import os os.environ.setdefault("HF_HOME", "/data/hf_cache") import gradio as gr from huggingface_hub import InferenceClient from src.config import ( MODEL_MAP, MODEL_CHOICES, DEFAULT_MODEL, DEFAULT_MAX_TOKENS, DEFAULT_Q, DEFAULT_M, ENTAILMENT_USE_API, ) # UI labels for the entailment backend toggle _BACKEND_LOCAL = "Local (GPU/CPU)" _BACKEND_API = "Inference API" _DEFAULT_BACKEND = _BACKEND_API if ENTAILMENT_USE_API else _BACKEND_LOCAL from src.generation import query_model_with_logprobs from src.pipeline import run_pipeline from src.rendering import LEGEND_HTML, TOOLTIP_CSS, render_token_probs_html, color_code_answer_html # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- def _make_client(): return InferenceClient(token=os.environ.get("HF_TOKEN")) def _result_ready_for(semantic_result, answer): """A result counts as ready only if it was computed for the CURRENT answer. Guards against a stale background pipeline (from a previous question) whose result lands after the user has already asked something new.""" return bool(semantic_result) and semantic_result.get("answer") == answer def _render_view(view, token_logprobs, semantic_result, answer): if view == "Token probability": return render_token_probs_html(token_logprobs or []) # Semantic entropy if _result_ready_for(semantic_result, answer): if semantic_result.get("error"): return ( "
⚠️ Semantic entropy could not be computed for " "this answer. Switch the toggle to Token probability to see " "per-token confidence.
" ) return color_code_answer_html( answer, semantic_result.get("sp_to_entropy", {}), semantic_result.get("sentence_parts", []), ) return "Semantic entropy is still being computed. Please try again in a moment, or switch to Token probability.
" # --------------------------------------------------------------------------- # Status banners # --------------------------------------------------------------------------- def _status_banner(icon, text, bg, border, fg): # Set the color directly on the text span with !important so Gradio's theme # CSS can't override it to a light color. return ( f"Please enter a question.
", gr.update(visible=False), gr.update(visible=False), "", gr.update(visible=False), ) client = _make_client() model_name = MODEL_MAP.get(model_key, MODEL_MAP[DEFAULT_MODEL]) answer, token_logprobs = query_model_with_logprobs( question, model_name, int(max_tokens_val), client ) safe_answer = ( answer .replace("&", "&") .replace("<", "<") .replace(">", ">") .replace("\n", "⏳ Semantic entropy is still being computed. " "Please wait for the ✅ banner above, then click Show probabilities again " "— or switch the toggle to Token probability to view that now.
" ) def on_show_probs(token_logprobs, semantic_result, answer, view): # If the semantic view is requested but the background pipeline is still # running, show a wait message and reveal the toggle so the user can view # Token probability immediately. if view == "Semantic entropy" and not _result_ready_for(semantic_result, answer): return _STILL_COMPUTING_MSG, gr.update(visible=True) html = _render_view(view, token_logprobs, semantic_result, answer) return html, gr.update(visible=True) def on_toggle_change(view, token_logprobs, semantic_result, answer): return _render_view(view, token_logprobs, semantic_result, answer) # --------------------------------------------------------------------------- # UI # --------------------------------------------------------------------------- with gr.Blocks(title="Probabot") as demo: # ---- Persistent state ---- answer_state = gr.State("") token_logprobs_state = gr.State([]) model_name_state = gr.State(DEFAULT_MODEL) semantic_result_state = gr.State(None) # ---- Header ---- gr.Markdown("# 🤖🔍 Probabot") gr.Markdown( "Ask a question and see how confident the model is — both at the token level " "and at the factoid level via semantic entropy." ) gr.HTML(LEGEND_HTML) # ---- Input row ---- with gr.Row(): user_input = gr.Textbox( label="Your question", placeholder="What is the capital of France?", scale=4, ) send_btn = gr.Button("Send", scale=1, variant="primary") # ---- Config row ---- with gr.Row(): model_choice = gr.Dropdown( choices=MODEL_CHOICES, value=DEFAULT_MODEL, label="Model", ) max_tokens = gr.Number(value=DEFAULT_MAX_TOKENS, label="Max tokens", precision=0) q_num = gr.Number(value=DEFAULT_Q, label="Q (questions per factoid)", precision=0) m_num = gr.Number(value=DEFAULT_M, label="M (answers per question)", precision=0) with gr.Row(): backend_choice = gr.Radio( choices=[_BACKEND_LOCAL, _BACKEND_API], value=_DEFAULT_BACKEND, label="Entailment backend", info="Local = this Space's GPU/CPU · Inference API = HF's hosted compute", ) # ---- Answer display ---- answer_display = gr.HTML(label="Model response") with gr.Row(): pipeline_status = gr.HTML(visible=False) cancel_btn = gr.Button("Cancel computation", variant="stop", visible=False, scale=0) # ---- Hidden group: results / visualization ---- with gr.Group(visible=False) as results_group: show_probs_btn = gr.Button( "Show probabilities", variant="secondary" ) view_toggle = gr.Radio( choices=["Semantic entropy", "Token probability"], value="Semantic entropy", label="Confidence view", visible=False, ) viz_html = gr.HTML(visible=False) # ---- Events ---- # Send / submit send_outputs = [ answer_state, token_logprobs_state, model_name_state, semantic_result_state, answer_display, pipeline_status, results_group, viz_html, cancel_btn, ] pipeline_outputs = [semantic_result_state, pipeline_status, cancel_btn] send_click = send_btn.click( fn=handle_send, inputs=[user_input, model_choice, max_tokens], outputs=send_outputs, ) pipe_from_click = send_click.then( fn=run_pipeline_bg, inputs=[answer_state, model_name_state, max_tokens, q_num, m_num, backend_choice], outputs=pipeline_outputs, ) submit_event = user_input.submit( fn=handle_send, inputs=[user_input, model_choice, max_tokens], outputs=send_outputs, ) pipe_from_submit = submit_event.then( fn=run_pipeline_bg, inputs=[answer_state, model_name_state, max_tokens, q_num, m_num, backend_choice], outputs=pipeline_outputs, ) # Cancel any in-flight semantic-entropy pipeline when a NEW question is sent, # so an old (slow) run can't keep burning GPU or land a stale result. _running_pipelines = [pipe_from_click, pipe_from_submit] send_btn.click(fn=None, inputs=None, outputs=None, cancels=_running_pipelines) user_input.submit(fn=None, inputs=None, outputs=None, cancels=_running_pipelines) # Manual cancel button: abort the running pipeline and reset the UI. _STATUS_CANCELLED = _status_banner( "🛑", "Semantic entropy computation cancelled. Token probability is still " "available via the toggle.", bg="#f0f0f0", border="#888", fg="#222", ) def _on_cancel(): return ( gr.update(value=_STATUS_CANCELLED, visible=True), # pipeline_status gr.update(visible=False), # cancel_btn ) cancel_btn.click( fn=_on_cancel, inputs=None, outputs=[pipeline_status, cancel_btn], cancels=_running_pipelines, ) # Show probabilities show_probs_btn.click( fn=on_show_probs, inputs=[ token_logprobs_state, semantic_result_state, answer_state, view_toggle, ], outputs=[viz_html, view_toggle], ).then( fn=lambda: gr.update(visible=True), inputs=None, outputs=[viz_html], ) # Toggle view view_toggle.change( fn=on_toggle_change, inputs=[view_toggle, token_logprobs_state, semantic_result_state, answer_state], outputs=[viz_html], ) demo.launch(css=TOOLTIP_CSS)