| """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 |
| from anchor import max_rel_err |
| from mamba3_ref import mimo_step_ref |
| from test_mamba3 import build, run_fwd |
|
|
| DEV = "cuda" |
| GEN, TC, STR = "generic", "tensor-core", "streamed" |
|
|
| |
| GRID = [ |
| |
| |
| ("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)), |
| |
| ("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 headdim 48", (1, 64, 4, 1, 48, 128, 4, 64, 16), torch.bfloat16, (GEN, GEN)), |
| |
| ("bf16 N=72", (1, 64, 4, 1, 64, 72, 4, 36, 16), torch.bfloat16, (GEN, TC)), |
| |
| ("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)), |
| |
| ("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)), |
| |
| ("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)), |
| ] |
|
|
| |
| |
| |
| 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) |
|
|