| import math |
| import gradio as gr |
| from transformers import AutoTokenizer, AutoModelForCausalLM |
| import torch |
|
|
| MODEL_NAME = "gpt2" |
| device = "cuda" if torch.cuda.is_available() else "cpu" |
|
|
| tok = AutoTokenizer.from_pretrained(MODEL_NAME) |
| model = AutoModelForCausalLM.from_pretrained(MODEL_NAME).to(device) |
| model.eval() |
|
|
| def next_seq_prob(context, candidate, assume_leading_space, show_topk): |
| if not context.strip(): |
| return "Please enter context.", "", "" |
| if not candidate.strip(): |
| return "Please enter a candidate next word/token.", "", "" |
|
|
| |
| cand_text = (" " + candidate) if assume_leading_space else candidate |
|
|
| with torch.no_grad(): |
| |
| ctx_ids = tok.encode(context, return_tensors="pt").to(device) |
| |
| cand_ids = tok.encode(cand_text, add_special_tokens=False) |
| if len(cand_ids) == 0: |
| return "Candidate tokenized to empty sequence (check spacing).", "", "" |
|
|
| total_logprob = 0.0 |
| step_details = [] |
|
|
| |
| input_ids = ctx_ids |
| for i, t_id in enumerate(cand_ids): |
| outputs = model(input_ids=input_ids) |
| logits = outputs.logits[:, -1, :] |
| logprobs = torch.log_softmax(logits, dim=-1) |
| token_logprob = logprobs[0, t_id].item() |
| total_logprob += token_logprob |
|
|
| |
| topk_vals, topk_idx = torch.topk(logprobs, k=min(show_topk, logprobs.shape[-1]), dim=-1) |
| topk_pairs = [ |
| (tok.decode(int(idx)), float(val)) |
| for idx, val in zip(topk_idx[0].tolist(), topk_vals[0].tolist()) |
| ] |
|
|
| step_details.append({ |
| "step": i+1, |
| "predicted_for": tok.decode([t_id]), |
| "logprob": token_logprob, |
| "prob": math.exp(token_logprob), |
| "topk": topk_pairs |
| }) |
|
|
| |
| input_ids = torch.cat([input_ids, torch.tensor([[t_id]], device=device)], dim=1) |
|
|
| seq_prob = math.exp(total_logprob) |
| |
| tokenized_candidate = [tok.decode([i]) for i in cand_ids] |
| summary = ( |
| f"Candidate tokenization: {tokenized_candidate}\n" |
| f"Total logprob (chain rule): {total_logprob:.6f}\n" |
| f"Sequence probability: {seq_prob:.6e}" |
| ) |
|
|
| |
| lines = [] |
| for d in step_details: |
| lines.append( |
| f"Step {d['step']}: token={repr(d['predicted_for'])} " |
| f"logprob={d['logprob']:.6f} prob={d['prob']:.6e}" |
| ) |
| if show_topk > 0: |
| tops = ", ".join([f"{repr(tok)}:{lp:.2f}" for tok, lp in d["topk"]]) |
| lines.append(f" top-{show_topk} logprobs: {tops}") |
| detail_text = "\n".join(lines) |
|
|
| return summary, detail_text, "" |
| |
| with gr.Blocks(title="Next-Token Probability (no fine-tuning)") as demo: |
| gr.Markdown("# Next-Token Probability\n" |
| "Compute the probability of a chosen next word/token sequence given a prior text segment.") |
| with gr.Row(): |
| context = gr.Textbox(label="Context (prompt)", lines=6) |
| with gr.Row(): |
| candidate = gr.Textbox(label="Candidate next word / token sequence") |
| with gr.Row(): |
| assume_space = gr.Checkbox(value=True, label="Assume leading space before candidate (useful for word starts in GPT-2 tokenization)") |
| topk = gr.Slider(0, 20, value=10, step=1, label="Show top-k alternatives (per step)") |
| btn = gr.Button("Compute probability") |
| summary = gr.Textbox(label="Summary", lines=4) |
| details = gr.Textbox(label="Step-by-step (per token)", lines=12) |
| _hidden = gr.Textbox(visible=False) |
|
|
| btn.click(fn=next_seq_prob, inputs=[context, candidate, assume_space, topk], outputs=[summary, details, _hidden]) |
|
|
| demo.launch() |