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)