"""Every path the forward launcher can select, against the sequential reference. launch_impl picks among a generic, a tensor-core and a row-block streamed kernel for chunk_state and for chunk_scan, on gates over dtype, C*R, the mimo rank, the head dimension and the state dimension. The tensor-core kernels take the rank and the chunk length from template arguments and discard the runtime values, so a geometry that reaches the wrong instantiation does not fall back to a slower kernel, it contracts with the wrong rank and returns a plausible tensor. Each row therefore does two things: it asks the extension which path the geometry takes, through the same gates the launcher selects with, and it checks the result against mimo_step_ref. The grid is then required to cover every combination the launcher can produce, so a path cannot stop being exercised without the run failing. run.cmd tests/test_dispatch.py """ import sys from pathlib import Path import torch ROOT = Path(__file__).resolve().parent.parent sys.path.insert(0, str(ROOT)) sys.path.insert(0, str(ROOT / "ref")) sys.path.insert(0, str(ROOT / "tests")) import load_local as mamba3 # noqa: E402 from anchor import max_rel_err # noqa: E402 from mamba3_ref import mimo_step_ref # noqa: E402 from test_mamba3 import build, run_fwd # noqa: E402 DEV = "cuda" GEN, TC, STR = "generic", "tensor-core", "streamed" # label, (B, S, H, G, P, N, R, Na, C), dtype, expected (state, scan) GRID = [ # float32 never reaches the bf16 state kernel, but its scan contracts on # three-term tf32 whenever C*R == 64 and the rank and head dim are in the set ("f32 baseline", (1, 64, 4, 1, 64, 128, 4, 64, 16), torch.float32, (GEN, TC)), ("f32 grouped query", (1, 64, 8, 2, 32, 64, 4, 32, 16), torch.float32, (GEN, TC)), ("f32 ragged S", (1, 70, 4, 1, 64, 128, 4, 64, 16), torch.float32, (GEN, TC)), ("f32 headdim 128", (1, 64, 4, 1, 128, 128, 4, 64, 16), torch.float32, (GEN, TC)), # both generic: C*R != 64, or a head dim outside the instantiated set ("f32 C*R=12, no tf32", (1, 48, 4, 1, 64, 128, 1, 64, 12), torch.float32, (GEN, GEN)), ("f32 headdim 48", (1, 64, 4, 1, 48, 128, 4, 64, 16), torch.float32, (GEN, GEN)), # bf16 with a head dimension outside the instantiated set ("bf16 headdim 48", (1, 64, 4, 1, 48, 128, 4, 64, 16), torch.bfloat16, (GEN, GEN)), # N a multiple of 8 but not of 16: the scan gate holds, the state gate does not ("bf16 N=72", (1, 64, 4, 1, 64, 72, 4, 36, 16), torch.bfloat16, (GEN, TC)), # C*R a multiple of 16 but not 64: state on tensor cores, scan generic ("bf16 C*R=32", (1, 64, 4, 1, 64, 128, 2, 64, 16), torch.bfloat16, (TC, GEN)), ("bf16 C*R=16", (1, 64, 4, 1, 64, 128, 1, 64, 16), torch.bfloat16, (TC, GEN)), # C*R = 64, the whole instantiated rank set and every head dimension ("bf16 R=1 C=64", (1, 64, 4, 1, 64, 128, 1, 64, 64), torch.bfloat16, (TC, TC)), ("bf16 R=2 C=32", (1, 64, 4, 1, 64, 128, 2, 64, 32), torch.bfloat16, (TC, TC)), ("bf16 R=4 C=16", (1, 64, 4, 1, 64, 128, 4, 64, 16), torch.bfloat16, (TC, TC)), ("bf16 R=8 C=8", (1, 64, 4, 1, 64, 128, 8, 64, 8), torch.bfloat16, (TC, TC)), ("bf16 headdim 16", (1, 64, 4, 1, 16, 128, 4, 64, 16), torch.bfloat16, (TC, TC)), ("bf16 headdim 32", (1, 64, 4, 1, 32, 128, 4, 64, 16), torch.bfloat16, (TC, TC)), ("bf16 headdim 128", (1, 64, 4, 1, 128, 128, 4, 64, 16), torch.bfloat16, (TC, TC)), # above C*R = 64 both stages stream, in either dtype ("f32 streamed", (1, 128, 4, 1, 64, 128, 4, 64, 32), torch.float32, (STR, STR)), ("bf16 streamed", (1, 128, 4, 1, 64, 128, 4, 64, 32), torch.bfloat16, (STR, STR)), ("bf16 streamed C*R=256", (1, 128, 4, 1, 64, 128, 8, 64, 32), torch.bfloat16, (STR, STR)), ] # Ranks the tensor-core scan has no instantiation for. C*R = 64 satisfies the # rest of its gate, so before the rank predicate was added these dispatched the # rank-8 kernel and returned garbage; they must be refused, not run. REFUSED = [(16, 4), (32, 2), (64, 1)] fails = [] covered = set() tf32_seen = set() def check(label, ok, detail=""): print(f" {label:<26} {'ok' if ok else 'FAIL':<5} {detail}") if not ok: fails.append(label) print(f"torch {torch.__version__} {torch.cuda.get_device_name(0)}\n") print("1. each dispatched path against mimo_step_ref") print(f" {'case':<26} {'state':<12} {'scan':<12} {'tf32':<5} {'rel':>10}") for label, geom, dtype, want in GRID: B, S, H, G, P, N, R, Na, C = geom is_bf = dtype is torch.bfloat16 d = mamba3.dispatch_paths(R, N, P, C, is_bf) state, scan, tf32 = d["state"], d["scan"], d["state_tf32"] covered.add((state, scan)) if state == GEN: tf32_seen.add(tf32) c32 = build(B, S, H, G, P, N, R, Na, DEV, seed=0) cin = build(B, S, H, G, P, N, R, Na, DEV, seed=0, dtype=dtype) if is_bf else c32 ref, _ = mimo_step_ref(c32["q"], c32["k"], c32["v"], c32["adt"], c32["dt"], c32["trap"], c32["q_bias"], c32["k_bias"], c32["angles"], c32["mimo_v"], c32["mimo_o"], D=c32["D"], z=c32["z"], mimo_z=c32["mimo_z"]) rel = max_rel_err(run_fwd(cin, C), ref) tol = 3e-2 if is_bf else 5e-5 print(f" {label:<26} {state:<12} {scan:<12} {str(tf32):<5} {rel:>10.3e}") if (state, scan) != want: check(label, False, f"took {(state, scan)}, grid says {want}") elif rel >= tol: check(label, False, f"rel {rel:.3e} against tol {tol:.0e}") print("\n2. ranks with no tensor-core instantiation are refused, not dispatched") for R, C in REFUSED: B, S, H, G, P, N, Na = 1, 64, 4, 1, 64, 128, 64 c = build(B, S, H, G, P, N, R, Na, DEV, seed=0, dtype=torch.bfloat16) try: run_fwd(c, C) torch.cuda.synchronize() check(f"R={R} C={C} C*R={C*R}", False, "ran instead of raising") except RuntimeError as exc: msg = str(exc) check(f"R={R} C={C} C*R={C*R}", "mimo_rank" in msg, msg.splitlines()[0][:56]) print("\n3. the grid covers every path the launcher can select") REACHABLE = {(GEN, GEN), (GEN, TC), (TC, GEN), (TC, TC), (STR, STR)} missing = REACHABLE - covered check("every (state, scan) pair exercised", not missing, "all five" if not missing else f"missing {sorted(missing)}") check("generic state seen with and without tf32", tf32_seen == {False, True}, f"tf32 flags {sorted(tf32_seen)}") print("\n" + ("PASS" if not fails else f"FAIL ({len(fails)}): " + "; ".join(fails))) sys.exit(0 if not fails else 1)