Raniahossam33/knowledge-drift-experiments
022
1#!/usr/bin/env python32"""3extract_models.py — Standardized Hidden State Extraction4=========================================================5Extracts last-token hidden states from any model in models.yaml.6Produces a standardized cache (.npz) with identical schema across all models.7 8Features:9 - Per-model chat template handling via tokenizer.apply_chat_template()10 - Saves lm_head + layer_norm weights for logit lens11 - Incremental checkpointing every 500 samples12 - Skip if cache already exists (--skip_if_cached)13 - Reuse existing Qwen2.5 cache with field remapping14 15Usage:16 # Extract Qwen2.5 (reuses existing cache)17 python extract_models.py --model qwen25 --skip_if_cached18 19 # Extract LLaMA-3.1 (fresh extraction, ~2-3h on 1 GPU)20 CUDA_VISIBLE_DEVICES=0 python extract_models.py --model llama3121 22 # Extract all models sequentially23 python extract_models.py --all --skip_if_cached24 25 # Migrate existing Qwen2.5 cache to v4 format26 python extract_models.py --model qwen25 --migrate_from data/experiments/v3/qwen25/cached_qwen25.npz27"""28 29import argparse30import json31import logging32import os33import sys34import time35import warnings36from datetime import datetime37from pathlib import Path38 39import numpy as np40import torch41import yaml42 43warnings.filterwarnings("ignore")44logging.basicConfig(45 level=logging.INFO,46 format="%(asctime)s [%(levelname)s] %(message)s",47 handlers=[logging.StreamHandler()])48logger = logging.getLogger(__name__)49 50 51# ─────────────────────────────────────────────────────────────────────────────52# CONFIG53# ─────────────────────────────────────────────────────────────────────────────54 55def load_config(config_path="models.yaml"):56 with open(config_path) as f:57 cfg = yaml.safe_load(f)58 return cfg59 60 61def get_model_cfg(cfg, model_key):62 if model_key not in cfg["models"]:63 raise ValueError(f"Unknown model '{model_key}'. "64 f"Available: {list(cfg['models'].keys())}")65 mcfg = cfg["models"][model_key]66 mcfg["key"] = model_key67 return mcfg68 69 70# ─────────────────────────────────────────────────────────────────────────────71# DATASET LOADING72# ─────────────────────────────────────────────────────────────────────────────73 74def load_dataset(dataset_path, model_key, drift_key):75 """Load and prepare dataset with model-specific drift labels."""76 logger.info(f"Loading dataset: {dataset_path}")77 with open(dataset_path) as f:78 raw = json.load(f)79 samples = raw.get("samples", raw)80 81 # Assign is_drifted from model-specific column82 for s in samples:83 val = s.get(drift_key, s.get("is_drifted", False))84 if isinstance(val, str):85 s["is_drifted"] = val.lower() in ("true", "1", "yes")86 else:87 s["is_drifted"] = bool(val)88 89 n_d = sum(1 for s in samples if s["is_drifted"])90 n_s = len(samples) - n_d91 logger.info(f" Total={len(samples)} Drifted={n_d} Stable={n_s}")92 93 if n_d == 0:94 logger.error(f"No drifted samples for drift_key='{drift_key}'. Aborting.")95 sys.exit(1)96 97 return samples98 99 100# ─────────────────────────────────────────────────────────────────────────────101# MODEL LOADING102# ─────────────────────────────────────────────────────────────────────────────103 104def load_model(model_name, device="auto"):105 from transformers import AutoModelForCausalLM, AutoTokenizer106 107 logger.info(f"Loading model: {model_name}")108 tok = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)109 if tok.pad_token is None:110 tok.pad_token = tok.eos_token111 112 mdl = AutoModelForCausalLM.from_pretrained(113 model_name,114 device_map=device,115 trust_remote_code=True,116 output_hidden_states=True,117 torch_dtype=torch.float16,118 )119 mdl.eval()120 n_layers = mdl.config.num_hidden_layers121 h_dim = mdl.config.hidden_size122 logger.info(f" Loaded: L={n_layers} D={h_dim}")123 return mdl, tok124 125 126# ─────────────────────────────────────────────────────────────────────────────127# LM HEAD EXTRACTION (for logit lens)128# ─────────────────────────────────────────────────────────────────────────────129 130def save_lm_head(model, out_dir, model_key):131 """Save lm_head weight + final layernorm for logit lens analysis."""132 lm_path = Path(out_dir) / f"lm_head_{model_key}.npz"133 if lm_path.exists():134 logger.info(f" lm_head already saved: {lm_path}")135 return136 137 lm_w = model.lm_head.weight.detach().float().cpu().numpy()138 139 # Find final layer norm — different names across architectures140 ln = None141 for attr in ["norm", "final_layernorm", "model.norm",142 "model.final_layernorm", "ln_f"]:143 parts = attr.split(".")144 obj = model145 try:146 for p in parts:147 obj = getattr(obj, p)148 ln = obj149 break150 except AttributeError:151 continue152 153 if ln is not None and hasattr(ln, "weight"):154 ln_w = ln.weight.detach().float().cpu().numpy()155 ln_b = (ln.bias.detach().float().cpu().numpy()156 if hasattr(ln, "bias") and ln.bias is not None157 else np.zeros_like(ln_w))158 else:159 logger.warning(" Could not find final layernorm — using identity")160 ln_w = np.ones(model.config.hidden_size, dtype=np.float32)161 ln_b = np.zeros(model.config.hidden_size, dtype=np.float32)162 163 np.savez_compressed(str(lm_path), lm_head=lm_w, ln_weight=ln_w, ln_bias=ln_b)164 logger.info(f" lm_head saved: {lm_path} (shape={lm_w.shape})")165 166 167# ─────────────────────────────────────────────────────────────────────────────168# TOKENIZATION (handles chat templates properly)169# ─────────────────────────────────────────────────────────────────────────────170 171def tokenize_query(tokenizer, query, model_cfg, max_length=512):172 """Tokenize query using proper chat template for each model."""173 is_instruct = model_cfg.get("is_instruct", True)174 175 if is_instruct and hasattr(tokenizer, "apply_chat_template"):176 try:177 messages = [{"role": "user", "content": query}]178 input_ids = tokenizer.apply_chat_template(179 messages, tokenize=True, add_generation_prompt=True,180 return_tensors="pt", max_length=max_length, truncation=True)181 attention_mask = torch.ones_like(input_ids)182 return {"input_ids": input_ids, "attention_mask": attention_mask}183 except Exception:184 pass # Fall through to raw tokenization185 186 # Raw tokenization for base models or if chat template fails187 return tokenizer(query, return_tensors="pt",188 truncation=True, max_length=max_length)189 190 191# ─────────────────────────────────────────────────────────────────────────────192# EXTRACTION193# ─────────────────────────────────────────────────────────────────────────────194 195def extract_single(model, tokenizer, sample, model_cfg, max_length=512):196 """Extract hidden states + logit info for a single sample."""197 query = sample.get("query", sample.get("question", ""))198 answer = sample.get("expected_answer", sample.get("answer", ""))199 200 inp = tokenize_query(tokenizer, query, model_cfg, max_length)201 inp = {k: v.to(model.device) for k, v in inp.items()}202 203 with torch.no_grad():204 out = model(**inp)205 206 n_layers = model.config.num_hidden_layers207 208 # Last-token hidden states at every layer209 hidden_states = {}210 for l in range(n_layers):211 h = out.hidden_states[l + 1][0, -1, :].float().cpu()212 h = torch.clamp(h, -1e6, 1e6)213 h[torch.isnan(h)] = 0.0214 hidden_states[l] = h.numpy()215 216 # Output logits217 logits = out.logits[0, -1, :].float().cpu()218 logits = torch.clamp(logits, -1e4, 1e4)219 logits[torch.isnan(logits)] = 0.0220 probs = torch.softmax(logits, dim=-1)221 222 top_prob = probs.max().item()223 top_idx = probs.argmax().item()224 top_token = tokenizer.decode([top_idx]).strip()225 entropy = -(probs * torch.log(probs + 1e-12)).sum().item()226 227 # Correctness: fuzzy match228 ans_lower = answer.lower().strip()229 tok_lower = top_token.lower().strip()230 correct = (ans_lower in tok_lower or tok_lower in ans_lower or231 any(w in tok_lower for w in ans_lower.split()[:3] if len(w) > 3))232 233 return {234 "hidden_states": hidden_states,235 "top_prob": top_prob,236 "top_token": top_token,237 "entropy": entropy,238 "correct": correct,239 }240 241 242def run_extraction(model, tokenizer, samples, model_cfg, out_dir, model_key,243 max_length=512, checkpoint_every=500):244 """Full extraction loop with incremental checkpointing."""245 out_dir = Path(out_dir)246 n_layers = model.config.num_hidden_layers247 t0 = time.time()248 249 # Save lm_head for logit lens250 save_lm_head(model, out_dir, model_key)251 252 results = []253 for idx, s in enumerate(samples):254 try:255 ext = extract_single(model, tokenizer, s, model_cfg, max_length)256 except Exception as e:257 logger.error(f" Sample {idx} error: {e}")258 continue259 260 result = {261 "idx": idx,262 "sample_id": s.get("sample_id", f"s_{idx}"),263 "query": s.get("query", ""),264 "expected_answer": s.get("expected_answer", ""),265 "is_drifted": s["is_drifted"],266 "relation": s.get("relation", "unknown"),267 "category": s.get("category", "unknown"),268 "entity": s.get("entity", ""),269 "knowledge_type": s.get("knowledge_type", ""),270 "drift_date": s.get("drift_date", ""),271 "year": s.get("year", ""),272 "dataset_source": s.get("dataset_source", ""),273 "hidden_states": ext["hidden_states"],274 "top_prob": ext["top_prob"],275 "top_token": ext["top_token"],276 "entropy": ext["entropy"],277 "correct": ext["correct"],278 }279 results.append(result)280 281 if (idx + 1) % 100 == 0:282 elapsed = time.time() - t0283 rate = (idx + 1) / elapsed284 eta = (len(samples) - idx - 1) / rate / 60285 logger.info(f" {idx+1}/{len(samples)} "286 f"({rate:.1f} samp/s, ETA {eta:.0f}m)")287 288 if (idx + 1) % checkpoint_every == 0:289 ckpt = out_dir / f"checkpoint_{model_key}_{idx+1}.npz"290 np.savez_compressed(str(ckpt),291 results=np.array(results, dtype=object))292 logger.info(f" Checkpoint: {ckpt}")293 294 # Final save295 cache_path = out_dir / f"cached_{model_key}.npz"296 logger.info(f"Saving final cache ({len(results)} samples)...")297 np.savez_compressed(str(cache_path),298 results=np.array(results, dtype=object))299 elapsed = time.time() - t0300 logger.info(f"Done: {cache_path} ({elapsed/60:.1f}m)")301 302 # Clean up checkpoints303 for ckpt in out_dir.glob(f"checkpoint_{model_key}_*.npz"):304 ckpt.unlink()305 logger.info(f" Removed checkpoint: {ckpt.name}")306 307 # Print summary308 n_correct = sum(1 for r in results if r["correct"])309 n_drifted = sum(1 for r in results if r["is_drifted"])310 logger.info(f"\n Summary for {model_key}:")311 logger.info(f" Samples: {len(results)}")312 logger.info(f" Drifted: {n_drifted}")313 logger.info(f" Stable: {len(results) - n_drifted}")314 logger.info(f" Correct: {n_correct} ({n_correct/len(results):.1%})")315 logger.info(f" Layers: {n_layers}")316 logger.info(f" H-dim: {model.config.hidden_size}")317 318 # Free GPU319 del model320 if torch.cuda.is_available():321 torch.cuda.empty_cache()322 323 return results324 325 326# ─────────────────────────────────────────────────────────────────────────────327# CACHE MIGRATION (reuse existing caches)328# ─────────────────────────────────────────────────────────────────────────────329 330def migrate_cache(src_path, dst_dir, model_key, dataset_path, drift_key):331 """332 Migrate an existing cache to v4 format.333 Adds missing fields (correct, drift_date, entity, etc.) by joining334 from the unified dataset.335 """336 logger.info(f"Migrating cache: {src_path} -> v4 format")337 338 # Load existing cache339 d = np.load(src_path, allow_pickle=True)340 results = d["results"].tolist()341 logger.info(f" Loaded {len(results)} cached samples")342 logger.info(f" Fields: {list(results[0].keys())}")343 344 # Load dataset for enrichment345 with open(dataset_path) as f:346 raw = json.load(f)347 samples = raw.get("samples", raw)348 lookup = {s.get("query", ""): s for s in samples}349 logger.info(f" Dataset: {len(samples)} samples, {len(lookup)} unique queries")350 351 # Fields to ensure exist352 required = ["correct", "is_drifted", "relation", "category", "entity",353 "knowledge_type", "drift_date", "year", "dataset_source",354 "sample_id", "expected_answer"]355 356 enriched = 0357 for r in results:358 # Fix correct field359 if "correct" not in r:360 r["correct"] = r.get("top_answer_matches", False)361 362 # Fix is_drifted from model-specific key363 src = lookup.get(r.get("query", ""))364 if src is not None:365 # Use model-specific drift label366 val = src.get(drift_key, src.get("is_drifted", False))367 if isinstance(val, str):368 r["is_drifted"] = val.lower() in ("true", "1", "yes")369 else:370 r["is_drifted"] = bool(val)371 372 # Enrich missing fields373 for field in required:374 if r.get(field) in (None, "", "None") and field in src:375 r[field] = src[field]376 enriched += 1377 378 logger.info(f" Enriched {enriched} field values")379 380 # Verify381 has_correct = sum(1 for r in results if "correct" in r)382 has_drift_date = sum(1 for r in results383 if r.get("drift_date") not in (None, "", "None"))384 n_drifted = sum(1 for r in results if r.get("is_drifted"))385 n_correct = sum(1 for r in results if r.get("correct"))386 387 logger.info(f" After migration:")388 logger.info(f" has_correct: {has_correct}/{len(results)}")389 logger.info(f" has_drift_date: {has_drift_date}/{len(results)}")390 logger.info(f" n_drifted: {n_drifted}")391 logger.info(f" n_correct: {n_correct}")392 393 # Save394 dst_dir = Path(dst_dir)395 dst_dir.mkdir(parents=True, exist_ok=True)396 dst_path = dst_dir / f"cached_{model_key}.npz"397 logger.info(f" Saving to {dst_path}...")398 np.savez_compressed(str(dst_path),399 results=np.array(results, dtype=object))400 logger.info(f" Done.")401 return results402 403 404# ─────────────────────────────────────────────────────────────────────────────405# MAIN406# ─────────────────────────────────────────────────────────────────────────────407 408def main():409 p = argparse.ArgumentParser(410 description="Extract hidden states from LLMs for drift detection",411 formatter_class=argparse.ArgumentDefaultsHelpFormatter)412 p.add_argument("--model", default="qwen25",413 help="Model key from models.yaml")414 p.add_argument("--config", default="models.yaml",415 help="Path to models.yaml config")416 p.add_argument("--dataset", default=None,417 help="Override dataset path from config")418 p.add_argument("--output_dir", default=None,419 help="Override output dir from config")420 p.add_argument("--device", default="auto",421 help="Device for model loading")422 p.add_argument("--skip_if_cached", action="store_true",423 help="Skip extraction if cache already exists")424 p.add_argument("--migrate_from", default=None,425 help="Migrate existing cache to v4 format")426 p.add_argument("--all", action="store_true",427 help="Extract all models sequentially")428 args = p.parse_args()429 430 cfg = load_config(args.config)431 defaults = cfg.get("defaults", {})432 dataset_path = args.dataset or defaults.get("dataset",433 "data/knowledge_drift_unified_tier1.json")434 output_base = args.output_dir or defaults.get("output_dir",435 "data/experiments/v4")436 437 models_to_run = (list(cfg["models"].keys()) if args.all438 else [args.model])439 440 for model_key in models_to_run:441 mcfg = get_model_cfg(cfg, model_key)442 drift_key = mcfg["drift_key"]443 model_out = Path(output_base) / model_key444 model_out.mkdir(parents=True, exist_ok=True)445 446 cache_path = model_out / f"cached_{model_key}.npz"447 448 # Migration mode449 if args.migrate_from and model_key == args.model:450 migrate_cache(args.migrate_from, str(model_out),451 model_key, dataset_path, drift_key)452 continue453 454 # Skip if cached455 if args.skip_if_cached and cache_path.exists():456 logger.info(f"[{model_key}] Cache exists: {cache_path} — skipping")457 continue458 459 # Load dataset460 samples = load_dataset(dataset_path, model_key, drift_key)461 462 # Load model and extract463 logger.info(f"\n{'='*60}")464 logger.info(f" Extracting: {model_key} ({mcfg['name']})")465 logger.info(f"{'='*60}")466 467 mdl, tok = load_model(mcfg["name"], args.device)468 469 max_length = defaults.get("max_length", 512)470 run_extraction(mdl, tok, samples, mcfg, str(model_out),471 model_key, max_length=max_length)472 473 logger.info("\nAll extractions complete.")474 475 476if __name__ == "__main__":477 main()