CoolFace
Datasetpublic

Raniahossam33/knowledge-drift-experiments

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes22downloads
run_experiments.py230 linesDownload Raw Back to root
1"""2Knowledge Drift Detection: Full Experiment Suite3==================================================4Master runner for all experiments. Run on GPU cluster.5 6Experiments (in order):7  1. Entropy & Confidence Analysis    → Does the model show uncertainty on drifted facts?8  2. Year-Token Attention Analysis     → Do attention heads attend differently to year tokens?9  3. Drift Neuron Discovery (L1)       → Can we find sparse "drift neurons" in MLP activations?10 11Each experiment produces:12  - Raw results (JSON)13  - Summary statistics (JSON) 14  - Console output with key findings15 16Usage:17    # Run everything18    python run_experiments.py --model Qwen/Qwen2.5-7B-Instruct --all19 20    # Run individual experiments21    python run_experiments.py --model Qwen/Qwen2.5-7B-Instruct --entropy22    python run_experiments.py --model Qwen/Qwen2.5-7B-Instruct --attention23    python run_experiments.py --model Qwen/Qwen2.5-7B-Instruct --neurons24 25    # Quick test (small sample)26    python run_experiments.py --model Qwen/Qwen2.5-7B-Instruct --all --max_samples 5027 28Paper References:29    - Entropy/Confidence: SEPs (Kossen et al., ICLR 2025), Calibration (Radharapu et al.)30    - Attention: D-LEAF (Yang et al., 2025) adapted for temporal queries31    - Neurons: SE Neurons (NeurIPS 2024), Neuron Circuits (Arora & Wu, 2026)32"""33 34import argparse35import json36import os37import sys38import time39import logging40from datetime import datetime41 42# Check dependencies early43try:44    import sklearn45except ImportError:46    print("Installing scikit-learn...")47    import subprocess48    subprocess.check_call([sys.executable, "-m", "pip", "install", "scikit-learn", "-q"])49 50logging.basicConfig(51    level=logging.INFO,52    format='%(asctime)s - %(levelname)s - %(message)s',53    handlers=[54        logging.StreamHandler(),55        logging.FileHandler(f"experiment_log_{datetime.now().strftime('%Y%m%d_%H%M%S')}.log")56    ]57)58logger = logging.getLogger(__name__)59 60 61def print_banner(text):62    width = 8063    print("\n" + "█" * width)64    print(f"█  {text:^{width-4}}  █")65    print("█" * width + "\n")66 67 68def run_entropy_analysis(args):69    """Experiment 1: Output entropy, top-k probs, logit lens."""70    print_banner("EXPERIMENT 1: ENTROPY & CONFIDENCE ANALYSIS")71    72    from analyze_drift_signals import load_model, run_analysis73    74    with open(args.dataset, 'r') as f:75        dataset = json.load(f)76    samples = dataset["samples"]77    78    if args.post_cutoff_only:79        samples = [s for s in samples if s.get("temporal_zone") == "post_cutoff"]80        logger.info(f"Filtered to {len(samples)} post-cutoff samples")81    82    model, tokenizer = load_model(args.model, args.device)83    84    import torch85    device = "cuda" if torch.cuda.is_available() else "cpu"86    87    output_dir = os.path.join(args.output_base, "entropy_analysis")88    summary = run_analysis(model, tokenizer, samples, output_dir, device, args.max_samples)89    90    # Cleanup91    del model92    if torch.cuda.is_available():93        torch.cuda.empty_cache()94    95    return summary96 97 98def run_attention_analysis(args):99    """Experiment 2: Year-token attention patterns (D-LEAF adapted)."""100    print_banner("EXPERIMENT 2: YEAR-TOKEN ATTENTION ANALYSIS")101    102    from year_attention_analysis import load_model, run_analysis103    104    with open(args.dataset, 'r') as f:105        dataset = json.load(f)106    samples = dataset["samples"]107    108    if args.post_cutoff_only:109        samples = [s for s in samples if s.get("temporal_zone") == "post_cutoff"]110    111    model, tokenizer = load_model(args.model, args.device)112    113    import torch114    device = "cuda" if torch.cuda.is_available() else "cpu"115    116    output_dir = os.path.join(args.output_base, "attention_analysis")117    summary = run_analysis(model, tokenizer, samples, output_dir, device, args.max_samples)118    119    del model120    if torch.cuda.is_available():121        torch.cuda.empty_cache()122    123    return summary124 125 126def run_neuron_discovery(args):127    """Experiment 3: L1-regularized drift neuron discovery."""128    print_banner("EXPERIMENT 3: DRIFT NEURON DISCOVERY (L1 PROBES)")129    130    from drift_neuron_discovery import run_full_analysis131    from transformers import AutoModelForCausalLM, AutoTokenizer132    import torch133    134    tokenizer = AutoTokenizer.from_pretrained(args.model, trust_remote_code=True)135    if tokenizer.pad_token is None:136        tokenizer.pad_token = tokenizer.eos_token137    model = AutoModelForCausalLM.from_pretrained(138        args.model, torch_dtype=torch.float16, device_map=args.device, trust_remote_code=True,139    )140    model.eval()141    142    with open(args.dataset, 'r') as f:143        dataset = json.load(f)144    samples = dataset["samples"]145    146    if args.post_cutoff_only:147        samples = [s for s in samples if s.get("temporal_zone") == "post_cutoff"]148    149    device = "cuda" if torch.cuda.is_available() else "cpu"150    output_dir = os.path.join(args.output_base, "drift_neurons")151    results = run_full_analysis(model, tokenizer, samples, output_dir, device, args.max_samples)152    153    del model154    if torch.cuda.is_available():155        torch.cuda.empty_cache()156    157    return results158 159 160def main():161    parser = argparse.ArgumentParser(description="Knowledge Drift Detection Experiments")162    parser.add_argument("--model", default="Qwen/Qwen2.5-7B-Instruct", help="Model name/path")163    parser.add_argument("--dataset", default="data/knowledge_drift_dataset.json", help="Dataset path")164    parser.add_argument("--output_base", default="data/experiments/", help="Base output directory")165    parser.add_argument("--device", default="auto", help="Device (auto/cuda/cpu)")166    parser.add_argument("--max_samples", type=int, default=None, help="Max samples per experiment")167    parser.add_argument("--post_cutoff_only", action="store_true", help="Only post-cutoff queries")168    169    # Experiment selection170    parser.add_argument("--all", action="store_true", help="Run all experiments")171    parser.add_argument("--entropy", action="store_true", help="Run entropy analysis")172    parser.add_argument("--attention", action="store_true", help="Run attention analysis")173    parser.add_argument("--neurons", action="store_true", help="Run neuron discovery")174    175    args = parser.parse_args()176    177    if not any([args.all, args.entropy, args.attention, args.neurons]):178        args.all = True179    180    os.makedirs(args.output_base, exist_ok=True)181    182    start_time = time.time()183    all_results = {}184    185    print_banner("KNOWLEDGE DRIFT DETECTION EXPERIMENT SUITE")186    print(f"  Model:   {args.model}")187    print(f"  Dataset: {args.dataset}")188    print(f"  Output:  {args.output_base}")189    print(f"  Max samples: {args.max_samples or 'all'}")190    print(f"  Post-cutoff only: {args.post_cutoff_only}")191    print()192    193    # Run experiments194    if args.all or args.entropy:195        t = time.time()196        all_results["entropy"] = run_entropy_analysis(args)197        logger.info(f"Entropy analysis completed in {time.time()-t:.1f}s")198    199    if args.all or args.attention:200        t = time.time()201        all_results["attention"] = run_attention_analysis(args)202        logger.info(f"Attention analysis completed in {time.time()-t:.1f}s")203    204    if args.all or args.neurons:205        t = time.time()206        all_results["neurons"] = run_neuron_discovery(args)207        logger.info(f"Neuron discovery completed in {time.time()-t:.1f}s")208    209    # Final summary210    total_time = time.time() - start_time211    print_banner("ALL EXPERIMENTS COMPLETE")212    print(f"  Total time: {total_time/60:.1f} minutes")213    print(f"  Results saved to: {args.output_base}")214    215    # Save experiment metadata216    metadata = {217        "model": args.model,218        "dataset": args.dataset,219        "max_samples": args.max_samples,220        "post_cutoff_only": args.post_cutoff_only,221        "total_time_seconds": total_time,222        "experiments_run": list(all_results.keys()),223        "timestamp": datetime.now().isoformat(),224    }225    with open(os.path.join(args.output_base, "experiment_metadata.json"), 'w') as f:226        json.dump(metadata, f, indent=2)227 228 229if __name__ == "__main__":230    main()