File size: 9,208 Bytes
6059b8c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7637da4
6059b8c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7637da4
6059b8c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7637da4
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
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
"""Gradio demo for gdiamos/amx-reasoning-v1-instruct.

A 7.5M-parameter (3.3M active) causal LM trained end-to-end on a single Intel
AMX CPU core. It does passage-grounded extractive QA: hand it a passage and a
question, get a one-or-two-word answer back.

The inference path here mirrors the model repo's own `example.py` exactly:
the `m2r` package that trained it, `render_prompt(..., thinking=False)` for the
prompt format, greedy argmax decoding stopped on EOT, and the 800-token
vocabulary mask from `generation.json` applied before the argmax.
"""

import json
import pathlib
import sys
import time

import gradio as gr
import torch
from huggingface_hub import snapshot_download
from safetensors.torch import load_file
from tokenizers import Tokenizer

MODEL_ID = "gdiamos/amx-reasoning-v1-instruct"

# The model's own source ships in the repo under m2r/ -- it is not a
# transformers architecture, so AutoModelForCausalLM will not load it.
LOCAL = pathlib.Path(
    snapshot_download(
        MODEL_ID,
        allow_patterns=[
            "m2r/**",
            "config.json",
            "training_config.yaml",
            "generation.json",
            "tokenizer.json",
            "model.safetensors",
        ],
    )
)
sys.path.insert(0, str(LOCAL))

from m2r.config import load  # noqa: E402
from m2r.data.templates import EOT, render_prompt  # noqa: E402
from m2r.model.torch_model import Model, swa_mask  # noqa: E402

torch.set_grad_enabled(False)

cfg = load(LOCAL / "training_config.yaml")
model = Model(cfg.model).to(torch.bfloat16)
model.load_state_dict(load_file(str(LOCAL / "model.safetensors")))
model.eval()

tok = Tokenizer.from_file(str(LOCAL / "tokenizer.json"))
MASK = swa_mask(cfg.model, dtype=torch.bfloat16)
GEN = json.loads((LOCAL / "generation.json").read_text())

# Required, not a knob: these 800 vocabulary rows never occur in the training
# corpus, so the sampled softmax never drew them as negatives and never pushed
# their logits down. They sit near 0 while trained-but-wrong tokens sit near
# -7.9, so they win the argmax whenever the model is unsure. Before masking,
# " ballo" and "Frequently" were 29% of all DROP answers.
BAN = torch.tensor(GEN["banned_token_ids"], dtype=torch.long)

PAD_TO = max(cfg.model.window, cfg.model.route_block or 1, 256)
EOT_ID = tok.encode(EOT, add_special_tokens=False).ids[0]
MAX_PASSAGE_TOKENS = 2000
DEFAULT_MAX_NEW_TOKENS = int(GEN.get("recommended", {}).get("max_new_tokens", 32))

EXAMPLES = json.loads((pathlib.Path(__file__).parent / "examples.json").read_text())

print(
    f"loaded {MODEL_ID}: "
    f"{sum(p.numel() for p in model.parameters()):,} stored parameters, "
    f"pad_to={PAD_TO}, eot_id={EOT_ID}, {len(BAN)} banned token ids",
    flush=True,
)


def _clip_passage(passage: str) -> str:
    """Trim a passage to MAX_PASSAGE_TOKENS on a token boundary."""
    enc = tok.encode(passage, add_special_tokens=False)
    if len(enc.ids) <= MAX_PASSAGE_TOKENS:
        return passage
    end = enc.offsets[MAX_PASSAGE_TOKENS - 1][1]
    return passage[:end].rstrip() + " ..."


def _next_token(ids: list[int], apply_vocab_mask: bool) -> int:
    n = len(ids)
    x = torch.tensor([ids + [0] * ((-n) % PAD_TO)])
    h = model.body(x, MASK)[:, n - 1]
    logits = (h @ model.emb.t().to(h.dtype)).float()[0]
    if apply_vocab_mask:
        logits[BAN] = -1e30
    return int(logits.argmax())


