CoolFace
Apppublic

davidbeaver/Next_token_probability

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
1likes
app.py204 linesDownload Raw Back to root
1import math2import gradio as gr3from transformers import AutoTokenizer, AutoModelForCausalLM4import torch5 6# ---- Model config ----7MODEL_NAME = "gpt2"  # 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> &nbsp;|&nbsp; "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 &gt; 1 ⇒ A more probable; &lt; 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()