"""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("", "").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)