multimodalart HF Staff commited on
Commit
6059b8c
·
verified ·
1 Parent(s): 4a61d49

Passage QA demo for amx-reasoning-v1-instruct

Browse files
Files changed (4) hide show
  1. README.md +37 -7
  2. app.py +253 -0
  3. examples.json +44 -0
  4. requirements.txt +4 -0
README.md CHANGED
@@ -1,13 +1,43 @@
1
  ---
2
- title: Amx Reasoning V1 Qa Demo
3
- emoji: 📉
4
- colorFrom: yellow
5
- colorTo: gray
6
  sdk: gradio
7
  sdk_version: 6.26.0
8
- python_version: '3.13'
9
  app_file: app.py
10
- pinned: false
 
 
 
 
 
11
  ---
12
 
13
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ title: amx-reasoning-v1-instruct QA
3
+ emoji: 🧮
4
+ colorFrom: green
5
+ colorTo: red
6
  sdk: gradio
7
  sdk_version: 6.26.0
 
8
  app_file: app.py
9
+ short_description: Passage QA with a 7.5M-param CPU-trained LM
10
+ python_version: "3.12"
11
+ startup_duration_timeout: 30m
12
+ license: apache-2.0
13
+ models:
14
+ - gdiamos/amx-reasoning-v1-instruct
15
  ---
16
 
