""" One-time precomputation: steering vectors + A_lin profile for all target cities. Results cached to ./cache/ and loaded by app.py at startup. Run: python precompute.py Run: python precompute.py --force (recompute even if cache exists) """ import json import os import sys import time import numpy as np import torch from transformers import AutoModelForCausalLM, AutoTokenizer from lap_compute import compute_alin, select_comparison_layer, select_optimal_layer from steering import ( CITY_DATA, TARGET_CITIES, compute_steering_vectors, num_layers, resolve_target_token, ) CACHE_DIR = "./cache" MODEL_ID = "allenai/OLMo-2-0425-1B-Instruct" FALLBACK_MODEL_ID = "EleutherAI/pythia-410m" def load_model(model_id: str): print(f"Loading model: {model_id}") tokenizer = AutoTokenizer.from_pretrained(model_id) if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token model = AutoModelForCausalLM.from_pretrained(model_id, dtype=torch.bfloat16) model.eval() return model, tokenizer def precompute_all(force: bool = False): os.makedirs(CACHE_DIR, exist_ok=True) meta_path = os.path.join(CACHE_DIR, "meta.json") if not force and os.path.exists(meta_path): with open(meta_path) as f: existing = json.load(f) cached_cities = set(existing.get("cities", {}).keys()) missing = [c for c in TARGET_CITIES if c not in cached_cities] if not missing: print("Cache exists and complete. Use --force to recompute.") return print(f"Cache exists but missing cities: {missing}. Computing those.") else: existing = None missing = TARGET_CITIES try: model, tokenizer = load_model(MODEL_ID) model_used = MODEL_ID except Exception as e: print(f"Could not load {MODEL_ID}: {e}\nFalling back to {FALLBACK_MODEL_ID}") model, tokenizer = load_model(FALLBACK_MODEL_ID) model_used = FALLBACK_MODEL_ID n = num_layers(model) device = "cpu" cities_meta = existing["cities"] if existing and "cities" in existing else {} for city in missing: print(f"\n{'='*50}") print(f"Processing: {city}") print(f"{'='*50}") t0 = time.time() target_id = resolve_target_token(tokenizer, city) print(f"Target token '{city}' → id {target_id} ('{tokenizer.decode([target_id])}')") print(f"\nComputing steering vectors (London → {city})...") steering_vecs = compute_steering_vectors(model, tokenizer, city, device) print(f"\nComputing A_lin for '{city}' prompts...") city_prompts = CITY_DATA[city]["prompts"] alin = compute_alin(model, tokenizer, city_prompts, city, device) optimal = select_optimal_layer(alin) comparison = select_comparison_layer(n) print(f"Optimal layer : {optimal} (A_lin={alin[optimal]:.3f})") print(f"Comparison layer: {comparison} (A_lin={alin[comparison]:.3f})") sv_path = os.path.join(CACHE_DIR, f"steering_{city.lower()}.npz") np.savez(sv_path, **{str(l): v.float().numpy() for l, v in steering_vecs.items()}) cities_meta[city] = { "target_token_id": target_id, "alin": alin, "optimal_layer": optimal, "comparison_layer": comparison, "steering_vectors_path": sv_path, } elapsed = time.time() - t0 print(f"Done in {elapsed:.1f}s") meta = { "model": model_used, "n_layers": n, "cities": cities_meta, } with open(meta_path, "w") as f: json.dump(meta, f, indent=2) print(f"\nAll done. Cache saved to {CACHE_DIR}/") def load_cache() -> dict: """Load precomputed data. Called by app.py at startup.""" meta_path = os.path.join(CACHE_DIR, "meta.json") if not os.path.exists(meta_path): raise FileNotFoundError( f"Cache not found at {CACHE_DIR}/. Run: python precompute.py" ) with open(meta_path) as f: meta = json.load(f) for city, city_meta in meta["cities"].items(): npz = np.load(city_meta["steering_vectors_path"]) city_meta["steering_vectors"] = {int(k): torch.tensor(npz[k]) for k in npz.files} return meta if __name__ == "__main__": precompute_all(force="--force" in sys.argv)