""" Interactive inference with steering vectors applied. Usage: python src/apply_steering.py python src/apply_steering.py --alpha 15 # Per-trait layers are read from config.yaml → steering.active_traits """ import argparse import torch from pathlib import Path from model_utils import load_config, load_model, build_chat_prompt, SteeringHook, HiddenStateCollector ROOT = Path(__file__).parent.parent def load_vector(trait: str, layer_idx: int) -> torch.Tensor: path = ROOT / "vectors" / f"{trait}_layer{layer_idx:03d}.pt" if not path.exists(): raise FileNotFoundError( f"Vector not found: {path}\n" f"Run: python src/extract_vectors.py --trait {trait}" ) return torch.load(path, weights_only=True) def generate(model, tokenizer, prompt: str, max_new_tokens: int = 256) -> str: inputs = tokenizer(prompt, return_tensors="pt").to(next(model.parameters()).device) with torch.no_grad(): output_ids = model.generate( **inputs, max_new_tokens=max_new_tokens, do_sample=False, temperature=None, top_p=None, pad_token_id=tokenizer.eos_token_id, repetition_penalty=1.3, ) new_ids = output_ids[0][inputs["input_ids"].shape[1]:] return tokenizer.decode(new_ids, skip_special_tokens=True) def diagnose_scale(model, tokenizer, trait_configs: list[dict], alpha: float): """Print hidden-state norms and calibration info for each trait's layer.""" test_text = "你好,有什么我可以帮你的?" prompt = build_chat_prompt(tokenizer, test_text) inputs = tokenizer(prompt, return_tensors="pt").to(next(model.parameters()).device) layers = list({tc["layer"] for tc in trait_configs}) norms = {} with HiddenStateCollector(model, layers) as c: with torch.no_grad(): model(**inputs) for l in layers: norms[l] = c.states[l].norm(dim=-1).mean().item() print("\n[Diagnostics]") for tc in trait_configs: l = tc["layer"] hn = norms[l] print(f" {tc['name']} @ layer {l} hidden-norm={hn:.1f} alpha={alpha} " f"perturbation ratio={alpha/hn:.2f}x (suggested: 0.2–0.5x → alpha {hn*0.2:.0f}–{hn*0.5:.0f})") print() def run_interactive(cfg: dict, trait_configs: list[dict], alpha: float): print("Loading model...") model, tokenizer = load_model(cfg) # Load vectors for tc in trait_configs: tc["vector"] = load_vector(tc["name"], tc["layer"]) print(f"Loaded vector: {tc['name']} @ layer {tc['layer']}") diagnose_scale(model, tokenizer, trait_configs, alpha) active_str = ", ".join(f"{tc['name']}(layer={tc['layer']})" for tc in trait_configs) print(f"Steering active — {active_str}, alpha={alpha}") print("Commands: 'baseline' toggle steering, 'quit' exit\n") steering_on = True while True: user_input = input("You: ").strip() if not user_input or user_input.lower() == "quit": break if user_input.lower() == "baseline": steering_on = not steering_on print(f"[Steering {'ON' if steering_on else 'OFF'}]") continue prompt = build_chat_prompt(tokenizer, user_input) if steering_on: hooks = [ SteeringHook(model, tc["layer"], tc["vector"], tc.get("alpha", alpha)) for tc in trait_configs ] response = generate(model, tokenizer, prompt) for h in hooks: h.remove() else: response = generate(model, tokenizer, prompt) print(f"Model: {response}\n") def main(): parser = argparse.ArgumentParser() parser.add_argument("--alpha", type=float, default=None) parser.add_argument("--max_new_tokens", type=int, default=256) args = parser.parse_args() cfg = load_config() alpha = args.alpha if args.alpha is not None else cfg["steering"]["alpha"] # Build trait_configs from config: [{name, layer}, ...] raw = cfg["steering"]["active_traits"] if raw and isinstance(raw[0], str): # Legacy format: list of strings, use a single apply_layer layer = cfg["steering"].get("apply_layer", 20) trait_configs = [{"name": t, "layer": layer} for t in raw] else: trait_configs = [{"name": t["name"], "layer": t["layer"], "alpha": t.get("alpha")} for t in raw] run_interactive(cfg, trait_configs, alpha) if __name__ == "__main__": main()