Jacobina / api /helpers.py
marinarosa's picture
initial commit
be82719
Raw
History Blame Contribute Delete
8.51 kB
"""Pure helpers shared across API handlers.
Ported from miru-tracer's ``ui/helpers.py`` and ``ui/lens_common.py`` minus
everything Gradio-specific. The frontend sends internal keys directly
("completion"/"chat"/"raw", "auto"/"off"/"prefill", "logit"/"jacobian"/"diff",
"adjusted"/"raw", "text"/"id") so no label mapping happens server-side.
"""
from __future__ import annotations
import json
from miru_tracer.core.interventions import Intervention
from miru_tracer.core.sampling import SamplingParams
LENS_MODES = ("logit", "jacobian", "diff")
GENERATION_MODES = ("completion", "chat", "raw")
THINKING_MODES = ("auto", "off", "prefill")
def ui_sampling_params(strategy, temperature, top_k, top_p) -> SamplingParams:
"""Build SamplingParams from raw widget values (clamped, not raising)."""
return SamplingParams(
strategy=strategy,
temperature=float(temperature),
top_k=int(top_k or 0),
top_p=min(max(float(top_p), 1e-3), 1.0),
)
class ChatValidationError(ValueError):
"""The chat JSON the user entered is not a valid message list."""
def parse_chat_messages(text: str) -> list[dict[str, str]]:
"""Parse and validate the chat-messages JSON from the UI.
Raises:
ChatValidationError: with a user-presentable message.
"""
try:
messages = json.loads(text)
except json.JSONDecodeError as e:
raise ChatValidationError(f"Invalid JSON: {e}") from e
if not isinstance(messages, list) or not messages:
raise ChatValidationError("Chat messages must be a non-empty JSON array")
for message in messages:
if (
not isinstance(message, dict)
or "role" not in message
or "content" not in message
):
raise ChatValidationError(
"Each message must have 'role' and 'content' fields"
)
return messages
def token_ref_to_id(ref: str, tokenizer, mode: str) -> int:
"""Resolve a user token reference under an explicit interpretation mode.
``mode == "id"``: parse a numeric token id; surrounding whitespace is
tolerated and non-numeric input is rejected. ``mode == "text"``: encode
the text verbatim — leading/trailing whitespace is significant (" Paris"
and "Paris" are different BPE tokens) and digits are NEVER treated as an
id — returning the first token id.
Raises:
ValueError: empty ref, unencodable text, non-numeric id, or an
out-of-range id.
"""
if mode == "id":
stripped = ref.strip()
if not stripped:
raise ValueError("Empty token reference")
if not stripped.lstrip("-").isdigit():
raise ValueError(f"Not a numeric token id: {ref!r}")
token_id = int(stripped)
if not 0 <= token_id < len(tokenizer):
raise ValueError(
f"Token id {token_id} out of range (vocab size {len(tokenizer)})"
)
return token_id
if not ref.strip():
raise ValueError("Empty token reference")
encoded = tokenizer.encode(ref, add_special_tokens=False)
if not encoded:
raise ValueError(f"Could not tokenize {ref!r}")
return int(encoded[0])
def parse_layer_refs(text: str) -> list[int]:
"""Comma-separated layers and inclusive ranges -> unique sorted layer list.
E.g. ``"11, 12-15, 18"`` -> ``[11, 12, 13, 14, 15, 18]``.
Raises:
ValueError: empty input, malformed entry, or a descending range.
"""
layers: list[int] = []
for part in (str(text) if text is not None else "").split(","):
part = part.strip()
if not part:
continue
lo, sep, hi = part.partition("-")
lo, hi = lo.strip(), hi.strip()
if not lo.isdigit() or (sep and not hi.isdigit()):
raise ValueError(
f"Bad layer reference {part!r}: use a number or range like 12-15"
)
start, end = int(lo), int(hi) if sep else int(lo)
if end < start:
raise ValueError(f"Descending layer range {part!r}")
layers.extend(range(start, end + 1))
if not layers:
raise ValueError("Empty layer reference")
return sorted(set(layers))
def layer_selection(n_layers: int, start, end, stride) -> list[int]:
"""Resolve UI layer-range inputs into a concrete layer list.
``end`` is inclusive; -1 (or blank) means the final layer. The final
selected layer is always included even if the stride skips it.
"""
start = int(start) if start is not None else 0
end = int(end) if end is not None else -1
stride = max(int(stride) if stride else 1, 1)
if end < 0:
end = n_layers - 1
start = max(0, min(start, n_layers - 1))
end = max(start, min(end, n_layers - 1))
layers = list(range(start, end + 1, stride))
if layers[-1] != end:
layers.append(end)
return layers
# ------------------------------------------------------------- interventions
def intervention_signature(iv: Intervention) -> tuple:
"""Stable key for deciding whether two UI interventions are duplicates."""
if iv.kind == "steer":
return (iv.kind, iv.layer, iv.token_id, float(iv.strength), iv.basis)
if iv.kind == "swap":
return (iv.kind, iv.layer, iv.token_id, iv.token_id_to, iv.basis)
return (iv.kind, iv.layer, iv.token_id, iv.basis)
def add_unique_intervention_rows(rows: list, candidates: list) -> tuple[list, list, int]:
"""Append candidates whose effective intervention parameters are new."""
updated = list(rows or [])
seen = {intervention_signature(row["intervention"]) for row in updated}
added = []
skipped = 0
for row in candidates:
signature = intervention_signature(row["intervention"])
if signature in seen:
skipped += 1
continue
updated.append(row)
added.append(row)
seen.add(signature)
return updated, added, skipped
def enabled_interventions(rows: list) -> list[Intervention]:
return [row["intervention"] for row in rows or [] if row.get("enabled", True)]
def describe_with_basis(iv: Intervention, tokenizer=None) -> str:
"""Human description of an intervention with its basis appended."""
return f"{iv.describe(tokenizer)} ({iv.basis})"
def intervened_layer_titles(
interventions: list[Intervention], tokenizer=None
) -> dict[int, str]:
"""Map each edited layer to a ``'; '``-joined description of its edits."""
titles: dict[int, str] = {}
for iv in interventions:
desc = describe_with_basis(iv, tokenizer)
titles[iv.layer] = f"{titles[iv.layer]}; {desc}" if iv.layer in titles else desc
return titles
def interventions_summary(
interventions: list[Intervention], tokenizer=None, *, limit: int = 4
) -> str:
"""One-line summary of the active interventions for the status area."""
parts = [describe_with_basis(iv, tokenizer) for iv in interventions[:limit]]
if len(interventions) > limit:
parts.append(f"+{len(interventions) - limit} more")
return "; ".join(parts)
def intervention_visibility_warning(
interventions: list[Intervention], mode: str, n_layers: int, tokenizer=None
) -> str | None:
"""Warn when an edit's basis differs from the current lens view mode.
A jacobian-basis edit moves the residual along pre-transport directions
(``J_ℓ v = û_t``), visible under the Jacobian lens but nearly invisible
under the Logit lens — and vice versa. The final layer is basis-independent
(both bases use ``û_t`` directly), and Diff renders both readouts, so
neither triggers a warning. For a fixed ``mode`` at most one basis can
mismatch, so this returns a single line (or None).
"""
if mode == "diff":
return None
final = n_layers - 1
mismatched = [iv for iv in interventions if iv.layer != final and iv.basis != mode]
if not mismatched:
return None
basis = mismatched[0].basis # only one basis can mismatch a given mode
shown = "; ".join(iv.describe(tokenizer) for iv in mismatched[:3])
if len(mismatched) > 3:
shown += f" (+{len(mismatched) - 3} more)"
verb = "uses" if len(mismatched) == 1 else "use"
pronoun = "its" if len(mismatched) == 1 else "their"
want = "Jacobian" if basis == "jacobian" else "Logit"
return (
f"⚠ {shown} {verb} {basis} basis — switch Lens to {want} (or Diff) "
f"to see {pronoun} effect in the readouts."
)