Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 8 additions & 5 deletions demo.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,13 +66,16 @@ def print_verdict_table(records, cross_gpu_results=None):
print(f"{'MODEL':<8}{'CONDITION':<14}{'REPRODUCIBLE':<14}{'CROSS-GPU':<12}{'FIRST-DIVERGENCE':<18}")
print("-" * 66)
for r in records:
repro = "-" if r["reproducible"] is None else ("PASS" if r["reproducible"] else "FAIL")
if r.get("status") == "UNSUPPORTED":
repro = "UNSUPPORTED"
else:
repro = "-" if r["reproducible"] is None else ("PASS" if r["reproducible"] else "FAIL")
fd = r["first_divergence_step"]
first_div = "-" if fd is None else f"step {fd}"
xg = "-"
key = (r["model"], r.get("condition"))
if key in cross:
xg = "SAME" if cross[key] == r["param_sha256"] else "DIFF"
if r.get("status") != "UNSUPPORTED" and key in cross:
xg = "SAME" if cross[key] == r.get("param_sha256") else "DIFF"
Comment thread
coderabbitai[bot] marked this conversation as resolved.
print(f"{r['model']:<8}{r['condition']:<14}{repro:<14}{xg:<12}{first_div:<18}")
print("-" * 66)
if not cross:
Expand All @@ -84,8 +87,8 @@ def show_debate_hook(records):
section("THE VERIFICATION BAR -- bitwise identity, with loss-tolerance as diagnostic only")
print("verify() compares losses at rel_tol=1e-6 as a diagnostic, but the VERDICT is the")
print("exact parameter hash: a run can match every loss to 1e-6 and still FAIL --\n")
hooks = [r for r in records if r["vs_fp32_bitwise"] is False and r["vs_fp32_losstol"] is True]
diffs = [r for r in records if r["vs_fp32_bitwise"] is False]
hooks = [r for r in records if r.get("vs_fp32_bitwise") is False and r.get("vs_fp32_losstol") is True and r.get("final_loss") is not None]
diffs = [r for r in records if r.get("vs_fp32_bitwise") is False and r.get("final_loss") is not None]
shown = hooks or diffs
if not shown:
print(" (No cell diverged from fp32 on this device. On CPU, TF32 is a no-op and bf16")
Expand Down
148 changes: 123 additions & 25 deletions src/experiment.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,58 @@ def _norm_bool(value):
return str(value).strip().lower() in ("1", "on", "true", "yes", "y")


class UnsupportedKernelError(RuntimeError):
"""Raised when an op/device/precision combination is not supported by the hardware backend."""


def check_kernel_support(model_name, precision, dev):
"""Check for known unsupported hardware/precision/operator combinations."""
if dev.type == "cpu" and model_name == "lstm" and precision == "bf16":
try:
probe_lstm = torch.nn.LSTM(4, 4, batch_first=True)
probe_x = torch.zeros(1, 1, 4)
with torch.autocast("cpu", dtype=torch.bfloat16):
probe_lstm(probe_x)
except RuntimeError as e:
err_msg = str(e).lower()
if "primitive descriptor" in err_msg:
raise UnsupportedKernelError("oneDNN on CPU has no LSTM bf16 forward primitive") from e
raise


def _unsupported_record(model_name, dataset_name, precision, deterministic, seed, dev, cfg,
reason, quiet=False):
record = {
"model": model_name,
"dataset": dataset_name,
"precision": precision,
"deterministic": deterministic,
"device": dev.type,
"device_name": device_name(dev),
"status": "UNSUPPORTED",
"unsupported_reason": reason,
"final_loss": None,
"param_sha256": None,
"merkle_root": None,
"merkle_chunk_count": None,
"artifact_size_bytes": None,
"first_divergence_step": None,
"reproducible": None,
"num_params": None,
"seed": seed,
"total_steps": cfg.get("total_steps", 0),
"batch_size": cfg.get("batch_size", 0),
"block_size": cfg.get("block_size", 0),
"torch": torch.__version__,
"precision_flags": precision_flags(),
"wall_time_s": 0.0,
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),
}
if not quiet:
_print_cell(record)
return record


