diff --git a/uv.lock b/uv.lock index 92cd53f3..e4bbfe46 100644 --- a/uv.lock +++ b/uv.lock @@ -903,7 +903,7 @@ wheels = [ [[package]] name = "nnsightful" version = "0.1.0" -source = { git = "https://github.com/AdamBelfki3/nnsightful.git#75e42670b64f2aa5fcb82421e910376368c934d3" } +source = { git = "https://github.com/AdamBelfki3/nnsightful.git#7123cc767447123a6199cba6f14d772b3439bce4" } dependencies = [ { name = "hf-xet" }, { name = "ipython" }, diff --git a/workbench/_api/_metadata_cache.json b/workbench/_api/_metadata_cache.json index 9190e715..388354e1 100644 --- a/workbench/_api/_metadata_cache.json +++ b/workbench/_api/_metadata_cache.json @@ -86,14 +86,14 @@ }, "zai-org/GLM-4.5-Air": { "name": "zai-org/GLM-4.5-Air", - "is_chat": false, + "is_chat": true, "n_layers": 46, "params": "110B", "gated": true }, "Qwen/Qwen3-32B": { "name": "Qwen/Qwen3-32B", - "is_chat": false, + "is_chat": true, "n_layers": 64, "params": "33B", "gated": true @@ -114,7 +114,7 @@ }, "microsoft/DialoGPT-small": { "name": "microsoft/DialoGPT-small", - "is_chat": false, + "is_chat": true, "n_layers": 12, "params": "176M", "gated": false @@ -163,7 +163,7 @@ }, "openai/gpt-oss-safeguard-120b": { "name": "openai/gpt-oss-safeguard-120b", - "is_chat": false, + "is_chat": true, "n_layers": 36, "params": "63B", "gated": true @@ -205,14 +205,14 @@ }, "Qwen/Qwen3-4B": { "name": "Qwen/Qwen3-4B", - "is_chat": false, + "is_chat": true, "n_layers": 36, "params": "4B", "gated": false }, "deepseek-ai/DeepSeek-R1-Distill-Llama-70B": { "name": "deepseek-ai/DeepSeek-R1-Distill-Llama-70B", - "is_chat": false, + "is_chat": true, "n_layers": 80, "params": "71B", "gated": true @@ -226,7 +226,7 @@ }, "microsoft/DialoGPT-medium": { "name": "microsoft/DialoGPT-medium", - "is_chat": false, + "is_chat": true, "n_layers": 24, "params": "unknown", "gated": false @@ -275,7 +275,7 @@ }, "deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B": { "name": "deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B", - "is_chat": false, + "is_chat": true, "n_layers": 28, "params": "2B", "gated": false @@ -352,7 +352,7 @@ }, "Qwen/Qwen3-0.6B": { "name": "Qwen/Qwen3-0.6B", - "is_chat": false, + "is_chat": true, "n_layers": 28, "params": "752M", "gated": false @@ -366,21 +366,21 @@ }, "deepseek-ai/DeepSeek-R1-Distill-Qwen-32B": { "name": "deepseek-ai/DeepSeek-R1-Distill-Qwen-32B", - "is_chat": false, + "is_chat": true, "n_layers": 64, "params": "33B", "gated": true }, "deepseek-ai/DeepSeek-R1-Distill-Llama-8B": { "name": "deepseek-ai/DeepSeek-R1-Distill-Llama-8B", - "is_chat": false, + "is_chat": true, "n_layers": 32, "params": "8B", "gated": false }, "deepseek-ai/DeepSeek-R1-Distill-Qwen-7B": { "name": "deepseek-ai/DeepSeek-R1-Distill-Qwen-7B", - "is_chat": false, + "is_chat": true, "n_layers": 28, "params": "8B", "gated": false @@ -429,14 +429,14 @@ }, "deepseek-ai/DeepSeek-R1-Distill-Qwen-14B": { "name": "deepseek-ai/DeepSeek-R1-Distill-Qwen-14B", - "is_chat": false, + "is_chat": true, "n_layers": 48, "params": "15B", "gated": true }, "openai/gpt-oss-20b": { "name": "openai/gpt-oss-20b", - "is_chat": false, + "is_chat": true, "n_layers": 24, "params": "12B", "gated": true @@ -464,7 +464,7 @@ }, "Qwen/QwQ-32B": { "name": "Qwen/QwQ-32B", - "is_chat": false, + "is_chat": true, "n_layers": 64, "params": "33B", "gated": true @@ -478,7 +478,7 @@ }, "openai/gpt-oss-120b": { "name": "openai/gpt-oss-120b", - "is_chat": false, + "is_chat": true, "n_layers": 36, "params": "63B", "gated": true @@ -506,7 +506,7 @@ }, "Qwen/Qwen3-8B": { "name": "Qwen/Qwen3-8B", - "is_chat": false, + "is_chat": true, "n_layers": 36, "params": "8B", "gated": false @@ -534,7 +534,7 @@ }, "deepseek-ai/DeepSeek-R1": { "name": "deepseek-ai/DeepSeek-R1", - "is_chat": false, + "is_chat": true, "n_layers": 61, "params": "685B", "gated": true @@ -569,7 +569,7 @@ }, "allenai/Olmo-3.1-32B-Think": { "name": "allenai/Olmo-3.1-32B-Think", - "is_chat": false, + "is_chat": true, "n_layers": 64, "params": "32B", "gated": true @@ -583,7 +583,7 @@ }, "Qwen/Qwen3-4B-Thinking-2507": { "name": "Qwen/Qwen3-4B-Thinking-2507", - "is_chat": false, + "is_chat": true, "n_layers": 36, "params": "4B", "gated": false @@ -625,14 +625,14 @@ }, "yujiepan/deepseek-llm-tiny-random": { "name": "yujiepan/deepseek-llm-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "410K", "gated": false }, "yujiepan/gpt-oss-tiny-random-mxfp4": { "name": "yujiepan/gpt-oss-tiny-random-mxfp4", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "7M", "gated": false @@ -646,7 +646,7 @@ }, "deepseek-ai/DeepSeek-V3": { "name": "deepseek-ai/DeepSeek-V3", - "is_chat": false, + "is_chat": true, "n_layers": 61, "params": "685B", "gated": true @@ -660,7 +660,7 @@ }, "yujiepan/llama-3-tiny-random": { "name": "yujiepan/llama-3-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "1M", "gated": false @@ -674,14 +674,14 @@ }, "yujiepan/mistral-v0.3-tiny-random": { "name": "yujiepan/mistral-v0.3-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "526K", "gated": false }, "yujiepan/llama-3.3-tiny-random": { "name": "yujiepan/llama-3.3-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "4M", "gated": false @@ -695,14 +695,14 @@ }, "yujiepan/mixtral-8xtiny-random": { "name": "yujiepan/mixtral-8xtiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "525K", "gated": false }, "yujiepan/gemma-tiny-random": { "name": "yujiepan/gemma-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "2M", "gated": false @@ -723,49 +723,49 @@ }, "yujiepan/deepseek-v2-tiny-random": { "name": "yujiepan/deepseek-v2-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "2M", "gated": false }, "yujiepan/glm-4.5-tiny-random": { "name": "yujiepan/glm-4.5-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "3M", "gated": false }, "yujiepan/phi-3.5-tiny-random": { "name": "yujiepan/phi-3.5-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "1M", "gated": false }, "yujiepan/dbrx-tiny-random": { "name": "yujiepan/dbrx-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "806K", "gated": false }, "yujiepan/phi-3-tiny-random": { "name": "yujiepan/phi-3-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "1M", "gated": false }, "yujiepan/QwQ-preview-tiny-random": { "name": "yujiepan/QwQ-preview-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "2M", "gated": false }, "yujiepan/mistral-tiny-random": { "name": "yujiepan/mistral-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "514K", "gated": false @@ -779,7 +779,7 @@ }, "yujiepan/gpt-oss-tiny-random-bf16": { "name": "yujiepan/gpt-oss-tiny-random-bf16", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "7M", "gated": false @@ -793,7 +793,7 @@ }, "yujiepan/dbrx-tiny256-random": { "name": "yujiepan/dbrx-tiny256-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "64M", "gated": false @@ -814,7 +814,7 @@ }, "yujiepan/glm-4-moe-tiny-random": { "name": "yujiepan/glm-4-moe-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "3M", "gated": false @@ -828,35 +828,35 @@ }, "yujiepan/deepseek-v3-tiny-random": { "name": "yujiepan/deepseek-v3-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "4M", "gated": false }, "yujiepan/ernie-4.5-tiny-random": { "name": "yujiepan/ernie-4.5-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "854K", "gated": false }, "yujiepan/gemma-2-tiny-random": { "name": "yujiepan/gemma-2-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "2M", "gated": false }, "yujiepan/deepseek-v2-0628-tiny-random": { "name": "yujiepan/deepseek-v2-0628-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "2M", "gated": false }, "yujiepan/llama-2-tiny-3layers-random": { "name": "yujiepan/llama-2-tiny-3layers-random", - "is_chat": false, + "is_chat": true, "n_layers": 3, "params": "515K", "gated": false @@ -884,21 +884,21 @@ }, "yujiepan/phi-3.5-moe-tiny-random": { "name": "yujiepan/phi-3.5-moe-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "1M", "gated": false }, "yujiepan/glm-4-tiny-random": { "name": "yujiepan/glm-4-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "5M", "gated": false }, "yujiepan/deepseek-v3.1-tiny-random": { "name": "yujiepan/deepseek-v3.1-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "5M", "gated": false @@ -919,7 +919,7 @@ }, "yujiepan/llama-3.1-tiny-random": { "name": "yujiepan/llama-3.1-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "2M", "gated": false @@ -933,7 +933,7 @@ }, "yujiepan/llama-3.3-tiny-random-dim64": { "name": "yujiepan/llama-3.3-tiny-random-dim64", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "8M", "gated": false @@ -947,7 +947,7 @@ }, "yujiepan/qwen3-tiny-random-tp": { "name": "yujiepan/qwen3-tiny-random-tp", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "2M", "gated": false @@ -961,7 +961,7 @@ }, "yujiepan/llama-2-tiny-random": { "name": "yujiepan/llama-2-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 1, "params": "513K", "gated": false @@ -975,21 +975,21 @@ }, "yujiepan/kimi-k2-tiny-random": { "name": "yujiepan/kimi-k2-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "11M", "gated": false }, "axolotl-ai-co/gemma-3-34M": { "name": "axolotl-ai-co/gemma-3-34M", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "34M", "gated": false }, "yujiepan/mixtral-8xtiny-random-openvino-8bit": { "name": "yujiepan/mixtral-8xtiny-random-openvino-8bit", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "unknown", "gated": false @@ -1003,7 +1003,7 @@ }, "yujiepan/baguettotron-tiny-random": { "name": "yujiepan/baguettotron-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "552K", "gated": false @@ -1038,14 +1038,14 @@ }, "yujiepan/ernie-4.5-moe-tiny-random": { "name": "yujiepan/ernie-4.5-moe-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "904K", "gated": false }, "yujiepan/mistral-nemo-2407-tiny-random": { "name": "yujiepan/mistral-nemo-2407-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "2M", "gated": false @@ -1059,56 +1059,56 @@ }, "yujiepan/llama-3.2-tiny-random": { "name": "yujiepan/llama-3.2-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "1M", "gated": false }, "yujiepan/phi-moe-tiny-random": { "name": "yujiepan/phi-moe-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "3M", "gated": false }, "yujiepan/gpt-oss-tiny-random": { "name": "yujiepan/gpt-oss-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "7M", "gated": false }, "yujiepan/smollm3-tiny-random": { "name": "yujiepan/smollm3-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "8M", "gated": false }, "yujiepan/qwen3-moe-tiny-random": { "name": "yujiepan/qwen3-moe-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "10M", "gated": false }, "yujiepan/phi-4-tiny-random": { "name": "yujiepan/phi-4-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "3M", "gated": false }, "yujiepan/QwQ-tiny-random": { "name": "yujiepan/QwQ-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "2M", "gated": false }, "yujiepan/mixtral-tiny-random": { "name": "yujiepan/mixtral-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "258K", "gated": false @@ -1122,7 +1122,7 @@ }, "yujiepan/qwen3-tiny-random": { "name": "yujiepan/qwen3-tiny-random", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "10M", "gated": false @@ -1136,14 +1136,14 @@ }, "yujiepan/qwq-tiny-random-dim64": { "name": "yujiepan/qwq-tiny-random-dim64", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "20M", "gated": false }, "yujiepan/llama-3-tiny-random-gptq-w4": { "name": "yujiepan/llama-3-tiny-random-gptq-w4", - "is_chat": false, + "is_chat": true, "n_layers": 2, "params": "1M", "gated": false diff --git a/workbench/_api/main.py b/workbench/_api/main.py index d3ffff10..282d8a39 100644 --- a/workbench/_api/main.py +++ b/workbench/_api/main.py @@ -4,7 +4,7 @@ import os import anyio -from .routes import lens, patch, models, logit_lens, j_lens, activation_patching, causal_mediation +from .routes import lens, patch, models, generate, logit_lens, j_lens, activation_patching, causal_mediation from .state import AppState from dotenv import load_dotenv; load_dotenv() @@ -62,6 +62,7 @@ def fastapi_app(): app.include_router(causal_mediation, prefix="/causal_mediation", tags=["causal_mediation"]) app.include_router(patch, prefix="/patch") app.include_router(models, prefix="/models") + app.include_router(generate, prefix="/generate") app.state.m = AppState() diff --git a/workbench/_api/metadata.py b/workbench/_api/metadata.py index 6ec50aef..a1a1b4b7 100644 --- a/workbench/_api/metadata.py +++ b/workbench/_api/metadata.py @@ -77,15 +77,26 @@ "git", } -# Chat/instruct classification is name-based, NOT tokenizer-chat-template-based. -# The tokenizer signal is unreliable: several families (notably Qwen) bundle a -# chat_template into their BASE tokenizers, so `chat_template is not None` would -# misclassify e.g. `Qwen2.5-7B` (base) as a chat model. Repo naming is the -# strongest cross-ecosystem signal — instruct/chat post-trained models almost -# always carry a marker in the name, base models don't. +# Chat/instruct classification combines two signals: # -# Tokens matched (case-insensitive) as hyphen/underscore/slash-delimited -# segments of the repo's last path component: +# 1. The presence of a chat template on the Hub — the ground truth for "this +# model expects a chat/instruct format". It lives in one of two places: a +# standalone ``chat_template.jinja`` (modern convention, e.g. gpt-oss) or a +# ``chat_template`` key inside ``tokenizer_config.json`` (classic). See +# ``_has_chat_template``. +# 2. Repo naming — post-trained models usually carry a marker (``instruct``, +# ``chat``, ``it``, ...). Kept as a fallback for gated repos whose template +# an unauthenticated probe can't read, and used exclusively for Qwen. +# +# The template supersedes the name check for every family EXCEPT Qwen: Qwen +# bundles a chat_template into its BASE tokenizers too (verified — ``Qwen2.5-7B`` +# and ``Qwen3-8B-Base`` both ship one), so a template is a false positive there. +# For Qwen we classify by name, with the twist that Qwen3 *inverts* the usual +# convention: the bare name (``Qwen3-8B``) is the chat checkpoint and the +# pretrained one carries a ``-Base`` suffix (``Qwen3-8B-Base``). +# +# Name markers matched (case-insensitive) as hyphen/underscore/slash/dot- +# delimited segments of the repo's last path component: # - generic post-training markers: instruct, chat, it (gemma), rlhf, dpo, # sft, orpo # - well-known marker-less chat families (so they don't fall to "base") @@ -97,16 +108,74 @@ re.IGNORECASE, ) +# Qwen bundles a chat_template into its base tokenizers, so the template signal +# is unusable for the family — classify by name instead. Qwen3 inverts the +# convention (bare = chat, ``-Base`` = pretrained); Qwen1.5/2/2.5 keep the +# classic ``-Instruct`` marker. +# +# Anchored at the START of the label so this only catches *official* Qwen repos +# (``Qwen3-8B``, ``Qwen2.5-7B``). Qwen-architecture derivatives from other orgs +# (e.g. ``deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B``) are genuine chat models +# that ship a real template — they must go through the template path, not this +# carve-out. ``QwQ`` (chat-only, no base) likewise doesn't match ``qwen`` and so +# is correctly template-classified. +_QWEN_RE = re.compile(r"^qwen", re.IGNORECASE) +_QWEN3_RE = re.compile(r"^qwen3(?:[-_/.]|$)", re.IGNORECASE) +_BASE_SUFFIX_RE = re.compile(r"[-_/.]base$", re.IGNORECASE) + + +def _has_chat_template(model_name: str) -> bool: + """Whether the repo ships a chat template, checking both Hub conventions. + + A chat template lives in a standalone ``chat_template.jinja`` (modern, e.g. + gpt-oss) or embedded under the ``chat_template`` key of + ``tokenizer_config.json`` (classic). Both files are small; we read them + directly rather than loading the whole tokenizer. + + Best-effort: a permanent read failure (e.g. a gated repo we can't access) is + treated as "no readable template" — the name check still catches marked + repos. Transient Hub failures are re-raised so the caller retries rather + than caching a wrong verdict. + """ + from huggingface_hub import hf_hub_download + from huggingface_hub.utils import EntryNotFoundError + + try: + try: + hf_hub_download(model_name, "chat_template.jinja") + return True + except EntryNotFoundError: + pass # not the standalone-file convention; try the embedded key + try: + path = hf_hub_download(model_name, "tokenizer_config.json") + except EntryNotFoundError: + return False # no tokenizer config at all → no template + with open(path, encoding="utf-8") as f: + return bool(json.load(f).get("chat_template")) + except Exception as e: + if _is_transient_metadata_error(e): + raise + logger.warning(f"Could not read chat template for {model_name}: {e}") + return False + -def _is_chat_model(model_name: str) -> bool: - """Return whether the repo name marks an instruction- or chat-tuned model. +def _is_chat_model(model_name: str, has_chat_template: bool) -> bool: + """Whether a repo is an instruction/chat model. - Classification is based on the repo's last path component (e.g. - ``Llama-3.1-8B-Instruct``), not on the tokenizer's chat template, which - is unreliable across ecosystems. + The chat template is authoritative for every family except Qwen, whose base + tokenizers also bundle a template (so the signal is a false positive there); + Qwen is classified by name convention instead. A name marker (``-Instruct`` + etc.) always counts as chat too, covering gated repos whose template the + probe couldn't read. """ label = model_name.split("/")[-1] - return _CHAT_NAME_RE.search(label) is not None + if _QWEN_RE.search(label): + if _QWEN3_RE.search(label): + # Qwen3 inverts: bare name is chat, `-Base` is the pretrained model. + return _BASE_SUFFIX_RE.search(label) is None + # Qwen1.5/2/2.5: classic convention — only explicit markers are chat. + return _CHAT_NAME_RE.search(label) is not None + return has_chat_template or _CHAT_NAME_RE.search(label) is not None class UnsupportedModel(Exception): @@ -122,7 +191,8 @@ class ModelMetadata(BaseModel): Attributes: name: Full repo ID (``org/model``). - is_chat: Whether the repo name indicates an instruct/chat variant. + is_chat: Whether the repo is an instruct/chat model — from its chat + template (or name, for Qwen and gated repos). See ``_is_chat_model``. n_layers: Transformer block count from ``AutoConfig``. params: Human-readable parameter count (e.g. ``"7B"``), or ``"unknown"`` when safetensors headers are unavailable. @@ -268,7 +338,7 @@ def fetch_model_metadata(model_name: str) -> ModelMetadata: return ModelMetadata( name=model_name, - is_chat=_is_chat_model(model_name), + is_chat=_is_chat_model(model_name, _has_chat_template(model_name)), n_layers=n_layers, params=_format_params(num_params) if num_params > 0 else "unknown", gated=gated, diff --git a/workbench/_api/routes/__init__.py b/workbench/_api/routes/__init__.py index 32a7057d..0d1b5e2d 100644 --- a/workbench/_api/routes/__init__.py +++ b/workbench/_api/routes/__init__.py @@ -1,6 +1,7 @@ from .lens import router as lens from .patch import router as patch from .models import router as models +from .generate import router as generate from .logit_lens import router as logit_lens from .j_lens import router as j_lens from .activation_patching import router as activation_patching @@ -14,6 +15,7 @@ "lens", "patch", "models", + "generate", "logit_lens", "j_lens", "activation_patching", diff --git a/workbench/_api/routes/generate.py b/workbench/_api/routes/generate.py new file mode 100644 index 00000000..3fa15c95 --- /dev/null +++ b/workbench/_api/routes/generate.py @@ -0,0 +1,116 @@ +from fastapi import APIRouter, Depends +from pydantic import BaseModel + +from ..state import AppState, get_state +from ..auth import require_user_email +from ..data_models import NDIFResponse + +router = APIRouter() + + +class GenerateRequest(BaseModel): + model: str + prompt: str + num_tokens: int = 25 # max new tokens to sample + temperature: float | None = None + top_k: int | None = None + top_p: float | None = None + stop_strings: list[str] | None = None + + +class GenerateData(BaseModel): + prompt: list[str] # the input (prompt) tokens, detokenized + completion: list[str] # the generated tokens, detokenized + + +class GenerateResponse(NDIFResponse): + data: GenerateData | None = None + + +def _sampling_kwargs(req: GenerateRequest) -> dict: + """Build the sampling kwargs forwarded to ``model.generate(...)``. + + A temperature/top_p/top_k value turns sampling on; the individual keys are + only included when set. ``do_sample`` is ALWAYS set explicitly: many + instruct/chat models ship a ``generation_config.json`` with + ``do_sample: true`` (e.g. Qwen3-8B has do_sample=true, temperature=0.6), so + omitting the flag lets the model's own config silently re-enable sampling — + a different completion every run even when the caller asked for greedy. + Passing ``do_sample=False`` forces deterministic decoding regardless. + """ + kwargs: dict = {} + sample = False + if req.temperature is not None: + kwargs["temperature"] = req.temperature + sample = True + if req.top_p is not None: + kwargs["top_p"] = req.top_p + sample = True + if req.top_k is not None: + kwargs["top_k"] = req.top_k + sample = True + kwargs["do_sample"] = sample + if req.stop_strings: + kwargs["stop_strings"] = req.stop_strings + return kwargs + + +def generate(model, req: GenerateRequest, state: AppState): + """Sample a completion. Returns the NDIF job id (remote) or a + ``(prompt_tokens, completion_tokens)`` pair of id tensors (local).""" + with model.generate( + req.prompt, + max_new_tokens=req.num_tokens, + remote=state.remote, + backend=state.make_backend(model=model), + **_sampling_kwargs(req), + ) as tracer: + prompt_tokens = model.inputs[1]['input_ids'].save() + completion_tokens = tracer.result[:, prompt_tokens.shape[-1]:].save() + + if state.remote: + return tracer.backend.job_id + + return prompt_tokens, completion_tokens + + +def get_remote_generation(job_id: str, state: AppState): + backend = state.make_backend(job_id=job_id) + results = backend() + return results["prompt_tokens"], results["completion_tokens"] + + +def process_generation(prompt_tokens, completion_tokens, tokenizer) -> GenerateData: + prompt = tokenizer.batch_decode(prompt_tokens[0]) + completion = tokenizer.batch_decode(completion_tokens[0]) + + return GenerateData(prompt=prompt, completion=completion) + + +@router.post("/start", response_model=GenerateResponse) +async def start_generate( + req: GenerateRequest, + state: AppState = Depends(get_state), + user_email: str = Depends(require_user_email), +): + model = state[req.model] + + output = generate(model, req, state) + + if state.remote: + return {"job_id": output} + + prompt_tokens, completion_tokens = output + return {"data": process_generation(prompt_tokens, completion_tokens, model.tokenizer)} + + +@router.post("/results/{job_id}", response_model=GenerateResponse) +async def collect_generate( + job_id: str, + req: GenerateRequest, + state: AppState = Depends(get_state), + user_email: str = Depends(require_user_email), +): + prompt_tokens, completion_tokens = get_remote_generation(job_id, state) + + return {"data": process_generation(prompt_tokens, completion_tokens, state[req.model].tokenizer)} diff --git a/workbench/_api/routes/j_lens.py b/workbench/_api/routes/j_lens.py index cc588ef2..64f17c74 100644 --- a/workbench/_api/routes/j_lens.py +++ b/workbench/_api/routes/j_lens.py @@ -13,8 +13,39 @@ class JLensRequest(BaseModel): model: str prompt: str - topk: int = 5 # Number of top-k predictions per cell + topk: int = 5 include_entropy: bool = True # Whether to include entropy data + max_new_tokens: int = 1 + # Sampling Params + temperature: float | None = None + top_p: float | None = None + top_k: int | None = None + + +def _generate_kwargs(req: JLensRequest) -> dict: + """Build the sampling kwargs forwarded to ``model.generate(...)`` via the + tool's ``generate_kwargs``. A temperature/top_p/top_k value turns sampling + on; the individual keys are only included when set. + + ``do_sample`` is ALWAYS set explicitly. This matters: many instruct/chat + models ship a ``generation_config.json`` with ``do_sample: true`` (e.g. + Qwen3-8B has do_sample=true, temperature=0.6), so omitting the flag lets the + model's own config silently re-enable sampling — producing a different + completion every run even when the user asked for greedy. Passing + ``do_sample=False`` forces deterministic decoding regardless of the config.""" + kwargs: dict = {} + sample = False + if req.temperature is not None: + kwargs["temperature"] = req.temperature + sample = True + if req.top_p is not None: + kwargs["top_p"] = req.top_p + sample = True + if req.top_k is not None: + kwargs["top_k"] = req.top_k + sample = True + kwargs["do_sample"] = sample + return kwargs class JLensResponse(NDIFResponse): @@ -30,7 +61,18 @@ async def start_j_lens( model = state[req.model] backend = state.make_backend(model=model) - output = j_lens._run(model, req.prompt, remote=state.remote, backend=backend, non_blocking=state.remote, raw=False, top_k=req.topk) + output = j_lens._run( + model, + req.prompt, + remote=state.remote, + backend=backend, + non_blocking=state.remote, + raw=False, + max_new_tokens=req.max_new_tokens, + generate_kwargs=_generate_kwargs(req), + top_k=req.topk, + include_entropy=req.include_entropy, + ) if not backend.blocking: return {"job_id": output} diff --git a/workbench/_api/routes/models.py b/workbench/_api/routes/models.py index cd97f59a..6b5d7db0 100644 --- a/workbench/_api/routes/models.py +++ b/workbench/_api/routes/models.py @@ -2,15 +2,12 @@ import time import requests -import torch as t from fastapi import APIRouter, Depends, HTTPException -from pydantic import BaseModel from nnsightful.tools.j_lens import j_lens -from ..auth import get_user_email, require_user_email, user_has_model_access -from ..data_models import NDIFResponse, Token, ModelHeat -from ..telemetry import TelemetryClient, RequestStatus +from ..auth import get_user_email +from ..data_models import ModelHeat from ..state import AppState, get_state logger = logging.getLogger(__name__) @@ -25,7 +22,7 @@ def _refresh_catalog(state: AppState) -> None: """Hit NDIF /status and rebuild the catalog of deployed models. Caches metadata for any model we haven't seen before; non-pinned models that fell out of the deployment set get unloaded (pinned ones stay loaded).""" - + ping_resp = requests.get(f"{state.ndif_backend_url}/ping", timeout=30) logger.info(f"Call NDIF_BACKEND/ping: {ping_resp.status_code}") if ping_resp.status_code != 200: @@ -110,7 +107,7 @@ async def get_models( # Local models are fully loaded on the dev backend, so they're effectively hot. for model in models: model['status'] = ModelHeat.HOT.value - + ## JLens supported models try: lens_models = j_lens.get_available_lenses() @@ -121,344 +118,3 @@ async def get_models( logger.warning(f"Failed to fetch Jacobian lens availability: {e}") return models - - -class LensCompletion(BaseModel): - model: str - prompt: str - token: Token - - -def prediction( - req: LensCompletion, state: AppState -) -> tuple[t.Tensor, t.Tensor] | str: - model = state[req.model] - idx = req.token.idx - - with model.trace( - req.prompt, - remote=state.remote, - backend=state.make_backend(model=model), - ) as tracer: - logits_BLV = model.logits - - # Get logits for the correct index - logits_LV = logits_BLV[0, [idx], :].softmax(dim=-1) - - # Sort logits by descending probability - values_LV_indices_LV = t.sort(logits_LV, dim=-1, descending=True) - - values_LV = values_LV_indices_LV[0].save() - indices_LV = values_LV_indices_LV[1].save() - - if state.remote: - return tracer.backend.job_id - - return values_LV, indices_LV - -def get_remote_prediction( - job_id: str, state: AppState -) -> tuple[t.Tensor, t.Tensor]: - backend = state.make_backend(job_id=job_id) - results = backend() - return results["values_LV"], results["indices_LV"] - - -class Prediction(BaseModel): - idx: int - ids: list[int] - probs: list[float] - texts: list[str] - - -class PredictionResponse(NDIFResponse): - data: Prediction | None = None - - -def process_prediction( - values_LV: t.Tensor, - indices_LV: t.Tensor, - req: LensCompletion, - state: AppState, -): - tok = state[req.model].tokenizer - idxs = [req.token.idx] - - # Round values to 2 decimal places - idx_values = t.round(values_LV[0] * 100) / 100 - nonzero = idx_values > 0 - - nonzero_values = idx_values[nonzero].tolist() - nonzero_indices = indices_LV[0][nonzero].tolist() - nonzero_texts = tok.batch_decode(nonzero_indices) - - prediction = Prediction( - idx=idxs[0], - ids=nonzero_indices, - probs=nonzero_values, - texts=nonzero_texts, - ) - - return prediction - - -@router.post("/start-prediction", response_model=PredictionResponse) -async def start_prediction( - prediction_request: LensCompletion, - state: AppState = Depends(get_state), - user_email: str = Depends(require_user_email) -): - if state.remote: - if not user_has_model_access(user_email, prediction_request.model, state): - message = f"User does not have access to {prediction_request.model}" - TelemetryClient.log_request( - RequestStatus.ERROR, - user_email, - method="PREDICTION", - type="NEXT_TOKEN", - msg=message, - ) - raise HTTPException(status_code=403, detail=message) - - TelemetryClient.log_request( - RequestStatus.STARTED, - user_email, - method="PREDICTION", - type="NEXT_TOKEN", - ) - - try: - result = prediction(prediction_request, state) - except Exception as e: - TelemetryClient.log_request( - RequestStatus.ERROR, - user_email, - method="PREDICTION", - type="NEXT_TOKEN", - msg=str(e), - ) - raise e - - if state.remote: - TelemetryClient.log_request( - RequestStatus.READY, - user_email, - method="PREDICTION", - type="NEXT_TOKEN", - job_id=result - ) - return {"job_id": result} - - values_LV, indices_LV = result - data = process_prediction(values_LV, indices_LV, prediction_request, state) - return {"data": data} - - -@router.post("/results-prediction/{job_id}", response_model=PredictionResponse) -async def results_prediction( - job_id: str, - prediction_request: LensCompletion, - state: AppState = Depends(get_state), - user_email: str = Depends(require_user_email) -): - - try: - values_LV, indices_LV = get_remote_prediction(job_id, state) - data = process_prediction(values_LV, indices_LV, prediction_request, state) - except Exception as e: - TelemetryClient.log_request( - RequestStatus.ERROR, - user_email, - job_id=job_id, - method="PREDICTION", - type="NEXT_TOKEN", - msg=str(e), - ) - raise e - - TelemetryClient.log_request( - RequestStatus.COMPLETE, - user_email, - job_id=job_id, - method="PREDICTION", - type="NEXT_TOKEN", - ) - - return {"data": data} - - -class Completion(BaseModel): - prompt: str - max_new_tokens: int - model: str - - -class Generation(BaseModel): - completion: list[Token] - last_token_prediction: Prediction - - -class GenerationResponse(NDIFResponse): - data: Generation | None = None - - -def generate(req: Completion, state: AppState): - model = state[req.model] - last_iter = req.max_new_tokens - 1 - with model.generate( - req.prompt, - max_new_tokens=req.max_new_tokens, - remote=state.remote, - backend=state.make_backend(model=model), - ) as tracer: - - with tracer.iter[last_iter]: - logits = model.logits - - probs_V = logits[0, -1, :].softmax(dim=-1) - values_V_indices_V = t.sort(probs_V, dim=-1, descending=True) - values_V = values_V_indices_V[0].save() - indices_V = values_V_indices_V[1].save() - - new_token_ids = model.generator.output[0].save() - - if state.remote: - return tracer.backend.job_id - - return values_V, indices_V, new_token_ids - - -def get_remote_generate( - job_id: str, state: AppState -) -> tuple[t.Tensor, t.Tensor, t.Tensor]: - backend = state.make_backend(job_id=job_id) - results = backend() - return results["values_V"], results["indices_V"], results["new_token_ids"] - - -def process_generation_results( - values_V: t.Tensor, - indices_V: t.Tensor, - new_token_ids: t.Tensor, - req: Completion, - state: AppState, -): - tok = state[req.model].tokenizer - new_token_text = tok.batch_decode(new_token_ids) - - tokens = [ - Token(idx=i, id=new_token_ids[i].item(), text=text, targetIds=[]) - for i, text in enumerate(new_token_text) - ] - - # Round values to 2 decimal places - idx_values = t.round(values_V * 100) / 100 - nonzero = idx_values > 0 - - nonzero_values = idx_values[nonzero].tolist() - nonzero_indices = indices_V[nonzero].tolist() - nonzero_texts = tok.batch_decode(nonzero_indices) - - last_token_prediction = Prediction( - idx=new_token_ids[-1], - ids=nonzero_indices, - probs=nonzero_values, - texts=nonzero_texts, - ).model_dump() - - return { - "completion": tokens, - "last_token_prediction": last_token_prediction, - } - - -@router.post("/start-generate", response_model=GenerationResponse) -async def start_generate( - req: Completion, - state: AppState = Depends(get_state), - user_email: str = Depends(require_user_email) -): - - if state.remote: - if not user_has_model_access(user_email, req.model, state): - message = f"User does not have access to {req.model}" - TelemetryClient.log_request( - RequestStatus.ERROR, - user_email, - method="GENERATE", - type="NEXT_TOKEN", - msg=message, - ) - raise HTTPException(status_code=403, detail=message) - - TelemetryClient.log_request( - RequestStatus.STARTED, - user_email, - method="GENERATE", - type="NEXT_TOKEN", - ) - - try: - result = generate(req, state) - except Exception as e: - TelemetryClient.log_request( - RequestStatus.ERROR, - user_email, - method="GENERATE", - type="NEXT_TOKEN", - msg=str(e), - ) - raise e - - if state.remote: - TelemetryClient.log_request( - RequestStatus.READY, - user_email, - method="GENERATE", - type="NEXT_TOKEN", - job_id=result - ) - return {"job_id": result} - - else: - values_V, indices_V, new_token_ids = result - - data = process_generation_results( - values_V, indices_V, new_token_ids, req, state - ) - return {"data": data} - - -@router.post("/results-generate/{job_id}", response_model=GenerationResponse) -async def results_generate( - job_id: str, - req: Completion, - state: AppState = Depends(get_state), - user_email: str = Depends(require_user_email) -): - - try: - values_V, indices_V, new_token_ids = get_remote_generate(job_id, state) - data = process_generation_results( - values_V, indices_V, new_token_ids, req, state - ) - except Exception as e: - TelemetryClient.log_request( - RequestStatus.ERROR, - user_email, - job_id=job_id, - method="GENERATE", - type="NEXT_TOKEN", - msg=str(e), - ) - raise e - - TelemetryClient.log_request( - RequestStatus.COMPLETE, - user_email, - job_id=job_id, - method="GENERATE", - type="NEXT_TOKEN", - ) - - return {"data": data} diff --git a/workbench/_api/state.py b/workbench/_api/state.py index ce892306..6c5c8a2d 100644 --- a/workbench/_api/state.py +++ b/workbench/_api/state.py @@ -183,14 +183,11 @@ def get_model(self, model_name: str) -> StandardizedTransformer: if model_name not in self.models: self._load_model(model_name) - if model_name not in self.pinned: - self._active_models[model_name] = None - self._evict_if_needed() - - elif model_name not in self.pinned: - # Already loaded; bump recency. - self._active_models.move_to_end(model_name) - + if model_name not in self.pinned: + self._active_models.pop(model_name, None) + self._active_models[model_name] = None + self._evict_if_needed() + return self.models[model_name] def _evict_if_needed(self) -> None: diff --git a/workbench/_web/bun.lock b/workbench/_web/bun.lock index a0688d36..c294328e 100644 --- a/workbench/_web/bun.lock +++ b/workbench/_web/bun.lock @@ -1498,7 +1498,7 @@ "next-themes": ["next-themes@0.4.6", "", { "peerDependencies": { "react": "^16.8 || ^17 || ^18 || ^19 || ^19.0.0-rc", "react-dom": "^16.8 || ^17 || ^18 || ^19 || ^19.0.0-rc" } }, "sha512-pZvgD5L0IEvX5/9GWyHMf3m8BKiVQwsCMHfoFosXtXBMnaS0ZnIJ9ST4b4NqLVKDEm8QBxoNNGNaBv2JNF6XNA=="], - "nnsightful": ["nnsightful@github:AdamBelfki3/nnsightful#75e4267", { "peerDependencies": { "react": "^18.2.0 || ^19.0.0", "react-dom": "^18.2.0 || ^19.0.0" } }, "AdamBelfki3-nnsightful-75e4267"], + "nnsightful": ["nnsightful@github:AdamBelfki3/nnsightful#7123cc7", { "peerDependencies": { "react": "^18.2.0 || ^19.0.0", "react-dom": "^18.2.0 || ^19.0.0" } }, "AdamBelfki3-nnsightful-7123cc7"], "node-abi": ["node-abi@3.75.0", "", { "dependencies": { "semver": "^7.3.5" } }, "sha512-OhYaY5sDsIka7H7AtijtI9jwGYLyl29eQn/W623DiN/MIv5sUqc4g7BIDThX+gb7di9f6xK02nkp8sdfFWZLTg=="], diff --git a/workbench/_web/src/app/workbench/[workspaceId]/[chartId]/components/lens/CompletionCard.tsx b/workbench/_web/src/app/workbench/[workspaceId]/[chartId]/components/lens/CompletionCard.tsx deleted file mode 100644 index fbf4112c..00000000 --- a/workbench/_web/src/app/workbench/[workspaceId]/[chartId]/components/lens/CompletionCard.tsx +++ /dev/null @@ -1,583 +0,0 @@ -"use client"; - -import { ChartLine, Grid3x3, Loader2, TriangleAlert, ChevronDown } from "lucide-react"; -import { Textarea } from "@/components/ui/textarea"; -import { TokenArea } from "./TokenArea"; -import { useState, useEffect, useRef } from "react"; -import { usePrediction } from "@/lib/api/modelsApi"; -import type { LensConfigData, LensHeatmapMetrics, LensLineMetrics } from "@/types/lens"; -import { Metrics } from "@/types/lens"; -import { - DropdownMenu, - DropdownMenuContent, - DropdownMenuItem, - DropdownMenuTrigger, - DropdownMenuLabel, - DropdownMenuSeparator, -} from "@/components/ui/dropdown-menu"; -import { Button } from "@/components/ui/button"; - -import { TargetTokenSelector } from "./TargetTokenSelector"; - -import { encodeText } from "@/actions/tok"; -import { TokenizerLoadError } from "@/actions/errors"; -import { useUpdateChartConfig } from "@/lib/api/configApi"; -import { useParams } from "next/navigation"; -import { useLensCharts } from "@/hooks/useLensCharts"; -import { cn } from "@/lib/utils"; - -import { LensConfig } from "@/db/schema"; -import GenerateButton from "./GenerateButton"; -import { DecoderSelector } from "./DecoderSelector"; -import { ChartType } from "@/types/charts"; -import { Token } from "@/types/models"; -import { Tooltip, TooltipContent, TooltipTrigger } from "@/components/ui/tooltip"; -import { toast } from "sonner"; - -interface CompletionCardProps { - initialConfig: LensConfig; - chartType: ChartType; - selectedModel: string; -} - -// Helper function to capitalize statistic type for display -const capitalizeStatistic = ( - statistic: LensHeatmapMetrics | LensLineMetrics | undefined, -): string => { - const stat = statistic || Metrics.PROBABILITY; - return stat.charAt(0).toUpperCase() + stat.slice(1); -}; - -// Helper function to get valid statistics for a chart type -const getValidStatistics = (chartType: ChartType): (LensHeatmapMetrics | LensLineMetrics)[] => { - if (chartType === "heatmap") { - return [Metrics.PROBABILITY, Metrics.RANK, Metrics.ENTROPY]; - } else { - return [Metrics.PROBABILITY, Metrics.RANK]; - } -}; - -// Helper function to check if a statistic is valid for a chart type -const isStatisticValid = ( - statistic: LensHeatmapMetrics | LensLineMetrics, - chartType: ChartType, -): boolean => { - const validStats = getValidStatistics(chartType); - return validStats.includes(statistic); -}; - -// Helper function to ensure the current statistic is valid for the chart type -const ensureValidStatistic = (config: LensConfigData, chartType: ChartType): LensConfigData => { - if (!isStatisticValid(config.statisticType, chartType)) { - // If current statistic is invalid for this chart type, default to PROBABILITY - return { - ...config, - statisticType: Metrics.PROBABILITY, - }; - } - return config; -}; - -export function CompletionCard({ initialConfig, chartType, selectedModel }: CompletionCardProps) { - const { workspaceId, chartId } = useParams<{ workspaceId: string; chartId: string }>(); - - const [tokenData, setTokenData] = useState([]); - - // creating the default config passed by the lensarea as initial config - const [config, setConfig] = useState(() => { - const baseConfig = { - ...initialConfig.data, - statisticType: initialConfig.data.statisticType || Metrics.PROBABILITY, - }; - return ensureValidStatistic(baseConfig, chartType); - }); - - // whether the chart has been generated? - const [editingText, setEditingText] = useState(initialConfig.data.prediction === undefined); - const [promptHasChangedState, setPromptHasChanged] = useState(false); - - // Track if we should auto-run: only if initial config has a prompt pre-filled - const shouldAutoRunRef = useRef( - initialConfig.data.prompt.length > 0 && !initialConfig.data.prediction, - ); - const hasAutoRunRef = useRef(false); - - const promptHasChanged = promptHasChangedState || config.model !== selectedModel; - - const { mutateAsync: getPrediction, isPending: isExecuting } = usePrediction(); - const { mutateAsync: updateChartConfigMutation } = useUpdateChartConfig(); - - const { handleCreateLineChart, handleCreateHeatmap, isCreatingLineChart, isCreatingHeatmap } = - useLensCharts({ configId: initialConfig.id }); - - // Reset promptHasChanged when config changes (e.g., when switching between different configs) - useEffect(() => { - setPromptHasChanged(false); - // Reset auto-run flags when switching configs - shouldAutoRunRef.current = - initialConfig.data.prompt.length > 0 && !initialConfig.data.prediction; - hasAutoRunRef.current = false; - }, [initialConfig.id, initialConfig.data.prompt.length, initialConfig.data.prediction]); - - // Ensure statistic is valid when chart type changes - useEffect(() => { - setConfig((prevConfig) => ensureValidStatistic(prevConfig, chartType)); - }, [chartType]); - - // Tokenize the prompt if the config changes and there's an existing prediction - useEffect(() => { - const fetchTokens = async () => { - if (config.prediction) { - const tokens = await encodeText(config.prompt, selectedModel); - setTokenData(tokens); - } - }; - fetchTokens(); - }, [initialConfig.id, config.prediction, config.prompt, selectedModel]); - - // Auto-run tokenization and heatmap generation ONLY on initial mount with pre-filled prompt - useEffect(() => { - const autoRunTokenization = async () => { - // Use pre-filled model from config if available, otherwise use selected model - const modelToUse = - initialConfig.data.model && initialConfig.data.model.length > 0 - ? initialConfig.data.model - : selectedModel; - - // Only auto-run if: - // 1. shouldAutoRunRef is true (prompt was pre-filled on mount) - // 2. We haven't auto-run before - // 3. A model is available (either pre-filled or selected) - // 4. Not currently executing - // 5. User hasn't manually edited the prompt - if ( - shouldAutoRunRef.current && - !hasAutoRunRef.current && - modelToUse && - modelToUse.length > 0 && - !isExecuting && - !promptHasChangedState - ) { - hasAutoRunRef.current = true; - shouldAutoRunRef.current = false; // Disable future auto-runs immediately - console.log( - "Auto-running tokenization and heatmap generation for pre-filled prompt:", - initialConfig.data.prompt, - ); - console.log("Using model:", modelToUse); - - try { - // Pass forceRun=true to bypass the promptHasChanged check, and pass modelToUse - await handleTokenize(true, modelToUse); - console.log("Auto-run completed successfully"); - } catch (error) { - console.error("Auto-run failed:", error); - // Don't reset flags - we only try once, even on error - // User can manually run if needed - } - } - }; - - // Small delay to ensure all dependencies are ready - const timer = setTimeout(autoRunTokenization, 800); - return () => clearTimeout(timer); - }, [ - selectedModel, - isExecuting, - promptHasChangedState, - initialConfig.data.prompt, - initialConfig.data.model, - config.model, - ]); - - // Toggle the TokenArea component to the TextArea component - const textareaRef = useRef(null); - const tokenContainerRef = useRef(null); - const settingsRef = useRef(null); - const escapeTokenArea = async () => { - setEditingText(true); - - // Focus the textarea and place cursor at the end after state updates - setTimeout(() => { - if (textareaRef.current) { - textareaRef.current.focus(); - const length = textareaRef.current.value.length; - textareaRef.current.setSelectionRange(length, length); - } - }, 0); - }; - - // Tokenize the prompt and run predictions - const handleTokenize = async (forceRun = false, modelOverride?: string) => { - const modelToUse = modelOverride || selectedModel; - let tokens: Token[]; - try { - tokens = await encodeText(config.prompt, modelToUse); - } catch (error) { - if (error instanceof TokenizerLoadError) { - toast.error( - `Could not load tokenizer for ${modelToUse}. The model may be gated and require authentication.`, - ); - } else { - toast.error("Failed to tokenize prompt."); - } - return; - } - - if (tokens.length <= 1) { - toast.error("Please enter a longer prompt."); - return; - } - - setTokenData(tokens); - // Set the token to the last token in the list - const temporaryConfig: LensConfigData = { - ...config, - model: modelToUse, - token: { idx: tokens[tokens.length - 1].idx, id: 0, text: "", targetIds: [] }, - }; - - if (!promptHasChanged && !forceRun) { - setEditingText(false); - return; - } - - // Run predictions - await runPredictions(temporaryConfig); - await handleCreateHeatmap(temporaryConfig); - setPromptHasChanged(false); - }; - - const handlePromptChange = (e: React.ChangeEvent) => { - setConfig({ - ...config, - prompt: e.target.value, - }); - if (!promptHasChanged) setPromptHasChanged(true); - }; - - const handleStatisticChange = async (value: LensHeatmapMetrics | LensLineMetrics) => { - const updatedConfig = { - ...config, - statisticType: value, - }; - setConfig(updatedConfig); - - // Update the config in the database - await updateChartConfigMutation({ - configId: initialConfig.id, - chartId: chartId, - config: { - data: updatedConfig, - workspaceId, - type: "lens", - }, - }); - - if ( - updatedConfig.prompt && - updatedConfig.prompt.trim().length > 0 && - updatedConfig.prediction - ) { - if (chartType === "heatmap") { - await handleCreateHeatmap(updatedConfig); - } else if (chartType === "line") { - await handleCreateLineChart(updatedConfig); - } - } - }; - - // Newline on shift + enter and tokenize on enter - const handleKeyDown = (e: React.KeyboardEvent) => { - if (e.key === "Enter" && !e.shiftKey && !isExecuting && config.prompt.length > 0) { - if (promptHasChanged) { - e.preventDefault(); - handleTokenize(); - console.log("wefaew", promptHasChanged); - } else { - console.log("promptHasChanged", promptHasChanged); - setEditingText(false); - } - } - }; - - // Auto-resize the textarea to fit its content - const autoResizeTextarea = () => { - if (textareaRef.current) { - textareaRef.current.style.height = "auto"; - textareaRef.current.style.height = `${textareaRef.current.scrollHeight}px`; - } - }; - useEffect(() => { - if (editingText) autoResizeTextarea(); - }, [config.prompt, editingText]); - - // Close editing when focus leaves to outside of textarea, token area, or settings - const handleTextareaBlur = (e: React.FocusEvent) => { - if (!config.prediction) return; // only exit editing once a prediction exists - - // Use setTimeout to allow click events to register first - setTimeout(() => { - const activeElement = document.activeElement; - const withinTextarea = activeElement && textareaRef.current?.contains(activeElement); - const withinToken = activeElement && tokenContainerRef.current?.contains(activeElement); - const withinSettings = activeElement && settingsRef.current?.contains(activeElement); - - // Check if a popover is open (Radix UI adds data-state="open" to popovers) - const popoverOpen = document.querySelector("[data-radix-popper-content-wrapper]"); - - // if (promptHasChanged) { - // handleTokenize(); - // } - - if (withinTextarea || withinToken || withinSettings || popoverOpen) return; - - setEditingText(false); - }, 0); - }; - - const runPredictions = async (temporaryConfig: LensConfigData) => { - // Run predictions for the selected token in the config - const prediction = await getPrediction(temporaryConfig); - const topThree = prediction.ids.slice(0, 3); - - // Update the config locally - temporaryConfig.prediction = prediction; - temporaryConfig.token.targetIds = topThree; - setConfig(temporaryConfig); - - // Update the config in the database - await updateChartConfigMutation({ - configId: initialConfig.id, - chartId: chartId, - config: { - data: temporaryConfig, - workspaceId, - type: "lens", - }, - }); - - // Exit the editing state - setEditingText(false); - }; - - const handleTokenClick = async (event: React.MouseEvent, idx: number) => { - // Prevent the editing state from activating - event.preventDefault(); - event.stopPropagation(); - - // Skip if the token is already selected - if (config.token.idx === idx) return; - - // Set the token to the last token in the list - const temporaryConfig: LensConfigData = { - ...config, - token: { idx, id: 0, text: "", targetIds: [] }, - }; - - // Run predictions - await runPredictions(temporaryConfig); - - setConfig(temporaryConfig); - await handleCreateLineChart(temporaryConfig); - }; - - return ( -
- {/* Content */} -
- {editingText ? ( -