trajectory_rag / webapp /components /tree-visualizer /ModelStreamPanel.tsx
gaaaaaaaaaaa's picture
Publish trajectory_rag
a414867 verified
Raw History Blame Contribute Delete
24.8 kB
"use client";
import { useEffect, useMemo, useRef, useState } from "react";
import { AnimatePresence, motion } from "framer-motion";
import { Bot, Check, Copy, Loader2, Play, Square, X } from "lucide-react";
import { Button } from "@/components/ui/Button";
import { Input } from "@/components/ui/Input";
import { Select } from "@/components/ui/Select";
import { Textarea } from "@/components/ui/Textarea";
import {
fetchModelStreamModels,
streamModelGeneration,
type ModelStreamEvent,
type ModelStreamModelsResponse,
type ModelStreamRequest
} from "@/lib/api";
import { cn } from "@/lib/utils";
interface ModelStreamPanelProps {
open: boolean;
onClose: () => void;
onStreamStart: () => void;
onStreamRawText: (rawText: string, done: boolean) => void;
onStreamError: (message: string) => void;
}
const defaultPrompt =
"If a store sells 3 notebooks for $12, how much would 8 notebooks cost at the same rate?";
const infiniteTokenCap = 32768;
function numberValue(value: string, fallback: number) {
const next = Number(value);
return Number.isFinite(next) ? next : fallback;
}
function boundedIntegerValue(value: string, fallback: number, min: number, max: number) {
return Math.min(max, Math.max(min, Math.round(numberValue(value, fallback))));
}
function numericDiagnostic(value: unknown) {
return typeof value === "number" && Number.isFinite(value) ? value : null;
}
function diagnosticValue(value: unknown) {
if (value == null || value === "") return "";
if (Array.isArray(value)) return value.map(String).join(", ");
if (typeof value === "object") {
const entries = Object.entries(value as Record<string, unknown>);
if (entries.length === 0) return "none";
return entries.map(([key, count]) => `${key}:${String(count)}`).join(", ");
}
if (typeof value === "number") return Number.isInteger(value) ? value.toLocaleString() : value.toFixed(2);
return String(value);
}
function generationDiagnosticItems(data: Record<string, unknown> | null) {
if (!data) return [];
const tokensPerSecond = numericDiagnostic(data.tokensPerSecond);
const generatedTokens = numericDiagnostic(data.generatedTokens);
const adapterState =
data.adapterMerged === true
? "merged"
: data.adapterPath
? data.adapterMergeError === "quantized_adapter_wrapper"
? "wrapped"
: "loaded"
: "base";
const items = [
["Device", diagnosticValue(data.device ?? data.parameterDevices)],
["Dtype", diagnosticValue(data.dtype ?? data.precision)],
["Attention", diagnosticValue(data.attention)],
["Prompt", data.promptTokens ? `${diagnosticValue(data.promptTokens)} tok` : ""],
["LoRA", adapterState],
["Offload", diagnosticValue(data.offloaded)]
];
if (generatedTokens != null) items.push(["Output", `${generatedTokens.toLocaleString()} tok`]);
if (tokensPerSecond != null) items.push(["Speed", `${tokensPerSecond.toFixed(1)} tok/s`]);
return items.filter(([, value]) => value);
}
export function ModelStreamPanel({
open,
onClose,
onStreamStart,
onStreamRawText,
onStreamError
}: ModelStreamPanelProps) {
const [prompt, setPrompt] = useState(defaultPrompt);
const [modelCatalog, setModelCatalog] = useState<ModelStreamModelsResponse | null>(null);
const [modelId, setModelId] = useState("Qwen/Qwen2.5-1.5B-Instruct");
const [adapterPath, setAdapterPath] = useState("");
const [systemPrompt, setSystemPrompt] = useState("");
const [systemPromptOpen, setSystemPromptOpen] = useState(false);
const [precision, setPrecision] = useState<ModelStreamRequest["precision"]>("bf16");
const [tokenMode, setTokenMode] = useState<"limited" | "infinite">("limited");
const [maxNewTokens, setMaxNewTokens] = useState("2048");
const [temperature, setTemperature] = useState("0.7");
const [topP, setTopP] = useState("0.9");
const [topK, setTopK] = useState("50");
const [repetitionPenalty, setRepetitionPenalty] = useState("1.05");
const [doSample, setDoSample] = useState(true);
const [rawGeneration, setRawGeneration] = useState("");
const [status, setStatus] = useState("idle");
const [error, setError] = useState("");
const [stage, setStage] = useState<"configure" | "streaming" | "done">("configure");
const [isStreaming, setIsStreaming] = useState(false);
const [copiedRaw, setCopiedRaw] = useState(false);
const [diagnostics, setDiagnostics] = useState<Record<string, unknown> | null>(null);
const abortRef = useRef<AbortController | null>(null);
const rawRef = useRef("");
const parsedLengthRef = useRef(0);
const modelChoices = useMemo(() => modelCatalog?.models ?? [], [modelCatalog]);
const selectedModel = useMemo(
() => modelChoices.find((choice) => choice.modelId === modelId || choice.id === modelId) ?? null,
[modelChoices, modelId]
);
const checkpointOptions = selectedModel?.checkpoints ?? [];
const diagnosticItems = useMemo(() => generationDiagnosticItems(diagnostics), [diagnostics]);
useEffect(() => {
if (!open || modelCatalog) return;
let cancelled = false;
setError("");
fetchModelStreamModels()
.then((catalog) => {
if (cancelled) return;
const defaultModel =
catalog.models.find((choice) => choice.modelId === catalog.defaultModelId) ?? catalog.models[0];
setModelCatalog(catalog);
setModelId(defaultModel?.modelId ?? catalog.defaultModelId);
setAdapterPath(defaultModel?.defaultCheckpointPath ?? catalog.defaultAdapterPath ?? "");
setSystemPrompt(defaultModel?.systemPrompt ?? "");
})
.catch((nextError) => {
if (!cancelled) setError(nextError instanceof Error ? nextError.message : String(nextError));
});
return () => {
cancelled = true;
};
}, [modelCatalog, open]);
const handleModelChange = (nextModelId: string) => {
setModelId(nextModelId);
const nextModel = modelChoices.find((choice) => choice.modelId === nextModelId || choice.id === nextModelId);
setAdapterPath(nextModel?.defaultCheckpointPath ?? "");
setSystemPrompt(nextModel?.systemPrompt ?? "");
};
const handleEvent = (streamEvent: ModelStreamEvent) => {
if (streamEvent.event === "status") {
setStatus(String(streamEvent.data.message ?? "streaming"));
setDiagnostics(streamEvent.data);
return;
}
if (streamEvent.event === "token") {
const text = typeof streamEvent.data.text === "string" ? streamEvent.data.text : "";
rawRef.current += text;
setRawGeneration(rawRef.current);
const completedIndex = Math.max(
rawRef.current.lastIndexOf("</trajectory>"),
rawRef.current.lastIndexOf("</|trajectory|>"),
rawRef.current.lastIndexOf("</|end_of_trajectory|>"),
rawRef.current.lastIndexOf("<|end_of_trajectory|>")
);
if (completedIndex > parsedLengthRef.current) {
parsedLengthRef.current = completedIndex;
onStreamRawText(rawRef.current, false);
}
return;
}
if (streamEvent.event === "trajectory" || streamEvent.event === "trajectory_update") {
parsedLengthRef.current = rawRef.current.length;
onStreamRawText(rawRef.current, false);
return;
}
if (streamEvent.event === "done") {
const rawText = typeof streamEvent.data.rawText === "string" ? streamEvent.data.rawText : rawRef.current;
rawRef.current = rawText;
setRawGeneration(rawText);
onStreamRawText(rawText, true);
setStatus("done");
setDiagnostics((current) => ({ ...(current ?? {}), ...streamEvent.data, message: "done" }));
setStage("done");
setIsStreaming(false);
return;
}
if (streamEvent.event === "error") {
const message = String(streamEvent.data.error ?? "Model stream failed.");
setError(message);
onStreamError(message);
setStatus("error");
setStage("done");
setIsStreaming(false);
}
};
const startStream = async () => {
if (!prompt.trim() || !modelId || isStreaming) return;
const controller = new AbortController();
abortRef.current = controller;
rawRef.current = "";
parsedLengthRef.current = 0;
setRawGeneration("");
setCopiedRaw(false);
setDiagnostics(null);
setError("");
setStatus("starting");
setStage("streaming");
setIsStreaming(true);
onStreamStart();
const defaultSystemPrompt = selectedModel?.systemPrompt?.trim() ?? "";
const nextSystemPrompt = systemPrompt.trim();
const systemPromptOverride =
nextSystemPrompt && nextSystemPrompt !== defaultSystemPrompt ? nextSystemPrompt : null;
try {
await streamModelGeneration(
{
prompt,
system_prompt: systemPromptOverride,
model_id: modelId,
adapter_path: adapterPath.trim() || null,
precision,
max_new_tokens:
tokenMode === "limited"
? boundedIntegerValue(maxNewTokens, 2048, 1, infiniteTokenCap)
: null,
temperature: numberValue(temperature, 0.7),
top_p: numberValue(topP, 0.9),
top_k: boundedIntegerValue(topK, 50, 0, 1000),
repetition_penalty: numberValue(repetitionPenalty, 1.05),
do_sample: doSample
},
handleEvent,
controller.signal
);
if (!controller.signal.aborted) {
setStatus("done");
setStage("done");
}
} catch (nextError) {
if (controller.signal.aborted) {
setStatus("stopped");
} else {
const message = nextError instanceof Error ? nextError.message : String(nextError);
setError(message);
onStreamError(message);
setStatus("error");
}
setStage("done");
} finally {
setIsStreaming(false);
abortRef.current = null;
}
};
const stopStream = () => {
abortRef.current?.abort();
setIsStreaming(false);
setStatus("stopped");
setStage("done");
};
const copyRawGeneration = async () => {
if (!rawGeneration) return;
try {
await navigator.clipboard.writeText(rawGeneration);
setCopiedRaw(true);
window.setTimeout(() => setCopiedRaw(false), 1400);
} catch (nextError) {
const message = nextError instanceof Error ? nextError.message : String(nextError);
setError(`Copy failed: ${message}`);
}
};
const handleClose = () => {
if (!isStreaming) setStage("configure");
onClose();
};
const openNewRun = () => {
if (isStreaming) return;
setError("");
setStatus("idle");
setRawGeneration("");
setCopiedRaw(false);
setDiagnostics(null);
rawRef.current = "";
parsedLengthRef.current = 0;
setStage("configure");
};
const renderConfigureModal = () => (
<motion.div
initial={{ opacity: 0 }}
animate={{ opacity: 1 }}
exit={{ opacity: 0 }}
transition={{ duration: 0.16 }}
className="fixed inset-0 z-50 flex items-center justify-center bg-ink-900/30 p-4 backdrop-blur-[2px]"
>
<motion.section
initial={{ opacity: 0, scale: 0.97, y: 14 }}
animate={{ opacity: 1, scale: 1, y: 0 }}
exit={{ opacity: 0, scale: 0.98, y: 10 }}
transition={{ duration: 0.18 }}
className="flex max-h-[calc(100vh-2rem)] w-[min(780px,calc(100vw-2rem))] flex-col overflow-hidden rounded-xl border border-paper-300 bg-white shadow-soft"
>
<div className="flex items-center justify-between border-b border-paper-300 px-5 py-4">
<div className="flex min-w-0 items-center gap-3">
<span className="flex h-10 w-10 shrink-0 items-center justify-center rounded-lg bg-claude-50 text-claude-700">
<Bot className="h-5 w-5" />
</span>
<div className="min-w-0">
<h2 className="truncate text-base font-semibold text-ink-900">Create tree by model</h2>
<p className="truncate text-xs text-ink-500">{status}</p>
</div>
</div>
<Button size="icon" variant="ghost" onClick={handleClose} title="Close">
<X className="h-4 w-4" />
</Button>
</div>
<div className="scrollbar-thin flex-1 space-y-5 overflow-auto p-5">
<label className="block text-xs font-semibold uppercase text-ink-500">
Question prompt
<Textarea
className="mt-2 min-h-32 resize-y"
value={prompt}
onChange={(event) => setPrompt(event.target.value)}
disabled={isStreaming}
/>
</label>
<details
open={systemPromptOpen}
onToggle={(event) => setSystemPromptOpen(event.currentTarget.open)}
className="rounded-lg border border-paper-300 bg-paper-50/70"
>
<summary className="flex cursor-pointer select-none items-center justify-between gap-3 px-3 py-2 text-xs font-semibold uppercase text-ink-600">
<span>System prompt</span>
<span className="truncate text-[11px] font-medium normal-case text-ink-400">
{selectedModel?.settingsPath ?? "Loaded from wrapper model_settings.yaml"}
</span>
</summary>
<div className="space-y-2 border-t border-paper-300 p-3">
<Textarea
className="min-h-44 resize-y font-mono text-xs leading-5"
value={systemPrompt}
onChange={(event) => setSystemPrompt(event.target.value)}
disabled={isStreaming}
placeholder="Wrapper system prompt"
/>
<div className="flex items-center justify-between gap-3 text-xs text-ink-500">
<span>
{systemPrompt.trim() === (selectedModel?.systemPrompt?.trim() ?? "")
? "Using wrapper prompt"
: "Using edited prompt for this run"}
</span>
<Button
size="sm"
variant="ghost"
onClick={() => setSystemPrompt(selectedModel?.systemPrompt ?? "")}
disabled={isStreaming}
>
Reset
</Button>
</div>
</div>
</details>
<div className="grid grid-cols-1 gap-3 md:grid-cols-2">
<label className="text-xs font-semibold uppercase text-ink-500">
Model
<Select
className="mt-2"
value={modelId}
onChange={(event) => handleModelChange(event.target.value)}
disabled={isStreaming || modelChoices.length === 0}
>
{modelChoices.length === 0 ? (
<option value={modelId}>{modelId}</option>
) : (
modelChoices.map((choice) => (
<option key={choice.key} value={choice.modelId}>
{choice.label}
</option>
))
)}
</Select>
</label>
<label className="text-xs font-semibold uppercase text-ink-500">
Checkpoint
<Select
className="mt-2"
value={adapterPath}
onChange={(event) => setAdapterPath(event.target.value)}
disabled={isStreaming}
>
<option value="">Base model</option>
{checkpointOptions.map((checkpoint) => (
<option key={checkpoint.path} value={checkpoint.path}>
{checkpoint.label}
</option>
))}
{adapterPath && !checkpointOptions.some((checkpoint) => checkpoint.path === adapterPath) && (
<option value={adapterPath}>{adapterPath}</option>
)}
</Select>
</label>
</div>
<label className="block text-xs font-semibold uppercase text-ink-500">
Custom checkpoint path
<Input
className="mt-2 font-mono text-xs"
value={adapterPath}
onChange={(event) => setAdapterPath(event.target.value)}
disabled={isStreaming}
placeholder="Optional LoRA adapter directory"
/>
</label>
<div className="grid grid-cols-1 gap-4 md:grid-cols-[1fr_220px]">
<label className="text-xs font-semibold uppercase text-ink-500">
Temperature
<div className="mt-2 flex items-center gap-3">
<input
type="range"
min="0"
max="2"
step="0.05"
value={temperature}
onChange={(event) => setTemperature(event.target.value)}
disabled={isStreaming}
className="h-2 w-full accent-claude-500"
/>
<Input
className="w-20"
inputMode="decimal"
value={temperature}
onChange={(event) => setTemperature(event.target.value)}
disabled={isStreaming}
/>
</div>
</label>
<label className="text-xs font-semibold uppercase text-ink-500">
Precision
<Select
className="mt-2"
value={precision}
onChange={(event) => setPrecision(event.target.value as ModelStreamRequest["precision"])}
disabled={isStreaming}
>
<option value="bf16">bf16</option>
<option value="fp16">fp16</option>
<option value="fp32">fp32</option>
<option value="4bit">4bit</option>
<option value="8bit">8bit</option>
</Select>
</label>
</div>
<div className="grid grid-cols-1 gap-3 md:grid-cols-[220px_1fr]">
<div className="text-xs font-semibold uppercase text-ink-500">
Token limit
<div className="mt-2 inline-flex h-9 rounded-lg border border-paper-300 bg-paper-50 p-1">
{(["limited", "infinite"] as const).map((mode) => (
<button
key={mode}
type="button"
onClick={() => setTokenMode(mode)}
disabled={isStreaming}
className={cn(
"h-7 rounded-md px-3 text-xs font-semibold capitalize transition",
tokenMode === mode ? "bg-white text-claude-700 shadow-sm" : "text-ink-500 hover:text-ink-900"
)}
>
{mode}
</button>
))}
</div>
</div>
<label className="text-xs font-semibold uppercase text-ink-500">
Max new tokens
<Input
className="mt-2"
inputMode="numeric"
value={tokenMode === "infinite" ? "Until EOS (32k cap)" : maxNewTokens}
onChange={(event) => setMaxNewTokens(event.target.value)}
disabled={isStreaming || tokenMode === "infinite"}
/>
</label>
</div>
<div className="grid grid-cols-2 gap-3 md:grid-cols-4">
<label className="text-xs font-semibold uppercase text-ink-500">
Top P
<Input className="mt-2" value={topP} onChange={(event) => setTopP(event.target.value)} disabled={isStreaming} />
</label>
<label className="text-xs font-semibold uppercase text-ink-500">
Top K
<Input className="mt-2" value={topK} onChange={(event) => setTopK(event.target.value)} disabled={isStreaming} />
</label>
<label className="text-xs font-semibold uppercase text-ink-500">
Repeat
<Input
className="mt-2"
value={repetitionPenalty}
onChange={(event) => setRepetitionPenalty(event.target.value)}
disabled={isStreaming}
/>
</label>
<label className="flex items-end gap-2 pb-2 text-xs font-semibold uppercase text-ink-500">
<input
type="checkbox"
checked={doSample}
onChange={(event) => setDoSample(event.target.checked)}
disabled={isStreaming}
className="h-4 w-4 rounded border-paper-300 text-claude-600"
/>
Sampling
</label>
</div>
{error && (
<div className="rounded-lg border border-red-200 bg-red-50 px-3 py-2 text-sm text-red-700">
{error}
</div>
)}
</div>
<div className="flex items-center justify-between gap-3 border-t border-paper-300 px-5 py-4">
<div className="truncate text-xs text-ink-500">
{selectedModel ? selectedModel.wrapperDir : "Loading model choices"}
</div>
<Button variant="primary" onClick={startStream} disabled={!prompt.trim() || !modelId || isStreaming}>
{isStreaming ? <Loader2 className="h-4 w-4 animate-spin" /> : <Play className="h-4 w-4" />}
Generate tree
</Button>
</div>
</motion.section>
</motion.div>
);
const renderMonitor = () => (
<motion.aside
initial={{ opacity: 0, y: 14 }}
animate={{ opacity: 1, y: 0 }}
exit={{ opacity: 0, y: 14 }}
transition={{ duration: 0.18 }}
className="absolute bottom-5 left-5 z-30 flex max-h-[min(430px,calc(100vh-7rem))] w-[min(560px,calc(100vw-2rem))] flex-col overflow-hidden rounded-xl border border-paper-300 bg-white/94 shadow-soft backdrop-blur-xl"
>
<div className="flex items-center justify-between border-b border-paper-300 px-4 py-3">
<div className="flex min-w-0 items-center gap-2">
<Bot className="h-4 w-4 shrink-0 text-claude-600" />
<div className="min-w-0">
<h2 className="truncate text-sm font-semibold text-ink-900">Model generation</h2>
<p className="truncate text-xs text-ink-500">{status}</p>
</div>
</div>
<div className="flex items-center gap-2">
{isStreaming ? (
<Button size="sm" variant="danger" onClick={stopStream}>
<Square className="h-4 w-4" />
Stop
</Button>
) : (
<Button size="sm" onClick={openNewRun}>
<Play className="h-4 w-4" />
New run
</Button>
)}
<Button
size="sm"
variant="ghost"
onClick={() => void copyRawGeneration()}
disabled={!rawGeneration}
title="Copy raw generation"
>
{copiedRaw ? <Check className="h-4 w-4" /> : <Copy className="h-4 w-4" />}
{copiedRaw ? "Copied" : "Copy raw"}
</Button>
<Button size="icon" variant="ghost" onClick={handleClose} title="Hide raw generation">
<X className="h-4 w-4" />
</Button>
</div>
</div>
<div className="flex-1 space-y-3 overflow-hidden p-4">
{diagnosticItems.length > 0 && (
<div className="grid grid-cols-2 gap-2 md:grid-cols-4">
{diagnosticItems.map(([label, value]) => (
<div key={label} className="min-w-0 rounded-md border border-paper-300 bg-paper-50 px-2 py-1.5">
<div className="text-[10px] font-semibold uppercase text-ink-400">{label}</div>
<div className="truncate text-xs font-medium text-ink-800" title={value}>
{value}
</div>
</div>
))}
</div>
)}
<Textarea
className="h-72 min-h-0 resize-none font-mono text-xs leading-5"
value={rawGeneration}
readOnly
placeholder="Streaming output"
/>
{error && (
<div className="rounded-lg border border-red-200 bg-red-50 px-3 py-2 text-sm text-red-700">
{error}
</div>
)}
</div>
<div className="flex items-center justify-between border-t border-paper-300 px-4 py-2 text-xs text-ink-500">
<span className="truncate">{modelId}</span>
<span>{rawGeneration.length.toLocaleString()} chars</span>
</div>
</motion.aside>
);
return (
<AnimatePresence>
{open && (stage === "configure" ? renderConfigureModal() : renderMonitor())}
</AnimatePresence>
);
}