def prepare_run(seed, precision, deterministic, warn_only=True):
"""Seed everything, then set determinism and precision for one run."""
os.environ.setdefault("CUBLAS_WORKSPACE_CONFIG", ":4096:8")
Expand All @@ -86,23 +138,31 @@ def _single_train(model_name, dataset_name, precision, deterministic, seed, dev,

bs, blk = cfg["batch_size"], cfg["block_size"]
losses, step_hashes = [], []
for _step in range(cfg["total_steps"]):
# FIRST statement in the loop: the batch draw consumes the global torch
# RNG, so it is captured by torch.get_rng_state() and stays replay- and
# run-to-run exact. Moving this out of the loop silently breaks both.
x, y = ds.get_batch(bs, blk, device=dev)
with autocast:
logits = model(x)
if vision:
loss = F.cross_entropy(logits, y)
else:
loss = F.cross_entropy(logits.reshape(-1, logits.size(-1)), y.reshape(-1))
optimizer.zero_grad()
loss.backward()
optimizer.step()
losses.append(loss.item())
if track_full:
step_hashes.append(model_parameters_sha256(model))
try:
for _step in range(cfg["total_steps"]):
# FIRST statement in the loop: the batch draw consumes the global torch
# RNG, so it is captured by torch.get_rng_state() and stays replay- and
# run-to-run exact. Moving this out of the loop silently breaks both.
x, y = ds.get_batch(bs, blk, device=dev)
with autocast:
logits = model(x)
if vision:
loss = F.cross_entropy(logits, y)
else:
loss = F.cross_entropy(logits.reshape(-1, logits.size(-1)), y.reshape(-1))
optimizer.zero_grad()
loss.backward()
optimizer.step()
losses.append(loss.item())
if track_full:
step_hashes.append(model_parameters_sha256(model))
except RuntimeError as e:
err_msg = str(e).lower()
if dev.type == "cpu" and model_name == "lstm" and precision == "bf16" and (
"primitive descriptor" in err_msg
):
raise UnsupportedKernelError("oneDNN on CPU has no LSTM bf16 forward primitive") from e
raise

return model, losses, tensor_mapping_sha256(_stable_cpu_state_dict(model)), step_hashes

Expand Down Expand Up @@ -150,21 +210,49 @@ def run_one(model_name, dataset_name="shakespeare", precision="fp32",
if cfg.get("total_steps", 0) < 1:
raise ValueError(f"total_steps must be at least 1, got {cfg.get('total_steps')}")

try:
check_kernel_support(model_name, precision, dev)
except UnsupportedKernelError as err:
return _unsupported_record(
model_name, dataset_name, precision, deterministic, seed, dev, cfg,
reason=str(err), quiet=quiet
)

tag = f"{model_name}_{dataset_name}_{precision}_det{'on' if deterministic else 'off'}_s{seed}"

t0 = time.time()
modelA, lossesA, hashA, stepA = _single_train(
model_name, dataset_name, precision, deterministic, seed, dev, cfg,
track_full=track_full)
def _handle_run_error(err):
err_msg = str(err).lower()
if isinstance(err, UnsupportedKernelError) or (
dev.type == "cpu" and model_name == "lstm" and precision == "bf16" and (
"primitive descriptor" in err_msg
)
):
reason = "oneDNN on CPU has no LSTM bf16 forward primitive" if isinstance(err, RuntimeError) else str(err)
return _unsupported_record(
model_name, dataset_name, precision, deterministic, seed, dev, cfg,
reason=reason, quiet=quiet
)
raise err

try:
modelA, lossesA, hashA, stepA = _single_train(
model_name, dataset_name, precision, deterministic, seed, dev, cfg,
track_full=track_full)
except (UnsupportedKernelError, RuntimeError) as err:
return _handle_run_error(err)

reproducible, first_div = None, None
if twin:
# Twin run: identical settings, same seed, same hardware -> tests (A).
_modelB, lossesB, hashB, stepB = _single_train(
model_name, dataset_name, precision, deterministic, seed, dev, cfg,
track_full=track_full)
reproducible = (hashA == hashB) and (lossesA == lossesB)
first_div = _first_divergence(lossesA, lossesB, stepA, stepB)
try:
_modelB, lossesB, hashB, stepB = _single_train(
model_name, dataset_name, precision, deterministic, seed, dev, cfg,
track_full=track_full)
reproducible = (hashA == hashB) and (lossesA == lossesB)
first_div = _first_divergence(lossesA, lossesB, stepA, stepB)
except (UnsupportedKernelError, RuntimeError) as err:
return _handle_run_error(err)

merkle_root, chunk_count, size_bytes = _merkle_from_model(modelA, tag, keep_artifact)

Expand All @@ -175,6 +263,7 @@ def run_one(model_name, dataset_name="shakespeare", precision="fp32",
"deterministic": deterministic,
"device": dev.type,
"device_name": device_name(dev),
"status": "PASS",
"final_loss": lossesA[-1],
"param_sha256": hashA,
"merkle_root": merkle_root,
Expand Down Expand Up @@ -202,6 +291,15 @@ def run_one(model_name, dataset_name="shakespeare", precision="fp32",


def _print_cell(r):
if r.get("status") == "UNSUPPORTED":
reason = r.get("unsupported_reason", "unsupported configuration")
print(
f" [{r['model']:>6} | {r['dataset']:>11} | {r['precision']:>4} | "
f"det {'on ' if r['deterministic'] else 'off'} | {r['device']:>4}] "
f"UNSUPPORTED: {reason}"
)
return

repro = "PASS" if r["reproducible"] else ("FAIL" if r["reproducible"] is not None else "-")
fd = r["first_divergence_step"]
print(
Expand Down
39 changes: 27 additions & 12 deletions sweep.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,12 +46,12 @@ def annotate_reference(records):
"""Tag each cell with agreement vs the per-(model,dataset) fp32+det reference."""
refs = {}
for r in records:
if r["precision"] == "fp32" and r["deterministic"]:
if r["precision"] == "fp32" and r["deterministic"] and r.get("status") != "UNSUPPORTED":
refs[(r["model"], r["dataset"])] = r
for r in records:
ref = refs.get((r["model"], r["dataset"]))
r["is_reference"] = ref is not None and r is ref
if not ref:
if not ref or r.get("status") == "UNSUPPORTED" or r.get("final_loss") is None:
r["vs_fp32_bitwise"] = None
r["vs_fp32_losstol"] = None
continue
Expand All @@ -63,14 +63,18 @@ def annotate_reference(records):


def _fmt_repro(r):
if r.get("status") == "UNSUPPORTED":
return "UNSUP"
if r["reproducible"] is None:
return " - "
return "PASS " if r["reproducible"] else "FAIL "


def _fmt_vs(bitwise, losstol):
if bitwise is None:
def _fmt_vs(bitwise, losstol, is_ref=False):
if is_ref:
return "ref "
if bitwise is None:
return " - "
if bitwise:
return "SAME"
return "DIFF"
Expand All @@ -95,23 +99,34 @@ def print_grid(records):
for r in cells:
bitwise, losstol = r["vs_fp32_bitwise"], r["vs_fp32_losstol"]
fd = r["first_divergence_step"]
loss_tol_str = "ref " if losstol is None else ("SAME" if losstol else "DIFF")
if r.get("is_reference"):
loss_tol_str = "ref "
elif losstol is None:
loss_tol_str = " - "
else:
loss_tol_str = "SAME" if losstol else "DIFF"
star = ""
if bitwise is False and losstol is True:
star = " <-- loss agrees, bits differ (DEBATE)"
debate_cells.append(r)
loss_val = r.get("final_loss")
loss_str = f"{loss_val:<14.6f}" if loss_val is not None else f"{'-':<14}"
chunk_cnt = r.get("merkle_chunk_count")
chunk_str = f"{chunk_cnt} chunks" if chunk_cnt is not None else "-"
print(f"{model:<8}{r['condition']:<14}{_fmt_repro(r):<7}"
f"{('-' if fd is None else str(fd)):<11}"
f"{_fmt_vs(bitwise, losstol):<14}{loss_tol_str:<11}"
f"{r['final_loss']:<14.6f}"
f"{str(r['merkle_chunk_count']) + ' chunks':<10}{star}")
f"{_fmt_vs(bitwise, losstol, r.get('is_reference')):<14}{loss_tol_str:<11}"
f"{loss_str}"
f"{chunk_str:<10}{star}")

print("\n" + "=" * 100)
n_pass = sum(1 for r in records if r["reproducible"] is True)
n_fail = sum(1 for r in records if r["reproducible"] is False)
n_diff = sum(1 for r in records if r["vs_fp32_bitwise"] is False)
n_pass = sum(1 for r in records if r.get("reproducible") is True)
n_fail = sum(1 for r in records if r.get("reproducible") is False)
n_unsup = sum(1 for r in records if r.get("status") == "UNSUPPORTED")
n_diff = sum(1 for r in records if r.get("vs_fp32_bitwise") is False)
unsup_info = f" | unsupported: {n_unsup}" if n_unsup else ""
print(f"Cells: {len(records)} | run-to-run reproducible: {n_pass} | "
f"run-to-run broken: {n_fail} | bit-differ from fp32 reference: {n_diff}")
f"run-to-run broken: {n_fail}{unsup_info} | bit-differ from fp32 reference: {n_diff}")
if debate_cells:
print(f"Debate-hook cells (loss within 1e-6 of fp32 but NOT bitwise equal): "
f"{len(debate_cells)} -> "
Expand Down
79 changes: 78 additions & 1 deletion tests/test_experiment.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,9 @@
import unittest
from pathlib import Path

SRC = Path(__file__).resolve().parents[1] / "src"
ROOT = Path(__file__).resolve().parents[1]
SRC = ROOT / "src"
sys.path.insert(0, str(ROOT))
sys.path.insert(0, str(SRC))

import signing # noqa: E402 (torch-free)
Expand Down Expand Up @@ -323,6 +325,81 @@ def test_t8_matrix_has_reference_and_divergent_cell(self):
self.assertTrue(any(r["is_reference"] for r in recs))
self.assertTrue(any(r["vs_fp32_bitwise"] is False for r in recs))

def test_lstm_cpu_bf16_unsupported_handled_gracefully(self):
from unittest.mock import patch
from experiment import UnsupportedKernelError, run_one
from sweep import annotate_reference, print_grid, _fmt_repro
from demo import print_verdict_table

with patch("experiment.check_kernel_support", side_effect=UnsupportedKernelError("oneDNN on CPU has no LSTM bf16 forward primitive")):
rec = run_one("lstm", "shakespeare", "bf16", True, device="cpu", overrides=SMOKE, quiet=True)

self.assertEqual(rec["status"], "UNSUPPORTED")
self.assertIn("oneDNN on CPU has no LSTM bf16", rec["unsupported_reason"])
self.assertIsNone(rec["final_loss"])
self.assertIsNone(rec["reproducible"])
self.assertEqual(_fmt_repro(rec), "UNSUP")

# Ensure grid and verdict table formatting succeed without exceptions
ref = run_one("mlp", "shakespeare", "fp32", True, device="cpu", overrides=SMOKE, quiet=True)
ref["condition"] = "fp32 det-on"
rec["condition"] = "bf16 det-on"
records = [ref, rec]

annotate_reference(records)
self.assertIsNone(rec["vs_fp32_bitwise"])
self.assertIsNone(rec["vs_fp32_losstol"])

# Validate that printing doesn't throw formatting errors on None values
# and that UNSUPPORTED records do not evaluate cross-GPU agreement
import io
from contextlib import redirect_stdout
with tempfile.NamedTemporaryFile("w+", suffix=".jsonl") as tmp:
tmp.write(json.dumps({"model": "lstm", "condition": "bf16 det-on", "param_sha256": "fakehash"}) + "\n")
tmp.flush()
buf = io.StringIO()
with redirect_stdout(buf):
print_grid(records)
print_verdict_table(records, cross_gpu_results=tmp.name)
out = buf.getvalue()
self.assertIn("UNSUP", out)
self.assertIn("UNSUPPORTED", out)
for line in out.splitlines():
if line.startswith("lstm") and "bf16 det-on" in line:
self.assertIn("-", line)
self.assertNotIn("DIFF", line)
self.assertNotIn("SAME", line)

def test_lstm_cpu_bf16_runtime_error_converted_in_single_train(self):
from unittest.mock import patch
from experiment import run_one

with patch("experiment.check_kernel_support"):
with patch("experiment._single_train", side_effect=RuntimeError("could not create a primitive descriptor for the LSTM forward propagation primitive")):
rec = run_one("lstm", "shakespeare", "bf16", True, device="cpu", overrides=SMOKE, quiet=True)

self.assertEqual(rec["status"], "UNSUPPORTED")
self.assertIn("oneDNN on CPU has no LSTM bf16", rec["unsupported_reason"])
self.assertIsNone(rec["final_loss"])

def test_unrelated_lstm_runtime_error_not_caught_as_unsupported(self):
from unittest.mock import patch
from experiment import run_one

with patch("experiment.check_kernel_support"):
with patch("experiment._single_train", side_effect=RuntimeError("lstm generic internal error")):
with self.assertRaises(RuntimeError) as cm:
run_one("lstm", "shakespeare", "bf16", True, device="cpu", overrides=SMOKE, quiet=True)
self.assertIn("lstm generic internal error", str(cm.exception))

def test_check_kernel_support_non_cpu_or_non_lstm(self):
from experiment import check_kernel_support
# Non-LSTM on CPU + bf16 is supported
check_kernel_support("mlp", "bf16", torch.device("cpu"))
# LSTM on CPU + fp32 is supported
check_kernel_support("lstm", "fp32", torch.device("cpu"))



# --------------------------------------------------------------------------- #
# CUDA-only (run on the pod) -- the headline failure exhibits
Expand Down