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