AQIT 0.1.0
Loading...
Searching...
No Matches
model_diff.py
Go to the documentation of this file.
1# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2# This file is part of the Aquin Engine. Unauthorized copying, modification,
3# distribution, or use of this file, via any medium, is strictly prohibited.
4# Proprietary and confidential. See LICENSE for terms.
5
6"""
7Ingested from inspection-backend/server.py (_run_model_diff function).
8Extracted into its own module so train_simulate.py and inspect tools can import it
9without pulling in FastAPI or the full server.
10"""
11from __future__ import annotations
12
13import torch
14
15from aquin.compute.model_loader import get_config, load_model, resolve_model_id
16from aquin.compute.causal_trace import DTYPE, DEVICE
17from aquin.compute.device import empty_device_cache
18
19
21 base_model,
22 ft_model,
23 prompts: list[str],
24 n_prompts: int = 5,
25) -> dict:
26 """Compare base vs fine-tuned HookedTransformer models (already loaded)."""
27 _default_prompts = [
28 "The capital of France is",
29 "Water boils at",
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",
34 "DNA stands for",
35 "The largest planet in our solar system is",
36 ]
37 eval_prompts = (prompts if prompts else _default_prompts)[:n_prompts]
38
39 def _gen_output(model, prompt: str, max_new: int = 60) -> str:
40 try:
41 if hasattr(model, "hf_model"):
42 from aquin.compute.causal_trace import run_chat
43 return run_chat(
44 prompt,
45 max_new_tokens=max_new,
46 temperature=0.0,
47 model=model,
48 )
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()
54 except Exception:
55 return ""
56
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]
59
60 def _word_overlap(a: str, b: str) -> float:
61 wa, wb = set(a.lower().split()), set(b.lower().split())
62 if not wa and not wb:
63 return 1.0
64 return len(wa & wb) / max(len(wa | wb), 1)
65
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)
70
71 len_ratios = [
72 min(max(len(ft.split()), 1) / max(len(base.split()), 1), 2.0)
73 for base, ft in zip(base_outs, ft_outs)
74 ]
75 mean_len_ratio = sum(len_ratios) / max(len(len_ratios), 1)
76 suppression_score = round(min(mean_len_ratio, 1.0), 4)
77
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)
82 else:
83 robustness_score = consistency_score
84
85 _DOMAIN_KEYWORDS = {
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"],
90 }
91
92 def _classify_prompt(p: str) -> str:
93 pl = p.lower()
94 for domain, kws in _DOMAIN_KEYWORDS.items():
95 if any(kw in pl for kw in kws):
96 return domain
97 return "general"
98
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)
107
108 BASELINE = 0.75
109 category_deltas = []
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)
116
117 max_drift_idx = max(range(len(per_prompt_drift)), key=lambda i: per_prompt_drift[i]) if per_prompt_drift else 0
118 max_drift_prompt = {
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,
123 }
124
125 return {
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,
134 }
135
136
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
139
140 _HF_TO_KEY = {
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",
145 }
146 try:
147 model_key = resolve_model_id(base_model_id)
148 except ValueError:
149 model_key = _HF_TO_KEY.get(base_model_id)
150 if model_key is None:
151 raise ValueError(
152 f"[model-diff] unsupported base model '{base_model_id}'. "
153 f"Supported: {list(_HF_TO_KEY.keys())}"
154 )
155
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"]
161
162 try:
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)
167 del hf_ft
168 except Exception:
169 ft_model = HookedTransformer.from_pretrained(hf_name, dtype=DTYPE)
170 ft_model.eval()
171 ft_model.to(DEVICE)
172 missing, unexpected = ft_model.load_state_dict(state, strict=False)
173 if unexpected:
174 print(f"[model-diff] {len(unexpected)} unexpected keys: {unexpected[:3]}", flush=True)
175 if missing:
176 print(f"[model-diff] {len(missing)} missing keys — ft weights partially applied", flush=True)
177
178 ft_model.eval()
179 result = _run_model_diff_tl(base_model, ft_model, prompts, n_prompts=n_prompts)
180 del ft_model
181 empty_device_cache()
182 return result
dict _run_model_diff(str base_model_id, str ft_ckpt_path, list[str] prompts, int n_prompts=5)
dict _run_model_diff_tl(base_model, ft_model, list[str] prompts, int n_prompts=5)
Definition model_diff.py:29