diff --git a/docs/.nav.yml b/docs/.nav.yml index d4ef49e8..7b2ac5a2 100644 --- a/docs/.nav.yml +++ b/docs/.nav.yml @@ -50,6 +50,7 @@ nav: - DirectionalAblation: "examples/notebooks/algorithms/directional_ablation.ipynb" - ITI: "examples/notebooks/algorithms/iti.ipynb" - PASTA: "examples/notebooks/algorithms/pasta.ipynb" + - VJPDelta: "examples/notebooks/algorithms/vjp_delta.ipynb" - Output control: - BestOfN: "examples/notebooks/algorithms/best_of_n.ipynb" - BudgetForcing: "examples/notebooks/algorithms/budget_forcing.ipynb" @@ -114,6 +115,7 @@ nav: - Directional Ablation: reference/algorithms/state_control/directional_ablation.md - ITI: reference/algorithms/state_control/iti.md - PASTA: reference/algorithms/state_control/pasta.md + - VJPDelta: reference/algorithms/state_control/vjp_delta.md - Output control: - Base classes: reference/algorithms/output_control/base_output_control.md - Common library: reference/algorithms/output_control/common.md diff --git a/docs/concepts/controls.md b/docs/concepts/controls.md index 128ed76e..d932c38d 100644 --- a/docs/concepts/controls.md +++ b/docs/concepts/controls.md @@ -152,6 +152,9 @@ patching. The toolkit implements: - `PASTA` ([API reference](../reference/algorithms/state_control/pasta.md), [notebook](../examples/notebooks/algorithms/pasta.ipynb)) - *Description*: post-hoc attention steering[@zhang2024tell], rescaling attention to targeted prompt substrings at selected layers and heads. The `head_config` argument takes a dict or list of layers and heads, or a `HeadProfile` recipe that runs the paper's head-profiling stage as a steer-time fit on the loaded model (scoring each candidate head by its paired lift over an unsteered baseline) and freezes the resolved head map. - *Backends*: HF with `attn_implementation` `"eager"` or `"sdpa"` (attention-map writes have no engine form). +- `VJPDelta` ([API reference](../reference/algorithms/state_control/vjp_delta.md), [notebook](../examples/notebooks/algorithms/vjp_delta.ipynb)) + - *Description*: fits a target-state contrast and uses vector-Jacobian products to derive one normalized additive direction per earlier residual layer. `VJPDeltaFit` uses raw prompts, excludes `skip_first` positions and each row's final real token from the VJP spans, and averages each class independently before subtraction. + - *Backends*: HF for fitting. The frozen form uses the existing `ActivationAdapter` additive intervention. Reusable building blocks shared across the residual-stream methods (estimators, gating, selectors, transforms, steering vectors, hook utilities) are located in diff --git a/docs/reference/algorithms/state_control/vjp_delta.md b/docs/reference/algorithms/state_control/vjp_delta.md new file mode 100644 index 00000000..2e6554e7 --- /dev/null +++ b/docs/reference/algorithms/state_control/vjp_delta.md @@ -0,0 +1,20 @@ +# VJPDelta + +::: steerability.algorithms.state_control.vjp_delta + handler: python + options: + show_if_no_docstring: true + show_source: true + show_root_heading: true + docstring_style: google + show_root_full_path: true + show_object_full_path: false + separate_signature: false + inherited_members: true + show_submodules: true + show_symbol_type_heading: true + show_symbol_type_toc: true + filters: + - "!.*Args$" + - "!^registry" + - "!^STEERING_METHOD" diff --git a/examples/index.md b/examples/index.md index c4269db8..7fc4422a 100644 --- a/examples/index.md +++ b/examples/index.md @@ -59,6 +59,8 @@ Algorithm notebooks demonstrate how each method (i.e., control) operates. The me :octicons-arrow-right-24: [PASTA](./notebooks/algorithms/pasta.ipynb) + :octicons-arrow-right-24: [VJPDelta](./notebooks/algorithms/vjp_delta.ipynb) + - __Output control__ --- diff --git a/examples/notebooks/algorithms/vjp_delta.ipynb b/examples/notebooks/algorithms/vjp_delta.ipynb new file mode 100644 index 00000000..89409742 --- /dev/null +++ b/examples/notebooks/algorithms/vjp_delta.ipynb @@ -0,0 +1,170 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "24c9ee9c", + "metadata": {}, + "source": [ + "# VJP-delta\n", + "\n", + "`VJPDelta` fits an additive direction at each chosen source layer. It reads a contrast at a later target layer, then uses a vector-Jacobian product (VJP) to map that contrast back to each source layer. The fitted vectors use the usual state-control additive intervention during generation.\n", + "\n", + "VJP-delta follows [Clark, Michael J. (2026), _vjp-steering: contrastive steering vectors from vector-Jacobian products_](https://github.com/wassname/vjp-steering), adapting the [Jacobian lens](https://transformer-circuits.pub/2026/workspace/)." + ] + }, + { + "cell_type": "markdown", + "id": "2fe4a263", + "metadata": {}, + "source": [ + "## Setup\n", + "\n", + "This CPU demonstration loads a small Hugging Face Llama checkpoint through the public pipeline API." + ] + }, + { + "cell_type": "code", + "execution_count": 1, + "id": "1c306fb6", + "metadata": {}, + "outputs": [ + { + "name": "stderr", + "output_type": "stream", + "text": [ + "[transformers] Model config: pad_token_id must be `None` or an integer within the vocabulary (between 0 and 31999), got -1. This may result in unexpected behavior.\n" + ] + }, + { + "name": "stderr", + "output_type": "stream", + "text": [ + "[transformers] The following generation flags are not valid and may be ignored: ['pad_token_id']. Set `TRANSFORMERS_VERBOSITY=info` for more details.\n" + ] + } + ], + "source": [ + "from pathlib import Path\n", + "import tempfile\n", + "\n", + "from steerability.algorithms.core.steering_pipeline import SteeringPipeline\n", + "from steerability.algorithms.state_control.vjp_delta import VJPDelta\n", + "from steerability.spipe import SPipe\n", + "from transformers import AutoModelForCausalLM, AutoTokenizer\n", + "\n", + "model_name = \"hf-internal-testing/tiny-random-LlamaForCausalLM\"\n", + "tokenizer = AutoTokenizer.from_pretrained(model_name)\n", + "tokenizer.pad_token = tokenizer.eos_token\n", + "model = AutoModelForCausalLM.from_pretrained(model_name)\n", + "fit_data = {\n", + " \"positives\": [\"the cat sat\", \"the dog ran\"],\n", + " \"negatives\": [\"dog ran fast\"],\n", + "}\n", + "control = VJPDelta(\n", + " data=fit_data,\n", + " target_layer=1,\n", + " source_layer_ids=[0],\n", + " skip_first=0,\n", + " strength=0.5,\n", + ")\n", + "pipeline = SteeringPipeline(\n", + " model=model,\n", + " tokenizer=tokenizer,\n", + " controls=[control],\n", + " model_name_or_path=model_name,\n", + ")" + ] + }, + { + "cell_type": "markdown", + "id": "5c87a2e7", + "metadata": {}, + "source": [ + "## Extract and steer\n", + "\n", + "The raw prompts have unequal positive and negative pool sizes. `steer()` extracts the vectors, then binds the standard additive intervention. The stored direction rows are unit norm." + ] + }, + { + "cell_type": "code", + "execution_count": 2, + "id": "ca2ba708", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "{'layers': [0], 'reply': \"'agreedָָ'\"}\n" + ] + } + ], + "source": [ + "pipeline.steer()\n", + "vector = control.export_state()[\"intervention_0/transform\"]\n", + "assert set(vector.directions) == {0}\n", + "assert all(abs(direction.norm().item() - 1.0) < 1e-5 for direction in vector.directions.values())\n", + "\n", + "reply = pipeline.generate(text=\"the cat\", max_new_tokens=3, do_sample=False)\n", + "print({\"layers\": sorted(vector.directions), \"reply\": repr(reply)})" + ] + }, + { + "cell_type": "markdown", + "id": "5b74b575", + "metadata": {}, + "source": [ + "## Freeze and reload\n", + "\n", + "The frozen form contains the fitted vectors as an `ActivationAdapter`. Reloading it resolves the stored additive artifact and does not run another VJP fit." + ] + }, + { + "cell_type": "code", + "execution_count": 3, + "id": "651ceb9c", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "{'bundle': '/tmp/tmpx73lemsj/vjp_delta_demo', 'reply_matches': True}\n" + ] + } + ], + "source": [ + "bundle = Path(tempfile.mkdtemp()) / \"vjp_delta_demo\"\n", + "saved = pipeline.to_spipe().save(bundle)\n", + "reloaded = SPipe.load(saved).pipeline()\n", + "assert type(reloaded.state_controls[0]).__name__ == \"ActivationAdapter\"\n", + "reloaded.model, reloaded.tokenizer = model, tokenizer\n", + "reloaded.steer()\n", + "reloaded_reply = reloaded.generate(text=\"the cat\", max_new_tokens=3, do_sample=False)\n", + "assert reloaded_reply == reply\n", + "print({\"bundle\": str(saved), \"reply_matches\": reloaded_reply == reply})" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.13.12" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/steerability/algorithms/state_control/vjp_delta/__init__.py b/steerability/algorithms/state_control/vjp_delta/__init__.py new file mode 100644 index 00000000..823e4bfa --- /dev/null +++ b/steerability/algorithms/state_control/vjp_delta/__init__.py @@ -0,0 +1,12 @@ +from .args import VJPDeltaArgs +from .control import VJPDelta +from .fit import VJPDeltaFit + +STEERING_METHOD = { + "category": "state_control", + "name": "vjp_delta", + "control": VJPDelta, + "args": VJPDeltaArgs, +} + +__all__ = ["VJPDelta", "VJPDeltaArgs", "VJPDeltaFit"] diff --git a/steerability/algorithms/state_control/vjp_delta/args.py b/steerability/algorithms/state_control/vjp_delta/args.py new file mode 100644 index 00000000..47bc4672 --- /dev/null +++ b/steerability/algorithms/state_control/vjp_delta/args.py @@ -0,0 +1,67 @@ +"""Arguments for VJP-delta steering.""" +from __future__ import annotations + +from dataclasses import dataclass +from typing import Sequence + +from steerability.algorithms.core.base_args import BaseArgs +from steerability.algorithms.core.internals.data import LabeledExamples, as_labeled_examples +from steerability.algorithms.state_control.common.sources import ArtifactSource +from steerability.algorithms.state_control.common.steering_vector import SteeringVector +from steerability.algorithms.state_control.common.token_scope import ScopeKind + + +@dataclass +class VJPDeltaArgs(BaseArgs): + """Arguments for `VJPDelta`. + + Args: + steering_vector: A precomputed `SteeringVector` skips VJP fitting. An `ArtifactSource` resolves its own artifact. + data: Independent positive and negative raw prompt pools for `VJPDeltaFit`. + target_layer: Target layer for the contrast. None selects `num_layers - 3`. + source_layer_ids: Source layers for VJPs. None selects every layer before the target. + skip_first: Prefix positions excluded during gradient extraction. + max_length: Maximum tokenized fit-prompt length. + batch_size: Fit prompts per differentiable forward. + strength: Multiplier used when applying the normalized directions. + token_scope: Positions that receive the additive intervention during generation. + last_k: Required with `token_scope="last_k"`. + from_position: Required with `token_scope="from_position"`. + """ + + steering_vector: SteeringVector | ArtifactSource | None = None + data: LabeledExamples | dict | None = None + target_layer: int | None = None + source_layer_ids: Sequence[int] | None = None + skip_first: int = 16 + max_length: int = 384 + batch_size: int = 8 + strength: float = 1.0 + token_scope: ScopeKind = "after_prompt" + last_k: int | None = None + from_position: int | None = None + + def __post_init__(self) -> None: + if (self.steering_vector is None) == (self.data is None): + raise ValueError("Provide exactly one of steering_vector or data.") + if isinstance(self.steering_vector, SteeringVector): + self.steering_vector.validate() + if self.data is not None and not isinstance(self.data, LabeledExamples): + self.data = as_labeled_examples(self.data) + if self.target_layer is not None and self.target_layer < 0: + raise ValueError("target_layer must be >= 0.") + if self.source_layer_ids is not None: + source_ids = tuple(int(layer_id) for layer_id in self.source_layer_ids) + if not source_ids or min(source_ids) < 0 or len(set(source_ids)) != len(source_ids): + raise ValueError("source_layer_ids must be a non-empty sequence of unique integers >= 0.") + self.source_layer_ids = source_ids + if self.skip_first < 0: + raise ValueError("skip_first must be >= 0.") + if self.max_length < 2: + raise ValueError("max_length must be >= 2.") + if self.batch_size < 1: + raise ValueError("batch_size must be >= 1.") + if self.token_scope == "last_k" and (self.last_k is None or self.last_k < 1): + raise ValueError("last_k must be >= 1 when token_scope is 'last_k'.") + if self.token_scope == "from_position" and (self.from_position is None or self.from_position < 0): + raise ValueError("from_position must be >= 0 when token_scope is 'from_position'.") diff --git a/steerability/algorithms/state_control/vjp_delta/control.py b/steerability/algorithms/state_control/vjp_delta/control.py new file mode 100644 index 00000000..ddafc094 --- /dev/null +++ b/steerability/algorithms/state_control/vjp_delta/control.py @@ -0,0 +1,62 @@ +"""VJP-delta control.""" +from __future__ import annotations + +from steerability.algorithms.state_control.base import InterventionControl +from steerability.algorithms.state_control.common.sources import _Precomputed +from steerability.algorithms.state_control.common.specs import CoveredLayers, Intervention, TokenScope +from steerability.algorithms.state_control.common.steering_vector import SteeringVector +from steerability.algorithms.state_control.common.transforms import AdditiveTransform + +from .args import VJPDeltaArgs +from .fit import VJPDeltaFit + + +class VJPDelta(InterventionControl): + """VJP-delta activation steering. + + The control fits one normalized additive direction per source layer during `steer()`. Its + target contrast is the positive-minus-negative final unpadded target state. The fit applies + that contrast as a cotangent at valid target tokens and averages valid source-token gradients + per prompt before separately averaging the positive and negative classes. At generation it + uses the standard additive intervention and token scopes. + + A precomputed `SteeringVector` skips VJP fitting. A supplied `ArtifactSource` resolves its own + artifact. The frozen form is `ActivationAdapter`, so a reloaded `.spipe` resolves the stored + vectors without a VJP fit. + + Reference: + + - Clark, Michael J. (2026). "vjp-steering: contrastive steering vectors from + vector-Jacobian products." + [https://github.com/wassname/vjp-steering](https://github.com/wassname/vjp-steering) + Adapts the [Jacobian lens](https://transformer-circuits.pub/2026/workspace/). + """ + + Args = VJPDeltaArgs + supports_batching = True + + def _configure(self) -> None: + if self.steering_vector is None: + source = VJPDeltaFit( + data=self.data, + target_layer=self.target_layer, + source_layer_ids=self.source_layer_ids, + skip_first=self.skip_first, + max_length=self.max_length, + batch_size=self.batch_size, + ) + elif isinstance(self.steering_vector, SteeringVector): + source = _Precomputed(self.steering_vector.clone()) + else: + source = self.steering_vector + self._template = ( + Intervention( + layers=CoveredLayers(), + transform=AdditiveTransform(source, strength=self.strength), + scope=TokenScope(self.token_scope, last_k=self.last_k, from_position=self.from_position), + ), + ) + + def cleanup(self) -> None: + """Drop fitted intervention tensors and their bound artifacts.""" + self.interventions = () diff --git a/steerability/algorithms/state_control/vjp_delta/fit.py b/steerability/algorithms/state_control/vjp_delta/fit.py new file mode 100644 index 00000000..990841fa --- /dev/null +++ b/steerability/algorithms/state_control/vjp_delta/fit.py @@ -0,0 +1,333 @@ +"""VJP-delta steering-vector extraction.""" +from __future__ import annotations + +import weakref +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Sequence + +import torch +from transformers import PreTrainedModel, PreTrainedTokenizerBase + +from steerability.algorithms.core.execution.access import ModelAccess +from steerability.algorithms.core.internals.data import LabeledExamples, as_labeled_examples +from steerability.algorithms.core.internals.model_layout import resolve_model_layout, text_config +from steerability.algorithms.state_control.common.steering_vector import SteeringVector + +if TYPE_CHECKING: + from steerability.algorithms.core.execution.backend import SteeringSession + + +def _hidden(output) -> torch.Tensor: + """Return the residual tensor from a decoder-layer output.""" + hidden = output[0] if isinstance(output, tuple) else output + if not isinstance(hidden, torch.Tensor): + raise TypeError(f"Decoder layer returned {type(hidden).__name__}, not a tensor or tuple headed by a tensor.") + return hidden + + +def _token_batch( + tokenizer: PreTrainedTokenizerBase, + texts: Sequence[str], + *, + max_length: int, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + """Tokenize raw texts into a right-padded batch without changing tokenizer state.""" + encoded = tokenizer( + list(texts), + add_special_tokens=True, + truncation=True, + max_length=max_length, + padding=False, + ) + rows = encoded["input_ids"] + if not rows or any(len(row) == 0 for row in rows): + raise ValueError("VJP-delta fitting requires every example to contain at least one token after tokenization.") + pad_id = tokenizer.pad_token_id if tokenizer.pad_token_id is not None else tokenizer.eos_token_id + if pad_id is None: + raise ValueError("VJP-delta fitting needs tokenizer.pad_token_id or tokenizer.eos_token_id for batching.") + width = max(len(row) for row in rows) + input_ids = torch.full((len(rows), width), int(pad_id), dtype=torch.long, device=device) + attention_mask = torch.zeros((len(rows), width), dtype=torch.long, device=device) + for index, row in enumerate(rows): + input_ids[index, :len(row)] = torch.tensor(row, dtype=torch.long, device=device) + attention_mask[index, :len(row)] = 1 + return input_ids, attention_mask + + +def _valid_token_mask(attention_mask: torch.Tensor, skip_first: int) -> torch.BoolTensor: + """Select valid positions, excluding the prefix and each row's final real token.""" + mask = attention_mask.to(torch.bool).clone() + lengths = mask.sum(dim=1) + if torch.any(lengths <= skip_first + 1): + short = torch.nonzero(lengths <= skip_first + 1, as_tuple=False).flatten().tolist() + raise ValueError( + "VJP-delta fitting needs a valid token span after skip_first and before the final real token; " + f"rows {short} have lengths {lengths[short].tolist()} with skip_first={skip_first}." + ) + positions = torch.arange(mask.size(1), device=mask.device).unsqueeze(0) + mask &= positions >= skip_first + mask[torch.arange(mask.size(0), device=mask.device), lengths - 1] = False + return mask + + +@dataclass +class VJPDeltaFit: + """Fit normalized VJP-delta directions from independent labeled prompt pools. + + Args: + data: Independent positive and negative raw prompts. Class sizes may differ. + target_layer: Target residual layer. None selects `num_layers - 3`. + source_layer_ids: Source residual layers. None selects every layer before the target. + skip_first: Prefix positions excluded from target and source token spans. + max_length: Maximum tokenized prompt length. + batch_size: Number of prompts per differentiable forward. + """ + + produces_positional = False + access = ModelAccess.MODULE + artifact_class = "direction" + + data: LabeledExamples | dict + target_layer: int | None = None + source_layer_ids: Sequence[int] | None = None + skip_first: int = 16 + max_length: int = 384 + batch_size: int = 8 + + _model_ref: weakref.ref | None = field(default=None, init=False, repr=False, compare=False) + _master: SteeringVector | None = field(default=None, init=False, repr=False, compare=False) + + def __post_init__(self) -> None: + if not isinstance(self.data, LabeledExamples): + self.data = as_labeled_examples(self.data) + if self.skip_first < 0: + raise ValueError(f"skip_first must be >= 0, got {self.skip_first}.") + if self.max_length < 2: + raise ValueError(f"max_length must be >= 2, got {self.max_length}.") + if self.batch_size < 1: + raise ValueError(f"batch_size must be >= 1, got {self.batch_size}.") + if self.target_layer is not None and self.target_layer < 0: + raise ValueError(f"target_layer must be >= 0, got {self.target_layer}.") + if self.source_layer_ids is not None: + source_ids = tuple(int(layer_id) for layer_id in self.source_layer_ids) + if not source_ids: + raise ValueError("source_layer_ids must not be empty.") + if len(set(source_ids)) != len(source_ids): + raise ValueError("source_layer_ids must not contain duplicates.") + if min(source_ids) < 0: + raise ValueError("source_layer_ids must all be >= 0.") + self.source_layer_ids = source_ids + + def _layers(self, model: PreTrainedModel) -> tuple[int, tuple[int, ...]]: + layout = resolve_model_layout(model) + target = layout.num_layers - 3 if self.target_layer is None else self.target_layer + if not 0 <= target < layout.num_layers: + raise ValueError(f"target_layer {target} is out of range for {layout.num_layers} decoder layers.") + source_ids = tuple(range(target)) if self.source_layer_ids is None else tuple(self.source_layer_ids) + if not source_ids: + raise ValueError("VJP-delta fitting needs at least one source layer before target_layer.") + invalid = [layer_id for layer_id in source_ids if not 0 <= layer_id < layout.num_layers] + if invalid: + raise ValueError(f"source_layer_ids {invalid} are out of range for {layout.num_layers} decoder layers.") + after_target = [layer_id for layer_id in source_ids if layer_id >= target] + if after_target: + raise ValueError( + f"VJP-delta source layers must precede target_layer {target}; got {after_target}." + ) + return target, source_ids + + def _target_mean( + self, + model: PreTrainedModel, + tokenizer: PreTrainedTokenizerBase, + texts: Sequence[str], + *, + target_layer: int, + target_module: torch.nn.Module, + device: torch.device, + ) -> torch.Tensor: + """Return the mean target-layer state at each row's final unpadded token.""" + total: torch.Tensor | None = None + count = 0 + for start in range(0, len(texts), self.batch_size): + input_ids, attention_mask = _token_batch( + tokenizer, texts[start:start + self.batch_size], max_length=self.max_length, device=device, + ) + captured: list[torch.Tensor] = [] + handle = target_module.register_forward_hook(lambda _m, _a, output: captured.append(_hidden(output))) + try: + with torch.no_grad(): + model(input_ids=input_ids, attention_mask=attention_mask) + finally: + handle.remove() + if len(captured) != 1: + raise RuntimeError(f"VJP-delta target hook fired {len(captured)} times; expected one forward output.") + lengths = attention_mask.sum(dim=1) - 1 + rows = captured[0][torch.arange(input_ids.size(0), device=device), lengths].float() + total = rows.sum(dim=0) if total is None else total + rows.sum(dim=0) + count += rows.size(0) + if total is None or count == 0: + raise ValueError("VJP-delta fitting needs at least one prompt in each class.") + return total / count + + def _class_gradients( + self, + model: PreTrainedModel, + tokenizer: PreTrainedTokenizerBase, + texts: Sequence[str], + *, + target_module: torch.nn.Module, + source_modules: dict[int, torch.nn.Module], + cotangent: torch.Tensor, + device: torch.device, + ) -> dict[int, torch.Tensor]: + """Return independent per-class means of per-prompt source-token VJPs.""" + totals: dict[int, torch.Tensor] = {} + count = 0 + for start in range(0, len(texts), self.batch_size): + input_ids, attention_mask = _token_batch( + tokenizer, texts[start:start + self.batch_size], max_length=self.max_length, device=device, + ) + valid = _valid_token_mask(attention_mask, self.skip_first) + source_states: dict[int, torch.Tensor] = {} + target_states: list[torch.Tensor] = [] + def source_capture(layer_id: int): + def hook(_module, _args, output): + source_states[layer_id] = _hidden(output) + return None + + return hook + + def target_capture(_module, _args, output): + target_states.append(_hidden(output)) + return None + + handles = [ + module.register_forward_hook(source_capture(layer_id)) + for layer_id, module in source_modules.items() + ] + handles.append(target_module.register_forward_hook(target_capture)) + try: + model(input_ids=input_ids, attention_mask=attention_mask) + if len(target_states) != 1 or set(source_states) != set(source_modules): + raise RuntimeError("VJP-delta extraction hooks did not capture one target and every source layer.") + target = target_states[0] + if not target.requires_grad or any(not state.requires_grad for state in source_states.values()): + raise RuntimeError( + "VJP-delta fitting needs differentiable decoder-layer outputs. " + "Use a standard torch attention implementation without inference_mode or no-grad wrappers." + ) + grad_outputs = torch.zeros_like(target) + grad_outputs[valid] = cotangent.to(dtype=target.dtype, device=target.device) + try: + gradients = torch.autograd.grad( + outputs=target, + inputs=tuple(source_states[layer_id] for layer_id in source_modules), + grad_outputs=grad_outputs, + allow_unused=False, + ) + except RuntimeError as error: + raise RuntimeError(f"VJP-delta autograd.grad failed: {error}") from error + finally: + for handle in handles: + handle.remove() + for layer_id, gradient in zip(source_modules, gradients, strict=True): + per_prompt = ( + gradient.float() * valid.unsqueeze(-1).to(dtype=gradient.dtype) + ).sum(dim=1) / valid.sum(dim=1, keepdim=True) + summed = per_prompt.sum(dim=0) + totals[layer_id] = summed if layer_id not in totals else totals[layer_id] + summed + count += input_ids.size(0) + return {layer_id: total / count for layer_id, total in totals.items()} + + def _fit(self, model: PreTrainedModel, tokenizer: PreTrainedTokenizerBase) -> SteeringVector: + if model is None: + raise ValueError("VJP-delta fitting requires a live model at steer time.") + if torch.is_inference_mode_enabled(): + raise RuntimeError("VJP-delta fitting cannot run while torch.inference_mode() is enabled.") + target_layer, source_ids = self._layers(model) + layout = resolve_model_layout(model) + target_module = model.get_submodule(layout.layer_names[target_layer]) + source_modules = {layer_id: model.get_submodule(layout.layer_names[layer_id]) for layer_id in source_ids} + parameters = tuple(model.parameters()) + if not parameters: + raise ValueError("VJP-delta fitting requires a model with parameters.") + devices = {tensor.device for tensor in (*parameters, *model.buffers())} + if len(devices) != 1: + raise ValueError("VJP-delta fitting requires a single-device model; sharded or offloaded models are unsupported.") + device = devices.pop() + if device.type == "meta": + raise ValueError("VJP-delta fitting requires materialized model parameters, not meta tensors.") + + requires_grad = [parameter.requires_grad for parameter in parameters] + parameter_grads = [ + None if parameter.grad is None else parameter.grad.detach().clone() + for parameter in parameters + ] + training_flags = {module: module.training for module in model.modules()} + try: + model.eval() + for parameter in parameters: + parameter.requires_grad_(True) + with torch.enable_grad(): + positive_target = self._target_mean( + model, tokenizer, self.data.positives, target_layer=target_layer, + target_module=target_module, device=device, + ) + negative_target = self._target_mean( + model, tokenizer, self.data.negatives, target_layer=target_layer, + target_module=target_module, device=device, + ) + cotangent = positive_target - negative_target + if not torch.isfinite(cotangent).all() or cotangent.norm() == 0: + raise ValueError("VJP-delta target contrast must have a finite, nonzero norm.") + positive = self._class_gradients( + model, tokenizer, self.data.positives, target_module=target_module, + source_modules=source_modules, cotangent=cotangent, device=device, + ) + negative = self._class_gradients( + model, tokenizer, self.data.negatives, target_module=target_module, + source_modules=source_modules, cotangent=cotangent, device=device, + ) + finally: + for module, training in training_flags.items(): + module.training = training + for parameter, flag, grad in zip(parameters, requires_grad, parameter_grads, strict=True): + parameter.requires_grad_(flag) + parameter.grad = grad + + directions: dict[int, torch.Tensor] = {} + for layer_id in source_ids: + direction = positive[layer_id] - negative[layer_id] + norm = direction.norm() + if not torch.isfinite(norm) or norm == 0: + raise ValueError(f"VJP-delta direction at source layer {layer_id} must have a finite, nonzero norm.") + directions[layer_id] = (direction / norm).detach().cpu().unsqueeze(0) + return SteeringVector( + model_type=text_config(model).model_type, + directions=directions, + meta={ + "method": "vjp_delta", + "target_layer": target_layer, + "source_layer_ids": list(source_ids), + "skip_first": self.skip_first, + "max_length": self.max_length, + }, + ) + + def resolve( + self, + model: PreTrainedModel, + tokenizer: PreTrainedTokenizerBase, + *, + session: "SteeringSession | None" = None, + ) -> SteeringVector: + """Fit once per model identity and return a defensive clone.""" + del session + if model is not None and self._model_ref is not None and self._model_ref() is model and self._master is not None: + return self._master.clone() + master = self._fit(model, tokenizer) + self._model_ref = weakref.ref(model) + self._master = master + return master.clone() diff --git a/tests/controls/test_vjp_delta.py b/tests/controls/test_vjp_delta.py new file mode 100644 index 00000000..aad0c249 --- /dev/null +++ b/tests/controls/test_vjp_delta.py @@ -0,0 +1,363 @@ +"""Regression tests for VJP-delta extraction and its frozen additive form.""" +from __future__ import annotations + +import gc +import json +import weakref + +import pytest +import torch + +from steerability.algorithms.core.steering_pipeline import SteeringPipeline +from steerability.algorithms.core.utils.assembly import collect_state_entries +from steerability.algorithms.state_control.activation_adapter import ActivationAdapter +from steerability.algorithms.state_control.common.steering_vector import SteeringVector +from steerability.algorithms.state_control.vjp_delta import VJPDelta, VJPDeltaFit +from steerability.spipe import SPipe, SpipeStaleError +from tests.utils.tiny_models import tiny_gpt2, tiny_llama, wordlevel_tokenizer + +HIDDEN = 32 + + +def _fit_data(): + return { + "positives": ["the cat sat", "the dog ran"], + "negatives": ["dog ran fast"], + } + + +def _source(): + return VJPDeltaFit(_fit_data(), target_layer=2, source_layer_ids=[0, 1], skip_first=0, batch_size=2) + + +def _capture_layer(model, pipeline, layer_id, input_ids): + entries = collect_state_entries( + pipeline.state_controls, input_ids, {}, hooks_in_process=True, + lowered_state=pipeline._lowered_state, model=pipeline.model, + ) + backend = pipeline._backend_for(pipeline._resolve_backend_spec(None)) + captured = {} + with backend.open_session() as session, session.entries_applied(entries): + def capture(_module, _args, _kwargs, output): + captured["hidden"] = (output[0] if isinstance(output, tuple) else output).detach().clone() + + handle = model.model.layers[layer_id].register_forward_hook(capture, with_kwargs=True) + try: + with torch.no_grad(): + model(input_ids=input_ids) + finally: + handle.remove() + return captured["hidden"] + + +def _replace_hidden(output, hidden): + """Replace the residual tensor while preserving a decoder layer's output container.""" + return (hidden, *output[1:]) if isinstance(output, tuple) else hidden + + +def _explicit_jacobian_direction(model, tokenizer, texts, cotangent, *, source_layer, target_layer, skip_first): + """Compute the VJP contraction from an explicit Jacobian, without production extraction code.""" + source_module = model.model.layers[source_layer] + target_module = model.model.layers[target_layer] + total = None + for text in texts: + batch = tokenizer(text, return_tensors="pt") + source_states = [] + + def capture_source(_module, _args, output): + source_states.append(output[0] if isinstance(output, tuple) else output) + return None + + source_handle = source_module.register_forward_hook(capture_source) + try: + with torch.no_grad(): + model(**batch) + finally: + source_handle.remove() + source_value = source_states.pop().detach() + valid = torch.ones(source_value.shape[:2], dtype=torch.bool) + valid[:, :skip_first] = False + valid[:, -1] = False + + def target_from_source(replacement): + target_states = [] + + def replace_source(_module, _args, output): + return _replace_hidden(output, replacement) + + def capture_target(_module, _args, output): + target_states.append(output[0] if isinstance(output, tuple) else output) + return None + + source_hook = source_module.register_forward_hook(replace_source) + target_hook = target_module.register_forward_hook(capture_target) + try: + model(**batch) + finally: + source_hook.remove() + target_hook.remove() + return target_states.pop()[valid].reshape(-1) + + jacobian = torch.autograd.functional.jacobian(target_from_source, source_value) + repeated_cotangent = cotangent.repeat(int(valid.sum())).to(jacobian.dtype) + gradient = torch.tensordot(repeated_cotangent, jacobian, dims=([0], [0])) + mean = gradient[valid].reshape(1, -1, gradient.size(-1)).mean(dim=1).squeeze(0) + total = mean if total is None else total + mean + return total / len(texts) + + +def _manual_target_mean(model, tokenizer, texts, target_layer): + module = model.model.layers[target_layer] + rows = [] + for text in texts: + captured = [] + handle = module.register_forward_hook( + lambda _m, _a, output: captured.append(output[0] if isinstance(output, tuple) else output) + ) + try: + with torch.no_grad(): + model(**tokenizer(text, return_tensors="pt")) + finally: + handle.remove() + rows.append(captured.pop()[0, -1].float()) + return torch.stack(rows).mean(dim=0) + + +def test_tiny_jacobian_matches_production_for_unequal_classes_and_cross_token_dependence(): + torch.manual_seed(7) + model = tiny_llama() + tokenizer = wordlevel_tokenizer() + source = _source() + vector = source.resolve(model, tokenizer) + + positive_mean = _manual_target_mean(model, tokenizer, _fit_data()["positives"], target_layer=2) + negative_mean = _manual_target_mean(model, tokenizer, _fit_data()["negatives"], target_layer=2) + cotangent = positive_mean - negative_mean + positive = _explicit_jacobian_direction( + model, tokenizer, _fit_data()["positives"], cotangent, + source_layer=0, target_layer=2, skip_first=0, + ) + negative = _explicit_jacobian_direction( + model, tokenizer, _fit_data()["negatives"], cotangent, + source_layer=0, target_layer=2, skip_first=0, + ) + expected = (positive - negative) / (positive - negative).norm() + assert torch.allclose(vector.directions[0].squeeze(0), expected, atol=1e-5) + + batch = tokenizer("the cat sat", return_tensors="pt") + source_module = model.model.layers[0] + target_module = model.model.layers[2] + captured = [] + capture_handle = source_module.register_forward_hook( + lambda _m, _a, output: captured.append(output[0] if isinstance(output, tuple) else output) + ) + try: + with torch.no_grad(): + model(**batch) + finally: + capture_handle.remove() + source_value = captured.pop().detach() + + def target_position_from_source(replacement): + target_states = [] + source_handle = source_module.register_forward_hook( + lambda _m, _a, output: _replace_hidden(output, replacement) + ) + target_handle = target_module.register_forward_hook( + lambda _m, _a, output: target_states.append(output[0] if isinstance(output, tuple) else output) + ) + try: + model(**batch) + finally: + source_handle.remove() + target_handle.remove() + return target_states.pop()[0, 2] + + jacobian = torch.autograd.functional.jacobian(target_position_from_source, source_value) + assert jacobian[:, 0, 1].abs().sum() > 0 + + +def test_tuple_decoder_outputs_are_not_replaced_by_extraction_hooks(): + model = tiny_gpt2() + vector = VJPDeltaFit(_fit_data(), target_layer=2, source_layer_ids=[0], skip_first=0).resolve( + model, wordlevel_tokenizer(), + ) + assert set(vector.directions) == {0} + + +def test_fit_restores_mode_flags_and_existing_parameter_grads_after_error(): + model = tiny_llama().train() + tokenizer = wordlevel_tokenizer() + parameters = list(model.parameters()) + for index, parameter in enumerate(parameters): + parameter.requires_grad_(index % 2 == 0) + parameter.grad = torch.full_like(parameter, float(index + 1)) + flags = [parameter.requires_grad for parameter in parameters] + grads = [parameter.grad.clone() for parameter in parameters] + weights = {name: value.detach().clone() for name, value in model.state_dict().items()} + + with pytest.raises(ValueError, match="valid token span"): + VJPDeltaFit(_fit_data(), target_layer=2, source_layer_ids=[0], skip_first=20).resolve(model, tokenizer) + + assert model.training is True + for parameter, flag, grad in zip(parameters, flags, grads, strict=True): + assert parameter.requires_grad is flag + assert torch.equal(parameter.grad, grad) + for name, value in weights.items(): + assert torch.equal(model.state_dict()[name], value) + + +def test_fit_restores_mode_flags_and_existing_parameter_grads_after_success(): + model = tiny_llama().train() + tokenizer = wordlevel_tokenizer() + parameters = list(model.parameters()) + for index, parameter in enumerate(parameters): + parameter.requires_grad_(index % 2 == 0) + parameter.grad = torch.full_like(parameter, float(index + 1)) + flags = [parameter.requires_grad for parameter in parameters] + grads = [parameter.grad.clone() for parameter in parameters] + weights = {name: value.detach().clone() for name, value in model.state_dict().items()} + + vector = _source().resolve(model, tokenizer) + + assert set(vector.directions) == {0, 1} + assert model.training is True + for parameter, flag, grad in zip(parameters, flags, grads, strict=True): + assert parameter.requires_grad is flag + assert torch.equal(parameter.grad, grad) + for name, value in weights.items(): + assert torch.equal(model.state_dict()[name], value) + + +def test_invalid_layer_order_and_zero_direction_fail_clearly(): + model = tiny_llama() + tokenizer = wordlevel_tokenizer() + with pytest.raises(ValueError, match="must precede"): + VJPDeltaFit(_fit_data(), target_layer=1, source_layer_ids=[1], skip_first=0).resolve(model, tokenizer) + repeated = {"positives": ["the cat sat"], "negatives": ["the cat sat"]} + with pytest.raises(ValueError, match="nonzero norm"): + VJPDeltaFit(repeated, target_layer=2, source_layer_ids=[0], skip_first=0).resolve(model, tokenizer) + with torch.inference_mode(), pytest.raises(RuntimeError, match="inference_mode"): + _source().resolve(model, tokenizer) + + +def test_autograd_error_keeps_the_original_diagnostic(monkeypatch): + def out_of_memory(*_args, **_kwargs): + raise RuntimeError("CUDA out of memory while allocating a gradient buffer") + + monkeypatch.setattr(torch.autograd, "grad", out_of_memory) + with pytest.raises(RuntimeError, match="VJP-delta autograd.grad failed: CUDA out of memory"): + _source().resolve(tiny_llama(), wordlevel_tokenizer()) + + +def test_additive_zero_strength_preserves_hidden_and_nonzero_adds_vector(): + torch.manual_seed(11) + tokenizer = wordlevel_tokenizer() + model = tiny_llama() + vector = SteeringVector(model_type="llama", directions={0: torch.ones(1, HIDDEN)}) + input_ids = tokenizer("the cat sat", return_tensors="pt")["input_ids"] + + baseline = SteeringPipeline(model=model, tokenizer=tokenizer, controls=[]) + baseline.steer() + zero = SteeringPipeline( + model=model, tokenizer=tokenizer, + controls=[VJPDelta(steering_vector=vector, strength=0.0, token_scope="all")], + ) + zero.steer() + nonzero = SteeringPipeline( + model=model, tokenizer=tokenizer, + controls=[VJPDelta(steering_vector=vector, strength=2.0, token_scope="all")], + ) + nonzero.steer() + + baseline_hidden = _capture_layer(model, baseline, 0, input_ids) + zero_hidden = _capture_layer(model, zero, 0, input_ids) + nonzero_hidden = _capture_layer(model, nonzero, 0, input_ids) + assert torch.equal(zero_hidden, baseline_hidden) + assert torch.allclose(nonzero_hidden - baseline_hidden, torch.full_like(baseline_hidden, 2.0), atol=1e-6) + + +def test_pipeline_freeze_reload_skips_vjp_and_strength_does_not_stale(tmp_path, monkeypatch): + model = tiny_llama() + tokenizer = wordlevel_tokenizer() + control = VJPDelta( + data=_fit_data(), target_layer=2, source_layer_ids=[0, 1], skip_first=0, strength=0.5, + ) + pipeline = SteeringPipeline(model=model, tokenizer=tokenizer, controls=[control], model_name_or_path="tiny-llama") + pipeline.steer() + reference = pipeline.generate(text="the cat", max_new_tokens=3, do_sample=False) + batch = pipeline.generate(text=["the cat", "the dog"], max_new_tokens=3, do_sample=False) + assert isinstance(batch, list) and len(batch) == 2 + saved = pipeline.to_spipe().save(tmp_path / "vjp") + + rebuilt = SPipe.load(saved).pipeline() + assert isinstance(rebuilt.state_controls[0], ActivationAdapter) + assert rebuilt.state_controls[0].steer_fits() == () + monkeypatch.setattr(VJPDeltaFit, "resolve", lambda *_a, **_k: pytest.fail("frozen reload ran a VJP fit")) + rebuilt.model, rebuilt.tokenizer = model, tokenizer + rebuilt.steer() + assert rebuilt.generate(text="the cat", max_new_tokens=3, do_sample=False) == reference + + manifest = json.loads((saved / "spipe.json").read_text()) + manifest["controls"][0]["args"]["strength"] = 2.0 + (saved / "spipe.json").write_text(json.dumps(manifest)) + assert isinstance(SPipe.load(saved).pipeline().state_controls[0], ActivationAdapter) + + +def test_recipe_edit_is_stale_and_frozen_vector_uses_existing_mismatch_policy(tmp_path): + model = tiny_llama() + tokenizer = wordlevel_tokenizer() + pipeline = SteeringPipeline( + model=model, tokenizer=tokenizer, + controls=[VJPDelta(data=_fit_data(), target_layer=2, source_layer_ids=[0], skip_first=0)], + model_name_or_path="tiny-llama", + ) + pipeline.steer() + saved = pipeline.to_spipe().save(tmp_path / "vjp") + manifest_path = saved / "spipe.json" + manifest = json.loads(manifest_path.read_text()) + manifest["controls"][0]["args"]["data"]["fields"]["positives"] = ["the cat sat", "the cat sat"] + manifest_path.write_text(json.dumps(manifest)) + with pytest.raises(SpipeStaleError): + SPipe.load(saved) + + vector = SteeringVector(model_type="llama", directions={0: torch.ones(1, HIDDEN)}) + precomputed = SteeringPipeline( + model=model, tokenizer=tokenizer, controls=[VJPDelta(steering_vector=vector)], model_name_or_path="tiny", + ) + precomputed.steer() + assert precomputed.state_controls[0].steer_fits() == () + saved_precomputed = precomputed.to_spipe().save(tmp_path / "precomputed") + changed = tiny_llama() + changed.load_state_dict(model.state_dict()) + with torch.no_grad(): + next(changed.parameters()).add_(0.01) + same_architecture = SPipe.load(saved_precomputed).pipeline() + same_architecture.model, same_architecture.tokenizer = changed, tokenizer + with pytest.warns(UserWarning, match="direction artifact"): + same_architecture.steer() + + frozen = SPipe.load(saved_precomputed).pipeline() + frozen.model, frozen.tokenizer = tiny_gpt2(), tokenizer + with pytest.raises(ValueError, match="model_type"): + frozen.steer() + + +def test_fit_holds_only_a_weak_model_reference_and_control_cleanup_drops_bound_tensors(): + model = tiny_llama() + tokenizer = wordlevel_tokenizer() + source = _source() + source.resolve(model, tokenizer) + ref = weakref.ref(model) + assert source._model_ref is ref or source._model_ref() is model + control = VJPDelta(data=_fit_data(), target_layer=2, source_layer_ids=[0], skip_first=0) + control.steer(model, tokenizer) + assert control.steer_fits() == (("VJPDeltaFit", "direction"),) + assert control.interventions + control.cleanup() + assert control.interventions == () + del model + gc.collect() + assert ref() is None + with pytest.raises(ValueError, match="requires a live model"): + source.resolve(None, tokenizer)