OpenTransformer's picture
AGILLM 4.3 GUI space: streaming NAT canvas, warm server / ZeroGPU in-process
d5558c9 verified
Raw
History Blame Contribute Delete
9.26 kB
"""AGILLM 4.3 GUI — Hugging Face Space (free CPU).
Mirrors the local Tk GUI: one warm `infer --server` child process stays
loaded between requests; streamed NAT decode renders as a canvas that
fills in confidence order ([STREAM_*] marker protocol).
"""
import json
import os
import re
import subprocess
import sys
import threading
import time
import gradio as gr
from huggingface_hub import hf_hub_download
HERE = os.path.dirname(os.path.abspath(__file__))
MODEL_REPO = "OpenTransformer/AGILLM-4.3"
DELTA_DIR = "checkpoints/recovery_fedC/artifacts/delta/pretrain_delta_step00374922_20260703T1248Z__sha256_f2d01389959f"
CKPT_FILE = "pretrain_delta_step00374922_20260703T1248Z.pt"
RUNTIME = os.path.join(HERE, "agillm41.py")
STAT_RE = re.compile(r"\[(?P<sec>[0-9.]+)s \| (?P<tok>[0-9]+) tokens \| (?P<tps>[0-9.]+) tok/s\]")
DEFAULT_PROMPT = "The quick brown fox jumps over the lazy dog and then"
print("[space] downloading checkpoint (5.4 GB, cached by HF hub)...", flush=True)
CKPT = hf_hub_download(MODEL_REPO, f"{DELTA_DIR}/{CKPT_FILE}")
TOKENIZER = hf_hub_download(MODEL_REPO, f"{DELTA_DIR}/{CKPT_FILE}.tokenizer.json")
print("[space] checkpoint ready:", CKPT, flush=True)
CPU_THREADS = str(max(1, os.cpu_count() or 2))
ENV = dict(os.environ)
ENV.update({
"PYTHONUNBUFFERED": "1",
"PYTHONUTF8": "1",
"AGILLM43_TOKENIZER_JSON": TOKENIZER,
"OMP_NUM_THREADS": CPU_THREADS,
"MKL_NUM_THREADS": CPU_THREADS,
})
lock = threading.Lock()
child = None
ready = False
def start_child():
global child, ready
cmd = [sys.executable, "-u", RUNTIME, "infer", "--server", "--device", "cpu",
"--cpu_threads", CPU_THREADS, "--ckpt", CKPT, "--mode", "nat",
"--max_new", "64", "--min_new", "0", "--temperature", "0.25",
"--top_p", "1.0", "--greedy", "--ignore_eos", "--plain-output",
"--repetition_penalty", "2.0", "--presence_penalty", "0.8",
"--frequency_penalty", "1.2", "--penalty_last_n", "0"]
child = subprocess.Popen(cmd, cwd=HERE, stdin=subprocess.PIPE, stdout=subprocess.PIPE,
stderr=subprocess.STDOUT, text=True, bufsize=1, env=ENV)
ready = False
for line in child.stdout:
print("[child]", line.rstrip(), flush=True)
if "[INFER_SERVER_READY]" in line:
ready = True
return True
if child.poll() is not None:
return False
return False
threading.Thread(target=start_child, daemon=True).start()
def generate(prompt: str, mode: str = "nat", max_new: int = 48, nat_passes: int = 1,
temperature: float = 0.25, top_k: int = 0, top_p: float = 1.0,
greedy: bool = True, ignore_eos: bool = True, rep_pen: float = 2.0,
presence: float = 0.8, frequency: float = 1.2, last_n: int = 0,
streaming: bool = True):
"""Generate text with the AGILLM 4.3 research LLM (1.2B params, CPU Space).
This function is also exposed as an MCP tool so agents can call the model.
Args:
prompt: Input text to continue.
mode: Decode head — "nat" (fast parallel mask-predict), "sat var",
"sat fixed", or "ar" (classic left-to-right). NAT is fastest.
max_new: Number of tokens to generate.
nat_passes: NAT refinement passes (1 is fast; more improves quality).
temperature: Sampling temperature (ignored when greedy=True).
top_k: Top-k filter, 0 = off.
top_p: Nucleus sampling threshold.
greedy: Take the argmax token each step (deterministic).
ignore_eos: Never stop early on the end-of-sequence token.
rep_pen: Repetition penalty (>1 discourages repeats).
presence: Presence penalty.
frequency: Frequency penalty.
last_n: Penalty window in tokens, 0 = unlimited (keep 0 for NAT).
streaming: Stream tokens as they commit.
Returns:
The prompt followed by the generated continuation. Note: at this
mid-pretraining checkpoint the output is word-salad by design.
"""
global child, ready
prompt = (prompt or "").strip() or DEFAULT_PROMPT
with lock:
if child is None or child.poll() is not None or not ready:
yield "(loading the 1.2B model — the first request after a Space restart takes a few minutes on free CPU...)"
if not start_child():
yield "ERROR: model process failed to start — check the Space logs."
return
base_mode = "sat" if str(mode).startswith("sat") else str(mode)
req = {"prompt": prompt, "mode": base_mode, "max_new": int(max_new), "min_new": 0,
"nat_passes": int(nat_passes), "temperature": float(temperature),
"top_k": int(top_k), "top_p": float(top_p), "greedy": bool(greedy),
"ignore_eos": bool(ignore_eos), "repetition_penalty": float(rep_pen),
"presence_penalty": float(presence), "frequency_penalty": float(frequency),
"penalty_last_n": int(last_n), "stream": bool(streaming)}
if mode == "sat var":
req["var"] = True
elif mode == "sat fixed":
req["var"] = False
child.stdin.write(json.dumps(req) + "\n")
child.stdin.flush()
slots = None
final = None
stats = ""
t0 = time.time()
for line in child.stdout:
s = line.strip()
if "[INFER_SERVER_RESULT_END]" in s:
break
if "[INFER_SERVER_ERROR]" in s:
final = s
break
if s.startswith("[STREAM_BEGIN] "):
try:
slots = [None] * int(json.loads(s[len("[STREAM_BEGIN] "):]).get("slots") or 0)
except Exception:
slots = None
continue
if s.startswith("[STREAM_NAT] ") or s.startswith("[STREAM_AR] "):
if slots is None:
continue
try:
d = json.loads(s.split("] ", 1)[1])
i = int(d.get("pos", d.get("i")))
if 0 <= i < len(slots):
slots[i] = str(d.get("text") or "")
except Exception:
pass
yield prompt + "".join(x if x is not None else " ·" for x in slots)
continue
if STAT_RE.search(s):
stats = s
continue
if s and not s.startswith("[") and not s.startswith("Generating"):
final = s
wall = time.time() - t0
out = final or "(no output)"
if stats:
out += f"\n\n{stats} | wall={wall:.2f}s (free CPU is slow; the ZeroGPU Space is much faster)"
yield out
with gr.Blocks(title="AGILLM 4.3 GUI (CPU)") as demo:
gr.Markdown(
"# AGILLM 4.3 — Local Inference GUI (CPU Space)\n"
"1.2B-param research model with **AR / SAT / NAT** decode heads, trained from scratch on rented GPUs. "
"NAT (mask-predict) **streams as a canvas filling in confidence order** — watch the diffusion-style decode live. "
"Mid-pretraining checkpoint: expect word-salad, not prose. Free CPU is slow (~1-2 tok/s); "
"first request after a restart loads the model (minutes)."
)
with gr.Row():
prompt = gr.Textbox(label="Prompt", value=DEFAULT_PROMPT, lines=2, scale=4)
with gr.Row():
mode = gr.Radio(["nat", "sat var", "sat fixed", "ar"], value="nat", label="Mode")
streaming = gr.Checkbox(value=True, label="Streaming (live canvas)")
greedy = gr.Checkbox(value=True, label="Greedy")
ignore_eos = gr.Checkbox(value=True, label="Ignore EOS")
with gr.Row():
max_new = gr.Slider(4, 256, value=48, step=4, label="Max tokens")
nat_passes = gr.Slider(1, 8, value=1, step=1, label="NAT passes")
temperature = gr.Slider(0.0, 1.5, value=0.25, step=0.05, label="Temperature")
with gr.Row():
top_k = gr.Slider(0, 200, value=0, step=1, label="Top-k (0=off)")
top_p = gr.Slider(0.1, 1.0, value=1.0, step=0.05, label="Top-p")
rep_pen = gr.Slider(1.0, 4.0, value=2.0, step=0.1, label="Repetition penalty")
with gr.Row():
presence = gr.Slider(0.0, 2.0, value=0.8, step=0.1, label="Presence penalty")
frequency = gr.Slider(0.0, 3.0, value=1.2, step=0.1, label="Frequency penalty")
last_n = gr.Slider(0, 1024, value=0, step=32, label="Penalty window (0=unlimited; keep 0 for NAT)")
out = gr.Textbox(label="Output", lines=12)
btn = gr.Button("Run Inference", variant="primary")
btn.click(generate,
inputs=[prompt, mode, max_new, nat_passes, temperature, top_k, top_p,
greedy, ignore_eos, rep_pen, presence, frequency, last_n, streaming],
outputs=out, api_name="generate")
gr.Markdown(
"**Agent / MCP access:** this Space is also an MCP server — point an MCP "
"client at `https://opentransformer-agillm43-gui-cpu.hf.space/gradio_api/mcp/sse` "
"and the `generate` tool becomes callable."
)
# AGILLM-MCP 20260703: mcp_server=True exposes generate() as an MCP tool.
demo.queue(default_concurrency_limit=1).launch(mcp_server=True)