Pragya / app.py
ArushBuilds's picture
Update app.py
5c67a8d
Raw History Blame
23.1 kB
"""
app.py — Pragya by Arush Kumar · Hugging Face Spaces (PyTorch backend)
Expects alongside app.py:
model.py – GPT / GPTConfig (training codebase, unchanged)
inference.py – generate(), load_model(), sampling helpers
tokenizer.py – TokenizerWrapper
tokenizer.model – SentencePiece vocab file
checkpoint.pt – trained checkpoint (from train.py)
Environment variables (Space Settings → Variables):
CHECKPOINT_PATH path to local checkpoint (default ./checkpoint.pt)
CHECKPOINT_REPO HF Hub repo id (e.g. yourname/pragya)
CHECKPOINT_FILE filename inside Hub repo (default checkpoint.pt)
TOKENIZER_PATH path to tokenizer.model (default ./tokenizer.model)
MAX_TOKENS_CAP hard cap on slider (default 512)
DEFAULT_MAX_TOKENS default slider value (default 128)
"""
from __future__ import annotations
import contextlib
import os
import sys
import time
import threading
from pathlib import Path
import torch
import gradio as gr
_HERE = Path(__file__).resolve().parent
if str(_HERE) not in sys.path:
sys.path.insert(0, str(_HERE))
from inference import ( # noqa: E402
load_model, pick_device, _resolve_dtype,
_apply_repetition_penalty, _block_repeated_ngrams,
_top_k_filter, _top_p_filter,
)
from tokenizer import TokenizerWrapper # noqa: E402
# ── env config ───────────────────────────────────────────────────────────────
_CHECKPOINT_LOCAL = os.environ.get("CHECKPOINT_PATH", str(_HERE / "pragya.pt"))
_HF_REPO = os.environ.get("CHECKPOINT_REPO", "")
_HF_FILE = os.environ.get("CHECKPOINT_FILE", "pragya.pt")
_TOKENIZER = os.environ.get("TOKENIZER_PATH", str(_HERE / "tokenizer.pkl"))
_MAX_TOKENS_CAP = int(os.environ.get("MAX_TOKENS_CAP", "512"))
_DEFAULT_TOKENS = int(os.environ.get("DEFAULT_MAX_TOKENS", "128"))
def _resolve_checkpoint() -> str:
if _HF_REPO and _HF_FILE:
try:
from huggingface_hub import hf_hub_download
print(f"[App] Pulling {_HF_REPO}/{_HF_FILE} from HF Hub …", flush=True)
return hf_hub_download(repo_id=_HF_REPO, filename=_HF_FILE)
except Exception as exc:
print(f"[App] Hub download failed ({exc}), trying local.", flush=True)
if Path(_CHECKPOINT_LOCAL).exists():
return _CHECKPOINT_LOCAL
raise FileNotFoundError(
f"No checkpoint at {_CHECKPOINT_LOCAL}. "
"Set CHECKPOINT_REPO + CHECKPOINT_FILE to pull from HF Hub."
)
# ── one-time model init ───────────────────────────────────────────────────────
print("[App] Initialising …", flush=True)
_device = pick_device(None)
_cast_dtype = _resolve_dtype("auto", _device)
_ckpt_path = _resolve_checkpoint()
_model, _cfg = load_model(_ckpt_path, _device)
_model.eval()
_tokenizer = TokenizerWrapper(_TOKENIZER)
_n_params_str = f"{_model.get_num_params() / 1e6:.1f}M"
_ctx_len = _cfg.block_size
_pattern = getattr(_cfg, "pattern", "dense").upper()
_n_layer = getattr(_cfg, "n_layer", "?")
_n_head = getattr(_cfg, "n_head", "?")
_kv_heads = getattr(_cfg, "num_kv_heads", _n_head)
_gen_lock = threading.Lock()
print(f"[App] Ready — {_n_params_str} params ctx={_ctx_len} device={_device}", flush=True)
# ── generation ────────────────────────────────────────────────────────────────
_ANCHOR = (
"The old lighthouse stood at the edge of the cliff, its beam sweeping "
"slowly across the dark water. Every night the keeper climbed the narrow "
"stairs to check the lamp, and every night the sea answered with the steady "
"sound of waves breaking against the rocks far below.\n\n"
)
def _cast_ctx():
if _cast_dtype is not None and _device.type in ("cuda", "cpu"):
return torch.autocast(device_type=_device.type, dtype=_cast_dtype)
return contextlib.nullcontext()
def respond(
message: str,
history: list,
max_new_tokens: int,
temperature: float,
top_k: int,
top_p: float,
repetition_penalty: float,
no_repeat_ngram_size: int,
use_anchor: bool,
):
if not message.strip():
yield "Please enter a prompt."
return
prompt = (_ANCHOR + message) if use_anchor else message
prompt_ids = _tokenizer.encode(prompt, add_bos=True, add_eos=False)
if len(prompt_ids) >= _ctx_len:
prompt_ids = prompt_ids[-(_ctx_len - 1):]
generated_ids : list[int] = []
eos_id = getattr(_tokenizer, "eos_id", None)
t_start = time.perf_counter()
ttft : float | None = None
def _stats(n: int) -> str:
elapsed = max(time.perf_counter() - t_start, 1e-6)
ttft_ms = (ttft * 1000) if ttft is not None else 0.0
return (
f"\n\n<span style='font-size:11px;font-family:\"IBM Plex Mono\",monospace;"
f"color:#9ca3af;letter-spacing:0.03em'>"
f"{n} tokens · TTFT {ttft_ms:.0f} ms · {n/elapsed:.1f} tok/s</span>"
)
idx = torch.tensor([prompt_ids], dtype=torch.long, device=_device)
with _gen_lock, torch.no_grad(), _cast_ctx():
for _ in range(max_new_tokens):
idx_cond = idx if idx.size(1) <= _ctx_len else idx[:, -_ctx_len:]
logits, _ = _model(idx_cond)
logits = logits[:, -1, :].float()
row = logits[0:1]
row = _apply_repetition_penalty(row, generated_ids, repetition_penalty)
row = _block_repeated_ngrams(row, generated_ids, no_repeat_ngram_size)
row = row / max(temperature, 1e-5)
if top_k > 0:
row = _top_k_filter(row, top_k)
if top_p < 1.0:
row = _top_p_filter(row, top_p)
probs = torch.softmax(row, dim=-1)
next_id = torch.multinomial(probs, num_samples=1).item()
if ttft is None:
ttft = time.perf_counter() - t_start
generated_ids.append(next_id)
idx = torch.cat([idx, torch.tensor([[next_id]], device=_device)], dim=1)
if eos_id is not None and next_id == eos_id:
break
full = _tokenizer.decode(generated_ids)
yield full + _stats(len(generated_ids))
final = _tokenizer.decode(generated_ids) or "(no output)"
yield final + _stats(len(generated_ids))
# ── UI ────────────────────────────────────────────────────────────────────────
CSS = """
/* ══════════════════════════════════════════════════════
Sarvam-style: warm white · near-black · saffron accent
Inter (body) + IBM Plex Mono (code / telemetry)
══════════════════════════════════════════════════════ */
@import url('https://fonts.googleapis.com/css2?family=Inter:wght@400;500;600;700&family=IBM+Plex+Mono:wght@400;500&display=swap');
:root {
/* surface */
--s-bg: #fafaf9;
--s-white: #ffffff;
--s-warm-50: #f5f5f4;
--s-warm-100: #e7e5e4;
--s-warm-200: #d6d3d1;
/* text */
--s-ink: #0a0a0a;
--s-ink-2: #292524;
--s-muted: #78716c;
--s-faint: #a8a29e;
/* accent — saffron */
--s-saf: #ea580c;
--s-saf-hover: #c2410c;
--s-saf-light: #fff7ed;
--s-saf-mid: #fed7aa;
/* Gradio token overrides */
--body-background-fill: var(--s-bg);
--background-fill-primary: var(--s-white);
--background-fill-secondary: var(--s-warm-50);
--block-background-fill: var(--s-white);
--block-border-color: var(--s-warm-100);
--block-title-text-color: var(--s-muted);
--block-label-text-color: var(--s-muted);
--body-text-color: var(--s-ink);
--body-text-color-subdued: var(--s-muted);
--border-color-primary: var(--s-warm-100);
--border-color-accent: var(--s-saf);
--color-accent: var(--s-saf);
--input-background-fill: var(--s-white);
--input-border-color: var(--s-warm-200);
--input-border-color-focus: var(--s-saf);
--input-placeholder-color: var(--s-faint);
--slider-color: var(--s-saf);
--checkbox-background-color-selected: var(--s-saf);
--checkbox-border-color-selected: var(--s-saf);
--checkbox-label-background-fill-selected:
color-mix(in srgb, var(--s-saf) 8%, var(--s-white));
--button-primary-background-fill: var(--s-saf);
--button-primary-background-fill-hover: var(--s-saf-hover);
--button-primary-text-color: #ffffff;
--button-secondary-background-fill: var(--s-white);
--button-secondary-border-color: var(--s-warm-200);
--button-secondary-background-fill-hover: var(--s-warm-50);
--link-text-color: var(--s-saf);
}
/* ── Base ──────────────────────────────────────────────── */
* { box-sizing: border-box; }
body, .gradio-container {
font-family: 'Inter', system-ui, -apple-system, sans-serif !important;
background: var(--s-bg) !important;
color: var(--s-ink) !important;
}
.gradio-container {
max-width: 820px !important;
margin: 0 auto !important;
padding: 0 20px !important;
}
/* ── Nav bar ──────────────────────────────────────────── */
#s-nav {
display: flex;
align-items: center;
padding: 18px 0 16px;
border-bottom: 1px solid var(--s-warm-100);
margin-bottom: 0;
gap: 12px;
}
#s-wordmark {
font-family: 'Inter', sans-serif;
font-weight: 700;
font-size: 1.05rem;
letter-spacing: -0.02em;
color: var(--s-ink);
display: flex;
align-items: center;
gap: 8px;
text-decoration: none;
}
#s-wordmark-dot {
width: 8px; height: 8px; border-radius: 50%;
background: var(--s-saf);
flex-shrink: 0;
}
#s-nav-tag {
font-size: 0.72rem;
font-weight: 500;
color: var(--s-saf);
background: var(--s-saf-light);
border: 1px solid var(--s-saf-mid);
padding: 3px 9px;
border-radius: 999px;
letter-spacing: 0.01em;
}
#s-nav-right {
margin-left: auto;
display: flex;
align-items: center;
gap: 20px;
}
.s-nav-link {
font-size: 0.82rem;
color: var(--s-muted);
text-decoration: none;
font-weight: 500;
transition: color 0.12s;
}
.s-nav-link:hover { color: var(--s-ink); }
/* ── Hero ─────────────────────────────────────────────── */
#s-hero {
padding: 56px 0 40px;
border-bottom: 1px solid var(--s-warm-100);
}
#s-hero-eyebrow {
font-size: 0.78rem;
font-weight: 600;
letter-spacing: 0.06em;
text-transform: uppercase;
color: var(--s-saf);
margin-bottom: 16px;
}
#s-hero h1 {
font-size: clamp(2rem, 5vw, 2.8rem);
font-weight: 700;
letter-spacing: -0.03em;
line-height: 1.12;
color: var(--s-ink);
margin: 0 0 16px;
max-width: 560px;
}
#s-hero-desc {
font-size: 1rem;
color: var(--s-muted);
line-height: 1.65;
max-width: 480px;
margin-bottom: 28px;
font-weight: 400;
}
#s-stat-row {
display: flex;
gap: 32px;
flex-wrap: wrap;
}
.s-stat {
display: flex;
flex-direction: column;
gap: 2px;
}
.s-stat-val {
font-size: 1.1rem;
font-weight: 700;
letter-spacing: -0.02em;
color: var(--s-ink);
font-variant-numeric: tabular-nums;
}
.s-stat-lbl {
font-size: 0.72rem;
color: var(--s-faint);
font-weight: 500;
text-transform: uppercase;
letter-spacing: 0.05em;
}
/* ── Section label ────────────────────────────────────── */
.s-section-label {
font-size: 0.72rem;
font-weight: 600;
letter-spacing: 0.08em;
text-transform: uppercase;
color: var(--s-faint);
padding: 28px 0 12px;
border-bottom: none;
}
/* ── Chat bubbles ─────────────────────────────────────── */
/* User: saffron-tinted */
.user-row {
background: var(--s-saf) !important;
border-radius: 12px 12px 2px 12px !important;
}
.user-row * { color: #fff !important; }
/* Bot: white card with warm border */
.bot-row {
background: var(--s-white) !important;
border: 1px solid var(--s-warm-100) !important;
border-radius: 12px 12px 12px 2px !important;
}
.message-bubble-border {
box-shadow: 0 1px 3px rgba(0,0,0,0.06) !important;
}
/* Code in bot replies */
.bot-row code, .bot-row pre {
font-family: 'IBM Plex Mono', monospace !important;
background: var(--s-warm-50) !important;
border: 1px solid var(--s-warm-100) !important;
border-radius: 6px !important;
font-size: 0.82rem !important;
}
/* ── Textbox & send button ────────────────────────────── */
textarea {
font-family: 'Inter', sans-serif !important;
font-size: 0.95rem !important;
border-radius: 8px !important;
border: 1.5px solid var(--s-warm-200) !important;
transition: border-color 0.15s, box-shadow 0.15s !important;
resize: none !important;
background: var(--s-white) !important;
color: var(--s-ink) !important;
}
textarea:focus {
border-color: var(--s-saf) !important;
box-shadow: 0 0 0 3px color-mix(in srgb, var(--s-saf) 12%, transparent) !important;
outline: none !important;
}
textarea::placeholder { color: var(--s-faint) !important; }
button.primary {
background: var(--s-saf) !important;
border: none !important;
border-radius: 8px !important;
font-weight: 600 !important;
font-size: 0.88rem !important;
letter-spacing: 0.01em !important;
color: #fff !important;
transition: background 0.12s, transform 0.1s !important;
}
button.primary:hover {
background: var(--s-saf-hover) !important;
transform: translateY(-1px) !important;
}
/* ── Settings accordion ───────────────────────────────── */
.accordion {
border: 1px solid var(--s-warm-100) !important;
border-radius: 10px !important;
background: var(--s-white) !important;
overflow: hidden;
}
.accordion > .label-wrap {
padding: 12px 16px !important;
font-size: 0.85rem !important;
font-weight: 600 !important;
color: var(--s-ink-2) !important;
background: var(--s-warm-50) !important;
border-bottom: 1px solid var(--s-warm-100) !important;
}
/* Sliders */
input[type="range"] { accent-color: var(--s-saf) !important; }
/* ── Example chips ────────────────────────────────────── */
.example {
font-size: 0.82rem !important;
font-weight: 500 !important;
border: 1px solid var(--s-warm-200) !important;
border-radius: 6px !important;
color: var(--s-ink-2) !important;
background: var(--s-white) !important;
padding: 6px 12px !important;
transition: border-color 0.12s, color 0.12s, background 0.12s !important;
}
.example:hover {
border-color: var(--s-saf) !important;
color: var(--s-saf) !important;
background: var(--s-saf-light) !important;
}
/* ── Footer ───────────────────────────────────────────── */
#s-footer {
border-top: 1px solid var(--s-warm-100);
padding: 28px 0 20px;
display: flex;
justify-content: space-between;
align-items: flex-start;
flex-wrap: wrap;
gap: 16px;
margin-top: 16px;
}
#s-footer-left {
display: flex;
flex-direction: column;
gap: 5px;
}
#s-footer-brand {
font-weight: 700;
font-size: 0.88rem;
color: var(--s-ink);
letter-spacing: -0.01em;
}
#s-footer-copy {
font-size: 0.72rem;
color: var(--s-faint);
}
#s-footer-chips {
display: flex;
gap: 6px;
flex-wrap: wrap;
align-items: center;
}
.s-badge {
font-size: 0.68rem;
font-weight: 500;
font-family: 'IBM Plex Mono', monospace;
color: var(--s-muted);
background: var(--s-warm-50);
border: 1px solid var(--s-warm-100);
padding: 3px 8px;
border-radius: 4px;
letter-spacing: 0.02em;
white-space: nowrap;
}
.s-badge-accent {
color: var(--s-saf);
background: var(--s-saf-light);
border-color: var(--s-saf-mid);
}
/* ── Misc fixes ───────────────────────────────────────── */
* { scrollbar-width: thin; scrollbar-color: var(--s-warm-200) transparent; }
::-webkit-scrollbar { width: 5px; }
::-webkit-scrollbar-thumb { background: var(--s-warm-200); border-radius: 4px; }
"""
def main():
# Build runtime stats for the hero section
device_label = str(_device).upper()
dtype_label = str(_cast_dtype).replace("torch.", "").upper() if _cast_dtype else "FP32"
layer_str = str(_n_layer)
head_str = f"{_n_head}/{_kv_heads}"
header_html = f"""
<div id="s-nav">
<span id="s-wordmark">
<span id="s-wordmark-dot"></span>
Pragya
</span>
<span id="s-nav-tag">by Arush Kumar</span>
<div id="s-nav-right">
<span class="s-nav-link">PyTorch</span>
<span class="s-nav-link">{device_label}</span>
</div>
</div>
<div id="s-hero">
<div id="s-hero-eyebrow">Arush Kumar · From-scratch language model</div>
<h1>Pragya — a model<br>built from the ground up.</h1>
<p id="s-hero-desc">
A custom GPT-style architecture with Mixture-of-Depths routing,
Multi-Token Prediction heads, and YaRN RoPE context scaling.
Every weight trained from scratch in PyTorch. Type a prompt and let it continue.
</p>
<div id="s-stat-row">
<div class="s-stat">
<span class="s-stat-val">{_n_params_str}</span>
<span class="s-stat-lbl">Parameters</span>
</div>
<div class="s-stat">
<span class="s-stat-val">{_ctx_len:,}</span>
<span class="s-stat-lbl">Context length</span>
</div>
<div class="s-stat">
<span class="s-stat-val">{layer_str}L · {head_str}H</span>
<span class="s-stat-lbl">Depth · GQA heads</span>
</div>
<div class="s-stat">
<span class="s-stat-val">{_pattern}</span>
<span class="s-stat-lbl">Attention pattern</span>
</div>
<div class="s-stat">
<span class="s-stat-val">{dtype_label}</span>
<span class="s-stat-lbl">Inference dtype</span>
</div>
</div>
</div>
"""
footer_html = f"""
<div id="s-footer">
<div id="s-footer-left">
<span id="s-footer-brand">Pragya &nbsp;·&nbsp; <span style="font-weight:400;color:var(--s-muted)">Arush Kumar</span></span>
<span id="s-footer-copy">
Continuation model · conversation history is shown in UI but not fed back to the model ·
TTFT and tok/s are live-measured each generation
</span>
</div>
<div id="s-footer-chips">
<span class="s-badge s-badge-accent">{_n_params_str} params</span>
<span class="s-badge">{device_label}</span>
<span class="s-badge">{dtype_label}</span>
<span class="s-badge">ctx {_ctx_len}</span>
<span class="s-badge">{_pattern}</span>
</div>
</div>
"""
with gr.Blocks(title="Pragya by Arush Kumar") as demo:
gr.HTML(header_html)
gr.HTML('<div class="s-section-label">Try it</div>')
gr.ChatInterface(
fn=respond,
additional_inputs=[
gr.Slider(16, _MAX_TOKENS_CAP, value=_DEFAULT_TOKENS, step=8,
label="Max new tokens"),
gr.Slider(0.1, 2.0, value=0.8, step=0.05,
label="Temperature"),
gr.Slider(0, 200, value=50, step=5,
label="Top-k",
info="0 disables top-k filtering."),
gr.Slider(0.5, 1.0, value=1.0, step=0.01,
label="Top-p (nucleus)",
info="1.0 disables nucleus sampling."),
gr.Slider(1.0, 2.0, value=1.3, step=0.05,
label="Repetition penalty",
info="Count-scaled: penalty^occurrences. 1.0 = off."),
gr.Slider(0, 6, value=3, step=1,
label="No-repeat n-gram",
info="Hard-bans any n-gram of this length from repeating. 0 = off."),
gr.Checkbox(value=False, label="Narrative anchor",
info="Prepends a short coherent prose passage to prime the context."),
],
additional_inputs_accordion=gr.Accordion("Generation settings", open=False),
chatbot=gr.Chatbot(
height=440,
show_label=False,
render_markdown=True,
),
textbox=gr.Textbox(
placeholder="Start a sentence and let Pragya complete it …",
show_label=False,
lines=2,
submit_btn="Send",
),
examples=[
["Hello"],
["Tell me a story."],
["Create a simple story."],
["who developed you?"],
],
cache_examples=False,
)
gr.HTML(footer_html)
demo.launch(
server_name="0.0.0.0",
server_port=7860,
share=False,
css=CSS,
theme=gr.themes.Base(
primary_hue=gr.themes.colors.orange,
neutral_hue=gr.themes.colors.stone,
font=gr.themes.GoogleFont("Inter"),
),
)
if __name__ == "__main__":
main()