| """Averaged input-output Jacobian estimator. |
| |
| This matches Anthropic's released Jacobian Lens estimator (Apache-2.0): each |
| cotangent row is injected at every valid target position, and source-position |
| gradients are averaged. ``dim_batch`` rows are computed in parallel. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import math |
| import os |
| from collections.abc import Sequence |
| from pathlib import Path |
|
|
| import torch |
| from tqdm import tqdm |
|
|
| from .corpus import CorpusItem |
| from .hooks import ActivationRecorder |
|
|
|
|
| def valid_position_mask(seq_len: int, skip_first: int = 16) -> torch.Tensor: |
| if skip_first < 0: |
| raise ValueError("skip_first must be nonnegative") |
| mask = torch.zeros(seq_len, dtype=torch.bool) |
| mask[skip_first:seq_len - 1] = True |
| if not mask.any(): |
| raise ValueError(f"sequence length {seq_len} leaves no positions after skip_first={skip_first}") |
| return mask |
|
|
|
|
| def jacobian_for_tokens( |
| model, |
| input_ids: torch.Tensor, |
| *, |
| source_layers: Sequence[int], |
| target_layer: int, |
| dim_batch: int = 8, |
| skip_first: int = 16, |
| ) -> tuple[dict[int, torch.Tensor], int]: |
| """Return per-prompt FP32 CPU Jacobians and valid-position count.""" |
| sources = sorted(set(source_layers)) |
| if not sources or sources[0] < 0 or sources[-1] >= target_layer: |
| raise ValueError("source layers must be nonempty and below target_layer") |
| if target_layer >= model.n_layers: |
| raise ValueError("target_layer is outside the model") |
| if dim_batch < 1: |
| raise ValueError("dim_batch must be positive") |
|
|
| width = model.d_model |
| mask = valid_position_mask(input_ids.shape[1], skip_first) |
| jacobians = {layer: torch.zeros(width, width, dtype=torch.float32) for layer in sources} |
| passes = math.ceil(width / dim_batch) |
|
|
| with ActivationRecorder(model.layers, at=[*sources, target_layer], start_graph_at=sources[0]) as recorder, torch.enable_grad(): |
| model.forward(input_ids.expand(dim_batch, -1)) |
| target = recorder.activations[target_layer] |
| source_activations = [recorder.activations[layer] for layer in sources] |
| positions = mask.nonzero(as_tuple=True)[0].to(target.device) |
| batch_indices = torch.arange(dim_batch, device=target.device) |
| cotangent = torch.zeros_like(target) |
|
|
| for pass_index, start in enumerate(range(0, width, dim_batch)): |
| row_count = min(dim_batch, width - start) |
| cotangent.zero_() |
| cotangent[ |
| batch_indices[:row_count, None], positions[None, :], |
| start + batch_indices[:row_count, None], |
| ] = 1 |
| gradients = torch.autograd.grad( |
| outputs=target, inputs=source_activations, grad_outputs=cotangent, |
| retain_graph=pass_index < passes - 1, |
| ) |
| for layer, gradient in zip(sources, gradients, strict=True): |
| local_positions = positions.to(gradient.device) |
| rows = gradient[:row_count, local_positions, :].float().mean(dim=1) |
| jacobians[layer][start:start + row_count] = rows.cpu() |
| return jacobians, int(mask.sum()) |
|
|
|
|
| def _atomic_save(state: dict, path: Path) -> None: |
| path.parent.mkdir(parents=True, exist_ok=True) |
| temporary = path.with_suffix(path.suffix + f".tmp.{os.getpid()}") |
| torch.save(state, temporary) |
| os.replace(temporary, path) |
|
|
|
|
| def fit( |
| model, |
| prompts: Sequence[CorpusItem], |
| *, |
| output_path: str, |
| checkpoint_path: str, |
| source_layers: Sequence[int] | None = None, |
| target_layer: int | None = None, |
| max_seq_len: int = 128, |
| dim_batch: int = 8, |
| skip_first: int = 16, |
| checkpoint_every: int = 1, |
| resume: bool = True, |
| export_dtype: torch.dtype = torch.bfloat16, |
| ) -> dict: |
| """Fit, checkpoint FP32 sums, and export BF16 averaged matrices.""" |
| target = model.n_layers - 1 if target_layer is None else target_layer |
| sources = list(range(target)) if source_layers is None else sorted(set(source_layers)) |
| checkpoint = Path(checkpoint_path) |
| metadata_keys = ("source_layers", "target_layer", "d_model", "max_seq_len", "skip_first") |
| expected = (sources, target, model.d_model, max_seq_len, skip_first) |
|
|
| if resume and checkpoint.exists(): |
| state = torch.load(checkpoint, map_location="cpu", weights_only=True) |
| actual = tuple(state[key] for key in metadata_keys) |
| if actual != expected: |
| raise ValueError(f"checkpoint configuration mismatch: {actual} != {expected}") |
| else: |
| state = { |
| "jacobian_sum": {layer: torch.zeros(model.d_model, model.d_model, dtype=torch.float32) for layer in sources}, |
| "n_done": 0, "next_index": 0, "source_layers": sources, |
| "target_layer": target, "d_model": model.d_model, |
| "max_seq_len": max_seq_len, "skip_first": skip_first, |
| } |
|
|
| progress = tqdm(range(state["next_index"], len(prompts)), desc="fitting prompts") |
| for index in progress: |
| item = prompts[index] |
| if isinstance(item, str): |
| input_ids = model.encode(item, max_seq_len) |
| else: |
| input_ids = torch.tensor( |
| [item[:max_seq_len]], |
| dtype=torch.long, |
| device=model.input_device, |
| ) |
| try: |
| per_prompt, n_valid = jacobian_for_tokens( |
| model, input_ids, source_layers=sources, target_layer=target, |
| dim_batch=dim_batch, skip_first=skip_first, |
| ) |
| except ValueError as exc: |
| progress.write(f"skipping prompt {index}: {exc}") |
| state["next_index"] = index + 1 |
| continue |
| for layer in sources: |
| state["jacobian_sum"][layer].add_(per_prompt[layer]) |
| state["n_done"] += 1 |
| state["next_index"] = index + 1 |
| progress.set_postfix(valid=n_valid, completed=state["n_done"]) |
| if checkpoint_every and state["n_done"] % checkpoint_every == 0: |
| _atomic_save(state, checkpoint) |
|
|
| _atomic_save(state, checkpoint) |
| if state["n_done"] == 0: |
| raise ValueError("no usable prompts were fitted") |
| exported = { |
| layer: (total / state["n_done"]).to(export_dtype) |
| for layer, total in state["jacobian_sum"].items() |
| } |
| |
| |
| exported[target] = torch.eye(model.d_model, dtype=export_dtype) |
| result = { |
| "J": exported, |
| "n_prompts": state["n_done"], "source_layers": sources, |
| "target_layer": target, "d_model": model.d_model, |
| "max_seq_len": max_seq_len, "skip_first": skip_first, |
| "dtype": str(export_dtype), "target_is_identity": True, |
| } |
| _atomic_save(result, Path(output_path)) |
| return result |
|
|