Spaces:
Runtime error
Runtime error
File size: 9,263 Bytes
8af3c88 d5558c9 8af3c88 d5558c9 8af3c88 d5558c9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 | """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)
|