def answer(
    question: str,
    passage: str,
    max_new_tokens: int = DEFAULT_MAX_NEW_TOKENS,
    apply_vocab_mask: bool = True,
) -> tuple[str, str]:
    """Answer a question about a passage with amx-reasoning-v1-instruct.

    Greedy decoding, stopped on the end-of-turn token, using the model's own
    prompt template. The answer is normally one or two words extracted from
    the passage.

    Args:
        question: the question to ask about the passage.
        passage: the passage the answer should be grounded in.
        max_new_tokens: hard cap on generated tokens (the model usually stops
            after one or two).
        apply_vocab_mask: apply the 800-token vocabulary mask from
            generation.json. Required for correct output; turn it off only to
            see the untrained-row failure mode the model card describes.

    Returns:
        The model's answer, and a one-line note about how it was produced.
    """
    question = (question or "").strip()
    passage = (passage or "").strip()
    if not question:
        return "", "Enter a question."
    if not passage:
        return "", "This model is extractive — it needs a passage to answer from."

    clipped = _clip_passage(passage)
    prompt = render_prompt([f"{question}\n\n{clipped}"], thinking=False)
    ids = tok.encode(prompt, add_special_tokens=False).ids
    n_prompt = len(ids)

    t0 = time.perf_counter()
    out: list[int] = []
    stopped = False
    for _ in range(int(max_new_tokens)):
        t = _next_token(ids, apply_vocab_mask)
        if t == EOT_ID:
            stopped = True
            break
        out.append(t)
        ids.append(t)
    dt = time.perf_counter() - t0

    text = tok.decode(out).replace("<think></think>", "").strip()
    if not text:
        text = "(empty)"

    note = (
        f"{n_prompt} prompt tokens → {len(out)} generated in {dt:.2f}s on CPU "
        f"({'stopped on EOT' if stopped else 'hit the token cap'})"
    )
    if not apply_vocab_mask:
        note += " · **vocabulary mask off**"
    if clipped is not passage:
        note += f" · passage clipped to {MAX_PASSAGE_TOKENS} tokens"
    return text, note


CSS = """
#col-container { max-width: 1040px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
#answer textarea { font-size: 20px; font-weight: 600; }
"""

INTRO = """# amx-reasoning-v1-instruct — passage QA

A **7,492,448-parameter** causal LM (3,315,552 active per token) trained end to
end on **one Intel Emerald Rapids CPU core**. Give it a passage and a question;
it extracts a one-or-two-word answer and stops. It is a research artifact —
18.2% exact match on held-out extractive QA — and the point is that a model
this small does passage-grounded retrieval at all.

It **retrieves and compares; it cannot calculate.** It also answers some
unanswerable questions anyway. The examples below include those failures on
purpose.

[Model card](https://huggingface.co/gdiamos/amx-reasoning-v1-instruct) ·
[Paper](https://huggingface.co/gdiamos/amx-reasoning-v1-instruct/blob/main/paper.pdf)
"""

with gr.Blocks(title="amx-reasoning-v1 QA") as demo:
    with gr.Column(elem_id="col-container"):
        gr.Markdown(INTRO)

        with gr.Row():
            question = gr.Textbox(
                label="Question",
                placeholder="What year was the company founded?",
                lines=1,
                scale=4,
            )
            run = gr.Button("Answer", variant="primary", scale=1)

        passage = gr.Textbox(
            label="Passage",
            placeholder="Paste the passage the answer should come from…",
            lines=9,
        )

        answer_box = gr.Textbox(label="Answer", elem_id="answer", lines=2)
        info = gr.Markdown()

        with gr.Accordion("Advanced", open=False):
            max_new_tokens = gr.Slider(
                1, 64, value=DEFAULT_MAX_NEW_TOKENS, step=1,
                label="Max new tokens",
                info="The model normally emits one or two and stops on EOT.",
            )
            apply_vocab_mask = gr.Checkbox(
                value=True,
                label="Apply the vocabulary mask from generation.json",
                info=(
                    "800 vocabulary rows never occurred in training, so their "
                    "logits were never pushed down and they win the argmax "
                    "whenever the model is unsure. Unchecking this reproduces "
                    "the ' ballo' failure mode from the model card."
                ),
            )

        gr.Examples(
            examples=[[e["question"], e["passage"]] for e in EXAMPLES],
            example_labels=[
                f"{e['source']} — {e['question']}" for e in EXAMPLES
            ],
            inputs=[question, passage],
            outputs=[answer_box, info],
            fn=answer,
            cache_examples=True,
            cache_mode="lazy",
        )

        gr.Markdown(
            "Example passages come from the datasets the model card reports on: "
            "[DROP](https://huggingface.co/datasets/ucinlp/drop) and "
            "[SQuAD v2](https://huggingface.co/datasets/rajpurkar/squad_v2) "
            "(CC BY-SA 4.0), and "
            "[databricks-dolly-15k](https://huggingface.co/datasets/databricks/databricks-dolly-15k) "
            "(CC BY-SA 3.0). The first example is the model repo's own "
            "`example.py`."
        )

    inputs = [question, passage, max_new_tokens, apply_vocab_mask]
    outputs = [answer_box, info]
    run.click(answer, inputs=inputs, outputs=outputs, api_name="answer")
    question.submit(answer, inputs=inputs, outputs=outputs, api_name=False)

if __name__ == "__main__":
    demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)