liautaud/Next_token_probability_quantized
0
1import math2import gradio as gr3from transformers import AutoTokenizer, AutoModelForCausalLM4import torch5 6# ---- Model config ----7MODEL_NAME = "Sreehariiii/BioGPT-4bit-quantized-version" # e.g., "distilgpt2", "gpt2-medium"8device = "cuda" if torch.cuda.is_available() else "cpu"9 10tok = AutoTokenizer.from_pretrained(MODEL_NAME)11model = AutoModelForCausalLM.from_pretrained(MODEL_NAME).to(device)12model.eval()13 14EPS = 1e-9 # tie tolerance15 16def safe_exp(x: float) -> str:17 # Pretty string even for big magnitudes18 try:19 return f"{math.exp(x):.6e}"20 except OverflowError:21 return "∞ (overflow)"22 except Exception:23 return "—"24 25def is_finite(x: float) -> bool:26 return x is not None and math.isfinite(x)27 28def seq_logprob(context: str, candidate: str, assume_leading_space: bool, show_topk: int):29 """30 Compute log P(candidate | context) via chain rule over tokens.31 Returns (total_logprob, detail_text, token_list, num_tokens).32 """33 cand_text = (" " + candidate) if assume_leading_space else candidate34 with torch.no_grad():35 ctx_ids = tok.encode(context, return_tensors="pt").to(device)36 cand_ids = tok.encode(cand_text, add_special_tokens=False)37 38 if len(cand_ids) == 0:39 return None, "Candidate tokenized to empty sequence (check spacing).", [], 040 41 total_logprob = 0.042 step_lines = []43 input_ids = ctx_ids44 token_texts = []45 46 for i, t_id in enumerate(cand_ids):47 outputs = model(input_ids=input_ids)48 logits = outputs.logits[:, -1, :]49 logprobs = torch.log_softmax(logits, dim=-1)50 token_lp = logprobs[0, t_id].item()51 total_logprob += token_lp52 53 tok_str = tok.decode([t_id])54 token_texts.append(tok_str)55 56 if show_topk > 0:57 topk_vals, topk_idx = torch.topk(logprobs, k=min(show_topk, logprobs.shape[-1]), dim=-1)58 tops = ", ".join([f"{repr(tok.decode([int(idx)]))}:{val.item():.2f}"59 for idx, val in zip(topk_idx[0], topk_vals[0])])60 step_lines.append(61 f"Step {i+1}: token={repr(tok_str)} logprob={token_lp:.6f} "62 f"prob={math.exp(token_lp):.6e}\n top-{show_topk}: {tops}"63 )64 else:65 step_lines.append(66 f"Step {i+1}: token={repr(tok_str)} logprob={token_lp:.6f} "67 f"prob={math.exp(token_lp):.6e}"68 )69 70 # teacher-forcing: append the true token to continue conditioning71 input_ids = torch.cat([input_ids, torch.tensor([[t_id]], device=device)], dim=1)72 73 return total_logprob, "\n".join(step_lines), token_texts, len(cand_ids)74 75def compare_candidates(context, candA, candB, assume_space, topk, use_len_norm):76 # Basic checks77 errors = []78 if not context.strip():79 errors.append("Please enter a context.")80 if not candA.strip():81 errors.append("Please enter Candidate A.")82 if not candB.strip():83 errors.append("Please enter Candidate B.")84 if errors:85 msg = " ".join(errors)86 return (f"<div style='color:#b00020;font-weight:600'>{msg}</div>",87 "", "", "", "", "")88 89 # Compute90 lpA, detA, toksA, nA = seq_logprob(context, candA, assume_space, topk)91 lpB, detB, toksB, nB = seq_logprob(context, candB, assume_space, topk)92 93 # Validate numbers94 if not (is_finite(lpA) and is_finite(lpB)):95 return ("<div style='color:#b00020;font-weight:600'>Numerical issue (NaN/Inf). "96 "Try shorter context, smaller model (e.g., distilgpt2), or disable length normalization.</div>",97 "", "", "", "", "")98 99 # Optionally length-normalize (per-token average log-prob)100 # Note: odds under length-normalization are "per-token odds", not whole-sequence odds.101 if use_len_norm:102 if nA == 0 or nB == 0:103 return ("<div style='color:#b00020;font-weight:600'>Empty tokenization. "104 "Check spacing or turn off 'assume leading space'.</div>",105 "", "", "", "", "")106 scoreA = lpA / nA107 scoreB = lpB / nB108 label_suffix = " (per-token)"109 else:110 scoreA = lpA111 scoreB = lpB112 label_suffix = ""113 114 diff = scoreA - scoreB # log-odds if unnormalized; log per-token odds otherwise115 116 # Winner logic with proper tie handling117 if abs(diff) <= EPS:118 winner = "Tie"119 elif diff > 0:120 winner = "Candidate A"121 else:122 winner = "Candidate B"123 124 ratio_str = safe_exp(diff)125 126 # Colors127 if winner == "Candidate A":128 win_color = "#166534" # green129 elif winner == "Candidate B":130 win_color = "#1d4ed8" # blue131 else:132 win_color = "#92400e" # amber133 134 # Headline135 headline = (136 f"<div style='padding:14px;border-radius:12px;background:#f8fafc;"137 f"border:1px solid #e2e8f0;margin-bottom:10px'>"138 f"<div style='font-size:20px;font-weight:800;color:{win_color};'>Winner: {winner}{label_suffix}</div>"139 f"<div style='margin-top:6px;font-size:16px;'>"140 f"Odds A/B{label_suffix} = <b>{ratio_str}</b> | "141 f"log-odds A−B{label_suffix} = <b>{diff:.6f}</b>"142 f"</div>"143 f"<div style='margin-top:6px;color:#475569'>"144 f"(Odds > 1 ⇒ A more probable; < 1 ⇒ B more probable. "145 f"{'Per-token uses average log-prob.' if use_len_norm else 'Whole-sequence comparison.'})"146 f"</div></div>"147 )148 149 def summarize(label, cand, lp, toks, n):150 return (151 f"**{label}**: {cand}\n\n"152 f"Tokenization: {toks}\n"153 f"Total logprob: {lp:.6f}\n"154 f"Sequence probability: {math.exp(lp):.6e}\n"155 f"Tokens: {n}"156 )157 158 sumA = summarize("Candidate A", candA, lpA, toksA, nA)159 sumB = summarize("Candidate B", candB, lpB, toksB, nB)160 161 return headline, sumA, detA, sumB, detB, ""162 163def swap(a, b):164 return b, a165 166with gr.Blocks(title="Two-Candidate Next-Token Comparator (Robust)") as demo:167 gr.Markdown(168 "# Two-Candidate Next-Word/Token Probability (Robust)\n"169 "Compare **P(A|context)** vs **P(B|context)** from a pretrained causal LM (no fine-tuning).\n"170 "- Proper tie handling and numerical guards.\n"171 "- Optional **length normalization** (per-token).\n"172 "- Use **Swap** to sanity-check symmetry."173 )174 with gr.Row():175 context = gr.Textbox(label="Context (prompt)", lines=6, placeholder="Paste prior text here...")176 with gr.Row():177 candA = gr.Textbox(label="Candidate A (follow-up)")178 candB = gr.Textbox(label="Candidate B (follow-up)")179 with gr.Row():180 assume_space = gr.Checkbox(value=True, label="Assume leading space before candidates (useful for GPT-2 tokenization)")181 topk = gr.Slider(0, 20, value=5, step=1, label="Show top-k alternatives (per token step)")182 use_len_norm = gr.Checkbox(value=False, label="Use length normalization (average log-prob per token)")183 with gr.Row():184 btn_compare = gr.Button("Compare", variant="primary")185 btn_swap = gr.Button("Swap A ↔ B")186 187 winner_html = gr.HTML()188 summaryA = gr.Markdown()189 detailsA = gr.Textbox(label="Candidate A — step-by-step", lines=10)190 summaryB = gr.Markdown()191 detailsB = gr.Textbox(label="Candidate B — step-by-step", lines=10)192 _hidden = gr.Textbox(visible=False)193 194 btn_compare.click(195 fn=compare_candidates,196 inputs=[context, candA, candB, assume_space, topk, use_len_norm],197 outputs=[winner_html, summaryA, detailsA, summaryB, detailsB, _hidden]198 )199 200 btn_swap.click(201 fn=swap, inputs=[candA, candB], outputs=[candA, candB]202 )203 204demo.launch()