Jacobina / api /lens_api.py
marinarosa's picture
Bundle fitted lenses and improve lens controls
abb8712
Raw
History Blame Contribute Delete
19.1 kB
"""Lens view API: layer-by-layer readouts (logit / Jacobian / diff) with
position and layer selection, aggregated readout browsing, multi-intervention
steering, pinned-token rank tracking, and fit-file management.
Semantics preserved from upstream: interventions are applied at generation
time. "Update readouts" re-slices the existing sequence under the
interventions it was generated with; changing the intervention list requires
Generate again (the status line says so). This keeps the displayed text and
the displayed readouts always consistent with each other.
The analysis (input_ids, intervention set, cached activations) and the last
computed bundle live server-side in ``state.get_lens_session()``; the position
selection and pinned-token list are plain data the frontend holds and sends
with each call.
"""
from __future__ import annotations
from collections.abc import Iterator
from api.helpers import (
ChatValidationError,
add_unique_intervention_rows,
enabled_interventions,
intervened_layer_titles,
intervention_visibility_warning,
interventions_summary,
layer_selection,
parse_layer_refs,
token_ref_to_id,
)
from api.interactive import _reset_tracer_for_mode
from api.lens_views import distribution_html, heatmap_html, readouts_table_html
from api.models import model_manager
from api.serialize import error_payload, fig_json, intervention_rows_payload
from api.state import (
get_intervention_rows,
get_lens_session,
set_intervention_rows,
)
from miru_tracer.core.interventions import Intervention
from miru_tracer.core.lens import (
LEGACY_LENS_FILENAME,
aggregate_readouts,
compute_lens_slice,
decode_token,
get_lens_store,
record_lens_activations,
)
from miru_tracer.core.lens_io import load_lens, save_lens
from miru_tracer.core.logging_config import get_logger
from miru_tracer.core.sampling import SamplingParams
from miru_tracer.core.tokenizer_utils import visible_whitespace
from miru_tracer.core.tracer import LLMTracer
from miru_tracer.visualization.plots import plot_pinned_token_ranks
logger = get_logger(__name__)
VIEWS = ("summary", "readouts", "heatmap", "pinned")
def _token_label(text: str) -> str:
shown = visible_whitespace(text)
return shown if shown.strip() else "·"
def _tokens_payload(position_texts: list[str]) -> list[dict]:
return [
{"index": i, "label": _token_label(text), "text": text}
for i, text in enumerate(position_texts)
]
def _compute_bundle(
analysis, positions, mode, l_start, l_end, l_stride, per_cell,
skip_non_words, pinned_ids,
):
"""One lens slice over the stored sequence -> (bundle, status)."""
if analysis is None:
return None, "Generate first."
model = model_manager.get_model()
tokenizer = model_manager.get_tokenizer()
if model is None or analysis["model_name"] != model_manager.get_model_name():
return None, (
"Error: the model changed since this sequence was generated. "
"Generate again."
)
jlens = get_lens_store().get(analysis["model_name"])
if mode in ("jacobian", "diff") and jlens is None:
return None, (
"Error: no fitted Jacobian lens for "
f"{analysis['model_name']}. Upload one below or run:\n"
f" miru-tracer-fit-lens {analysis['model_name']}"
)
try:
layers = layer_selection(analysis["n_layers"], l_start, l_end, l_stride)
if mode in ("jacobian", "diff") and jlens is not None:
fitted = set(jlens.source_layers) | {analysis["n_layers"] - 1}
dropped = [layer for layer in layers if layer not in fitted]
layers = [layer for layer in layers if layer in fitted]
if not layers:
return None, (
f"Error: none of the selected layers are fitted "
f"(lens covers {jlens.source_layers})."
)
else:
dropped = []
seq_len = int(analysis["input_ids"].shape[1])
selected = [p for p in positions or [] if 0 <= p < seq_len] or None
# The residuals only depend on (sequence, interventions), both frozen
# per analysis — record once, then every Update is unembed + top-k
# with no model forward.
activations = analysis.get("activations")
if activations is None:
activations = record_lens_activations(
model,
tokenizer,
analysis["input_ids"],
interventions=analysis["iset"],
)
analysis["activations"] = activations
readout_limit = max(1, int(per_cell))
slice_ = compute_lens_slice(
model,
tokenizer,
analysis["input_ids"],
layers=layers,
positions=selected,
mode=mode,
jlens=jlens,
top_k=readout_limit,
skip_non_words=bool(skip_non_words),
pinned_token_ids=[int(t) for t in pinned_ids or []],
interventions=analysis["iset"],
activations=activations,
)
rows = aggregate_readouts(slice_, limit=readout_limit)
n_cells = len(slice_.layers) * len(slice_.positions)
where = (
f"{len(slice_.positions)} selected positions"
if selected is not None
else f"all {len(slice_.positions)} positions"
)
status = (
f"{slice_.mode} lens: {len(slice_.layers)} layers × {where} "
f"= {n_cells} cells, {len(rows)} distinct readout tokens."
)
if dropped:
status += f" (skipped unfitted layers: {dropped})"
intervened: dict[int, str] = {}
if analysis["iset"] is not None:
ivs = analysis["iset"].interventions
intervened = intervened_layer_titles(ivs, tokenizer)
status += f" Interventions: {interventions_summary(ivs, tokenizer)}."
warning = intervention_visibility_warning(
ivs, mode, analysis["n_layers"], tokenizer
)
if warning:
status += f"\n{warning}"
return {"slice": slice_, "rows": rows, "intervened": intervened}, status
except Exception as e:
logger.error(f"Lens readout error: {e}", exc_info=True)
return None, f"Error: {e}"
def _render_view(bundle, view: str) -> dict:
"""One result view from the cached bundle: html or a Plotly figure."""
if bundle is None:
return {"view": view, "html": "", "figure": None}
slice_ = bundle["slice"]
intervened = bundle.get("intervened") or None
# JSON round-trips string keys; the views expect int layer keys.
if intervened:
intervened = {int(k): v for k, v in intervened.items()}
if view == "readouts":
return {
"view": view,
"html": distribution_html(bundle["rows"], slice_.layers, intervened=intervened),
"figure": None,
}
if view == "heatmap":
return {"view": view, "html": heatmap_html(slice_, intervened), "figure": None}
if view == "pinned":
fig = plot_pinned_token_ranks(slice_, model_manager.get_tokenizer())
return {"view": view, "html": "", "figure": fig_json(fig)}
return {
"view": "summary",
"html": readouts_table_html(bundle["rows"], intervened, layers=slice_.layers),
"figure": None,
}
def lens_generate(
mode: str, prompt: str, chat_json: str, raw_text: str,
thinking: str, think_prefill: str,
n_tokens: int, strategy: str, temperature: float,
lens_mode: str, layer_start: int, layer_end: int, layer_stride: int,
per_cell: int, skip_non_words: bool, pinned_ids: list[int],
active_view: str,
) -> Iterator[dict]:
"""Generate (optionally with the enabled interventions) then analyze."""
lens = get_lens_session()
model = model_manager.get_model()
tokenizer = model_manager.get_tokenizer()
device = model_manager.get_device()
if model is None or tokenizer is None:
yield error_payload("No model loaded. Use the Model view.")
return
try:
with lens.lock:
tracer = LLMTracer(model, tokenizer, device)
jlens = get_lens_store().get(model_manager.get_model_name())
active_interventions = enabled_interventions(get_intervention_rows())
try:
tracer.set_interventions(active_interventions or None, jlens=jlens)
except ValueError as e:
yield error_payload(
f"Error in interventions: {e}\n"
"(jacobian basis needs a fitted lens covering that layer "
"— upload one below, or use the logit basis)"
)
return
_reset_tracer_for_mode(
tracer, mode, prompt, chat_json, raw_text, thinking, think_prefill
)
params = SamplingParams(strategy=strategy, temperature=float(temperature))
if n_tokens and int(n_tokens) > 0:
for _step in tracer.generate_stream(
max_new_tokens=int(n_tokens), params=params
):
yield {
"ok": True,
"type": "progress",
"status": f"Generating... {len(tracer.history)}/{int(n_tokens)}",
"text": tracer.get_full_text(),
}
position_texts = [
decode_token(tokenizer, int(t)) for t in tracer.input_ids[0]
]
analysis = {
"input_ids": tracer.input_ids.clone(),
"model_name": model_manager.get_model_name(),
"iset": tracer._intervention_set,
"n_layers": model.config.get_text_config().num_hidden_layers,
"prompt_len": tracer._prompt_len,
"position_texts": position_texts,
}
bundle, status = _compute_bundle(
analysis, [], lens_mode, layer_start, layer_end, layer_stride,
per_cell, skip_non_words, pinned_ids,
)
lens.analysis = analysis
lens.bundle = bundle
yield {
"ok": True,
"type": "final",
"status": status,
"text": tracer.get_full_text(),
"tokens": _tokens_payload(position_texts),
"prompt_len": analysis["prompt_len"],
"n_layers": analysis["n_layers"],
"positions": [],
**_render_view(bundle, active_view),
}
except ChatValidationError as e:
yield error_payload(str(e))
except Exception as e:
logger.error(f"Lens generate error: {e}", exc_info=True)
yield error_payload(str(e), trace=True)
def lens_update(
positions: list[int],
lens_mode: str, layer_start: int, layer_end: int, layer_stride: int,
per_cell: int, skip_non_words: bool, pinned_ids: list[int],
active_view: str,
) -> dict:
"""Re-slice the stored sequence; interventions only change on Generate."""
lens = get_lens_session()
with lens.lock:
bundle, status = _compute_bundle(
lens.analysis, positions, lens_mode, layer_start, layer_end,
layer_stride, per_cell, skip_non_words, pinned_ids,
)
if bundle is not None:
lens.bundle = bundle
payload = {"ok": bundle is not None, "status": status}
if bundle is None:
payload["error"] = status
return payload
return {**payload, **_render_view(bundle, active_view)}
def lens_view(view: str) -> dict:
"""Render one result view from the cached bundle (result-tab switch)."""
if view not in VIEWS:
return error_payload(f"Unknown view: {view}")
lens = get_lens_session()
with lens.lock:
return {"ok": True, **_render_view(lens.bundle, view)}
# ------------------------------------------------------------- interventions
def _rows_status(rows: list[dict], action: str) -> dict:
tokenizer = model_manager.get_tokenizer()
return {
"ok": True,
"rows": intervention_rows_payload(rows, tokenizer),
"status": (
f"{action} {len(enabled_interventions(rows))} enabled "
"intervention(s) — regenerate to apply."
),
}
def lens_add_intervention(
kind: str, token_ref: str, token_mode: str,
swap_to_ref: str, swap_to_mode: str,
layer_refs: str, strength: float, basis: str,
) -> dict:
tokenizer = model_manager.get_tokenizer()
if tokenizer is None:
return error_payload("No model loaded.")
try:
token_id = token_ref_to_id(token_ref, tokenizer, token_mode)
token_id_to = (
token_ref_to_id(swap_to_ref, tokenizer, swap_to_mode)
if kind == "swap"
else None
)
effective_strength = float(strength) if kind == "steer" else 0.0
candidates = [
{
"enabled": True,
"intervention": Intervention(
kind=kind,
layer=layer,
token_id=token_id,
strength=effective_strength,
token_id_to=token_id_to,
basis=basis,
),
}
for layer in parse_layer_refs(layer_refs)
]
rows = get_intervention_rows()
updated, added, skipped = add_unique_intervention_rows(rows, candidates)
set_intervention_rows(updated)
if added:
described = added[0]["intervention"].describe(tokenizer)
if len(added) > 1:
described += f" … (+{len(added) - 1} more layers)"
action = f"Added: {described}."
if skipped:
action += f" Skipped {skipped} duplicate intervention(s)."
else:
action = f"Skipped {skipped} duplicate intervention(s)."
return _rows_status(updated, action)
except ValueError as e:
return error_payload(str(e))
def lens_intervention_action(action: str, index: int, enabled: bool) -> dict:
"""Toggle or delete one intervention row."""
rows = get_intervention_rows()
try:
index = int(index)
except (TypeError, ValueError):
return _rows_status(rows, "Ignored invalid intervention action.")
if index < 0 or index >= len(rows):
return _rows_status(rows, "Ignored invalid intervention action.")
if action == "toggle":
rows[index] = {**rows[index], "enabled": bool(enabled)}
state = "enabled" if rows[index]["enabled"] else "disabled"
set_intervention_rows(rows)
return _rows_status(rows, f"Intervention {index} {state}.")
if action == "delete":
rows = [row for i, row in enumerate(rows) if i != index]
set_intervention_rows(rows)
return _rows_status(rows, f"Deleted intervention {index}.")
return _rows_status(rows, "Ignored invalid intervention action.")
def lens_clear_interventions() -> dict:
set_intervention_rows([])
return _rows_status([], "Cleared all interventions.")
def lens_list_interventions() -> dict:
"""Current rows (restores the table after a page load)."""
rows = get_intervention_rows()
tokenizer = model_manager.get_tokenizer()
return {"ok": True, "rows": intervention_rows_payload(rows, tokenizer)}
# ------------------------------------------------------------- pinned tokens
def lens_resolve_pinned(token_ref: str, token_mode: str, pinned_ids: list[int]) -> dict:
"""Resolve one pinned-token reference and append it (deduped, order kept)."""
tokenizer = model_manager.get_tokenizer()
if tokenizer is None:
return error_payload("No model loaded.")
try:
token_id = token_ref_to_id(token_ref, tokenizer, token_mode)
updated = [int(t) for t in pinned_ids or []]
if token_id not in updated:
updated.append(token_id)
pinned = [
{
"token_id": t,
"token": tokenizer.convert_ids_to_tokens([t])[0],
"decoded": tokenizer.decode([t]),
}
for t in updated
]
return {"ok": True, "pinned": pinned, "status": f"{len(updated)} pinned token(s)."}
except ValueError as e:
return error_payload(str(e))
# ------------------------------------------------------------------ fit file
def lens_fit_status() -> dict:
"""Describe the fit-file situation for the loaded model."""
model_name = model_manager.get_model_name()
if model_name is None:
return {"ok": True, "status": "No model loaded."}
store = get_lens_store()
lens = store.get(model_name)
if lens is None:
return {
"ok": True,
"fitted": False,
"status": (
f"No fitted lens for {model_name}.\n"
f"Expected at: {store.lens_path(model_name)}\n"
f"Fit one on a GPU instance: miru-tracer-fit-lens {model_name}"
),
}
return {
"ok": True,
"fitted": True,
"status": (
f"Fitted lens loaded for {model_name}: "
f"{len(lens.source_layers)} layers "
f"(L{lens.source_layers[0]}..L{lens.source_layers[-1]}), "
f"averaged over {lens.n_prompts} prompts.\n"
f"Path: {store.existing_lens_path(model_name)}"
),
}
def install_fit_file(filepath: str) -> dict:
"""Validate an uploaded fit file and install it for the loaded model.
Called from the plain FastAPI upload route in app.py.
"""
model = model_manager.get_model()
model_name = model_manager.get_model_name()
if model is None:
return error_payload("Load a model first — the fit file is stored per model.")
try:
lens = load_lens(filepath)
except Exception as e:
return error_payload(f"Not a valid lens file: {e}")
d_model = model.config.get_text_config().hidden_size
n_layers = model.config.get_text_config().num_hidden_layers
if lens.d_model != d_model:
return error_payload(
f"Fit file has d_model={lens.d_model}, but {model_name} has "
f"d_model={d_model}. This lens was fitted for a different model."
)
if lens.source_layers[-1] >= n_layers:
return error_payload(
f"Fit file covers layer {lens.source_layers[-1]}, but "
f"{model_name} only has {n_layers} layers."
)
path = get_lens_store().lens_path(model_name)
path.parent.mkdir(parents=True, exist_ok=True)
save_lens(lens, path)
# A stale legacy pickle next to the fresh safetensors would keep the
# unsafe copy around — drop it.
path.with_name(LEGACY_LENS_FILENAME).unlink(missing_ok=True)
logger.info(f"Installed fit file for {model_name} at {path}")
return {"ok": True, "status": f"Installed.\n{lens_fit_status()['status']}"}