Jacobina / api /interactive.py
marinarosa's picture
initial commit
be82719
Raw
History Blame Contribute Delete
14.4 kB
"""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)."
),
}