CoolFace
Datasetpublic

Raniahossam33/knowledge-drift-experiments

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes22downloads
extract_models.py477 linesDownload Raw Back to root
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()