""" 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" f"{n} tokens · TTFT {ttft_ms:.0f} ms · {n/elapsed:.1f} tok/s" ) 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"""
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.