traj-mc / code /smoke_gpu.py
ttishere's picture
Publish code/smoke_gpu.py
5e9c921 verified
Raw History Blame Contribute Delete
3 kB
"""
smoke_gpu.py -- Step 2 self-checks that need the real model on a GPU.
Single model load covers:
A. target-linear enumeration: all=224, attn=128, mlp=96 (32 blocks x 7).
B. lm_head hard defense passes on the real dense model.
C. 3-item REAL-eval schema check: eval path dumps valid per-item JSONL that
mcnemar.py can read (item_id unique, fields complete). Scores NOT inspected.
"""
import os
import sys
import json
import tempfile
import importlib.util
import torch
HERE = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, HERE)
import common as C
from analysis import mcnemar
spec = importlib.util.spec_from_file_location("tm_eval", os.path.join(HERE, "eval", "eval.py"))
E = importlib.util.module_from_spec(spec)
spec.loader.exec_module(E)
ok = True
# ── A. counts ─────────────────────────────────────────────────────────────────
print("[A] loading dense model...")
model, tok = C.load_model()
n_all = C.count_target_linears(model, "all")
n_attn = C.count_target_linears(model, "attn")
n_mlp = C.count_target_linears(model, "mlp")
a_ok = (n_all == 224 and n_attn == 128 and n_mlp == 96)
ok &= a_ok
print(f"[A] target linears all={n_all} attn={n_attn} mlp={n_mlp} "
f"(expect 224/128/96) -> {'PASS' if a_ok else 'FAIL'}")
# ── B. head guard on real model ───────────────────────────────────────────────
try:
C.assert_head_dense(model)
head = C.get_output_head(model)
b_ok = True
print(f"[B] head is {type(head).__name__} (dense) -> PASS")
except Exception as e:
b_ok = False
print(f"[B] head guard FAILED on dense model: {e}")
ok &= b_ok
# ── C. 3-item real-eval schema dump (MMLU, ref arm) ───────────────────────────
from datasets import load_dataset
rows = load_dataset("cais/mmlu", "all", split="test")
rows = [rows[i] for i in range(3)]
fd, path = tempfile.mkstemp(suffix="_smoke_items.jsonl")
with os.fdopen(fd, "w") as f:
c, cn, tot = E._run_mcq(model, tok, f, "mmlu", rows, E._mmlu_task,
mc_num=32, batch_size=16, cfg=0.5,
shard=0, num_shards=1, limit=3)
read = mcnemar.read_items(path)
required = {"item_id", "prompt_hash", "correct", "correct_norm", "pred", "gold"}
unique = len(read) == 3
complete = all(required <= set(v) for v in read.values())
typed = all(isinstance(v["correct"], bool) for v in read.values())
c_ok = unique and complete and typed and tot == 3
ok &= c_ok
print(f"[C] 3-item MMLU dump: n={tot} unique={unique} complete={complete} "
f"bool={typed} -> {'PASS' if c_ok else 'FAIL'}")
print(f" sample record: {json.dumps(list(read.values())[0])}")
os.unlink(path)
print(f"\n=== GPU SMOKE {'ALL PASS' if ok else 'FAIL'} ===")
sys.exit(0 if ok else 1)