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