Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions strategy/runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@

import os
import time
from numbers import Integral

import numpy as np

from .backend import Backend
Expand Down Expand Up @@ -51,6 +53,12 @@ def run(n: int, cfg: Config, fill: str = "lowrank", verify: bool = False,
Default fill is 'lowrank' because that is the regime the strategy targets;
verify (small n) reports the reconstruction error vs a float64 reference.
"""
# Reject a degenerate n before touching the GPU, mirroring matmul.runner.run:
# storage's block-size math divides by n, so n=0 raised ZeroDivisionError and
# n<0 a numpy "negative dimensions" error -- both only after a Backend (and a
# real GPU) had already been demanded.
if isinstance(n, bool) or not isinstance(n, Integral) or n < 1:
raise ValueError(f"n must be a positive integer, got {n!r}")
backend = Backend(cfg.device, cfg.verbose)
dt = cfg.np_dtype
on_disk = storage.should_use_disk(
Expand Down Expand Up @@ -139,6 +147,8 @@ def compare(n: int, cfg: Config, fill: str = "lowrank",
spectral_alpha: float = 1.0) -> dict:
"""Run the exact baseline and the subspace strategy on the SAME inputs and
report time, throughput and the smart strategy's reconstruction error."""
if isinstance(n, bool) or not isinstance(n, Integral) or n < 1:
raise ValueError(f"n must be a positive integer, got {n!r}")
backend = Backend(cfg.device, cfg.verbose)
dt = cfg.np_dtype
on_disk = storage.should_use_disk(
Expand Down
63 changes: 63 additions & 0 deletions tests/test_runner_n_parity.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,63 @@
"""matmul and strategy must agree on the runner `n` boundary.

The two packages are deliberately standalone copies of the same orchestration,
so a boundary check added to one silently drifts out of the other. `matmul`
rejected a degenerate n up front while `strategy` accepted it and only failed
later -- after a Backend (a real GPU) had already been demanded, and then with
a raw ZeroDivisionError from storage's `// (n * 8)` block-size math.

These are parity tests: each public runner entry point is exercised through the
SAME table, so the two packages cannot drift apart again. A stub Backend that
raises on construction proves the check runs *before* any device work, which is
what lets this run on CPU with no GPU/torch.
"""
from __future__ import annotations

import pytest

import matmul.runner as matmul_runner
import strategy.runner as strategy_runner
from matmul.config import Config as MatmulConfig
from strategy.config import Config as StrategyConfig

# n values that are not a positive integer. `True` is included because bool is a
# subclass of int -- matmul.run(True) must not be read as n=1.
BAD_N = (0, -1, 1.5, True)

# (label, module, callable-name, config-factory) for every runner entry point.
ENTRY_POINTS = [
("matmul.run", matmul_runner, "run", lambda: MatmulConfig(verbose=False)),
("strategy.run", strategy_runner, "run", lambda: StrategyConfig(verbose=False)),
("strategy.compare", strategy_runner, "compare", lambda: StrategyConfig(verbose=False)),
]


class _ExplodingBackend:
"""Stands in for Backend so any device work fails loudly instead of being
skipped for lack of a GPU."""

def __init__(self, *args, **kwargs):
raise AssertionError("n must be validated before Backend construction")


@pytest.mark.parametrize("label,module,fn_name,make_cfg", ENTRY_POINTS,
ids=[e[0] for e in ENTRY_POINTS])
@pytest.mark.parametrize("n", BAD_N, ids=[repr(v) for v in BAD_N])
def test_runner_rejects_bad_n_before_backend(monkeypatch, label, module, fn_name, make_cfg, n):
monkeypatch.setattr(module, "Backend", _ExplodingBackend)
with pytest.raises(ValueError, match="n must be a positive integer"):
getattr(module, fn_name)(n, make_cfg())


@pytest.mark.parametrize("label,module,fn_name,make_cfg", ENTRY_POINTS,
ids=[e[0] for e in ENTRY_POINTS])
def test_runner_accepts_smallest_valid_n(monkeypatch, label, module, fn_name, make_cfg):
"""n=1 is valid, so it must get past the check and reach Backend (our stub),
guarding the boundary against an off-by-one."""
monkeypatch.setattr(module, "Backend", _ExplodingBackend)
with pytest.raises(AssertionError, match="before Backend construction"):
getattr(module, fn_name)(1, make_cfg())


if __name__ == "__main__":
raise SystemExit(pytest.main([__file__, "-q"]))
Loading