liautaud's picture
fkn moron
69ddaf5 verified
Raw
History Blame Contribute Delete
17.7 kB
import math
import os
import time
import gradio as gr
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
# ---- Model config ----
MODEL_NAME = "microsoft/biogpt" # e.g. "distilgpt2", "HuggingFaceTB/SmolLM2-135M"
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
FULL_BATCH_CONTEXT_THRESHOLD = 64 # short contexts are often faster as one batched full forward
EPS = 1e-9
# Set TORCH_COMPILE=1 to compile.
ENABLE_TORCH_COMPILE = os.environ.get("TORCH_COMPILE", "0") != "0"
COMPILE_STATUS = "not attempted"
# T4-friendly defaults: fp16 on CUDA, no gradients, eval mode.
torch.set_grad_enabled(False)
if DEVICE == "cuda":
torch.backends.cudnn.benchmark = True
torch.backends.cuda.matmul.allow_tf32 = True # harmless on T4; useful on newer GPUs
try:
torch.set_float32_matmul_precision("high")
except Exception:
pass
model_kwargs = {}
if DEVICE == "cuda":
model_kwargs["torch_dtype"] = torch.float16
tok = AutoTokenizer.from_pretrained(MODEL_NAME)
model = AutoModelForCausalLM.from_pretrained(MODEL_NAME, **model_kwargs).to(DEVICE)
model.eval()
model.config.use_cache = True
# Causal LMs such as GPT-2 often have no pad token. Right padding is safe here
# because we score only real positions and pass an attention mask.
if tok.pad_token_id is None:
if tok.eos_token_id is not None:
tok.pad_token = tok.eos_token
else:
tok.add_special_tokens({"pad_token": "<|pad|>"})
model.resize_token_embeddings(len(tok))
model.config.pad_token_id = tok.pad_token_id
PAD_ID = tok.pad_token_id
# Pre-tokenized hot candidates. Exact string match only.
# Add more dictation commands/ambiguities here as they become common.
PRETOKENIZED_CANDIDATES = {
text: tok.encode(text, add_special_tokens=False)
for text in ("column", "colon", ":")
}
# Optional torch.compile. This can reduce Python/dispatch overhead after warmup, but
# may not help all GPU/model/shape combinations, so it is deliberately best-effort.
if ENABLE_TORCH_COMPILE and hasattr(torch, "compile"):
try:
model = torch.compile(model, mode="reduce-overhead", fullgraph=False)
COMPILE_STATUS = "torch.compile attempted: enabled"
except Exception as exc:
COMPILE_STATUS = f"torch.compile attempted: failed ({type(exc).__name__})"
elif not ENABLE_TORCH_COMPILE:
COMPILE_STATUS = "torch.compile disabled by TORCH_COMPILE=0"
else:
COMPILE_STATUS = "torch.compile unavailable in this PyTorch"
def cuda_sync() -> None:
if DEVICE == "cuda":
torch.cuda.synchronize()
def now_ms() -> float:
return time.perf_counter() * 1000.0
def safe_exp(x: float) -> str:
try:
return f"{math.exp(x):.6e}"
except OverflowError:
return "inf (overflow)"
except Exception:
return "-"
def is_finite(x: float) -> bool:
return x is not None and math.isfinite(x)
def encode_candidate(candidate: str) -> list[int]:
cached = PRETOKENIZED_CANDIDATES.get(candidate)
if cached is not None:
# Return a copy so downstream code can treat all candidate id lists normally.
return list(cached)
return tok.encode(candidate, add_special_tokens=False)
def encode_inputs(context: str, candidates: list[str]):
ctx_ids_cpu = tok.encode(context, return_tensors="pt").squeeze(0)
cand_ids_list = [encode_candidate(c) for c in candidates]
return ctx_ids_cpu, cand_ids_list
def token_strings(cand_ids: list[int]) -> list[str]:
return [tok.decode([token_id]) for token_id in cand_ids]
def empty_result(cand_ids: list[int], message: str):
return {
"total_score": None,
"score_kind": "not_scored",
"token_list": token_strings(cand_ids),
"num_tokens": len(cand_ids),
"per_token_scores": [],
"message": message,
}
def make_result(cand_ids: list[int], per_token_scores: list[float], score_kind: str):
return {
"total_score": float(sum(per_token_scores)),
"score_kind": score_kind, # "raw_logit" for one-token fast path, "logprob" otherwise
"token_list": token_strings(cand_ids),
"num_tokens": len(cand_ids),
"per_token_scores": [float(x) for x in per_token_scores],
"message": "",
}
def repeat_cache_to_batch(cache, batch_size: int):
"""Repeat a KV cache from batch size 1 to batch_size."""
if cache is None:
return None
# Newer Transformers Cache objects may expose batch_repeat_interleave.
if hasattr(cache, "batch_repeat_interleave"):
maybe_returned = cache.batch_repeat_interleave(batch_size)
return cache if maybe_returned is None else maybe_returned
if torch.is_tensor(cache):
# expand avoids work, contiguous makes it safe for all attention implementations.
return cache.expand(batch_size, *([-1] * (cache.dim() - 1))).contiguous()
if isinstance(cache, tuple):
return tuple(repeat_cache_to_batch(x, batch_size) for x in cache)
if isinstance(cache, list):
return [repeat_cache_to_batch(x, batch_size) for x in cache]
raise TypeError(f"Unsupported cache type: {type(cache)}")
def score_single_token_fast(ctx_ids_cpu: torch.Tensor, cand_ids_list: list[list[int]]):
"""
Fastest path: one context forward, then gather raw next-token logits.
No softmax/log_softmax is needed. For candidates scored from the same next-token
distribution, logprob(A) - logprob(B) == logit(A) - logit(B) because the shared
normalization denominator cancels.
"""
ctx_ids = ctx_ids_cpu.unsqueeze(0).to(DEVICE)
cuda_sync()
t0 = now_ms()
with torch.inference_mode():
outputs = model(input_ids=ctx_ids, use_cache=False)
next_logits = outputs.logits[:, -1, :].float().squeeze(0)
token_ids = torch.tensor([ids[0] for ids in cand_ids_list], dtype=torch.long, device=DEVICE)
raw_logits = next_logits.index_select(0, token_ids)
cuda_sync()
model_ms = now_ms() - t0
scores = raw_logits.detach().cpu().tolist()
results = [make_result(ids, [score], "raw_logit") for ids, score in zip(cand_ids_list, scores)]
return results, model_ms, "single context forward for 1-token candidates; raw logits, no softmax"
def score_full_batch(ctx_ids_cpu: torch.Tensor, cand_ids_list: list[list[int]]):
"""Score context+candidate sequences in one right-padded batched forward."""
ctx_len = int(ctx_ids_cpu.numel())
seqs = [torch.cat([ctx_ids_cpu, torch.tensor(ids, dtype=torch.long)]) for ids in cand_ids_list]
max_len = max(int(seq.numel()) for seq in seqs)
batch_size = len(seqs)
input_ids = torch.full((batch_size, max_len), PAD_ID, dtype=torch.long)
attention_mask = torch.zeros((batch_size, max_len), dtype=torch.long)
for row, seq in enumerate(seqs):
seq_len = int(seq.numel())
input_ids[row, :seq_len] = seq
attention_mask[row, :seq_len] = 1
input_ids = input_ids.to(DEVICE)
attention_mask = attention_mask.to(DEVICE)
cuda_sync()
t0 = now_ms()
with torch.inference_mode():
outputs = model(input_ids=input_ids, attention_mask=attention_mask, use_cache=False)
logits = outputs.logits.float()
cuda_sync()
model_ms = now_ms() - t0
results = []
for row, cand_ids in enumerate(cand_ids_list):
per_token_lps = []
for i, token_id in enumerate(cand_ids):
pred_pos = ctx_len + i - 1
step_logits = logits[row, pred_pos, :]
lp = float((step_logits[token_id] - torch.logsumexp(step_logits, dim=-1)).item())
per_token_lps.append(lp)
results.append(make_result(cand_ids, per_token_lps, "logprob"))
return results, model_ms, "one full batched forward pass"
def score_with_prefix_cache(ctx_ids_cpu: torch.Tensor, cand_ids_list: list[list[int]]):
"""Score long-context multi-token candidates using one context pass + one batched continuation pass."""
ctx_ids = ctx_ids_cpu.unsqueeze(0).to(DEVICE)
ctx_len = int(ctx_ids.shape[1])
batch_size = len(cand_ids_list)
max_cand_len = max(len(ids) for ids in cand_ids_list)
cuda_sync()
t0 = now_ms()
with torch.inference_mode():
prefix_outputs = model(input_ids=ctx_ids, use_cache=True)
prefix_logits = prefix_outputs.logits[:, -1, :].float()
suffix_logits = None
if max_cand_len > 1:
past = repeat_cache_to_batch(prefix_outputs.past_key_values, batch_size)
suffix_len = max_cand_len - 1
suffix_input_ids = torch.full((batch_size, suffix_len), PAD_ID, dtype=torch.long, device=DEVICE)
suffix_attention = torch.zeros((batch_size, suffix_len), dtype=torch.long, device=DEVICE)
for row, ids in enumerate(cand_ids_list):
prefix_ids = ids[:-1]
if prefix_ids:
n = len(prefix_ids)
suffix_input_ids[row, :n] = torch.tensor(prefix_ids, dtype=torch.long, device=DEVICE)
suffix_attention[row, :n] = 1
full_attention = torch.cat(
[torch.ones((batch_size, ctx_len), dtype=torch.long, device=DEVICE), suffix_attention],
dim=1,
)
suffix_outputs = model(
input_ids=suffix_input_ids,
attention_mask=full_attention,
past_key_values=past,
use_cache=False,
)
suffix_logits = suffix_outputs.logits.float()
cuda_sync()
model_ms = now_ms() - t0
results = []
prefix_log_denom = torch.logsumexp(prefix_logits[0], dim=-1)
for row, cand_ids in enumerate(cand_ids_list):
per_token_lps = []
first_token_id = cand_ids[0]
first_lp = float((prefix_logits[0, first_token_id] - prefix_log_denom).item())
per_token_lps.append(first_lp)
for i in range(1, len(cand_ids)):
step_logits = suffix_logits[row, i - 1, :]
token_id = cand_ids[i]
lp = float((step_logits[token_id] - torch.logsumexp(step_logits, dim=-1)).item())
per_token_lps.append(lp)
results.append(make_result(cand_ids, per_token_lps, "logprob"))
return results, model_ms, "shared-prefix cache + batched candidate pass"
def score_candidates(context: str, candidates: list[str]):
"""
Compute scores for all candidates.
Candidates are scored exactly as typed. No leading space is added.
One-token candidates use raw logits without softmax/log_softmax; multi-token
candidates use summed log probabilities.
"""
total_t0 = now_ms()
ctx_ids_cpu, cand_ids_list = encode_inputs(context, candidates)
if any(len(ids) == 0 for ids in cand_ids_list):
results = [
empty_result(ids, "Candidate tokenized to an empty sequence. Type the candidate exactly as you want it scored.")
if len(ids) == 0
else empty_result(ids, "Not scored because another candidate tokenized to an empty sequence.")
for ids in cand_ids_list
]
return results, 0.0, now_ms() - total_t0, "not scored"
max_cand_len = max(len(ids) for ids in cand_ids_list)
ctx_len = int(ctx_ids_cpu.numel())
try:
if max_cand_len == 1:
results, model_ms, mode = score_single_token_fast(ctx_ids_cpu, cand_ids_list)
elif ctx_len <= FULL_BATCH_CONTEXT_THRESHOLD:
results, model_ms, mode = score_full_batch(ctx_ids_cpu, cand_ids_list)
else:
results, model_ms, mode = score_with_prefix_cache(ctx_ids_cpu, cand_ids_list)
except Exception:
# Robust fallback for any model/Transformers cache/compile incompatibility.
results, model_ms, mode = score_full_batch(ctx_ids_cpu, cand_ids_list)
mode = f"fallback: {mode}"
total_ms = now_ms() - total_t0
return results, model_ms, total_ms, mode
def compare_candidates(context, candA, candB, use_len_norm):
request_t0 = now_ms()
errors = []
if not context.strip():
errors.append("Please enter a context.")
if not candA.strip():
errors.append("Please enter Candidate A.")
if not candB.strip():
errors.append("Please enter Candidate B.")
if errors:
msg = " ".join(errors)
return f"<div style='color:#b00020;font-weight:600'>{msg}</div>", "", ""
scored, model_ms, scoring_ms, inference_mode = score_candidates(context, [candA, candB])
resA, resB = scored[0], scored[1]
rawA = resA["total_score"]
rawB = resB["total_score"]
nA = resA["num_tokens"]
nB = resB["num_tokens"]
if not (is_finite(rawA) and is_finite(rawB)):
return (
"<div style='color:#b00020;font-weight:600'>Numerical or tokenization issue. "
"Try shorter context, a smaller model, or check the candidate text.</div>",
summarize_candidate("Candidate A", candA, resA),
summarize_candidate("Candidate B", candB, resB),
)
if use_len_norm:
scoreA = rawA / nA
scoreB = rawB / nB
label_suffix = " (per-token)"
else:
scoreA = rawA
scoreB = rawB
label_suffix = ""
diff = scoreA - scoreB
if abs(diff) <= EPS:
winner = "Tie"
win_color = "#92400e"
elif diff > 0:
winner = "Candidate A"
win_color = "#166534"
else:
winner = "Candidate B"
win_color = "#1d4ed8"
request_ms = now_ms() - request_t0
ratio_str = safe_exp(diff)
score_kind = resA["score_kind"] if resA["score_kind"] == resB["score_kind"] else "mixed"
ratio_label = "exp(raw-logit difference)" if score_kind == "raw_logit" else "odds A/B"
headline = (
f"<div style='padding:14px;border-radius:12px;background:#f8fafc;"
f"border:1px solid #e2e8f0;margin-bottom:10px'>"
f"<div style='font-size:20px;font-weight:800;color:{win_color};'>Winner: {winner}{label_suffix}</div>"
f"<div style='margin-top:6px;font-size:16px;'>"
f"{ratio_label}{label_suffix} = <b>{ratio_str}</b> &nbsp;|&nbsp; "
f"score diff A-B{label_suffix} = <b>{diff:.6f}</b>"
f"</div>"
f"<div style='margin-top:6px;font-size:16px;'>"
f"Model inference = <b>{model_ms:.2f} ms</b> &nbsp;|&nbsp; "
f"Scoring total = <b>{scoring_ms:.2f} ms</b> &nbsp;|&nbsp; "
f"Request function = <b>{request_ms:.2f} ms</b>"
f"</div>"
f"<div style='margin-top:6px;color:#475569'>"
f"Mode: {inference_mode} on {DEVICE}. Compile: {COMPILE_STATUS}. "
f"Pre-tokenized exact candidates: {', '.join(repr(k) for k in PRETOKENIZED_CANDIDATES.keys())}. "
f"Candidates are scored exactly as typed; no leading space is added. "
f"{'Per-token uses average score.' if use_len_norm else 'Whole-sequence comparison.'}"
f"</div></div>"
)
return headline, summarize_candidate("Candidate A", candA, resA), summarize_candidate("Candidate B", candB, resB)
def summarize_candidate(label: str, cand: str, res: dict) -> str:
if res["total_score"] is None:
return (
f"**{label}**: {repr(cand)}\n\n"
f"Tokenization: {res['token_list']}\n"
f"Tokens: {res['num_tokens']}\n"
f"{res['message']}"
)
per_token = ", ".join(f"{x:.4f}" for x in res["per_token_scores"])
score_kind = res["score_kind"]
if score_kind == "raw_logit":
score_lines = (
f"Raw logit score: {res['total_score']:.6f}\n"
f"Per-token raw logits: [{per_token}]\n"
"Sequence probability: not computed in one-token fast path\n"
)
else:
score_lines = (
f"Total logprob: {res['total_score']:.6f}\n"
f"Sequence probability: {math.exp(res['total_score']):.6e}\n"
f"Per-token logprobs: [{per_token}]\n"
)
return (
f"**{label}**: {repr(cand)}\n\n"
f"Tokenization: {res['token_list']}\n"
f"{score_lines}"
f"Tokens: {res['num_tokens']}"
)
def swap(a, b):
return b, a
with gr.Blocks(title="Ultra-Fast Two-Candidate Next-Token Comparator") as demo:
gr.Markdown(
"# Ultra-Fast Two-Candidate Next-Word/Token Comparator\n"
"Compare candidate continuations from a pretrained causal LM.\n"
"- One-token candidates use raw logits only: no softmax/log_softmax.\n"
"- Exact candidates `column`, `colon`, and `:` are pre-tokenized at startup.\n"
"- The app doesn't attempt `torch.compile(..., mode='reduce-overhead')` unless `TORCH_COMPILE=1`.\n"
"- Candidates are scored exactly as typed; no leading space is automatically added.\n"
"- Multi-token candidates still use summed log probabilities."
)
with gr.Row():
context = gr.Textbox(label="Context (prompt)", lines=6, placeholder="Paste prior text here...")
with gr.Row():
candA = gr.Textbox(label="Candidate A", value="colon")
candB = gr.Textbox(label="Candidate B", value=":")
with gr.Row():
use_len_norm = gr.Checkbox(value=False, label="Use length normalization (average score per token)")
with gr.Row():
btn_compare = gr.Button("Compare", variant="primary")
btn_swap = gr.Button("Swap A <-> B")
winner_html = gr.HTML()
with gr.Row():
summaryA = gr.Markdown()
summaryB = gr.Markdown()
btn_compare.click(
fn=compare_candidates,
inputs=[context, candA, candB, use_len_norm],
outputs=[winner_html, summaryA, summaryB],
)
btn_swap.click(fn=swap, inputs=[candA, candB], outputs=[candA, candB])
demo.launch()