|
| 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