OpenTransformer's picture
Add AGILLM 4.3 Gradio GUI Space
b78b894 verified
Raw
History Blame Contribute Delete
17 kB
import json
import os
import re
import subprocess
import sys
import threading
import time
from pathlib import Path
import gradio as gr
from huggingface_hub import hf_hub_download
try:
import spaces
except Exception:
class _SpacesFallback:
def GPU(self, *args, **kwargs):
if args and callable(args[0]) and len(args) == 1 and not kwargs:
return args[0]
def deco(fn):
return fn
return deco
spaces = _SpacesFallback()
APP_DIR = Path(__file__).resolve().parent
RUNTIME = APP_DIR / "agillm41.py"
MODEL_REPO = "OpenTransformer/AGILLM-4.3"
CKPT_FILE = (
"checkpoints/recovery_fedC/artifacts/delta/"
"pretrain_delta_step00363424_20260703T1105Z__sha256_3e3f65ca7784/"
"pretrain_delta_step00363424_20260703T1105Z.pt"
)
TOKENIZER_FILE = (
"checkpoints/recovery_fedC/artifacts/full/"
"pretrain_step00002127_from00243186_20260701T0647Z__sha256_760874aadf59/"
"pretrain_step00002127_from00243186_20260701T0647Z.pt.tokenizer.json"
)
PROFILE = os.environ.get("AGILLM_SPACE_PROFILE", "cpu").strip().lower()
SPACE_REPO_NAME = os.environ.get("SPACE_REPO_NAME", "").strip().lower()
ACCELERATOR = os.environ.get("ACCELERATOR", "").strip().lower()
ZERO_GPU = (
PROFILE in {"zero", "zerogpu", "zero-gpu", "gpu"}
or "zerogpu" in SPACE_REPO_NAME
or ACCELERATOR.startswith("zero")
)
STAT_RE = re.compile(r"\[(?P<sec>[0-9.]+)s \| (?P<tok>[0-9]+) tokens \| (?P<tps>[0-9.]+) tok/s\]")
SERVER_LOCK = threading.RLock()
SERVER_PROC = None
SERVER_KEY = None
def _space_threads(default=2):
raw = os.environ.get("CPU_CORES") or os.cpu_count() or default
try:
return max(1, min(8, int(float(raw))))
except Exception:
return default
def _materialize_files():
local_dir = APP_DIR / "checkpoints"
local_dir.mkdir(parents=True, exist_ok=True)
ckpt = Path(hf_hub_download(MODEL_REPO, CKPT_FILE, repo_type="model", local_dir=local_dir))
tokenizer = Path(hf_hub_download(MODEL_REPO, TOKENIZER_FILE, repo_type="model", local_dir=local_dir))
return ckpt, tokenizer
def _runtime_env(tokenizer, threads):
env = os.environ.copy()
env["PYTHONUNBUFFERED"] = "1"
env["PYTHONUTF8"] = "1"
env["AGILLM43_TOKENIZER_JSON"] = str(tokenizer)
env["OMP_NUM_THREADS"] = str(max(1, int(threads)))
env["MKL_NUM_THREADS"] = str(max(1, int(threads)))
return env
def _mode_parts(mode_label):
if mode_label == "sat fixed":
return "sat", False
if mode_label == "sat var":
return "sat", True
return mode_label, None
def _payload(
prompt,
mode_label,
output_mode,
max_new,
min_new,
nat_passes,
temperature,
top_p,
top_k,
greedy,
ignore_eos,
repetition_penalty,
presence_penalty,
frequency_penalty,
penalty_last_n,
):
mode, sat_var = _mode_parts(mode_label)
data = {
"prompt": str(prompt or ""),
"mode": mode,
"max_new": int(max_new),
"min_new": int(min_new),
"nat_passes": int(nat_passes),
"temperature": float(temperature),
"top_p": float(top_p),
"top_k": int(top_k),
"greedy": bool(greedy),
"ignore_eos": bool(ignore_eos),
"repetition_penalty": float(repetition_penalty),
"presence_penalty": float(presence_penalty),
"frequency_penalty": float(frequency_penalty),
"penalty_last_n": int(penalty_last_n),
"stream": output_mode == "Streaming",
}
if sat_var is not None:
data["var"] = bool(sat_var)
return data
def _command_from_payload(ckpt, data, device, threads):
cmd = [
sys.executable,
"-u",
str(RUNTIME),
"infer",
"--ckpt",
str(ckpt),
"--prompt",
data["prompt"],
"--mode",
data["mode"],
"--max_new",
str(data["max_new"]),
"--min_new",
str(data["min_new"]),
"--temperature",
str(data["temperature"]),
"--top_p",
str(data["top_p"]),
"--top_k",
str(data["top_k"]),
"--repetition_penalty",
str(data["repetition_penalty"]),
"--presence_penalty",
str(data["presence_penalty"]),
"--frequency_penalty",
str(data["frequency_penalty"]),
"--penalty_last_n",
str(data["penalty_last_n"]),
"--plain-output",
"--device",
device,
]
if device == "cpu":
cmd.extend(["--cpu_threads", str(max(1, int(threads))), "--infer_dtype", "fp32"])
else:
cmd.extend(["--infer_dtype", "fp16", "--attn_backend", "sdpa"])
if data.get("stream"):
cmd.append("--stream")
if data.get("greedy"):
cmd.append("--greedy")
if data.get("ignore_eos"):
cmd.append("--ignore_eos")
if data["mode"] == "nat":
cmd.extend(["--nat_passes", str(data["nat_passes"])])
if data["mode"] == "sat" and "var" in data:
cmd.append("--var" if data["var"] else "--no-var")
return cmd
def _server_command(ckpt, threads):
return [
sys.executable,
"-u",
str(RUNTIME),
"infer",
"--server",
"--device",
"cpu",
"--cpu_threads",
str(max(1, int(threads))),
"--ckpt",
str(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",
"--infer_dtype",
"fp32",
]
def _alive(proc):
return proc is not None and proc.poll() is None
def _ensure_cpu_server(threads):
global SERVER_PROC, SERVER_KEY
ckpt, tokenizer = _materialize_files()
key = (str(ckpt), str(tokenizer), int(threads))
if _alive(SERVER_PROC) and SERVER_KEY == key:
return SERVER_PROC, ckpt
if _alive(SERVER_PROC):
try:
SERVER_PROC.stdin.write('{"cmd":"quit"}\n')
SERVER_PROC.stdin.flush()
except Exception:
pass
try:
SERVER_PROC.terminate()
except Exception:
pass
env = _runtime_env(tokenizer, threads)
proc = subprocess.Popen(
_server_command(ckpt, threads),
cwd=str(APP_DIR),
env=env,
text=True,
encoding="utf-8",
errors="replace",
stdin=subprocess.PIPE,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
bufsize=1,
)
boot = []
deadline = time.time() + 900
while time.time() < deadline:
line = proc.stdout.readline()
if line:
boot.append(line.rstrip("\n"))
if "[INFER_SERVER_READY]" in line:
SERVER_PROC = proc
SERVER_KEY = key
return proc, ckpt
if proc.poll() is not None:
tail = "\n".join(boot[-40:])
raise RuntimeError(f"runtime exited during warm load\n{tail}")
raise TimeoutError("warm load timed out before runtime was ready")
def _strip_prompt(text, prompt):
text = (text or "").strip()
prompt = (prompt or "").strip()
if prompt and text.startswith(prompt):
return text[len(prompt):].lstrip()
return text
def _stats_status(kind, started, stats, ckpt_name):
elapsed = max(0.001, time.time() - started)
if not stats:
return f"{kind} | button_to_done={elapsed:.2f}s | checkpoint={ckpt_name}"
tokens = int(stats.get("tokens") or 0)
button_tps = tokens / elapsed if tokens else 0.0
return (
f"{kind} | button_to_done={elapsed:.2f}s | "
f"button_to_done_tok_s={button_tps:.2f} | "
f"generation={stats.get('gen_s', '?')}s | "
f"generation_tok_s={stats.get('tok_s', '?')} | "
f"tokens={tokens} | checkpoint={ckpt_name}"
)
def _read_result_lines(proc, prompt, streaming, started, ckpt_name):
slots = None
stats = None
final_lines = []
saw_start = False
while True:
line = proc.stdout.readline()
if not line:
if proc.poll() is not None:
raise RuntimeError("runtime exited mid-generation")
continue
s = line.rstrip("\n")
if "[INFER_SERVER_RESULT_START]" in s:
saw_start = True
continue
if "[INFER_SERVER_RESULT_END]" in s:
break
if "[INFER_SERVER_ERROR]" in s:
raise RuntimeError(s)
if not saw_start:
continue
if s.startswith("[STREAM_BEGIN] "):
try:
info = json.loads(s.split("] ", 1)[1])
slots = [""] * int(info.get("slots") or 0)
if streaming:
yield "".join("." for _ in slots), "streaming..."
except Exception:
pass
continue
if s.startswith("[STREAM_NAT] ") or s.startswith("[STREAM_AR] ") or s.startswith("[STREAM_SAT] "):
try:
event = json.loads(s.split("] ", 1)[1])
idx = event.get("pos", event.get("i"))
if slots is not None and idx is not None:
idx = int(idx)
if 0 <= idx < len(slots):
slots[idx] = str(event.get("text") or "")
if streaming and slots is not None:
yield "".join(piece if piece else "." for piece in slots), "streaming..."
except Exception:
pass
continue
match = STAT_RE.search(s)
if match:
stats = {
"gen_s": float(match.group("sec")),
"tokens": int(match.group("tok")),
"tok_s": float(match.group("tps")),
}
continue
if s.startswith("[infer]") or s.startswith("Generating") or s.startswith("["):
continue
final_lines.append(s)
if streaming and slots is None:
yield _strip_prompt(s, prompt), "streaming..."
final = _strip_prompt("\n".join(final_lines), prompt)
yield final, _stats_status("done", started, stats, ckpt_name)
def _read_one_shot(proc, prompt, streaming, started, ckpt_name):
slots = None
stats = None
final_lines = []
while True:
line = proc.stdout.readline()
if not line:
if proc.poll() is not None:
break
continue
s = line.rstrip("\n")
if s.startswith("[STREAM_BEGIN] "):
try:
info = json.loads(s.split("] ", 1)[1])
slots = [""] * int(info.get("slots") or 0)
if streaming:
yield "".join("." for _ in slots), "streaming..."
except Exception:
pass
continue
if s.startswith("[STREAM_NAT] ") or s.startswith("[STREAM_AR] ") or s.startswith("[STREAM_SAT] "):
try:
event = json.loads(s.split("] ", 1)[1])
idx = event.get("pos", event.get("i"))
if slots is not None and idx is not None:
idx = int(idx)
if 0 <= idx < len(slots):
slots[idx] = str(event.get("text") or "")
if streaming and slots is not None:
yield "".join(piece if piece else "." for piece in slots), "streaming..."
except Exception:
pass
continue
match = STAT_RE.search(s)
if match:
stats = {
"gen_s": float(match.group("sec")),
"tokens": int(match.group("tok")),
"tok_s": float(match.group("tps")),
}
continue
if s.startswith("[infer]") or s.startswith("Generating") or s.startswith("["):
continue
final_lines.append(s)
if streaming and slots is None:
yield _strip_prompt(s, prompt), "streaming..."
rc = proc.wait()
if rc != 0:
raise RuntimeError(f"runtime exited with rc={rc}")
final = _strip_prompt("\n".join(final_lines), prompt)
yield final, _stats_status("done", started, stats, ckpt_name)
def _generate_cpu(data, threads):
streaming = bool(data.get("stream"))
started = time.time()
yield "", "loading warm CPU runtime..."
with SERVER_LOCK:
proc, ckpt = _ensure_cpu_server(threads)
proc.stdin.write(json.dumps(data) + "\n")
proc.stdin.flush()
yield from _read_result_lines(proc, data["prompt"], streaming, started, ckpt.name)
def _generate_once(data, device, threads):
streaming = bool(data.get("stream"))
started = time.time()
yield "", f"loading {device} runtime..."
ckpt, tokenizer = _materialize_files()
env = _runtime_env(tokenizer, threads)
proc = subprocess.Popen(
_command_from_payload(ckpt, data, device, threads),
cwd=str(APP_DIR),
env=env,
text=True,
encoding="utf-8",
errors="replace",
stdin=subprocess.DEVNULL,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
bufsize=1,
)
yield from _read_one_shot(proc, data["prompt"], streaming, started, ckpt.name)
def _collect_inputs(*args):
return _payload(*args[:-1]), int(args[-1])
def generate_cpu(*args):
data, threads = _collect_inputs(*args)
yield from _generate_cpu(data, threads)
def _gpu_duration(*args):
try:
max_new = int(args[3])
except Exception:
max_new = 16
return max(60, min(240, 70 + max_new * 4))
@spaces.GPU(duration=_gpu_duration)
def generate_zerogpu(*args):
data, threads = _collect_inputs(*args)
yield from _generate_once(data, "cuda", threads)
def warm_load(threads):
if ZERO_GPU:
return "ZeroGPU warms inside each GPU call."
started = time.time()
with SERVER_LOCK:
_proc, ckpt = _ensure_cpu_server(int(threads))
return f"CPU runtime ready in {time.time() - started:.2f}s | checkpoint={ckpt.name}"
def default_status():
hw = "ZeroGPU" if ZERO_GPU else "CPU"
accelerator = os.environ.get("ACCELERATOR", "none")
return f"{hw} Space | accelerator={accelerator} | profile={PROFILE}"
with gr.Blocks(title="AGILLM 4.3 Inference") as demo:
with gr.Row():
prompt = gr.Textbox(
value="The quick brown fox jumps over the lazy dog and then",
label="Prompt",
lines=2,
scale=5,
)
with gr.Row():
mode = gr.Dropdown(["nat", "sat fixed", "sat var", "ar"], value="nat", label="Mode")
output_mode = gr.Dropdown(["Streaming", "Full result"], value="Streaming", label="Output")
max_new = gr.Slider(1, 256, value=16 if ZERO_GPU else 8, step=1, label="Max")
min_new = gr.Slider(0, 256, value=0, step=1, label="Min")
nat_passes = gr.Slider(1, 128, value=1, step=1, label="NAT passes")
threads = gr.Slider(1, 8, value=_space_threads(), step=1, label="Threads")
with gr.Row():
temperature = gr.Number(value=0.25, label="Temp")
top_p = gr.Number(value=1.0, label="Top-p")
top_k = gr.Number(value=0, label="Top-k")
greedy = gr.Checkbox(value=True, label="Greedy")
ignore_eos = gr.Checkbox(value=True, label="Ignore EOS")
with gr.Row():
repetition_penalty = gr.Number(value=2.0, label="Repeat pen")
presence_penalty = gr.Number(value=0.8, label="Presence")
frequency_penalty = gr.Number(value=1.2, label="Frequency")
penalty_last_n = gr.Number(value=0, precision=0, label="Last N")
with gr.Row():
run = gr.Button("Run Inference", variant="primary")
warm = gr.Button("Warm Load")
output = gr.Textbox(label="Output", lines=14, show_copy_button=True)
status = gr.Textbox(value=default_status(), label="Status", lines=3)
inputs = [
prompt,
mode,
output_mode,
max_new,
min_new,
nat_passes,
temperature,
top_p,
top_k,
greedy,
ignore_eos,
repetition_penalty,
presence_penalty,
frequency_penalty,
penalty_last_n,
threads,
]
run.click(
fn=generate_zerogpu if ZERO_GPU else generate_cpu,
inputs=inputs,
outputs=[output, status],
show_progress="minimal",
concurrency_limit=1,
)
prompt.submit(
fn=generate_zerogpu if ZERO_GPU else generate_cpu,
inputs=inputs,
outputs=[output, status],
show_progress="minimal",
concurrency_limit=1,
)
warm.click(fn=warm_load, inputs=[threads], outputs=[status], show_progress="minimal", concurrency_limit=1)
if __name__ == "__main__":
demo.queue(max_size=8, default_concurrency_limit=1).launch()