"""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() } # The transport from the target block output to itself is exactly identity. # Include it so consumers have one matrix for every block-output boundary. 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