From 6f0e605c8bd696e3aab691587251532ef8005676 Mon Sep 17 00:00:00 2001 From: Hrishikesh Yadav Date: Sat, 3 Oct 2026 22:23:11 +0530 Subject: [PATCH 1/3] fix(cpu): handle missing oneDNN LSTM bf16 primitive on CPU gracefully --- demo.py | 11 +-- src/experiment.py | 148 ++++++++++++++++++++++++++++++++------- sweep.py | 39 +++++++---- tests/test_experiment.py | 56 +++++++++++++++ 4 files changed, 213 insertions(+), 41 deletions(-) diff --git a/demo.py b/demo.py index 71e2f5e4..12f98f6b 100644 --- a/demo.py +++ b/demo.py @@ -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" + xg = "SAME" if cross[key] == r.get("param_sha256") else "DIFF" print(f"{r['model']:<8}{r['condition']:<14}{repro:<14}{xg:<12}{first_div:<18}") print("-" * 66) if not cross: @@ -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") diff --git a/src/experiment.py b/src/experiment.py index 96d80f92..b398c590 100644 --- a/src/experiment.py +++ b/src/experiment.py @@ -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 or "onednn" in err_msg or "lstm" 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") @@ -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 or "onednn" in err_msg or "lstm" 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 @@ -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 or "onednn" in err_msg or "lstm" 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) @@ -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, @@ -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( diff --git a/sweep.py b/sweep.py index 849b595d..058f2a10 100644 --- a/sweep.py +++ b/sweep.py @@ -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 @@ -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" @@ -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)} -> " diff --git a/tests/test_experiment.py b/tests/test_experiment.py index b9f659f0..960b6111 100644 --- a/tests/test_experiment.py +++ b/tests/test_experiment.py @@ -323,6 +323,62 @@ 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 + import io + from contextlib import redirect_stdout + buf = io.StringIO() + with redirect_stdout(buf): + print_grid(records) + print_verdict_table(records) + out = buf.getvalue() + self.assertIn("UNSUP", out) + self.assertIn("UNSUPPORTED", out) + + 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_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 From 825feebaaab854e1847760331da2747be1eb20b6 Mon Sep 17 00:00:00 2001 From: Hrishikesh Yadav Date: Sat, 3 Oct 2026 22:32:13 +0530 Subject: [PATCH 2/3] fix(demo,experiment): address CodeRabbit review feedback on error classification and cross-gpu display --- demo.py | 2 +- src/experiment.py | 6 +++--- tests/test_experiment.py | 27 +++++++++++++++++++++++---- 3 files changed, 27 insertions(+), 8 deletions(-) diff --git a/demo.py b/demo.py index 12f98f6b..d18f088e 100644 --- a/demo.py +++ b/demo.py @@ -74,7 +74,7 @@ def print_verdict_table(records, cross_gpu_results=None): first_div = "-" if fd is None else f"step {fd}" xg = "-" key = (r["model"], r.get("condition")) - if key in cross: + if r.get("status") != "UNSUPPORTED" and key in cross: xg = "SAME" if cross[key] == r.get("param_sha256") else "DIFF" print(f"{r['model']:<8}{r['condition']:<14}{repro:<14}{xg:<12}{first_div:<18}") print("-" * 66) diff --git a/src/experiment.py b/src/experiment.py index b398c590..4f3606fc 100644 --- a/src/experiment.py +++ b/src/experiment.py @@ -76,7 +76,7 @@ def check_kernel_support(model_name, precision, dev): probe_lstm(probe_x) except RuntimeError as e: err_msg = str(e).lower() - if "primitive descriptor" in err_msg or "onednn" in err_msg or "lstm" in err_msg: + if "primitive descriptor" in err_msg: raise UnsupportedKernelError("oneDNN on CPU has no LSTM bf16 forward primitive") from e raise @@ -159,7 +159,7 @@ def _single_train(model_name, dataset_name, precision, deterministic, seed, dev, 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 or "onednn" in err_msg or "lstm" in err_msg + "primitive descriptor" in err_msg ): raise UnsupportedKernelError("oneDNN on CPU has no LSTM bf16 forward primitive") from e raise @@ -225,7 +225,7 @@ 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 or "onednn" in err_msg or "lstm" in err_msg + "primitive descriptor" in err_msg ) ): reason = "oneDNN on CPU has no LSTM bf16 forward primitive" if isinstance(err, RuntimeError) else str(err) diff --git a/tests/test_experiment.py b/tests/test_experiment.py index 960b6111..1fffd736 100644 --- a/tests/test_experiment.py +++ b/tests/test_experiment.py @@ -349,15 +349,24 @@ def test_lstm_cpu_bf16_unsupported_handled_gracefully(self): 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 - buf = io.StringIO() - with redirect_stdout(buf): - print_grid(records) - print_verdict_table(records) + 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 @@ -371,6 +380,16 @@ def test_lstm_cpu_bf16_runtime_error_converted_in_single_train(self): 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 From fdfed62bf11e2ce8e490d44811c5c86af659ed66 Mon Sep 17 00:00:00 2001 From: Hrishikesh Yadav Date: Sat, 3 Oct 2026 23:48:08 +0530 Subject: [PATCH 3/3] test: ensure repository root is in sys.path for test runners --- tests/test_experiment.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tests/test_experiment.py b/tests/test_experiment.py index 1fffd736..04282d55 100644 --- a/tests/test_experiment.py +++ b/tests/test_experiment.py @@ -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)