"""Model loading and hook utilities for steering vector extraction/application.""" import torch import yaml from pathlib import Path from typing import Optional from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig ROOT = Path(__file__).parent.parent def load_config() -> dict: with open(ROOT / "config.yaml", encoding="utf-8") as f: return yaml.safe_load(f) def load_model(cfg: Optional[dict] = None, quantization: Optional[str] = None): """Load Qwen3.5-9B with optional quantization. Returns (model, tokenizer).""" if cfg is None: cfg = load_config() model_id = cfg["model"]["local_path"] or cfg["model"]["id"] q = quantization or cfg["model"]["quantization"] dtype = getattr(torch, cfg["model"]["torch_dtype"]) bnb_config = None if q == "4bit": bnb_config = BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_compute_dtype=dtype, bnb_4bit_use_double_quant=True, bnb_4bit_quant_type="nf4", ) elif q == "8bit": bnb_config = BitsAndBytesConfig(load_in_8bit=True) model = AutoModelForCausalLM.from_pretrained( model_id, quantization_config=bnb_config, torch_dtype=dtype if bnb_config is None else None, device_map="auto", trust_remote_code=True, ) model.eval() tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True) if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token return model, tokenizer def get_transformer_layers(model): """Return the list of transformer decoder layers.""" # Qwen3 uses model.model.layers return model.model.layers def build_chat_prompt(tokenizer, user_text: str, response_prefix: str = "") -> str: """Format a single-turn chat. Thinking mode is suppressed via system prompt.""" # enable_thinking=False is ignored by this tokenizer version; # use a system message to suppress CoT instead. messages = [ {"role": "system", "content": "/no_think"}, {"role": "user", "content": user_text}, ] try: prompt = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True, enable_thinking=False, ) except TypeError: prompt = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True, ) if response_prefix: prompt = prompt + response_prefix return prompt class HiddenStateCollector: """ Registers forward hooks on transformer layers and collects the output hidden states (residual stream after the full layer). """ def __init__(self, model, layers: list[int]): self.layers = layers self.states: dict[int, torch.Tensor] = {} self._hooks = [] transformer_layers = get_transformer_layers(model) for idx in layers: layer = transformer_layers[idx] handle = layer.register_forward_hook(self._make_hook(idx)) self._hooks.append(handle) def _make_hook(self, idx: int): def hook(module, input, output): # Qwen3 layer output is a tuple; first element is hidden states hidden = output[0] if isinstance(output, tuple) else output # Store mean over sequence length, keep batch dim → (batch, hidden) self.states[idx] = hidden.detach().mean(dim=1) return hook def clear(self): self.states.clear() def remove(self): for h in self._hooks: h.remove() self._hooks.clear() def __enter__(self): return self def __exit__(self, *_): self.remove() class SteeringHook: """ Injects a steering vector into the residual stream at a specific layer during generation. """ def __init__(self, model, layer_idx: int, vector: torch.Tensor, alpha: float): transformer_layers = get_transformer_layers(model) self.alpha = alpha self.vector = vector # shape: (hidden_size,) self._handle = transformer_layers[layer_idx].register_forward_hook( self._hook ) def _hook(self, module, input, output): hidden = output[0] if isinstance(output, tuple) else output v = self.alpha * self.vector.to(hidden.device, hidden.dtype) hidden = hidden + v if isinstance(output, tuple): return (hidden,) + output[1:] return hidden def remove(self): self._handle.remove() def __enter__(self): return self def __exit__(self, *_): self.remove() def num_layers(model) -> int: return len(get_transformer_layers(model))