Spaces:
Sleeping
Sleeping
| """Interactive mode: step-through generation with session-based state. | |
| The frontend carries only the session id; the live tracer lives in the core | |
| SessionManager. Every handler resolves the session, takes its lock, mutates | |
| the tracer, and returns the canonical payload via ``serialize.render_state``. | |
| """ | |
| from __future__ import annotations | |
| from collections.abc import Iterator | |
| from api.helpers import ( | |
| ChatValidationError, | |
| layer_selection, | |
| parse_chat_messages, | |
| ui_sampling_params, | |
| ) | |
| from api.models import model_manager | |
| from api.serialize import error_payload, fig_json, render_state | |
| from api.state import get_active_interventions | |
| from miru_tracer.core.lens import compute_lens_slice, get_lens_store | |
| from miru_tracer.core.logging_config import get_logger | |
| from miru_tracer.core.session_manager import get_session_manager | |
| from miru_tracer.visualization.plots import plot_lens_heatmap | |
| logger = get_logger(__name__) | |
| def _resolve_session(session_id: str | None): | |
| """Common session lookup; returns (session, error_payload_or_None).""" | |
| if not session_id: | |
| return None, error_payload("Not initialized. Click 'Initialize' first.") | |
| session = get_session_manager().get_session(session_id) | |
| if session is None: | |
| return None, error_payload("Session not found. Please reinitialize.") | |
| return session, None | |
| def _reset_tracer_for_mode( | |
| tracer, mode: str, prompt: str, chat_json: str, raw_text: str, | |
| thinking: str, think_prefill: str, | |
| ) -> None: | |
| """Shared reset dispatch (raises ChatValidationError / ValueError).""" | |
| if mode == "chat": | |
| tracer.reset( | |
| messages=parse_chat_messages(chat_json), | |
| mode="chat", | |
| thinking=thinking or "auto", | |
| think_prefill=think_prefill or "", | |
| ) | |
| elif mode == "raw": | |
| tracer.reset(prompt=raw_text, mode="raw") | |
| else: | |
| tracer.reset(prompt=prompt, mode="completion") | |
| def interactive_init( | |
| mode: str, prompt: str, chat_json: str, raw_text: str, | |
| thinking: str, think_prefill: str, | |
| strategy: str, temperature: float, top_k: int, top_p: float, | |
| log_top_k: int, | |
| ) -> dict: | |
| model = model_manager.get_model() | |
| tokenizer = model_manager.get_tokenizer() | |
| device = model_manager.get_device() | |
| if model is None or tokenizer is None: | |
| return error_payload("No model loaded") | |
| try: | |
| session_manager = get_session_manager() | |
| session_id = session_manager.create_session(model, tokenizer, device) | |
| session = session_manager.get_session(session_id) | |
| with session.lock: | |
| _reset_tracer_for_mode( | |
| session.tracer, mode, prompt, chat_json, raw_text, | |
| thinking, think_prefill, | |
| ) | |
| params = ui_sampling_params(strategy, temperature, top_k, top_p) | |
| logger.info(f"Interactive session initialized: {session_id} (mode={mode})") | |
| return render_state( | |
| session_id, session.tracer, f"Initialized in {mode} mode", | |
| params, log_top_k, | |
| ) | |
| except ChatValidationError as e: | |
| return error_payload(str(e)) | |
| except Exception as e: | |
| logger.error(f"Initialize error: {e}", exc_info=True) | |
| return error_payload(str(e), trace=True) | |
| def interactive_reset(session_id: str) -> dict: | |
| if session_id: | |
| get_session_manager().delete_session(session_id) | |
| logger.info(f"Interactive session reset: {session_id}") | |
| return { | |
| "ok": True, | |
| "status": "Reset complete. Click 'Initialize' to start a new generation.", | |
| "session_id": None, | |
| "text": "", | |
| "step": 0, | |
| "candidates": [], | |
| "preview_id": None, | |
| "eos": False, | |
| } | |
| def interactive_step( | |
| session_id: str, | |
| strategy: str, temperature: float, top_k: int, top_p: float, | |
| selected_token_id: int | None, | |
| override_enabled: bool, override_id: int | None, | |
| log_top_k: int, log_full_probs: bool, stop_at_eos: bool, | |
| ) -> dict: | |
| session, error = _resolve_session(session_id) | |
| if error: | |
| return error | |
| with session.lock: | |
| tracer = session.tracer | |
| try: | |
| params = ui_sampling_params(strategy, temperature, top_k, top_p) | |
| if override_enabled: | |
| if override_id is None: | |
| return error_payload("Override enabled but no token ID provided") | |
| token_id = int(override_id) | |
| if not 0 <= token_id < len(tracer.tokenizer): | |
| return error_payload( | |
| f"Token ID {token_id} is out of range " | |
| f"(vocab size: {len(tracer.tokenizer)})" | |
| ) | |
| else: | |
| token_id = ( | |
| int(selected_token_id) if selected_token_id is not None else None | |
| ) | |
| step_data = tracer.step( | |
| params, | |
| token_id=token_id, | |
| log_top_k=max(int(log_top_k or 10), 1), | |
| log_full_probs=bool(log_full_probs), | |
| ) | |
| if stop_at_eos and tracer.is_eos(step_data.token_id): | |
| logger.info( | |
| f"EOS reached: session={session_id}, steps={len(tracer.history)}" | |
| ) | |
| return { | |
| "ok": True, | |
| "status": ( | |
| f"Generation complete (EOS reached)\n" | |
| f"Total steps: {len(tracer.history)}" | |
| ), | |
| "session_id": session_id, | |
| "text": tracer.get_full_text(), | |
| "step": len(tracer.history), | |
| "candidates": [], | |
| "preview_id": None, | |
| "eos": True, | |
| } | |
| status = ( | |
| f"Step {len(tracer.history)} complete\n" | |
| f"Generated: {step_data.token_text} (p={step_data.probability:.4f})" | |
| ) | |
| return render_state(session_id, tracer, status, params, log_top_k) | |
| except Exception as e: | |
| logger.error(f"Step error: {e}", exc_info=True) | |
| return error_payload(str(e), trace=True) | |
| def interactive_undo( | |
| session_id: str, | |
| strategy: str, temperature: float, top_k: int, top_p: float, log_top_k: int, | |
| ) -> dict: | |
| session, error = _resolve_session(session_id) | |
| if error: | |
| return error | |
| with session.lock: | |
| tracer = session.tracer | |
| try: | |
| if not tracer.undo(): | |
| return error_payload("No steps to undo") | |
| params = ui_sampling_params(strategy, temperature, top_k, top_p) | |
| return render_state( | |
| session_id, tracer, | |
| f"Undone last step. Current steps: {len(tracer.history)}", | |
| params, log_top_k, | |
| ) | |
| except Exception as e: | |
| logger.error(f"Undo error: {e}", exc_info=True) | |
| return error_payload(str(e), trace=True) | |
| def interactive_goto( | |
| session_id: str, target_step: int, | |
| strategy: str, temperature: float, top_k: int, top_p: float, log_top_k: int, | |
| ) -> dict: | |
| session, error = _resolve_session(session_id) | |
| if error: | |
| return error | |
| with session.lock: | |
| tracer = session.tracer | |
| current_steps = len(tracer.history) | |
| try: | |
| if target_step is None or target_step < 0: | |
| return error_payload( | |
| f"Target step must be 0 or greater. Current step: {current_steps}" | |
| ) | |
| if target_step > current_steps: | |
| return error_payload( | |
| f"Target step {int(target_step)} is beyond current step " | |
| f"{current_steps}" | |
| ) | |
| tracer.goto_step(int(target_step)) | |
| params = ui_sampling_params(strategy, temperature, top_k, top_p) | |
| undone = current_steps - int(target_step) | |
| status = ( | |
| f"Already at step {int(target_step)}" | |
| if undone == 0 | |
| else f"Went back to step {int(target_step)} (undid {undone} steps)" | |
| ) | |
| return render_state(session_id, tracer, status, params, log_top_k) | |
| except Exception as e: | |
| logger.error(f"Go-to-step error: {e}", exc_info=True) | |
| return error_payload(str(e), trace=True) | |
| def interactive_continue( | |
| session_id: str, | |
| strategy: str, temperature: float, top_k: int, top_p: float, | |
| n_tokens: int, log_top_k: int, log_full_probs: bool, stop_at_eos: bool, | |
| ) -> Iterator[dict]: | |
| """Run N steps, streaming progress; stops cooperatively via request_stop.""" | |
| session, error = _resolve_session(session_id) | |
| if error: | |
| yield error | |
| return | |
| if n_tokens is None or n_tokens < 1: | |
| yield error_payload("Number of tokens must be at least 1") | |
| return | |
| with session.lock: | |
| tracer = session.tracer | |
| try: | |
| params = ui_sampling_params(strategy, temperature, top_k, top_p) | |
| tracer.clear_stop_flag() | |
| logger.info(f"Continue generation: session={session_id}, n_tokens={n_tokens}") | |
| stopped_reason = None | |
| for i in range(int(n_tokens)): | |
| if tracer._stop_requested: | |
| stopped_reason = f"Generation stopped by user after {i} tokens" | |
| break | |
| step_data = tracer.step( | |
| params, | |
| log_top_k=max(int(log_top_k or 10), 1), | |
| log_full_probs=bool(log_full_probs), | |
| ) | |
| if stop_at_eos and tracer.is_eos(step_data.token_id): | |
| stopped_reason = "Generation complete (EOS reached)" | |
| break | |
| yield { | |
| "ok": True, | |
| "type": "progress", | |
| "status": ( | |
| f"Generating... Step {len(tracer.history)} " | |
| f"({i + 1}/{int(n_tokens)})" | |
| ), | |
| "text": tracer.get_full_text(), | |
| "step": len(tracer.history), | |
| } | |
| status = stopped_reason or "Continue complete" | |
| status += f"\nTotal steps: {len(tracer.history)}" | |
| final = render_state(session_id, tracer, status, params, log_top_k) | |
| final["type"] = "final" | |
| yield final | |
| except Exception as e: | |
| logger.error(f"Continue error: {e}", exc_info=True) | |
| payload = error_payload(str(e), trace=True) | |
| payload["type"] = "final" | |
| yield payload | |
| def interactive_stop(session_id: str) -> dict: | |
| """Request stop; the running Continue stream finalizes on its own.""" | |
| if session_id: | |
| session = get_session_manager().get_session(session_id) | |
| if session is not None: | |
| session.tracer.request_stop() | |
| logger.info(f"Stop requested for session {session_id}") | |
| return {"ok": True} | |
| def interactive_export( | |
| session_id: str, | |
| strategy: str, temperature: float, top_k: int, top_p: float, | |
| ) -> dict: | |
| """The full session log; the frontend saves it as a JSON download.""" | |
| session, error = _resolve_session(session_id) | |
| if error: | |
| return error | |
| with session.lock: | |
| params = ui_sampling_params(strategy, temperature, top_k, top_p) | |
| return {"ok": True, "export": session.tracer.export_to_dict(params)} | |
| def interactive_lens( | |
| session_id: str, mode: str, stride: int, top_k: int, | |
| ) -> dict: | |
| """Per-layer lens readout of the next-token position (Plotly heatmap).""" | |
| session, error = _resolve_session(session_id) | |
| if error: | |
| return error | |
| with session.lock: | |
| tracer = session.tracer | |
| if tracer.input_ids is None: | |
| return error_payload("Initialize a prompt first.") | |
| model_name = model_manager.get_model_name() | |
| jlens = get_lens_store().get(model_name) | |
| if mode in ("jacobian", "diff") and jlens is None: | |
| return error_payload( | |
| f"No fitted Jacobian lens for {model_name}. Upload one in the " | |
| f"Lens view, or fit one on a GPU box: miru-tracer-fit-lens {model_name}" | |
| ) | |
| try: | |
| n_layers = tracer.model.config.get_text_config().num_hidden_layers | |
| layers = layer_selection(n_layers, 0, -1, stride) | |
| if mode in ("jacobian", "diff") and jlens is not None: | |
| fitted = set(jlens.source_layers) | {n_layers - 1} | |
| layers = [layer for layer in layers if layer in fitted] | |
| slice_ = compute_lens_slice( | |
| tracer.model, | |
| tracer.tokenizer, | |
| tracer.input_ids, | |
| layers=layers, | |
| positions=[tracer.seq_len - 1], | |
| mode=mode, | |
| jlens=jlens, | |
| top_k=int(top_k), | |
| interventions=tracer._intervention_set, | |
| ) | |
| active = len(tracer.interventions) | |
| status = ( | |
| f"{mode} lens over {len(layers)} layers at position " | |
| f"{tracer.seq_len - 1}." | |
| ) | |
| if active: | |
| status += f" {active} intervention(s) active on this session." | |
| return {"ok": True, "figure": fig_json(plot_lens_heatmap(slice_)), "status": status} | |
| except Exception as e: | |
| logger.error(f"Interactive lens error: {e}", exc_info=True) | |
| return error_payload(str(e)) | |
| def interactive_apply_interventions(session_id: str) -> dict: | |
| """Apply the Lens view's active interventions to this session.""" | |
| session, error = _resolve_session(session_id) | |
| if error: | |
| return error | |
| interventions = get_active_interventions() | |
| with session.lock: | |
| jlens = get_lens_store().get(model_manager.get_model_name()) | |
| try: | |
| session.tracer.set_interventions(interventions or None, jlens=jlens) | |
| except ValueError as e: | |
| return error_payload(str(e)) | |
| if not interventions: | |
| return { | |
| "ok": True, | |
| "status": "No active interventions in the Lens view — session cleared.", | |
| } | |
| return { | |
| "ok": True, | |
| "status": ( | |
| f"Applied {len(interventions)} intervention(s) to this session. " | |
| "They affect all subsequent steps (KV cache was rebuilt)." | |
| ), | |
| } | |