one-for-all / app.py
frankyy03's picture
Custom gr.Server frontend + 6-teacher model (Nemotron), token streaming, arena
9409aaa verified
Raw
History Blame Contribute Delete
19.1 kB
"""
space/app.py — One for All HuggingFace ZeroGPU Gradio Space.
Run locally (from space/ dir, with a local viz_data.json):
cd space && VIZ_DATA_PATH=/path/to/viz_data.json python app.py
"""
from __future__ import annotations
import os
import html as _html_stdlib
import gradio as gr
import spaces
import _fig
import _glb
import _html
# ── Startup: shared runtime (viz + UMAP + student) lives in _boot ─────────
import _boot
RT = _boot.load_runtime()
HF_TOKEN = RT.hf_token
VIZ = RT.viz
REDUCER = RT.reducer
COORDS3D = RT.coords3d
BACKEND = RT.backend
_MODEL_READY = RT.model_ready
_INIT_GLB = _glb.build_glb(VIZ, COORDS3D, [])
print(f"[ofa-space] GLB path: {_INIT_GLB}")
if _INIT_GLB:
print(f"[ofa-space] GLB exists: {os.path.exists(_INIT_GLB)}, size: {os.path.getsize(_INIT_GLB)} bytes")
# Camera: pull back far enough to see all points at startup
if COORDS3D is not None:
import numpy as _np
_span = float(_np.linalg.norm(COORDS3D.max(axis=0) - COORDS3D.min(axis=0)))
_CAM = (45, 30, _span * 1.8)
else:
_CAM = (45, 30, 10)
def _response_html(text: str, title: str = "MODEL RESPONSE", accent: str = "#8b949e") -> str:
safe = _html_stdlib.escape(text).replace("\n", "<br>")
return (
'<div style="background:#0d1117;border:1px solid #30363d;border-radius:6px;'
'padding:14px;margin-top:8px;">'
f'<div style="font-size:10px;color:{accent};font-family:monospace;'
f'margin-bottom:8px;letter-spacing:0.04em;">{title}</div>'
f'<div style="font-size:13px;color:#e6edf3;line-height:1.65;">{safe}</div>'
"</div>"
)
# ── Backend dispatch: same handler code drives torch and llama.cpp ────────
def _to_device():
_boot.to_device(RT)
def _stream(text: str):
return _boot.stream(RT, text)
def _stream_pair(text: str):
return _boot.stream_pair(RT, text)
def _final_probe(text: str):
return _boot.final_probe(RT, text)
# ── ZeroGPU probe handler (streaming generator) ───────────────────────────
@spaces.GPU
def probe_fn(text: str, probe_points: list):
if not text.strip():
yield gr.skip(), probe_points, "", "", ""
return
if not _MODEL_READY or REDUCER is None:
msg = _html.gate_html([0.2] * 5, VIZ["teacher_names"] or ["—"] * 5)
yield _glb.build_glb(VIZ, COORDS3D, probe_points), probe_points, "", msg, ""
return
_to_device()
names = VIZ["teacher_names"]
# Stream tokens + live gate bars; skip the 3D plot until the end.
partial = ""
for partial, gates in _stream(text):
yield (
gr.skip(), probe_points,
_response_html(partial),
_html.gate_html(gates, names, ranked=False),
gr.skip(),
)
# Final pass: pooled probe point in soul space + dominant-teacher badge.
new_pt, gate_weights = _final_probe(text)
updated = probe_points + [new_pt]
yield (
_glb.build_glb(VIZ, COORDS3D, updated), updated,
_response_html(partial),
_html.gate_html(gate_weights, names),
_html.task_html(gate_weights, names),
)
# ── ZeroGPU arena handler: base (LoRA off) vs deku, same prompt ───────────
@spaces.GPU
def arena_fn(text: str):
if not text.strip():
yield "", "", ""
return
if not _MODEL_READY:
yield "", _response_html("model not loaded", "DEKU · DISTILLED"), ""
return
_to_device()
names = VIZ["teacher_names"]
for base_text, deku_text, gates in _stream_pair(text):
yield (
_response_html(base_text, "BASE · QWEN2.5-0.5B", accent="#8b949e"),
_response_html(deku_text, "DEKU · DISTILLED", accent="#7c3aed"),
_html.gate_html(gates, names, ranked=False),
)
# ── CSS ───────────────────────────────────────────────────────────────────
CSS = """
/* ── Variables ─────────────────────────────────────────── */
:root {
--bg: #080b10;
--panel: #0d1117;
--panel2: #111620;
--border: #1c2129;
--border-hi: #30363d;
--indigo: #7c3aed;
--indigo-dim: rgba(124,58,237,0.18);
--cyan: #06b6d4;
--amber: #f59e0b;
--green: #34d399;
--pink: #f472b6;
--text: #e6edf3;
--text-dim: #8b949e;
--text-faint: #484f58;
--mono: "JetBrains Mono", ui-monospace, SFMono-Regular, monospace;
--radius: 10px;
}
/* ── Base ───────────────────────────────────────────────── */
* { box-sizing: border-box; }
.gradio-container {
background: var(--bg) !important;
background-image:
radial-gradient(ellipse 90% 60% at 10% -5%, rgba(124,58,237,0.07) 0%, transparent 55%),
radial-gradient(ellipse 70% 50% at 90% 105%, rgba(6,182,212,0.05) 0%, transparent 55%)
!important;
font-family: system-ui, -apple-system, BlinkMacSystemFont, sans-serif;
min-height: 100vh;
}
footer { display: none !important; }
/* ── All blocks: remove default boxy look ───────────────── */
.block {
background: var(--panel) !important;
border: 1px solid var(--border) !important;
border-radius: var(--radius) !important;
box-shadow: 0 1px 3px rgba(0,0,0,0.4), 0 0 0 0 transparent !important;
transition: border-color 0.25s, box-shadow 0.25s !important;
padding: 0 !important;
}
.block:hover {
border-color: var(--border-hi) !important;
}
/* Remove the double-border that Gradio adds */
.block .block { border: none !important; background: transparent !important; }
/* ── Label text ─────────────────────────────────────────── */
label > span, .label-wrap span, .block-label span {
font-family: var(--mono) !important;
font-size: 10px !important;
font-weight: 600 !important;
letter-spacing: 0.09em !important;
text-transform: uppercase !important;
color: var(--text-faint) !important;
}
/* ── Textarea / input ───────────────────────────────────── */
textarea, input[type="text"], input[type="number"] {
background: #060a10 !important;
border: 1px solid var(--border-hi) !important;
color: var(--text) !important;
border-radius: 8px !important;
font-size: 13px !important;
line-height: 1.6 !important;
transition: border-color 0.2s, box-shadow 0.2s !important;
resize: vertical !important;
}
textarea:focus, input[type="text"]:focus {
border-color: var(--indigo) !important;
box-shadow: 0 0 0 3px rgba(124,58,237,0.12),
0 0 18px rgba(124,58,237,0.18) !important;
outline: none !important;
}
textarea::placeholder { color: var(--text-faint) !important; }
/* ── Primary button ─────────────────────────────────────── */
button.primary, button[variant="primary"] {
background: linear-gradient(135deg, #6d28d9, var(--indigo)) !important;
border: 1px solid rgba(124,58,237,0.45) !important;
border-radius: 8px !important;
color: #fff !important;
font-family: var(--mono) !important;
font-size: 12px !important;
font-weight: 700 !important;
letter-spacing: 0.07em !important;
text-transform: uppercase !important;
padding: 10px 22px !important;
transition: all 0.2s ease !important;
box-shadow: 0 2px 10px rgba(124,58,237,0.22) !important;
cursor: pointer !important;
}
button.primary:hover {
background: linear-gradient(135deg, var(--indigo), #8b5cf6) !important;
box-shadow: 0 4px 20px rgba(124,58,237,0.45),
0 0 0 1px rgba(124,58,237,0.35) !important;
transform: translateY(-1px) !important;
}
button.primary:active { transform: translateY(0) !important; }
/* Secondary buttons */
button.secondary {
background: var(--panel2) !important;
border: 1px solid var(--border-hi) !important;
color: var(--text-dim) !important;
border-radius: 8px !important;
transition: all 0.2s !important;
}
button.secondary:hover {
border-color: var(--indigo) !important;
color: var(--text) !important;
}
/* ── Tabs: underline style (Linear / GitHub / Vercel) ───── */
.tabs > .tab-nav,
div[role="tablist"] {
background: transparent !important;
border-bottom: 1px solid var(--border) !important;
gap: 0 !important;
padding: 0 2px !important;
}
.tab-nav button, div[role="tab"] {
background: transparent !important;
border: none !important;
border-bottom: 2px solid transparent !important;
border-radius: 0 !important;
color: var(--text-dim) !important;
font-family: var(--mono) !important;
font-size: 11px !important;
font-weight: 600 !important;
letter-spacing: 0.07em !important;
text-transform: uppercase !important;
padding: 10px 18px 9px !important;
margin-bottom: -1px !important;
transition: color 0.18s, border-color 0.18s !important;
box-shadow: none !important;
}
.tab-nav button:hover {
color: var(--text) !important;
border-bottom-color: rgba(124,58,237,0.4) !important;
}
.tab-nav button.selected {
color: var(--indigo) !important;
border-bottom: 2px solid var(--indigo) !important;
background: transparent !important;
box-shadow: none !important;
}
/* ── Plotly / charts: transparent background ────────────── */
.plot-container, .plot-container > div, .js-plotly-plot {
background: transparent !important;
}
/* ── Model3D container ──────────────────────────────────── */
div[data-testid="model3d"], .model3D-component {
border-radius: var(--radius) !important;
overflow: hidden !important;
border: 1px solid var(--border) !important;
box-shadow: 0 0 40px rgba(124,58,237,0.08) inset !important;
}
/* ── Scrollbars ─────────────────────────────────────────── */
::-webkit-scrollbar { width: 5px; height: 5px; }
::-webkit-scrollbar-track { background: var(--bg); }
::-webkit-scrollbar-thumb { background: var(--border-hi); border-radius: 99px; }
::-webkit-scrollbar-thumb:hover { background: var(--text-faint); }
/* ── Animated LIVE badge ─────────────────────────────────── */
@keyframes pulse-dot {
0%, 100% { opacity: 1; box-shadow: 0 0 6px var(--cyan); }
50% { opacity: 0.5; box-shadow: 0 0 2px var(--cyan); }
}
.live-dot { animation: pulse-dot 1.8s ease-in-out infinite; }
/* ── Fade-in on load ─────────────────────────────────────── */
@keyframes fadeUp {
from { opacity: 0; transform: translateY(10px); }
to { opacity: 1; transform: translateY(0); }
}
.gradio-container > .main > .wrap { animation: fadeUp 0.45s ease; }
/* ── Row/col gaps ────────────────────────────────────────── */
.gap { gap: 14px !important; }
.row { gap: 14px !important; }
"""
# ── Layout ────────────────────────────────────────────────────────────────
with gr.Blocks(css=CSS, theme=gr.themes.Base(), title="One for All") as demo:
gr.HTML(_html.header_html(
n_teachers=len(VIZ["teacher_names"]) or 6,
backend=BACKEND,
))
probe_state = gr.State([])
with gr.Tabs():
# ── Tab 1: Almas ──────────────────────────────────────────────────
with gr.TabItem("Souls"):
with gr.Row():
with gr.Column(scale=6):
umap_plot = gr.Model3D(
value=_INIT_GLB,
display_mode="solid",
clear_color=[0.031, 0.043, 0.063, 1.0],
height=500,
label=None,
camera_position=_CAM,
)
gr.HTML(_glb.build_legend_html(VIZ))
with gr.Column(scale=4):
gr.HTML(
'<div style="display:flex;align-items:center;gap:8px;'
'font-size:14px;font-weight:600;color:#e6edf3;margin-bottom:8px;">'
'<span style="color:#06b6d4;">⚡</span>Probe the student'
'<span style="font-family:monospace;font-size:10px;color:#06b6d4;'
'border:1px solid rgba(6,182,212,0.4);border-radius:4px;padding:2px 7px;">LIVE</span>'
'</div>'
)
prompt_box = gr.Textbox(
lines=4,
placeholder="Ask anything — code, math, language…",
label="",
)
run_btn = gr.Button("Run", variant="primary")
resp_out = gr.HTML()
gate_out = gr.HTML()
task_out = gr.HTML()
gr.HTML(
'<div style="font-size:11px;color:#8b949e;margin-top:8px;'
'font-family:monospace;">↑ new probe point will appear in soul space</div>'
)
# ── Tab 2: Arena — base vs deku, same prompt ──────────────────────
with gr.TabItem("Arena"):
gr.HTML(
'<div style="display:flex;align-items:center;gap:8px;'
'font-size:14px;font-weight:600;color:#e6edf3;margin:8px 0;">'
'<span style="color:#f59e0b;">⚔</span>One prompt, two models'
'<span style="font-family:monospace;font-size:10px;color:#8b949e;">'
'same 0.5B weights — LoRA adapter off vs on</span>'
'</div>'
)
arena_box = gr.Textbox(
lines=3,
placeholder="Ask something — watch base and distilled race side by side…",
label="",
)
arena_btn = gr.Button("Run both", variant="primary")
with gr.Row():
with gr.Column():
base_out = gr.HTML()
with gr.Column():
deku_out = gr.HTML()
arena_gate_out = gr.HTML()
gr.Examples(
examples=[
"Natalia sold clips to 48 of her friends in April, and then she "
"sold half as many clips in May. How many clips did Natalia sell "
"altogether in April and May?",
"Which property of a mineral can be determined just by looking "
"at it? (A) luster (B) mass (C) weight (D) hardness",
"Write a Python function that checks if a word is a palindrome.",
"Explain why the sky is blue in two sentences.",
],
inputs=[arena_box],
)
# ── Tab 3: Geometria ──────────────────────────────────────────────
with gr.TabItem("Geometry"):
with gr.Row():
with gr.Column(scale=7):
gr.Plot(
value=_fig.build_cka_fig(VIZ["cka"]),
label="CKA geometry alignment",
)
with gr.Column(scale=3):
cka_matrix = VIZ["cka"].get("matrix", [])
if cka_matrix:
import numpy as _np
mat = _np.array(cka_matrix)
n = mat.shape[0]
mask = ~_np.eye(n, dtype=bool)
mean_off = float(mat[mask].mean())
masked = mat.copy()
_np.fill_diagonal(masked, 1.0)
min_idx = _np.unravel_index(masked.argmin(), masked.shape)
hard_pair = (VIZ["cka"]["models"][min_idx[0]],
VIZ["cka"]["models"][min_idx[1]])
hard_val = float(masked[min_idx])
gr.HTML(
f'<div style="background:#161b22;border:1px solid #30363d;'
f'border-radius:6px;padding:18px;margin-top:8px;">'
f'<div style="font-size:28px;font-family:monospace;color:#06b6d4;'
f'font-weight:700;">{mean_off:.3f}</div>'
f'<div style="font-size:11px;color:#8b949e;margin-top:4px;">'
f'mean off-diagonal CKA</div>'
f'<div style="margin-top:16px;font-size:11px;color:#8b949e;">hardest pair</div>'
f'<div style="font-family:monospace;font-size:12px;color:#f59e0b;margin-top:4px;">'
f'{hard_pair[0]}{hard_pair[1]}'
f' <span style="color:#8b949e;">{hard_val:.2f}</span></div>'
f'</div>'
)
# ── Tab 4: Treino ─────────────────────────────────────────────────
with gr.TabItem("Training"):
with gr.Row():
gr.Plot(
value=_fig.build_curves_fig(VIZ["curves"]),
label="Loss curves",
)
gr.Plot(
value=_fig.build_gate_area_fig(VIZ["curves"]),
label="Gate evolution",
)
# ── Event wiring ──────────────────────────────────────────────────────
run_btn.click(
probe_fn,
inputs=[prompt_box, probe_state],
outputs=[umap_plot, probe_state, resp_out, gate_out, task_out],
)
arena_btn.click(
arena_fn,
inputs=[arena_box],
outputs=[base_out, deku_out, arena_gate_out],
)
if __name__ == "__main__":
demo.launch()