17
+ # amx-reasoning-v1-instruct — passage QA
18
+
19
+ Demo for [`gdiamos/amx-reasoning-v1-instruct`](https://huggingface.co/gdiamos/amx-reasoning-v1-instruct):
20
+ a 7,492,448-parameter causal LM (3,315,552 active per token) trained end to end
21
+ on a single Intel Emerald Rapids CPU core. Give it a passage and a question and
22
+ it extracts a one-or-two-word answer.
23
+
24
+ The demo runs the model's own reference path, unmodified:
25
+
26
+ - the `m2r` package that trained it (the architecture is not a `transformers`
27
+ one — `AutoModelForCausalLM` will not load it),
28
+ - `render_prompt(..., thinking=False)` for the exact prompt format,
29
+ - greedy argmax decoding, stopped on the `<SPECIAL_12>` end-of-turn token,
30
+ - the 800-token vocabulary mask from `generation.json` applied before the
31
+ argmax (required for correct output; the demo exposes a toggle so you can see
32
+ what happens without it).
33
+
34
+ Inference runs on CPU, which is the hardware the model was designed and trained
35
+ for — a forward pass at these dimensions is a few GFLOP.
36
+
37
+ ## Example attribution
38
+
39
+ Example passages are drawn from the datasets the model card reports on:
40
+ [DROP](https://huggingface.co/datasets/ucinlp/drop) and
41
+ [SQuAD v2](https://huggingface.co/datasets/rajpurkar/squad_v2) (CC BY-SA 4.0),
42
+ and [databricks-dolly-15k](https://huggingface.co/datasets/databricks/databricks-dolly-15k)
43
+ (CC BY-SA 3.0). The first example is the model repo's own `example.py`.
app.py ADDED
@@ -0,0 +1,253 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Gradio demo for gdiamos/amx-reasoning-v1-instruct.
2
+
3
+ A 7.5M-parameter (3.3M active) causal LM trained end-to-end on a single Intel
4
+ AMX CPU core. It does passage-grounded extractive QA: hand it a passage and a
5
+ question, get a one-or-two-word answer back.
6
+
7
+ The inference path here mirrors the model repo's own `example.py` exactly:
8
+ the `m2r` package that trained it, `render_prompt(..., thinking=False)` for the
9
+ prompt format, greedy argmax decoding stopped on EOT, and the 800-token
10
+ vocabulary mask from `generation.json` applied before the argmax.
11
+ """
12
+
13
+ import json
14
+ import pathlib
15
+ import sys
16
+ import time
17
+
18
+ import gradio as gr
19
+ import torch
20
+ from huggingface_hub import snapshot_download
21
+ from safetensors.torch import load_file
22
+ from tokenizers import Tokenizer
23
+
24
+ MODEL_ID = "gdiamos/amx-reasoning-v1-instruct"
25
+
26
+ # The model's own source ships in the repo under m2r/ -- it is not a
27
+ # transformers architecture, so AutoModelForCausalLM will not load it.
28
+ LOCAL = pathlib.Path(
29
+ snapshot_download(
30
+ MODEL_ID,
31
+ allow_patterns=[
32
+ "m2r/**",
33
+ "config.json",
34
+ "training_config.yaml",
35
+ "generation.json",
36
+ "tokenizer.json",
37
+ "model.safetensors",
38
+ ],
39
+ )
40
+ )
41
+ sys.path.insert(0, str(LOCAL))
42
+
43
+ from m2r.config import load # noqa: E402
44
+ from m2r.data.templates import EOT, render_prompt # noqa: E402
45
+ from m2r.model.torch_model import Model, swa_mask # noqa: E402
46
+
47
+ torch.set_grad_enabled(False)
48
+
49
+ cfg = load(LOCAL / "training_config.yaml")
50
+ model = Model(cfg.model).to(torch.bfloat16)
51
+ model.load_state_dict(load_file(str(LOCAL / "model.safetensors")))
52
+ model.eval()
53
+
54
+ tok = Tokenizer.from_file(str(LOCAL / "tokenizer.json"))
55
+ MASK = swa_mask(cfg.model, dtype=torch.bfloat16)
56
+ GEN = json.loads((LOCAL / "generation.json").read_text())
57
+
58
+ # Required, not a knob: these 800 vocabulary rows never occur in the training
59
+ # corpus, so the sampled softmax never drew them as negatives and never pushed
60
+ # their logits down. They sit near 0 while trained-but-wrong tokens sit near
61
+ # -7.9, so they win the argmax whenever the model is unsure. Before masking,
62
+ # " ballo" and "Frequently" were 29% of all DROP answers.
63
+ BAN = torch.tensor(GEN["banned_token_ids"], dtype=torch.long)
64
+
65
+ PAD_TO = max(cfg.model.window, cfg.model.route_block or 1, 256)
66
+ EOT_ID = tok.encode(EOT, add_special_tokens=False).ids[0]
67
+ MAX_PASSAGE_TOKENS = 2000
68
+ DEFAULT_MAX_NEW_TOKENS = int(GEN.get("recommended", {}).get("max_new_tokens", 32))
69
+
70
+ EXAMPLES = json.loads((pathlib.Path(__file__).parent / "examples.json").read_text())
71
+
72
+ print(
73
+ f"loaded {MODEL_ID}: "
74
+ f"{sum(p.numel() for p in model.parameters()):,} stored parameters, "
75
+ f"pad_to={PAD_TO}, eot_id={EOT_ID}, {len(BAN)} banned token ids",
76
+ flush=True,
77
+ )
78
+
79
+
80
+ def _clip_passage(passage: str) -> str:
81
+ """Trim a passage to MAX_PASSAGE_TOKENS on a token boundary."""
82
+ enc = tok.encode(passage, add_special_tokens=False)
83
+ if len(enc.ids) <= MAX_PASSAGE_TOKENS:
84
+ return passage
85
+ end = enc.offsets[MAX_PASSAGE_TOKENS - 1][1]
86
+ return passage[:end].rstrip() + " ..."
87
+
88
+
89
+ def _next_token(ids: list[int], apply_vocab_mask: bool) -> int:
90
+ n = len(ids)
91
+ x = torch.tensor([ids + [0] * ((-n) % PAD_TO)])
92
+ h = model.body(x, MASK)[:, n - 1]
93
+ logits = (h @ model.emb.t().to(h.dtype)).float()[0]
94
+ if apply_vocab_mask:
95
+ logits[BAN] = -1e30
96
+ return int(logits.argmax())
97
+
98
+
99
+ def answer(
100
+ question: str,
101
+ passage: str,
102
+ max_new_tokens: int = DEFAULT_MAX_NEW_TOKENS,
103
+ apply_vocab_mask: bool = True,
104
+ ) -> tuple[str, str]:
105
+ """Answer a question about a passage with amx-reasoning-v1-instruct.
106
+
107
+ Greedy decoding, stopped on the end-of-turn token, using the model's own
108
+ prompt template. The answer is normally one or two words extracted from
109
+ the passage.
110
+
111
+ Args:
112
+ question: the question to ask about the passage.
113
+ passage: the passage the answer should be grounded in.
114
+ max_new_tokens: hard cap on generated tokens (the model usually stops
115
+ after one or two).
116
+ apply_vocab_mask: apply the 800-token vocabulary mask from
117
+ generation.json. Required for correct output; turn it off only to
118
+ see the untrained-row failure mode the model card describes.
119
+
120
+ Returns:
121
+ The model's answer, and a one-line note about how it was produced.
122
+ """
123
+ question = (question or "").strip()
124
+ passage = (passage or "").strip()
125
+ if not question:
126
+ return "", "Enter a question."
127
+ if not passage:
128
+ return "", "This model is extractive — it needs a passage to answer from."
129
+
130
+ clipped = _clip_passage(passage)
131
+ prompt = render_prompt([f"{question}\n\n{clipped}"], thinking=False)
132
+ ids = tok.encode(prompt, add_special_tokens=False).ids
133
+ n_prompt = len(ids)
134
+
135
+ t0 = time.perf_counter()
136
+ out: list[int] = []
137
+ stopped = False
138
+ for _ in range(int(max_new_tokens)):
139
+ t = _next_token(ids, apply_vocab_mask)
140
+ if t == EOT_ID:
141
+ stopped = True
142
+ break
143
+ out.append(t)
144
+ ids.append(t)
145
+ dt = time.perf_counter() - t0
146
+
147
+ text = tok.decode(out).replace("<think></think>", "").strip()
148
+ if not text:
149
+ text = "(empty)"
150
+
151
+ note = (
152
+ f"{n_prompt} prompt tokens → {len(out)} generated in {dt:.2f}s on CPU "
153
+ f"({'stopped on EOT' if stopped else 'hit the token cap'})"
154
+ )
155
+ if not apply_vocab_mask:
156
+ note += " · **vocabulary mask off**"
157
+ if clipped is not passage:
158
+ note += f" · passage clipped to {MAX_PASSAGE_TOKENS} tokens"
159
+ return text, note
160
+
161
+
162
+ CSS = """
163
+ #col-container { max-width: 1040px; margin: 0 auto; }
164
+ .dark .gradio-container { color: var(--body-text-color); }
165
+ #answer textarea { font-size: 20px; font-weight: 600; }
166
+ """
167
+
168
+ INTRO = """# amx-reasoning-v1-instruct — passage QA
169
+
170
+ A **7,492,448-parameter** causal LM (3,315,552 active per token) trained end to
171
+ end on **one Intel Emerald Rapids CPU core**. Give it a passage and a question;
172
+ it extracts a one-or-two-word answer and stops. It is a research artifact —
173
+ 18.2% exact match on held-out extractive QA — and the point is that a model
174
+ this small does passage-grounded retrieval at all.
175
+
176
+ It **retrieves and compares; it cannot calculate.** It also answers some
177
+ unanswerable questions anyway. The examples below include those failures on
178
+ purpose.
179
+
180
+ [Model card](https://huggingface.co/gdiamos/amx-reasoning-v1-instruct) ·
181
+ [Paper](https://huggingface.co/gdiamos/amx-reasoning-v1-instruct/blob/main/paper.pdf)
182
+ """
183
+
184
+ with gr.Blocks(theme=gr.themes.Citrus(), css=CSS, title="amx-reasoning-v1 QA") as demo:
185
+ with gr.Column(elem_id="col-container"):
186
+ gr.Markdown(INTRO)
187
+
188
+ with gr.Row():
189
+ question = gr.Textbox(
190
+ label="Question",
191
+ placeholder="What year was the company founded?",
192
+ lines=1,
193
+ scale=4,
194
+ )
195
+ run = gr.Button("Answer", variant="primary", scale=1)
196
+
197
+ passage = gr.Textbox(
198
+ label="Passage",
199
+ placeholder="Paste the passage the answer should come from…",
200
+ lines=9,
201
+ )
202
+
203
+ answer_box = gr.Textbox(
204
+ label="Answer", elem_id="answer", lines=2, show_copy_button=True
205
+ )
206
+ info = gr.Markdown()
207
+
208
+ with gr.Accordion("Advanced", open=False):
209
+ max_new_tokens = gr.Slider(
210
+ 1, 64, value=DEFAULT_MAX_NEW_TOKENS, step=1,
211
+ label="Max new tokens",
212
+ info="The model normally emits one or two and stops on EOT.",
213
+ )
214
+ apply_vocab_mask = gr.Checkbox(
215
+ value=True,
216
+ label="Apply the vocabulary mask from generation.json",
217
+ info=(
218
+ "800 vocabulary rows never occurred in training, so their "
219
+ "logits were never pushed down and they win the argmax "
220
+ "whenever the model is unsure. Unchecking this reproduces "
221
+ "the ' ballo' failure mode from the model card."
222
+ ),
223
+ )
224
+
225
+ gr.Examples(
226
+ examples=[[e["question"], e["passage"]] for e in EXAMPLES],
227
+ example_labels=[
228
+ f"{e['source']} — {e['question']}" for e in EXAMPLES
229
+ ],
230
+ inputs=[question, passage],
231
+ outputs=[answer_box, info],
232
+ fn=answer,
233
+ cache_examples=True,
234
+ cache_mode="lazy",
235
+ )
236
+
237
+ gr.Markdown(
238
+ "Example passages come from the datasets the model card reports on: "
239
+ "[DROP](https://huggingface.co/datasets/ucinlp/drop) and "
240
+ "[SQuAD v2](https://huggingface.co/datasets/rajpurkar/squad_v2) "
241
+ "(CC BY-SA 4.0), and "
242
+ "[databricks-dolly-15k](https://huggingface.co/datasets/databricks/databricks-dolly-15k) "
243
+ "(CC BY-SA 3.0). The first example is the model repo's own "
244
+ "`example.py`."
245
+ )
246
+
247
+ inputs = [question, passage, max_new_tokens, apply_vocab_mask]
248
+ outputs = [answer_box, info]
249
+ run.click(answer, inputs=inputs, outputs=outputs, api_name="answer")
250
+ question.submit(answer, inputs=inputs, outputs=outputs, api_name=False)
251
+
252
+ if __name__ == "__main__":
253
+ demo.launch(mcp_server=True)
examples.json ADDED
@@ -0,0 +1,44 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [
2
+ {
3
+ "question": "In what year was the Royal Dutch Petroleum Company founded?",
4
+ "passage": "In February 1907, the Royal Dutch Shell Group was created through the amalgamation of two rival companies: the Royal Dutch Petroleum Company, founded in 1890, and the Shell Transport and Trading Company of the United Kingdom.",
5
+ "source": "the model's own example.py",
6
+ "reference": "1890"
7
+ },
8
+ {
9
+ "question": "Which field goals did Neil Rackers make?",
10
+ "passage": "The Texans' fifteenth game was an AFC duel with the Broncos. The Texans commanded the first half with RB Arian Foster getting a 3-yard TD run, followed by QB Matt Schaub getting a 3-yard TD pass to TE Owen Daniels, then with kicker Neil Rackers hitting a 34-yard field goal. The Broncos got on the board with RB Correll Buckhalter getting a 3-yard TD run, but the Texans scored again with Rackers nailing a 54-yard field goal. The Broncos replied as kicker Steven Hauschka got a 27-yard field goal, but the Texans extended their lead with Rackers hitting a 57-yard field goal. However, they failed to maintain this lead after QB Tim Tebow completed a 23-yard TD pass to Buckhalter, followed by Tebow scrambling 6-yards for a touchdown.",
11
+ "source": "DROP",
12
+ "reference": "34-yard, 54-yard, 57-yard"
13
+ },
14
+ {
15
+ "question": "How many years did Manipur raid the Upper Chindwin region?",
16
+ "passage": "Manipur was a tributary to Burma in the 16th century but had gone its own way since. It raided Upper Chindwin region in 1647 and 1692. However, in 1704, the raja of Manipur presented his daughter to Ava. Starting in the 1720s, Manipur under the leadership of Pamheiba became a thorn to Upper Burma. In early 1724, the Manipuris raided Upper Burma. In response, an expedition force of 3,000 men marched to Manipur in November 1724. The army was ambushed in the swamps at Heirok, and retreated in haste. The Manipuris then returned ten years later. From 1735 to 1741, Manipuris raided the Upper Chindwin regions, increasingly deeper with each raid. Burmese defences were simply bypassed the Manipuris on their horseback. In December 1739, they reached as far as Sagaing, and looted and burned everything insight. The Burmese defences finally stopped them at Myedu in early 1741, with each side agreeing to an uneasy truce. But the Manipuris had annexed the Kabaw valley. The truce did not last. Another raid came all the way down to Ava in 1744. The last raid came in 1749. Upon arrival at Ava, the Manipuri chief found a large Burmese army, and presented his 12-year-old daughter instead, and left.",
17
+ "source": "DROP",
18
+ "reference": "46"
19
+ },
20
+ {
21
+ "question": "Where did researchers study chimps in heavily forested regions?",
22
+ "passage": "In 2006–07, researchers from the Wildlife Conservation Society studied gorillas in heavily forested regions centered on the Ouesso district of the Sangha Region. They suggest a population on the order of 125,000 Western Lowland Gorillas, whose isolation from humans has been largely preserved by inhospitable swamps.",
23
+ "source": "SQuAD v2 (unanswerable)",
24
+ "reference": "The passage does not say."
25
+ },
26
+ {
27
+ "question": "Which species avoided the polar areas?",
28
+ "passage": "The Early Cretaceous spans from 145 million to 100 million years ago. The Early Cretaceous saw the expansion of seaways, and as a result, the decline and extinction of sauropods (except in South America). Many coastal shallows were created, and that caused Ichthyosaurs to die out. Mosasaurs evolved to replace them as head of the seas. Some island-hopping dinosaurs, like Eustreptospondylus, evolved to cope with the coastal shallows and small islands of ancient Europe. Other dinosaurs rose up to fill the empty space that the Jurassic-Cretaceous extinction left behind, such as Carcharodontosaurus and Spinosaurus. Of the most successful would be the Iguanodon which spread to every continent. Seasons came back into effect and the poles got seasonally colder, but dinosaurs still inhabited this area like the Leaellynasaura which inhabited the polar forests year-round, and many dinosaurs migrated there during summer like Muttaburrasaurus. Since it was too cold for crocodiles, it was the last stronghold for large amphibians, like Koolasuchus. Pterosaurs got larger as species like Tapejara and Ornithocheirus evolved.",
29
+ "source": "SQuAD v2 (unanswerable)",
30
+ "reference": "The passage does not say."
31
+ },
32
+ {
33
+ "question": "What is the Kentucky Derby Trophy",
34
+ "passage": "The Kentucky Derby Trophy is a set of four trophies that are awarded to the winning connections of America's most famous race: the grade one $3,000,000 Kentucky Derby. The owner receives a gold trophy while the trainer, the jockey and the breeder win a silver half size replica of the main gold trophy. The trophy itself has been run for since the 50th running of the Kentucky Derby in 1924. Churchill Downs Race Course of Louisville, Kentucky has annually presented a gold trophy to the winning owner of the famed \"Run for the Roses.\"",
35
+ "source": "Dolly 15k",
36
+ "reference": "The Kentucky Derby Trophy is a set of four trophies that are awarded to the winning connections of America's most famous race: the grade one $3,000,000 Kentucky Derby. The owner receives a gold trophy while the trainer, the jockey and the breeder win a silver half size replica of the main gold trophy. The trophy itself has been run for since the 50th running of the Kentucky Derby in 1924. Churchill Downs Race Course of Louisville, Kentucky has annually presented a gold trophy to the winning owner of the famed \"Run for the Roses.\""
37
+ },
38
+ {
39
+ "question": "From the passage note down the name of the countries which have most voting power. List the results in comma separated format.",
40
+ "passage": "The World Bank is an international financial institution that provides loans and grants to the governments of low- and middle-income countries for the purpose of pursuing capital projects. The World Bank is the collective name for the International Bank for Reconstruction and Development (IBRD) and International Development Association (IDA), two of five international organizations owned by the World Bank Group. It was established along with the International Monetary Fund at the 1944 Bretton Woods Conference. After a slow start, its first loan was to France in 1947. In the 1970s, it focused on loans to developing world countries, shifting away from that mission in the 1980s. For the last 30 years, it has included NGOs and environmental groups in its loan portfolio. Its loan strategy is influenced by the United Nations' Sustainable Development Goals, as well as environmental and social safeguards.\n\nAs of 2022, the World Bank is run by a president and 25 executive directors, as well as 29 various vice presidents. IBRD and IDA have 189 and 174 member countries, respectively. The U.S., Japan, China, Germany and the U.K. have the most voting power. The bank aims loans at developing countries to help reduce poverty. The bank is engaged in several global partnerships and initiatives, and takes a role in working toward addressing climate change. The World Bank operates a number of training wings and it works with the Clean Air Initiative and the UN Development Business. It works within the Open Data Initiative and hosts an Open Knowledge Repository.\n\nThe World Bank has been criticized as promoting inflation and harming economic development, causing protests in 1988 and 2000. There has also been criticism of the bank's governance and response to the COVID-19 pandemic.",
41
+ "source": "Dolly 15k",
42
+ "reference": "U.S., Japan, China, Germany, U.K."
43
+ }
44
+ ]
requirements.txt ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ torch
2
+ safetensors
3
+ tokenizers
4
+ pyyaml