"""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']}"}