import gradio as gr
import difflib
import pandas as pd
from samples import BASE_INSTRUCTION, DRIFT_SCENARIOS
from drift_env import evaluate_instruction_drift
from drift_agent import generate_response as flan_generate
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
import torch
# =====================================================
# MODELS
# =====================================================
ALT_MODEL = "google/flan-t5-small"
alt_tokenizer = AutoTokenizer.from_pretrained(ALT_MODEL)
alt_model = AutoModelForSeq2SeqLM.from_pretrained(ALT_MODEL)
alt_model.eval()
def alt_generate(prompt: str) -> str:
inputs = alt_tokenizer(prompt, return_tensors="pt", truncation=True)
with torch.no_grad():
outputs = alt_model.generate(**inputs, max_new_tokens=160)
return alt_tokenizer.decode(outputs[0], skip_special_tokens=True)
MODELS = {
"FLAN-T5-Base": flan_generate,
"FLAN-T5-Small": alt_generate
}
SAMPLE_QUESTION = "Explain how the system should respond when API latency increases."
# =====================================================
# HELPERS
# =====================================================
def build_prompt(extra):
return BASE_INSTRUCTION + "\n\n" + extra + "\n\nQuestion: " + SAMPLE_QUESTION
def diff_highlight(a, b):
diff = difflib.ndiff(a.split(), b.split())
html = ""
for t in diff:
if t.startswith("+"):
html += f" {t}"
elif t.startswith("-"):
html += f" {t}"
else:
html += f"{t} "
return html
# =====================================================
# CORE
# =====================================================
def run_eval(model_name):
gen = MODELS[model_name]
steps = []
scores = []
baseline = gen(build_prompt(""))
for i, s in enumerate(DRIFT_SCENARIOS):
resp = gen(build_prompt(s["extra_context"]))
_, score = evaluate_instruction_drift(resp)
steps.append(i)
scores.append(score)
df = pd.DataFrame({"Step": steps, "Drift": scores})
final_resp = gen(build_prompt(DRIFT_SCENARIOS[-1]["extra_context"]))
verdict = "🟢 Stable" if scores[-1] == 0 else "🔴 Drift Detected"
return df, baseline, final_resp, diff_highlight(baseline, final_resp), verdict
# =====================================================
# COMPARISON MODE
# =====================================================
def compare_models(m1, m2):
df1, base1, drift1, diff1, v1 = run_eval(m1)
df2, base2, drift2, diff2, v2 = run_eval(m2)
return df1, df2, diff1, diff2, v1, v2
# =====================================================
# UI
# =====================================================
with gr.Blocks(css="""
body {background:#020617;color:#e5e7eb;font-family:Inter}
button {background:linear-gradient(135deg,#0ea5e9,#22c55e);color:white}
a {color:#38bdf8}
""") as app:
gr.Markdown("# 🚀 Instruction Drift Simulator")
gr.Markdown("### Compare how models lose instruction fidelity under pressure")
with gr.Tabs():
# ---------------- SINGLE MODEL ----------------
with gr.Tab("Single Model"):
model = gr.Dropdown(list(MODELS.keys()), value="FLAN-T5-Base")
run = gr.Button("Run Evaluation")
chart = gr.LinePlot(label="Drift Over Time")
baseline = gr.Textbox(label="Baseline", lines=4)
drifted = gr.Textbox(label="Drifted", lines=4)
diff = gr.HTML(label="Diff")
verdict = gr.Textbox(label="Verdict")
run.click(
run_eval,
model,
[chart, baseline, drifted, diff, verdict]
)
# ---------------- COMPARISON ----------------
with gr.Tab("Model Comparison"):
m1 = gr.Dropdown(list(MODELS.keys()), value="FLAN-T5-Base", label="Model A")
m2 = gr.Dropdown(list(MODELS.keys()), value="FLAN-T5-Small", label="Model B")
compare = gr.Button("Compare")
chart1 = gr.LinePlot(label="Model A Drift")
chart2 = gr.LinePlot(label="Model B Drift")
diff1 = gr.HTML(label="Diff A")
diff2 = gr.HTML(label="Diff B")
v1 = gr.Textbox(label="Verdict A")
v2 = gr.Textbox(label="Verdict B")
compare.click(
compare_models,
[m1, m2],
[chart1, chart2, diff1, diff2, v1, v2]
)
# FOOTER
gr.Markdown("""
---
### Built by Aditi Khare
🌐
AditiKhare.com — Enterprise AI Product Ecosystem | Decision Intelligence
""")
app.launch()