fix(muh_dispatch): add verify_against_headers() to close the loop
Adds verification that hand-written values in muh_dispatch.py (reduce_threads=512, reduce_items=16, etc.) match the C++ headers (bi100_float32_plus_o4 in tuning_reduce.cuh). Previously: muh_dispatch.py had hand-coded values with no link to the C++ source of truth. gen_patch.py reads from C++ headers, but muh_dispatch.py was a separate copy that could diverge. Now: verify_against_headers() calls gen_patch.extract_bi100_structs() and compares. Self-test prints mismatches if any exist.
This commit is contained in:
@@ -224,6 +224,60 @@ def qwen36_config() -> AttentionConfig:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
# ============================================================
|
||||||
|
# Verification: check hand-written values match C++ headers
|
||||||
|
# Closes the loop: muh_dispatch.py values must come from tuning_*.cuh
|
||||||
|
# ============================================================
|
||||||
|
|
||||||
|
def verify_against_headers(header_dir: str = "muh/include/muh/tuning") -> list:
|
||||||
|
"""Verify that muh_dispatch.py's hand-written values match C++ headers.
|
||||||
|
|
||||||
|
Returns list of mismatches. Empty list = all values verified.
|
||||||
|
This is the closed-loop check that prevents the gen_patch pipeline
|
||||||
|
from diverging from the runtime dispatch values.
|
||||||
|
"""
|
||||||
|
mismatches = []
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Import from sibling module
|
||||||
|
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||||
|
from gen_patch import extract_bi100_structs
|
||||||
|
except ImportError:
|
||||||
|
return [("IMPORT_ERROR", "Cannot import gen_patch.extract_bi100_structs")]
|
||||||
|
|
||||||
|
header_map = {
|
||||||
|
"tuning_reduce.cuh": {
|
||||||
|
"bi100_float32_plus_o4": {"items": 16, "threads": 512}, # must match reduce_items, reduce_threads
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
reduce_header = os.path.join(header_dir, "tuning_reduce.cuh")
|
||||||
|
if os.path.exists(reduce_header):
|
||||||
|
structs = extract_bi100_structs(reduce_header)
|
||||||
|
for name, fields in structs:
|
||||||
|
if name == "bi100_float32_plus_o4":
|
||||||
|
cpp_threads = fields.get("threads", fields.get("threads_per_block"))
|
||||||
|
cpp_items = fields.get("items", fields.get("items_per_thread"))
|
||||||
|
|
||||||
|
# Compare against AttentionConfig defaults
|
||||||
|
cfg = qwen36_config()
|
||||||
|
if cfg.reduce_threads != cpp_threads:
|
||||||
|
mismatches.append((
|
||||||
|
"reduce_threads",
|
||||||
|
f"muh_dispatch={cfg.reduce_threads} vs tuning_reduce.cuh={cpp_threads}"
|
||||||
|
))
|
||||||
|
if cfg.reduce_items != cpp_items:
|
||||||
|
mismatches.append((
|
||||||
|
"reduce_items",
|
||||||
|
f"muh_dispatch={cfg.reduce_items} vs tuning_reduce.cuh={cpp_items}"
|
||||||
|
))
|
||||||
|
else:
|
||||||
|
mismatches.append(("HEADER_MISSING", reduce_header))
|
||||||
|
|
||||||
|
return mismatches
|
||||||
|
|
||||||
# ============================================================
|
# ============================================================
|
||||||
# Self-test
|
# Self-test
|
||||||
# ============================================================
|
# ============================================================
|
||||||
@@ -246,3 +300,13 @@ if __name__ == "__main__":
|
|||||||
print(f" decode: partition={cfg.partition_size} v1_thresh={cfg.v1_v2_threshold}")
|
print(f" decode: partition={cfg.partition_size} v1_thresh={cfg.v1_v2_threshold}")
|
||||||
print(f" reduce: threads={cfg.reduce_threads} items={cfg.reduce_items} vec={cfg.vec_size}")
|
print(f" reduce: threads={cfg.reduce_threads} items={cfg.reduce_items} vec={cfg.vec_size}")
|
||||||
print()
|
print()
|
||||||
|
# Verification: check values match C++ headers
|
||||||
|
print("\n=== Verification against C++ headers ===")
|
||||||
|
mismatches = verify_against_headers()
|
||||||
|
if mismatches:
|
||||||
|
for field, msg in mismatches:
|
||||||
|
print(f" ✗ MISMATCH {field}: {msg}")
|
||||||
|
print(f"\n {len(mismatches)} mismatches found — update muh_dispatch.py!")
|
||||||
|
else:
|
||||||
|
print(" ✓ All hand-written values match C++ headers")
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user