WaterScopeAI / app.py
starfriend's picture
Update app.py
f592c34 verified
Raw
History Blame Contribute Delete
42.6 kB
import os
import re
import time
import threading
import torch
import pandas as pd
import gradio as gr
from difflib import SequenceMatcher
from peft import PeftModel
from huggingface_hub import snapshot_download
from transformers import AutoModelForCausalLM, AutoTokenizer
# --- 1. Configuration ---
BASE_MODEL_ID = "unsloth/Meta-Llama-3.1-8B-Instruct"
MAX_OPTIONS = 8
# Retained models after the literature-agent assessment:
# Base, DPO-only, DA-DPO, and TuluCore.
ADAPTER_REPO_ID = "starfriend/WaterScopeAI-Adapters"
ADAPTER_SUBFOLDERS = {
"da_it": "DA-IT",
"dpo": "DPO",
"da_dpo": "DA-DPO",
"tulucore": "DA/Tulucore",
}
MODEL_DISPLAY_NAMES = {
"base": "Base",
"da_it": "DA-IT",
"dpo": "DPO-only",
"da_dpo": "DA-DPO",
"tulucore": "TuluCore",
}
MCQA_MODEL_ORDER = ["da_it", "da_dpo"]
CHAT_MODEL_ORDER = ["base", "dpo", "da_dpo", "tulucore"]
AGENT_MODEL_ORDER = ["base", "dpo", "da_dpo", "tulucore"]
DEFAULT_CHAT_MODEL = "da_dpo"
DEFAULT_AGENT_MODEL = "base"
DATA_PATH = os.path.join("Testing MCQA data", "Decarbonization_MCQA.csv")
# --- 2. Load dataset ---
try:
MCQA_DF = pd.read_csv(DATA_PATH, encoding="utf-8")
except UnicodeDecodeError:
MCQA_DF = pd.read_csv(DATA_PATH, encoding="latin1")
# Ensure only Question + A-D columns
MCQA_DF = MCQA_DF[["Question", "A", "B", "C", "D"]]
# --- 3. Lazy Loading for Models ---
_model = None
_tokenizer = None
# Generation is serialized because PEFT adapters are model-global state.
_generation_lock = threading.Lock()
def download_required_adapters():
"""Download only the four required PEFT adapters from the Hub."""
hf_token = os.getenv("HF_TOKEN")
allow_patterns = []
for subfolder in ADAPTER_SUBFOLDERS.values():
allow_patterns.extend(
[
f"{subfolder}/adapter_config.json",
f"{subfolder}/adapter_model.safetensors",
]
)
local_repo_path = snapshot_download(
repo_id=ADAPTER_REPO_ID,
repo_type="model",
revision="main",
token=hf_token,
allow_patterns=allow_patterns,
)
adapter_paths = {}
for adapter_name, subfolder in ADAPTER_SUBFOLDERS.items():
adapter_path = os.path.join(local_repo_path, *subfolder.split("/"))
config_path = os.path.join(adapter_path, "adapter_config.json")
weights_path = os.path.join(adapter_path, "adapter_model.safetensors")
if not os.path.isfile(config_path):
raise FileNotFoundError(
f"{adapter_name} adapter_config.json was not found at "
f"{config_path}. Check the repository structure and HF_TOKEN."
)
if not os.path.isfile(weights_path):
raise FileNotFoundError(
f"{adapter_name} adapter weights were not found at "
f"{weights_path}."
)
adapter_paths[adapter_name] = adapter_path
return adapter_paths
def load_model_and_tokenizer():
"""Load the base model and the three retained adapters."""
global _model, _tokenizer
if _model is not None and _tokenizer is not None:
return _model, _tokenizer
print("Initializing WaterScope-AI models...")
if not torch.cuda.is_available():
raise RuntimeError("No CUDA GPU detected. This Space requires a GPU.")
_tokenizer = AutoTokenizer.from_pretrained(
BASE_MODEL_ID,
use_fast=True,
)
if _tokenizer.pad_token_id is None:
_tokenizer.pad_token = _tokenizer.eos_token
base_model = AutoModelForCausalLM.from_pretrained(
BASE_MODEL_ID,
torch_dtype=torch.bfloat16,
device_map="auto",
low_cpu_mem_usage=True,
)
base_model.eval()
print("Base model loaded.")
adapter_paths = download_required_adapters()
first_adapter = "da_it"
_model = PeftModel.from_pretrained(
base_model,
adapter_paths[first_adapter],
adapter_name=first_adapter,
is_trainable=False,
)
print(f"DA-IT adapter loaded from: {adapter_paths[first_adapter]}")
for adapter_name in ("dpo", "da_dpo", "tulucore"):
_model.load_adapter(
adapter_paths[adapter_name],
adapter_name=adapter_name,
is_trainable=False,
)
print(
f"{MODEL_DISPLAY_NAMES[adapter_name]} adapter loaded from: "
f"{adapter_paths[adapter_name]}"
)
_model.set_adapter(DEFAULT_CHAT_MODEL)
_model.eval()
print("Available adapters:", list(_model.peft_config.keys()))
print("MCQA models:", [MODEL_DISPLAY_NAMES[name] for name in MCQA_MODEL_ORDER])
print("Chat models:", [MODEL_DISPLAY_NAMES[name] for name in CHAT_MODEL_ORDER])
print("Agent models:", [MODEL_DISPLAY_NAMES[name] for name in AGENT_MODEL_ORDER])
return _model, _tokenizer
def activate_model(model_name):
"""Activate a retained adapter or temporarily disable adapters for Base."""
all_models = ["base"] + list(ADAPTER_SUBFOLDERS.keys())
if model_name not in all_models:
raise ValueError(
f"Unknown model '{model_name}'. Choose from: {', '.join(all_models)}"
)
if model_name == "base":
return _model.disable_adapter()
_model.set_adapter(model_name)
return None
# --- 4. Utility Functions ---
def extract_letter(raw_answer: str) -> str:
"""Extract predicted option letter from model output"""
# Priority 1: Look for explicit phrases like "answer is B"
match = re.search(r"(?:answer|option) is\s+([A-H])", raw_answer, re.IGNORECASE)
if match:
return match.group(1).upper()
# Priority 2: Look for formats like "B." or "B)" at the start
match = re.search(r"^\s*([A-H])[\.\):]", raw_answer)
if match:
return match.group(1).upper()
# Priority 3: Look for the first standalone letter in the text
match = re.search(r"\b([A-H])\b", raw_answer)
if match:
return match.group(1).upper()
return "N/A"
def clean_repetitions(text: str) -> str:
lines = [l.strip() for l in text.strip().splitlines() if l.strip()]
if not lines:
return ""
# split into words (keep punctuation as part of word)
def tokenize(line):
return re.findall(r"\S+", line)
result = tokenize(lines[0])
for line in lines[1:]:
tokens = tokenize(line)
# find overlap
i = 0
while i < len(result) and i < len(tokens) and result[i].rstrip(".,!?") == tokens[i].rstrip(".,!?"):
i += 1
# append only the non-overlapping part
result.extend(tokens[i:])
return " ".join(result)
# Global variable to track cancellation
cancellation_requested = False
def run_mcqa_comparison(
question,
opt_a,
opt_b,
opt_c,
opt_d,
opt_e,
opt_f,
opt_g,
opt_h,
generate_explanation,
):
"""Run the original MCQA comparison with DA-IT and DA-DPO."""
global _model, _tokenizer, cancellation_requested
cancellation_requested = False
if _model is None or _tokenizer is None:
gr.Info("Initializing models for the first time...")
load_model_and_tokenizer()
options = [opt_a, opt_b, opt_c, opt_d, opt_e, opt_f, opt_g, opt_h]
active_options = [opt for opt in options if opt and opt.strip()]
if not question or len(active_options) < 2:
yield (
"Error",
"Please enter a question and at least two options.",
"Error",
"Please enter a question and at least two options.",
)
return
option_labels = "ABCDEFGHIJKLMNOPQRSTUVWXYZ"
option_text = "\n".join(
f"{option_labels[i]}. {value}"
for i, value in enumerate(active_options)
)
if generate_explanation:
instruction = (
"Provide the letter first, followed by a concise expert explanation "
"in the form: 'The answer is [LETTER]. Because ...'"
)
max_tokens = 220
else:
instruction = "Return only the letter of the best answer."
max_tokens = 30
messages = [
{
"role": "system",
"content": (
"You are an expert in water and wastewater treatment, "
"decarbonization, emissions, resource recovery, and "
"environmental sustainability. Answer the multiple-choice "
f"question accurately. {instruction}"
),
},
{
"role": "user",
"content": f"Question: {question}\nCandidate options:\n{option_text}",
},
]
chat_input = _tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
)
inputs = safe_tokenize(
chat_input,
_tokenizer,
_model,
max_new_tokens=max_tokens,
)
def generate_for(model_name):
if cancellation_requested:
raise gr.Error("Processing cancelled by user")
generation_kwargs = dict(
**inputs,
max_new_tokens=max_tokens,
eos_token_id=_tokenizer.eos_token_id,
pad_token_id=_tokenizer.pad_token_id,
do_sample=False,
use_cache=True,
)
if model_name == "da_dpo":
generation_kwargs.update(
repetition_penalty=1.10,
no_repeat_ngram_size=4,
)
with torch.inference_mode():
_model.set_adapter(model_name)
outputs = _model.generate(**generation_kwargs)
generated_ids = outputs[0][inputs["input_ids"].shape[1]:]
response = _tokenizer.decode(
generated_ids,
skip_special_tokens=True,
clean_up_tokenization_spaces=False,
).strip()
return clean_repetitions(response)
try:
with _generation_lock:
yield "", "Running DA-IT...", "", ""
da_it_raw = generate_for("da_it")
da_it_letter = extract_letter(da_it_raw)
yield da_it_letter, da_it_raw, "", "Running DA-DPO..."
da_dpo_raw = generate_for("da_dpo")
da_dpo_letter = extract_letter(da_dpo_raw)
yield da_it_letter, da_it_raw, da_dpo_letter, da_dpo_raw
except gr.Error as exc:
if "cancelled" in str(exc).lower():
gr.Info("Processing cancelled by user")
return
raise
# Function to handle cancellation
def cancel_processing():
global cancellation_requested
cancellation_requested = True
return "Cancellation requested"
# Safe tokenization wrapper
def safe_tokenize(chat_input, _tokenizer, _model, max_new_tokens=1600):
# 1. Validate input type
if not isinstance(chat_input, str) or len(chat_input.strip()) == 0:
raise ValueError("chat_input must be a non-empty string")
# 2. Sanitize weird characters (e.g., emojis, zero-width spaces)
clean_input = re.sub(r"[^\x00-\x7F]+", " ", chat_input)
# 3. Tokenize with truncation to avoid position limit issues
max_input_tokens = (
_model.config.max_position_embeddings
- max_new_tokens
- 100
)
tokens = _tokenizer(
clean_input,
return_tensors="pt",
truncation=True,
max_length=max_input_tokens
)
# 4. Validate token IDs
vocab_size = _model.get_input_embeddings().weight.shape[0]
max_id = tokens["input_ids"].max().item()
min_id = tokens["input_ids"].min().item()
print(f"[DEBUG] chat_input: {repr(chat_input)}")
print(f"[DEBUG] sanitized_input: {repr(clean_input)}")
print(f"[DEBUG] token IDs min: {min_id}, max: {max_id}, vocab size: {vocab_size}")
if max_id >= vocab_size or min_id < 0:
raise ValueError(f"Token IDs out of range: min {min_id}, max {max_id}, vocab size {vocab_size}")
# 5. Move tokens to model device
tokens = {k: v.to(_model.device) for k, v in tokens.items() if isinstance(v, torch.Tensor)}
return tokens
# General chat function
def chat_with_model(
message,
history,
selected_model=DEFAULT_CHAT_MODEL,
max_new_tokens=900,
):
"""General conversational QA while preserving multi-turn history."""
global _model, _tokenizer
if _model is None or _tokenizer is None:
gr.Info("Initializing models for the first time...")
load_model_and_tokenizer()
if not isinstance(message, str) or not message.strip():
return "Please provide a non-empty message."
selected_model = selected_model or DEFAULT_CHAT_MODEL
if selected_model not in CHAT_MODEL_ORDER:
raise gr.Error(
"Chat supports Base, DPO-only, DA-DPO, and TuluCore."
)
max_new_tokens = max(64, min(int(max_new_tokens), 1600))
system_prompt = (
"""
You are an expert AI assistant in water and wastewater engineering.
When answering a question:
- First, write down all relevant facts or values.
- Next, identify which one is correct based on those facts for factual
comparisons or multiple options.
- Next, provide a clear description for conceptual definition questions.
- Finally, clearly state your conclusion in this format:
[Main answer]. [one or two sentences explaining the reasoning].
[appropriate values, equations to support the reasoning].
"""
)
messages = [{"role": "system", "content": system_prompt}]
messages.extend(normalize_chat_history(history))
messages.append({"role": "user", "content": message.strip()})
return generate_chat_response(
messages=messages,
selected_model=selected_model,
max_new_tokens=max_new_tokens,
)
AGENT_APPLICATIONS = {
"general_qa": {
"label": "General environmental QA",
"instruction": (
"Answer the technical question directly. Distinguish facts, "
"assumptions, and recommendations."
),
},
"document_qa": {
"label": "Document-grounded QA",
"instruction": (
"Use the supplied context as the only source for document-specific "
"claims. State clearly when the context does not support an answer."
),
},
"paper_synthesis": {
"label": "Multi-paper synthesis",
"instruction": (
"Compare the supplied papers with explicit source attribution. "
"Preserve key methods and quantitative findings and do not mix "
"evidence among papers."
),
},
"method_recommendation": {
"label": "Method recommendation",
"instruction": (
"Recommend a method using explicit evidence, implementation "
"constraints, limitations, and uncertainty. Do not claim superiority "
"unless the supplied evidence directly supports it."
),
},
"research_gaps": {
"label": "Research-gap identification",
"instruction": (
"Identify evidence-grounded research gaps. Link every gap to a "
"specific limitation or unresolved issue in the supplied context."
),
},
"claim_check": {
"label": "Unsupported-claim check",
"instruction": (
"Evaluate whether the requested claim is supported. Refuse to invent "
"a paper, method, value, causal relationship, or percentage."
),
},
}
def normalize_chat_history(history):
"""Convert Gradio messages or legacy tuples into chat-template messages."""
normalized = []
for item in history or []:
if isinstance(item, dict):
role = item.get("role")
content = item.get("content")
if isinstance(content, list):
parts = []
for part in content:
if isinstance(part, dict) and part.get("type") == "text":
if part.get("text"):
parts.append(str(part["text"]))
elif isinstance(part, str):
parts.append(part)
content = "\n".join(parts)
if role in {"user", "assistant"} and content:
normalized.append({"role": role, "content": str(content)})
elif isinstance(item, (list, tuple)) and len(item) >= 2:
user_message, assistant_message = item[0], item[1]
if user_message:
normalized.append(
{"role": "user", "content": str(user_message)}
)
if assistant_message:
normalized.append(
{"role": "assistant", "content": str(assistant_message)}
)
return normalized
def generate_chat_response(messages, selected_model, max_new_tokens):
"""Shared generation path for Chat and Agent."""
chat_input = _tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
)
inputs = safe_tokenize(
chat_input,
_tokenizer,
_model,
max_new_tokens=max_new_tokens,
)
generation_kwargs = dict(
**inputs,
max_new_tokens=max_new_tokens,
do_sample=False,
eos_token_id=_tokenizer.eos_token_id,
pad_token_id=_tokenizer.pad_token_id,
use_cache=True,
)
# The benchmark showed that DA-DPO and TuluCore benefit from modest
# repetition control for longer, open-ended generations.
if selected_model in {"da_dpo", "tulucore"}:
generation_kwargs.update(
repetition_penalty=1.10,
no_repeat_ngram_size=4,
)
with _generation_lock:
try:
if selected_model == "base":
with _model.disable_adapter():
outputs = _model.generate(**generation_kwargs)
else:
_model.set_adapter(selected_model)
outputs = _model.generate(**generation_kwargs)
except Exception as exc:
print(
f"[ERROR] Generation failed for {selected_model}: {exc}",
flush=True,
)
raise gr.Error(
f"{MODEL_DISPLAY_NAMES.get(selected_model, selected_model)} "
f"could not generate a response: {exc}"
)
input_length = inputs["input_ids"].shape[1]
generated_ids = outputs[0][input_length:]
return _tokenizer.decode(
generated_ids,
skip_special_tokens=True,
clean_up_tokenization_spaces=False,
).strip()
def agent_run(
selected_model,
application,
context,
question,
max_new_tokens=900,
):
"""Run a selectable WaterScope agent application."""
global _model, _tokenizer
if _model is None or _tokenizer is None:
gr.Info("Initializing models for the first time...")
load_model_and_tokenizer()
if selected_model not in AGENT_MODEL_ORDER:
raise gr.Error(
"Agent applications support Base, DPO-only, DA-DPO, and TuluCore."
)
if application not in AGENT_APPLICATIONS:
raise gr.Error("Select a valid agent application.")
if not isinstance(question, str) or not question.strip():
raise gr.Error("Enter a question or task.")
max_new_tokens = max(64, min(int(max_new_tokens), 1600))
application_config = AGENT_APPLICATIONS[application]
system_prompt = f"""
You are WaterScope-AI, an expert scientific agent for water, wastewater,
environmental engineering, and sustainability.
Application: {application_config["label"]}
Task behavior: {application_config["instruction"]}
Core requirements:
- Answer the user's actual task directly.
- Preserve important quantitative values and units.
- Attribute document-specific evidence clearly.
- Separate evidence from inference and recommendation.
- Never invent papers, citations, numerical values, standards, or findings.
- When evidence is insufficient, say exactly what cannot be concluded.
- Avoid repeated sentences or sections and stop when complete.
""".strip()
context = (context or "").strip()
user_content = question.strip()
if context:
user_content = (
"SUPPLIED CONTEXT\n"
"================\n"
f"{context}\n\n"
"USER TASK\n"
"=========\n"
f"{question.strip()}"
)
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_content},
]
started = time.time()
response = generate_chat_response(
messages=messages,
selected_model=selected_model,
max_new_tokens=max_new_tokens,
)
elapsed = time.time() - started
diagnostics = (
f"Model: {MODEL_DISPLAY_NAMES[selected_model]}\n"
f"Application: {application_config['label']}\n"
f"Context characters: {len(context):,}\n"
f"Response characters: {len(response):,}\n"
f"Generation time: {elapsed:.1f} s"
)
return response, diagnostics
# Backward-compatible API endpoint used by prior local scripts.
def agent_chat_with_model(
message,
chat_history=None,
max_new_tokens=600,
selected_model=DEFAULT_AGENT_MODEL,
):
messages = [
{
"role": "system",
"content": (
"You are WaterScope-AI, an expert assistant in water and "
"wastewater engineering. Answer accurately, remain grounded, "
"and do not invent evidence."
),
}
]
messages.extend(normalize_chat_history(chat_history))
messages.append({"role": "user", "content": str(message)})
return generate_chat_response(
messages=messages,
selected_model=selected_model,
max_new_tokens=max(64, min(int(max_new_tokens), 1200)),
)
# Custom CSS for website-like appearance with lighter blue header
custom_css = """
.gradio-container {
max-width: 1200px !important;
margin: 0 auto !important;
font-family: 'Segoe UI', Tahoma, Geneva, Verdana, sans-serif !important;
}
.header {
text-align: center;
padding: 20px;
background: linear-gradient(135deg, #6eb1ff 0%, #88d3fe 100%);
color: white;
border-radius: 8px;
margin-bottom: 20px;
}
.header h1 {
margin: 0;
font-size: 2.5em;
font-weight: 600;
}
.header p {
margin: 10px 0 0;
font-size: 1.2em;
opacity: 0.9;
}
.section {
background: white;
padding: 20px;
border-radius: 8px;
box-shadow: 0 4px 6px rgba(0, 0, 0, 0.1);
margin-bottom: 20px;
}
.nav-bar {
margin-bottom: 20px;
display: flex;
justify-content: center;
gap: 10px;
}
.footer {
text-align: center;
padding: 15px;
margin-top: 30px;
color: #666;
font-size: 0.9em;
border-top: 1px solid #eee;
}
.dataframe-container {
margin-top: 20px;
}
.model-output {
background: #f8f9fa;
padding: 15px;
border-radius: 8px;
border-left: 4px solid #6eb1ff;
}
.model-output h4 {
margin-top: 0;
color: #6eb1ff;
}
.option-controls {
margin-top: 15px;
display: flex;
gap: 10px;
}
.cancel-btn {
background: #f39c12 !important;
color: white !important;
}
.cancel-btn:hover {
background: #e67e22 !important;
}
.status-message {
padding: 10px;
border-radius: 4px;
margin: 10px 0;
}
.status-info {
background-color: #e3f2fd;
border-left: 4px solid #2196f3;
}
.status-warning {
background-color: #fff3e0;
border-left: 4px solid #ff9800;
}
.status-error {
background-color: #ffebee;
border-left: 4px solid #f44336;
}
.status-success {
background-color: #e8f5e9;
border-left: 4px solid #4caf50;
}
/* Chat specific styles */
.chat-container {
display: flex;
flex-direction: column;
height: 500px;
}
.chat-messages {
flex: 1;
overflow-y: auto;
padding: 15px;
background: var(--light);
border-radius: 6px;
margin-bottom: 15px;
display: flex;
flex-direction: column;
gap: 15px;
}
.message {
display: flex;
max-width: 80%;
}
.user-message {
align-self: flex-end;
}
.bot-message {
align-self: flex-start;
}
.message-content {
padding: 12px 16px;
border-radius: 18px;
line-height: 1.4;
}
.user-message .message-content {
background: var(--accent);
color: white;
border-bottom-right-radius: 4px;
}
.bot-message .message-content {
background: var(--light-gray);
color: var(--dark);
border-bottom-left-radius: 4px;
}
.chat-input-container {
display: flex;
gap: 10px;
}
.chat-input-container textarea {
flex: 1;
padding: 12px;
border: 1px solid var(--border);
border-radius: 6px;
resize: vertical;
font-family: inherit;
font-size: 14px;
}
"""
# --- 5. Gradio UI ---
with gr.Blocks(
title="WaterScope-AI",
fill_width=True,
) as demo:
# Custom Header with lighter blue
with gr.Column(elem_classes="header"):
gr.Markdown("WaterScope-AI")
gr.Markdown("Domain-Specific Language Models and Agent Applications for Water Sustainability")
# Navigation Bar
with gr.Row(elem_classes="nav-bar"):
gr.Button("Home", variant="secondary", size="sm")
gr.Button("About", variant="secondary", size="sm")
gr.Button("Documentation", variant="secondary", size="sm")
gr.Button("Contact", variant="secondary", size="sm")
# Create tabs for different functionalities
with gr.Tabs():
# MCQA Demo Tab
with gr.TabItem("MCQA Demo"):
# Status message area
status_message = gr.HTML("", elem_classes="status-message")
# Main content in a styled section
with gr.Column(elem_classes="section"):
# State for tracking number of visible options
num_options_state = gr.State(4)
# Top row with input and output panels
with gr.Row():
# Left panel with inputs
with gr.Column(scale=1):
with gr.Group():
question_box = gr.Textbox(label="Question", lines=2, interactive=True)
gr.Markdown("#### Options")
# Create option boxes using a list (like in the working version)
option_boxes = []
for i in range(MAX_OPTIONS):
option_boxes.append(gr.Textbox(
label=f"Option {chr(ord('A') + i)}",
visible=(i < 4),
interactive=True
))
with gr.Row():
add_option_btn = gr.Button("Add Option")
clear_btn = gr.Button("Clear")
explanation_checkbox = gr.Checkbox(label="Generate Explanation", value=False)
with gr.Row():
run_btn = gr.Button("Run Comparison", variant="primary")
cancel_btn = gr.Button("Cancel", variant="stop", visible=False, elem_classes="cancel-btn")
# Right panel with outputs
with gr.Column(scale=1):
gr.Markdown("### Model Outputs")
with gr.Row():
with gr.Column(elem_classes="model-output"):
gr.Markdown("#### DA-IT Model")
da_it_letter_box = gr.Textbox(
label="Predicted Letter",
interactive=False,
)
da_it_raw_box = gr.Textbox(
label="Raw Answer",
lines=3,
interactive=False,
)
with gr.Column(elem_classes="model-output"):
gr.Markdown("#### DA-DPO Model")
da_dpo_letter_box = gr.Textbox(
label="Predicted Letter",
interactive=False,
)
da_dpo_raw_box = gr.Textbox(
label="Raw Answer",
lines=3,
interactive=False,
)
# Table section with custom styling
with gr.Column(elem_classes="section dataframe-container"):
gr.Markdown("### Browse 777 MCQAs (click a row to autofill)")
mcqa_table = gr.Dataframe(
value=MCQA_DF.values.tolist(),
headers=["Question", "A", "B", "C", "D"],
datatype=["str"]*5,
interactive=False,
wrap=True,
max_height=400
)
# Chat Tab
with gr.TabItem("Chat"):
with gr.Column(elem_classes="section"):
gr.Markdown(
"### General Chat\n"
"Use any retained model for conversational water and "
"environmental engineering questions."
)
with gr.Row():
chat_model = gr.Dropdown(
choices=[
(MODEL_DISPLAY_NAMES[name], name)
for name in CHAT_MODEL_ORDER
],
value=DEFAULT_CHAT_MODEL,
label="Chat model",
)
chat_max_tokens = gr.Slider(
minimum=128,
maximum=1600,
value=900,
step=64,
label="Maximum new tokens",
)
with gr.Row():
system_status = gr.Textbox(
value="Ready",
label="System status",
interactive=False,
)
api_status = gr.Textbox(
value="Ready",
label="Generation status",
interactive=False,
)
chatbot = gr.Chatbot(
label="Conversation",
elem_classes="chat-messages",
height=430,
)
with gr.Row():
msg = gr.Textbox(
label="Your message",
placeholder="Ask a water or environmental engineering question...",
lines=3,
scale=5,
)
send_btn = gr.Button(
"Send",
variant="primary",
scale=1,
)
clear_chat = gr.Button("Clear conversation")
# Agent Applications Tab
with gr.TabItem("Agent Applications"):
with gr.Column(elem_classes="section"):
gr.Markdown(
"### Test WaterScope-AI across agent applications\n"
"Select a model and application, provide optional source "
"material, and submit a task."
)
with gr.Row():
agent_model = gr.Dropdown(
choices=[
(MODEL_DISPLAY_NAMES[name], name)
for name in AGENT_MODEL_ORDER
],
value=DEFAULT_AGENT_MODEL,
label="Agent model",
)
agent_application = gr.Dropdown(
choices=[
(config["label"], key)
for key, config in AGENT_APPLICATIONS.items()
],
value="document_qa",
label="Application",
)
agent_max_tokens = gr.Slider(
minimum=128,
maximum=1600,
value=900,
step=64,
label="Maximum new tokens",
)
agent_context = gr.Textbox(
label="Optional context or documents",
placeholder=(
"Paste one or more paper summaries, regulatory passages, "
"process data, or other source material here."
),
lines=14,
)
agent_question = gr.Textbox(
label="Question or task",
placeholder=(
"Example: Compare the treatment methods and recommend "
"the most defensible option based only on the context."
),
lines=4,
)
with gr.Row():
agent_run_button = gr.Button(
"Run Agent",
variant="primary",
)
agent_clear_button = gr.Button("Clear")
agent_output = gr.Textbox(
label="Agent response",
lines=18,
interactive=False,
)
agent_diagnostics = gr.Textbox(
label="Run diagnostics",
lines=5,
interactive=False,
)
# Backward-compatible hidden endpoint: /agent_chat
agent_api_message = gr.Textbox(visible=False)
agent_api_history = gr.JSON(value=[], visible=False)
agent_api_max_tokens = gr.Number(value=600, precision=0, visible=False)
agent_api_model = gr.Dropdown(
choices=AGENT_MODEL_ORDER,
value=DEFAULT_AGENT_MODEL,
visible=False,
)
agent_api_output = gr.Textbox(visible=False)
agent_api_trigger = gr.Button("Agent Chat API", visible=False)
agent_api_trigger.click(
fn=agent_chat_with_model,
inputs=[
agent_api_message,
agent_api_history,
agent_api_max_tokens,
agent_api_model,
],
outputs=agent_api_output,
api_name="agent_chat",
)
# Structured agent endpoint: /agent_run
agent_run_button.click(
fn=agent_run,
inputs=[
agent_model,
agent_application,
agent_context,
agent_question,
agent_max_tokens,
],
outputs=[agent_output, agent_diagnostics],
api_name="agent_run",
)
agent_clear_button.click(
fn=lambda: ("", "", "", ""),
inputs=None,
outputs=[
agent_context,
agent_question,
agent_output,
agent_diagnostics,
],
queue=False,
)
# Footer
with gr.Column(elem_classes="footer"):
gr.Markdown("© 2025 WaterScope-AI | Built with Gradio")
# Function to add more options
def add_option(current_count):
if current_count < MAX_OPTIONS:
current_count += 1
updates = [gr.update(visible=i < current_count) for i in range(MAX_OPTIONS)]
return current_count, *updates
# Function to clear all inputs and outputs (from working version)
def clear_all():
"""Clear all MCQA inputs, outputs, and option visibility states."""
option_visibility_updates = [
gr.update(visible=(i < 4), value="")
for i in range(MAX_OPTIONS)
]
return (
4, # Reset the visible-option count
"", # Clear the question
*[""] * MAX_OPTIONS, # Clear option values
False, # Uncheck explanation
"", "", "", "", # Clear DA-IT and DA-DPO outputs
*option_visibility_updates,
)
# Fixed function to load row data
def load_row(evt: gr.SelectData):
"""Load a selected row from the dataframe into the input fields"""
if evt.index[0] >= len(MCQA_DF):
return ["", ""] + [""] * MAX_OPTIONS
row = MCQA_DF.iloc[evt.index[0]]
# Return question and first 4 options (A-D), and empty for the rest
return_values = [
row["Question"] if pd.notna(row["Question"]) else "",
row["A"] if pd.notna(row["A"]) else "",
row["B"] if pd.notna(row["B"]) else "",
row["C"] if pd.notna(row["C"]) else "",
row["D"] if pd.notna(row["D"]) else ""
]
# Add empty values for any additional options
return_values += [""] * (MAX_OPTIONS - 4)
return return_values
# Function to toggle cancel button visibility
def toggle_cancel_button():
return gr.update(visible=True)
# Function to hide cancel button
def hide_cancel_button():
return gr.update(visible=False)
# Function to update status message
def update_status(message, type="info"):
if type == "info":
cls = "status-info"
elif type == "warning":
cls = "status-warning"
elif type == "error":
cls = "status-error"
elif type == "success":
cls = "status-success"
else:
cls = "status-info"
return f'<div class="status-message {cls}">{message}</div>'
# Connect the table selection event
mcqa_table.select(
fn=load_row,
inputs=None,
outputs=[question_box, *option_boxes]
)
# Connect the add option button
add_option_btn.click(
fn=add_option,
inputs=[num_options_state],
outputs=[num_options_state, *option_boxes]
)
# Define the MCQA components cleared by the Clear button.
outputs_to_clear = [
num_options_state,
question_box,
*option_boxes,
explanation_checkbox,
da_it_letter_box,
da_it_raw_box,
da_dpo_letter_box,
da_dpo_raw_box,
*option_boxes,
]
# Connect the clear button (from working version)
clear_btn.click(
fn=clear_all,
inputs=None,
outputs=outputs_to_clear,
queue=False
).then(
fn=lambda: update_status("Form cleared", "info"),
inputs=None,
outputs=[status_message],
queue=False
)
# Connect the run button
run_event = run_btn.click(
fn=lambda: update_status("Initializing processing...", "info"),
inputs=None,
outputs=[status_message],
queue=False
).then(
fn=toggle_cancel_button,
inputs=None,
outputs=[cancel_btn],
queue=False
).then(
fn=run_mcqa_comparison,
inputs=[question_box, *option_boxes, explanation_checkbox],
outputs=[
da_it_letter_box,
da_it_raw_box,
da_dpo_letter_box,
da_dpo_raw_box,
]
).then(
fn=lambda: update_status("Processing completed successfully", "success"),
inputs=None,
outputs=[status_message],
queue=False
).then(
fn=hide_cancel_button,
inputs=None,
outputs=[cancel_btn],
queue=False
)
# Connect the cancel button
cancel_btn.click(
fn=cancel_processing,
inputs=None,
outputs=None,
queue=False
).then(
fn=lambda: update_status("Processing cancelled by user", "warning"),
inputs=None,
outputs=[status_message],
queue=False
).then(
fn=hide_cancel_button,
inputs=None,
outputs=[cancel_btn],
queue=False
)
# Chat functionality
def respond(
message,
chat_history,
selected_model,
max_new_tokens,
):
chat_history = list(chat_history or [])
if not isinstance(message, str) or not message.strip():
return "", chat_history, "Ready", "No message submitted"
message = message.strip()
try:
bot_message = chat_with_model(
message=message,
history=chat_history,
selected_model=selected_model,
max_new_tokens=max_new_tokens,
)
chat_history.extend(
[
{"role": "user", "content": message},
{"role": "assistant", "content": bot_message},
]
)
return "", chat_history, "Ready", (
f"Response generated with "
f"{MODEL_DISPLAY_NAMES[selected_model]}"
)
except Exception as exc:
error_message = f"Sorry, I encountered an error: {exc}"
chat_history.extend(
[
{"role": "user", "content": message},
{"role": "assistant", "content": error_message},
]
)
return "", chat_history, "Error", str(exc)
# Connect the chat send button
chat_inputs = [msg, chatbot, chat_model, chat_max_tokens]
chat_outputs = [msg, chatbot, system_status, api_status]
msg.submit(
respond,
chat_inputs,
chat_outputs,
api_name="respond",
)
send_btn.click(
respond,
chat_inputs,
chat_outputs,
api_name="respond_button",
)
# Connect the clear chat button
def clear_chat_func():
system_status.value = "Ready"
api_status.value = "Ready"
return []
clear_chat.click(clear_chat_func, None, chatbot, queue=False)
# Pre-load the model when the app starts
print("Pre-loading models...")
load_model_and_tokenizer()
print("Models loaded successfully!")
demo.queue(default_concurrency_limit=1).launch(
debug=True,
show_error=True,
theme=gr.themes.Glass(primary_hue="blue"),
css=custom_css,
)