J-space / math_jlens /fitting.py
ayh015's picture
Upload folder using huggingface_hub
f6158c7 verified
Raw
History Blame Contribute Delete
6.86 kB
"""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