Skip to content

Commit 28b212c

Browse files
committed
test(§TB1): z3 proof driver for batched B2 GEMM rewrite (3 positives + 2 non-vacuity negatives)
1 parent 8344356 commit 28b212c

1 file changed

Lines changed: 95 additions & 0 deletions

File tree

scratch/proof_b2_batched_driver.py

Lines changed: 95 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,95 @@
1+
"""z3/TLA proof driver for the batched B2 GEMM rewrite (§TB1).
2+
3+
Discharges the 3 GEMM-able batched contractions (dchunk_states transpose_A +
4+
decay-fold, dC_off dense, dC_diag lower-tri mask) at prod dims L=N=P=64,
5+
HEADS_PER_CTA=4, AND proves non-vacuity with two negative controls:
6+
(1) overlap-band tiling (m_stride < tile_m) -> single_writer MUST be False
7+
(2) transpose-bugged operand map -> operand_maps_match MUST be False
8+
RAISES (RULE #1) if any positive fails or any negative spuriously passes.
9+
"""
10+
import json
11+
12+
from cppmega_mlx.nn._tilelang import _gemm_rewrite_proof as grp
13+
14+
grp._msl_transform.ensure_libz3_preloaded()
15+
import z3 # noqa: E402
16+
17+
L = N = P = 64
18+
HPC = 4
19+
out = {"positives": {}, "negatives": {}}
20+
21+
# ---------------- POSITIVE: the 3 batched GEMM-able contractions ------------
22+
dchunk = grp.b2_batched_dchunk_contraction(z3, chunk_size=L, headdim=P, dstate=N)
23+
t_dchunk = grp.b2_batched_tiling(tile_m=P, tile_n=N, tile_k=L, heads_per_cta=HPC)
24+
p_dchunk = grp.require_gemm_rewrite_proof(dchunk, t_dchunk)
25+
26+
dcoff = grp.b2_batched_dcoff_contraction(z3, chunk_size=L, headdim=P, dstate=N)
27+
t_dcoff = grp.b2_batched_tiling(tile_m=L, tile_n=N, tile_k=P, heads_per_cta=HPC)
28+
p_dcoff = grp.require_gemm_rewrite_proof(dcoff, t_dcoff)
29+
30+
dcdiag = grp.b2_batched_dcdiag_contraction(z3, chunk_size=L, dstate=N)
31+
t_dcdiag = grp.b2_batched_tiling(tile_m=L, tile_n=N, tile_k=L, heads_per_cta=HPC)
32+
p_dcdiag = grp.require_gemm_rewrite_proof(dcdiag, t_dcdiag)
33+
34+
for nm, pr in (("dchunk", p_dchunk), ("dcoff", p_dcoff), ("dcdiag", p_dcdiag)):
35+
out["positives"][nm] = {
36+
"z3_used": pr.z3_used,
37+
"z3_proved": pr.z3_proved,
38+
"operand_maps_match": pr.operand_maps_match,
39+
"mask_equiv": pr.mask_equiv,
40+
"scale_equiv": pr.scale_equiv,
41+
"single_writer": pr.single_writer,
42+
"k_covered": pr.k_covered,
43+
}
44+
assert pr.z3_proved, f"POSITIVE {nm} not proved: {pr.reason}"
45+
assert pr.single_writer, f"POSITIVE {nm} single_writer False: {pr.reason}"
46+
47+
# ---------------- NEGATIVE 1: overlapping head-bands (race) -----------------
48+
# m_stride < tile_m => band b and b+1 share rows => single_writer MUST be False.
49+
t_overlap = grp.GemmTiling(
50+
tile_m=P, tile_n=N, tile_k=L, m_blocks=HPC, n_blocks=1, k_steps=1,
51+
m_stride=P // 2, n_stride=N,
52+
)
53+
p_overlap = grp.prove_gemm_rewrite(dchunk, t_overlap)
54+
out["negatives"]["overlap_band"] = {
55+
"z3_proved": p_overlap.z3_proved,
56+
"single_writer": p_overlap.single_writer,
57+
"reason": p_overlap.reason[:160],
58+
}
59+
assert not p_overlap.single_writer, "NEG overlap_band spuriously single_writer=True (VACUOUS)"
60+
assert not p_overlap.z3_proved, "NEG overlap_band spuriously z3_proved=True (VACUOUS)"
61+
62+
# ---------------- NEGATIVE 2: transpose-bugged operand map ------------------
63+
# Inject a transpose bug: make the GEMM A-address disagree with the serial
64+
# A-address (gemm reads k*M+i where serial reads i*K+k). operand_maps_match
65+
# MUST be False (the rewrite would compute the wrong thing).
66+
src = grp.b2_batched_dchunk_contraction(z3, chunk_size=L, headdim=P, dstate=N)
67+
orig_a_gemm = src.a_addr_gemm
68+
M_, K_ = src.m_extent, src.k_extent
69+
def bugged_a_gemm(i, k): # transposed flattened address (row<->col swap)
70+
return k * M_ + i
71+
bad = grp.GemmContraction(
72+
name=src.name + "_transpose_bug",
73+
m_extent=src.m_extent, n_extent=src.n_extent, k_extent=src.k_extent,
74+
a_addr_serial=src.a_addr_serial,
75+
a_addr_gemm=bugged_a_gemm,
76+
b_addr_serial=src.b_addr_serial,
77+
b_addr_gemm=src.b_addr_gemm,
78+
mask_serial=src.mask_serial,
79+
mask_gemm=src.mask_gemm,
80+
scale_serial=src.scale_serial,
81+
scale_gemm=src.scale_gemm,
82+
)
83+
p_bad = grp.prove_gemm_rewrite(bad, t_dchunk)
84+
out["negatives"]["transpose_bug"] = {
85+
"z3_proved": p_bad.z3_proved,
86+
"operand_maps_match": p_bad.operand_maps_match,
87+
"reason": p_bad.reason[:160],
88+
}
89+
assert not p_bad.operand_maps_match, "NEG transpose_bug spuriously operand_maps_match=True (VACUOUS)"
90+
assert not p_bad.z3_proved, "NEG transpose_bug spuriously z3_proved=True (VACUOUS)"
91+
92+
out["VERDICT"] = "ALL_POSITIVES_PROVED_AND_NON_VACUOUS"
93+
print("PROOF_RESULT_JSON_BEGIN")
94+
print(json.dumps(out, indent=2))
95+
print("PROOF_RESULT_JSON_END")

0 commit comments

Comments
 (0)