File size: 14,409 Bytes
be82719
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
"""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)."
        ),
    }