26 """Compare base vs fine-tuned HookedTransformer models (already loaded)."""
28 "The capital of France is",
30 "The square root of 144 is",
31 "Albert Einstein was born in",
32 "The speed of light is approximately",
33 "The chemical formula for water is",
35 "The largest planet in our solar system is",
37 eval_prompts = (prompts
if prompts
else _default_prompts)[:n_prompts]
39 def _gen_output(model, prompt: str, max_new: int = 60) -> str:
41 if hasattr(model,
"hf_model"):
45 max_new_tokens=max_new,
49 tokens = model.generate(prompt, max_new_tokens=max_new, do_sample=
False)
50 if isinstance(tokens, str):
51 return tokens[len(prompt):].strip()
52 input_len = model.to_tokens(prompt).shape[1]
53 return model.tokenizer.decode(tokens[0][input_len:].tolist(), skip_special_tokens=
True).strip()
57 base_outs = [_gen_output(base_model, p)
for p
in eval_prompts]
58 ft_outs = [_gen_output(ft_model, p)
for p
in eval_prompts]
60 def _word_overlap(a: str, b: str) -> float:
61 wa, wb = set(a.lower().split()), set(b.lower().split())
64 return len(wa & wb) / max(len(wa | wb), 1)
66 per_prompt_overlap = [_word_overlap(a, b)
for a, b
in zip(base_outs, ft_outs)]
67 per_prompt_drift = [1.0 - s
for s
in per_prompt_overlap]
68 mean_drift = sum(per_prompt_drift) / max(len(per_prompt_drift), 1)
69 consistency_score = round(max(0.0, 1.0 - mean_drift), 4)
72 min(max(len(ft.split()), 1) / max(len(base.split()), 1), 2.0)
73 for base, ft
in zip(base_outs, ft_outs)
75 mean_len_ratio = sum(len_ratios) / max(len(len_ratios), 1)
76 suppression_score = round(min(mean_len_ratio, 1.0), 4)
78 if len(per_prompt_drift) > 1:
79 mean_d = sum(per_prompt_drift) / len(per_prompt_drift)
80 variance = sum((d - mean_d) ** 2
for d
in per_prompt_drift) / len(per_prompt_drift)
81 robustness_score = round(max(0.0, 1.0 - variance * 4), 4)
83 robustness_score = consistency_score
86 "factual": [
"capital",
"born",
"formula",
"stands for",
"invented",
"wrote",
"founded"],
87 "science": [
"boils",
"speed of light",
"square root",
"photosynthesis",
"dna",
"atom",
"planet",
"chemical"],
88 "reasoning": [
"if",
"therefore",
"implies",
"deduce",
"conclude",
"given that",
"factorial",
"function"],
89 "language": [
"write",
"explain",
"describe",
"list",
"name",
"define",
"summarise",
"translate"],
92 def _classify_prompt(p: str) -> str:
94 for domain, kws
in _DOMAIN_KEYWORDS.items():
95 if any(kw
in pl
for kw
in kws):
99 from collections
import defaultdict
as _dd
100 cat_data: dict = _dd(list)
101 for i, (base_o, ft_o)
in enumerate(zip(base_outs, ft_outs)):
102 cat = _classify_prompt(eval_prompts[i])
103 drift = per_prompt_drift[i]
104 length_r = len_ratios[i]
105 score = round((1.0 - drift) * 0.7 + min(length_r, 1.0) * 0.3, 4)
106 cat_data[cat].append(score)
110 for cat, scores
in cat_data.items():
111 mean_score = round(sum(scores) / len(scores), 4)
112 delta = round(mean_score - BASELINE, 4)
113 direction =
"improved" if delta > 0.05
else "degraded" if delta < -0.05
else "unchanged"
114 category_deltas.append({
"category": cat,
"score": mean_score,
"delta": delta,
"direction": direction,
"n": len(scores)})
115 category_deltas.sort(key=
lambda x: abs(x[
"delta"]), reverse=
True)
117 max_drift_idx = max(range(len(per_prompt_drift)), key=
lambda i: per_prompt_drift[i])
if per_prompt_drift
else 0
119 "prompt": eval_prompts[max_drift_idx]
if eval_prompts
else "",
120 "base": base_outs[max_drift_idx]
if base_outs
else "",
121 "ft": ft_outs[max_drift_idx]
if ft_outs
else "",
122 "drift": round(per_prompt_drift[max_drift_idx], 4)
if per_prompt_drift
else 0.0,
126 "consistencyScore": consistency_score,
127 "suppressionScore": suppression_score,
128 "robustnessScore": robustness_score,
129 "categoryDeltas": category_deltas,
130 "maxDriftPrompt": max_drift_prompt,
131 "baseOutputs": [{
"prompt": p,
"output": base_outs[i]}
for i, p
in enumerate(eval_prompts)],
132 "ftOutputs": [{
"prompt": p,
"output": ft_outs[i]}
for i, p
in enumerate(eval_prompts)],
133 "promptsUsed": eval_prompts,
137def _run_model_diff(base_model_id: str, ft_ckpt_path: str, prompts: list[str], n_prompts: int = 5) -> dict:
138 from transformer_lens
import HookedTransformer
141 "meta-llama/Llama-3.2-1B":
"llama-3.2-1b",
142 "meta-llama/Llama-3.2-1B-Instruct":
"llama-3.2-1b",
143 "meta-llama/Llama-3.2-3B":
"llama-3.2-3b",
144 "meta-llama/Llama-3.2-3B-Instruct":
"llama-3.2-3b",
147 model_key = resolve_model_id(base_model_id)
149 model_key = _HF_TO_KEY.get(base_model_id)
150 if model_key
is None:
152 f
"[model-diff] unsupported base model '{base_model_id}'. "
153 f
"Supported: {list(_HF_TO_KEY.keys())}"
156 base_model = load_model(model_key)
157 hf_name = get_config(model_key)[
"hf_name"]
158 state = torch.load(ft_ckpt_path, map_location=DEVICE, weights_only=
True)
159 if isinstance(state, dict)
and "model_state_dict" in state:
160 state = state[
"model_state_dict"]
163 from transformers
import AutoModelForCausalLM
164 hf_ft = AutoModelForCausalLM.from_pretrained(hf_name, torch_dtype=DTYPE, device_map=DEVICE)
165 hf_ft.load_state_dict(state, strict=
False)
166 ft_model = HookedTransformer.from_pretrained(hf_name, hf_model=hf_ft, dtype=DTYPE, device=DEVICE)
169 ft_model = HookedTransformer.from_pretrained(hf_name, dtype=DTYPE)
172 missing, unexpected = ft_model.load_state_dict(state, strict=
False)
174 print(f
"[model-diff] {len(unexpected)} unexpected keys: {unexpected[:3]}", flush=
True)
176 print(f
"[model-diff] {len(missing)} missing keys — ft weights partially applied", flush=
True)