multimodalart's picture
multimodalart HF Staff
Fix Gradio 6 API: theme/css on launch, drop show_copy_button
7637da4 verified
Raw History Blame Contribute Delete
9.21 kB
"""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)