Skip to main content

taconite_qwen35/
lib.rs

1// SPDX-FileCopyrightText: Copyright (C) 2026 Brishen Hawkins
2// SPDX-License-Identifier: Apache-2.0
3
4//! Qwen3.5-2B (`Qwen/Qwen3.5-2B`, its text model) chatting on an AMD XDNA
5//! NPU.
6//!
7//! The model is a hybrid: 18 Gated DeltaNet (linear attention) layers and
8//! 6 gated softmax-attention layers, each followed by a SwiGLU MLP. The
9//! bundle `iron/applications/qwen3_5/export_qwen35.py` writes holds every
10//! compiled IRON kernel, the weights (the NPU's pre-packed) and the
11//! tokenizer; this crate replays the forward the Python app runs:
12//!
13//! | | NPU | host (here) |
14//! |---|---|---|
15//! | prompt | every projection as an `flm.GEMM` over 256-row chunks, one hardware context | tokenizer, embedding, norms, the DeltaNet's conv + recurrence, RoPE, attention |
16//! | a generated token | every projection and the LM head as a `GEMVbfp16`, a second context | the same, one row |
17//! | an image (the vision tower) | every projection as an `flm.GEMM` (K = 1024), a third context | preprocessing (`image.rs`), LayerNorms, 2D RoPE, attention |
18//!
19//! [`Qwen35::chat`] answers one user turn (greedy), with or without
20//! images; [`Qwen35::check`] verifies a bundle against the references its
21//! exporter recorded.
22
23use std::fmt;
24use std::path::Path;
25use std::time::Instant;
26
27pub use taconite::Timing;
28use taconite_bundle::{Manifest, Store};
29
30pub mod image;
31pub mod model;
32pub mod npu;
33pub mod tokenizer;
34pub mod vision;
35
36use model::{ImageInput, Model};
37use tokenizer::Tokenizer;
38
39pub const VERSION: u32 = 1;
40
41#[cfg(not(any(feature = "xrt", feature = "direct")))]
42compile_error!("taconite-qwen35 needs the `xrt` or the `direct` feature");
43
44#[derive(Debug)]
45pub enum Error {
46    Bundle(String),
47    Npu(String),
48    Input(String),
49}
50
51impl fmt::Display for Error {
52    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
53        match self {
54            Error::Bundle(m) => write!(f, "bundle: {m}"),
55            Error::Npu(m) => write!(f, "NPU: {m}"),
56            Error::Input(m) => write!(f, "input: {m}"),
57        }
58    }
59}
60
61impl std::error::Error for Error {}
62
63impl From<taconite_bundle::Error> for Error {
64    fn from(e: taconite_bundle::Error) -> Self {
65        Error::Bundle(e.to_string())
66    }
67}
68
69impl From<taconite::Error> for Error {
70    fn from(e: taconite::Error) -> Self {
71        Error::Npu(e.to_string())
72    }
73}
74
75/// An 8-bit RGB image, rows top to bottom, `[height, width, 3]`.
76#[derive(Debug, Clone)]
77pub struct RgbImage {
78    pub width: usize,
79    pub height: usize,
80    pub rgb: Vec<u8>,
81}
82
83/// How to answer a turn.
84#[derive(Debug, Clone)]
85pub struct ChatOptions {
86    pub system: Option<String>,
87    /// let the model reason in a `<think>` block first
88    pub thinking: bool,
89    pub max_new: usize,
90    /// images the turn shows, before its text
91    pub images: Vec<RgbImage>,
92}
93
94impl Default for ChatOptions {
95    fn default() -> Self {
96        ChatOptions { system: None, thinking: false, max_new: 512, images: Vec::new() }
97    }
98}
99
100/// What a generation did.
101#[derive(Debug, Clone, Default)]
102pub struct Stats {
103    /// the vision tower's time over the turn's images
104    pub vision_s: f64,
105    pub prompt_tokens: usize,
106    pub new_tokens: usize,
107    pub prefill_s: f64,
108    pub decode_s: f64,
109}
110
111impl Stats {
112    /// Tokens a second after the first (which the prefill produces).
113    pub fn decode_tok_s(&self) -> f64 {
114        if self.new_tokens > 1 && self.decode_s > 0.0 { (self.new_tokens - 1) as f64 / self.decode_s } else { 0.0 }
115    }
116}
117
118pub struct Qwen35 {
119    pub model: Model,
120    pub tok: Tokenizer,
121    manifest: Manifest,
122    store: Store,
123    pub timing: Timing,
124}
125
126fn argmax(x: &[f32]) -> u32 {
127    let mut best = 0;
128    for (i, &v) in x.iter().enumerate() {
129        if v > x[best] {
130            best = i;
131        }
132    }
133    best as u32
134}
135
136fn unhex(s: &str) -> Result<String, Error> {
137    if s == "-" {
138        return Ok(String::new());
139    }
140    let bytes: Option<Vec<u8>> =
141        (0..s.len()).step_by(2).map(|i| s.get(i..i + 2).and_then(|h| u8::from_str_radix(h, 16).ok())).collect();
142    bytes.and_then(|b| String::from_utf8(b).ok()).ok_or_else(|| Error::Bundle(format!("bad hex string {s}")))
143}
144
145/// Streams text from token bytes: holds back an incomplete UTF-8 sequence
146/// until the token that completes it.
147#[derive(Default)]
148struct Utf8Stream {
149    pending: Vec<u8>,
150}
151
152impl Utf8Stream {
153    fn push(&mut self, bytes: &[u8]) -> String {
154        self.pending.extend_from_slice(bytes);
155        let valid = match std::str::from_utf8(&self.pending) {
156            Ok(_) => self.pending.len(),
157            // an invalid sequence (not just a truncated one) goes out
158            // replaced, as the final decode would show it
159            Err(e) if e.error_len().is_some() => self.pending.len(),
160            Err(e) => e.valid_up_to(),
161        };
162        let out = String::from_utf8_lossy(&self.pending[..valid]).into_owned();
163        self.pending.drain(..valid);
164        out
165    }
166}
167
168fn greedy(
169    model: &mut Model,
170    store: &Store,
171    timing: &mut Timing,
172    ids: &[u32],
173    image: Option<&ImageInput>,
174    max_new: usize,
175    mut on_token: impl FnMut(u32),
176) -> Result<(Vec<u32>, Stats), Error> {
177    let t0 = Instant::now();
178    let mut logits = model.prefill(store, ids, image, timing)?;
179    let prefill_s = t0.elapsed().as_secs_f64();
180    let t1 = Instant::now();
181    let budget = max_new.min(model.max_ctx.saturating_sub(ids.len()));
182    let mut out = Vec::new();
183    while out.len() < budget {
184        let t = argmax(&logits);
185        if model.cfg.stop.contains(&t) {
186            break;
187        }
188        out.push(t);
189        on_token(t);
190        if out.len() == budget {
191            break;
192        }
193        logits = model.decode(store, t, timing)?;
194    }
195    let stats = Stats {
196        vision_s: 0.0,
197        prompt_tokens: ids.len(),
198        new_tokens: out.len(),
199        prefill_s,
200        decode_s: t1.elapsed().as_secs_f64(),
201    };
202    Ok((out, stats))
203}
204
205impl Qwen35 {
206    /// Loads a bundle: every kernel into the NPU, the weights into device
207    /// buffers. `max_ctx` bounds prompt + answer (the KV cache's length).
208    pub fn load(dir: &Path, max_ctx: usize) -> Result<Self, Error> {
209        let manifest = Manifest::load(dir, VERSION)?;
210        let store = Store::load(dir)?;
211        let json = std::fs::read_to_string(dir.join("tokenizer.json"))
212            .map_err(|e| Error::Bundle(format!("tokenizer.json: {e}")))?;
213        let tok = Tokenizer::from_hf_json(&json).map_err(Error::Bundle)?;
214        let model = Model::load(&manifest, &store, max_ctx)?;
215        Ok(Qwen35 { model, tok, manifest, store, timing: Timing::default() })
216    }
217
218    /// One user turn through the checkpoint's chat template, ready for
219    /// the answer. `image_tokens`: each shown image's token count (they
220    /// come before the text, `<|image_pad|>` repeated that often).
221    pub fn prompt_ids(&self, message: &str, system: Option<&str>, thinking: bool, image_tokens: &[usize]) -> Vec<u32> {
222        let mut text = String::new();
223        if let Some(s) = system {
224            text += &format!("<|im_start|>system\n{}<|im_end|>\n", s.trim());
225        }
226        let images: String = image_tokens
227            .iter()
228            .map(|&n| format!("<|vision_start|>{}<|vision_end|>", "<|image_pad|>".repeat(n)))
229            .collect();
230        let content = format!("{images}{message}");
231        text += &format!("<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n<think>\n", content.trim());
232        if !thinking {
233            text += "\n</think>\n\n";
234        }
235        self.tok.encode(&text)
236    }
237
238    /// Greedy continuation of `ids` up to a stop token or `max_new`
239    /// tokens; `on_token` sees each new token as it is chosen.
240    pub fn generate(
241        &mut self,
242        ids: &[u32],
243        max_new: usize,
244        on_token: impl FnMut(u32),
245    ) -> Result<(Vec<u32>, Stats), Error> {
246        greedy(&mut self.model, &self.store, &mut self.timing, ids, None, max_new, on_token)
247    }
248
249    /// An image's patches and grid, preprocessed as the checkpoint's
250    /// processor does (see [`image`]).
251    pub fn preprocess(&self, img: &RgbImage) -> Result<(Vec<f32>, (usize, usize)), Error> {
252        let v = self.model.vision.as_ref().ok_or_else(|| Error::Input("the bundle has no vision tower".into()))?;
253        let c = &v.cfg;
254        if img.rgb.len() != img.width * img.height * 3 {
255            return Err(Error::Input("image data is not width x height x 3 bytes".into()));
256        }
257        Ok(image::preprocess(&img.rgb, img.width, img.height, c.patch, c.merge, c.min_pixels, c.max_pixels))
258    }
259
260    /// Images -> (their token embeddings, concatenated; each one's grid).
261    fn encode(&mut self, images: &[RgbImage]) -> Result<(Vec<f32>, Vec<(usize, usize)>), Error> {
262        let (mut emb, mut grids) = (Vec::new(), Vec::new());
263        for img in images {
264            let (patches, grid) = self.preprocess(img)?;
265            emb.extend(self.model.encode_image(&patches, grid, &mut self.timing)?);
266            grids.push(grid);
267        }
268        Ok((emb, grids))
269    }
270
271    /// The prefill inputs of a prompt with images: M-RoPE positions.
272    fn image_input<'a>(
273        &self,
274        ids: &[u32],
275        emb: &'a [f32],
276        grids: &[(usize, usize)],
277        pos: &'a mut Vec<[usize; 3]>,
278    ) -> Result<ImageInput<'a>, Error> {
279        let v = self.model.vision.as_ref().ok_or_else(|| Error::Input("the bundle has no vision tower".into()))?;
280        let (p, next) = model::positions(ids, grids, v.cfg.image_token, v.cfg.merge)?;
281        *pos = p;
282        Ok(ImageInput { image_token: v.cfg.image_token, embeddings: emb, positions: pos, next })
283    }
284
285    /// Answers `message`; `stream` gets the answer's text as it grows.
286    pub fn chat(
287        &mut self,
288        message: &str,
289        opts: &ChatOptions,
290        mut stream: impl FnMut(&str),
291    ) -> Result<(String, Stats), Error> {
292        let t0 = Instant::now();
293        let (emb, grids) = self.encode(&opts.images)?;
294        let vision_s = t0.elapsed().as_secs_f64();
295        let merge2 = self.model.vision.as_ref().map_or(1, |v| v.cfg.merge * v.cfg.merge);
296        let counts: Vec<usize> = grids.iter().map(|g| g.0 * g.1 / merge2).collect();
297        let ids = self.prompt_ids(message, opts.system.as_deref(), opts.thinking, &counts);
298        let mut pos = Vec::new();
299        let image = if grids.is_empty() { None } else { Some(self.image_input(&ids, &emb, &grids, &mut pos)?) };
300        let mut utf8 = Utf8Stream::default();
301        let tok = &self.tok;
302        let (out, mut stats) =
303            greedy(&mut self.model, &self.store, &mut self.timing, &ids, image.as_ref(), opts.max_new, |t| {
304                let bytes = tok.decode_bytes(&[t], true);
305                let s = utf8.push(&bytes);
306                if !s.is_empty() {
307                    stream(&s);
308                }
309            })?;
310        stats.vision_s = vision_s;
311        Ok((self.tok.decode(&out, true), stats))
312    }
313
314    /// Checks the bundle against what its exporter recorded: the
315    /// tokenizer on its test strings, then each reference run -- its
316    /// prompt through the chat template, every step teacher-forced on the
317    /// reference's tokens (the next token must agree wherever the float32
318    /// reference's top two logits are 0.5 or more apart), and a free
319    /// greedy run. Prints a report; returns whether everything passed.
320    pub fn check(&mut self, mut say: impl FnMut(&str)) -> Result<bool, Error> {
321        let mut ok = true;
322        let cases: Vec<(String, Vec<u32>)> = self
323            .manifest
324            .tagged("tokcase")
325            .map(|r| {
326                Ok((
327                    unhex(r.field(0)?)?,
328                    r.field(1)?.split(',').filter(|s| !s.is_empty()).map(|s| s.parse().unwrap_or(u32::MAX)).collect(),
329                ))
330            })
331            .collect::<Result<_, Error>>()?;
332        let bad: Vec<&String> = cases.iter().filter(|(t, ids)| &self.tok.encode(t) != ids).map(|(t, _)| t).collect();
333        say(&format!("tokenizer: {}/{} strings encode as HF does", cases.len() - bad.len(), cases.len()));
334        for t in &bad {
335            say(&format!("  MISMATCH {t:?}: {:?}", self.tok.encode(t)));
336        }
337        ok &= bad.is_empty();
338
339        let refs: Vec<(String, String, bool, bool)> = self
340            .manifest
341            .tagged("ref")
342            .map(|r| {
343                let image = r.has("image") && r.get::<u8>("image")? == 1;
344                Ok((r.field(0)?.to_string(), unhex(r.str("prompt")?)?, r.get::<u8>("thinking")? == 1, image))
345            })
346            .collect::<Result<_, Error>>()?;
347        for (name, prompt, thinking, has_image) in refs {
348            let p = |k: &str| format!("ref.{name}.{k}");
349            let ids: Vec<u32> = self.store.i32(&p("ids"))?.iter().map(|&x| x as u32).collect();
350            let toks: Vec<u32> = self.store.i32(&p("tokens"))?.iter().map(|&x| x as u32).collect();
351            let top_ids = self.store.i32(&p("top_ids"))?.to_vec();
352            let top = self.store.f32(&p("top_logits"))?.to_vec();
353            let k = top.len() / toks.len();
354            self.timing.clear();
355            let (mut emb, mut grids, mut counts) = (Vec::new(), Vec::new(), Vec::new());
356            if has_image {
357                // the exporter's decoded pixels -> our patches (against the
358                // app's) -> the vision tower (against the float32 model's)
359                let shape = self.store.shape(&p("rgb"))?.to_vec();
360                let img = RgbImage { height: shape[0], width: shape[1], rgb: self.store.u8(&p("rgb"))?.to_vec() };
361                let (patches, grid) = self.preprocess(&img)?;
362                let want = self.store.f32(&p("patches"))?;
363                let dp = patches.iter().zip(want).fold(0f32, |m, (a, b)| m.max((a - b).abs()));
364                let want_grid = self.store.i32(&p("grid"))?;
365                let same_grid = grid == (want_grid[0] as usize, want_grid[1] as usize) && patches.len() == want.len();
366                let t0 = Instant::now();
367                emb = self.model.encode_image(&patches, grid, &mut self.timing)?;
368                let vs = t0.elapsed().as_secs_f64();
369                let r = self.store.f32(&p("vision"))?;
370                let (mut dot, mut nn, mut rr, mut dd) = (0f64, 0f64, 0f64, 0f64);
371                for (&a, &b) in emb.iter().zip(r) {
372                    let (a, b) = (a as f64, b as f64);
373                    dot += a * b;
374                    nn += a * a;
375                    rr += b * b;
376                    dd += (a - b) * (a - b);
377                }
378                let cos = dot / (nn.sqrt() * rr.sqrt());
379                say(&format!(
380                    "{name} image: {}x{} -> grid {grid:?} ({}), patches max |d| vs the app's {dp:.2e}; vision vs float32: cosine {cos:.4}, relative RMS {:.3}; {vs:.2} s",
381                    img.width,
382                    img.height,
383                    if same_grid { "matches" } else { "MISMATCH" },
384                    (dd / rr).sqrt(),
385                ));
386                ok &= same_grid && dp < 1e-6 && cos > 0.99 && emb.len() == r.len();
387                counts.push(grid.0 * grid.1 / 4);
388                grids.push(grid);
389            }
390            let same_prompt = self.prompt_ids(&prompt, None, thinking, &counts) == ids;
391            ok &= same_prompt;
392            let mut pos = Vec::new();
393            let image = if has_image { Some(self.image_input(&ids, &emb, &grids, &mut pos)?) } else { None };
394
395            let t0 = Instant::now();
396            let mut logits = self.model.prefill(&self.store, &ids, image.as_ref(), &mut self.timing)?;
397            let prefill_s = t0.elapsed().as_secs_f64();
398            let (mut agree, mut flips, mut real, mut max_d) = (0, Vec::new(), 0, 0f32);
399            let t1 = Instant::now();
400            for (j, &want) in toks.iter().enumerate() {
401                let (ti, tl) = (&top_ids[j * k..(j + 1) * k], &top[j * k..(j + 1) * k]);
402                for (&id, &l) in ti.iter().zip(tl) {
403                    max_d = max_d.max((logits[id as usize] - l).abs());
404                }
405                if argmax(&logits) == want {
406                    agree += 1;
407                } else {
408                    let margin = tl[0] - tl[1];
409                    flips.push(format!("{j} ({margin:.3})"));
410                    if margin > 0.5 {
411                        real += 1;
412                    }
413                }
414                if j + 1 < toks.len() {
415                    logits = self.model.decode(&self.store, want, &mut self.timing)?;
416                }
417            }
418            let dec_ms = t1.elapsed().as_secs_f64() * 1e3 / toks.len().saturating_sub(1).max(1) as f64;
419            say(&format!(
420                "{name}: {} prompt tokens (template {}), {} steps: top-1 agrees {agree}/{}, max |dlogit| over the reference's top {k} {max_d:.3}, flips (ref margin) [{}]; prefill {prefill_s:.2} s, {dec_ms:.0} ms/token",
421                ids.len(),
422                if same_prompt { "matches" } else { "MISMATCH" },
423                toks.len(),
424                toks.len(),
425                flips.join(", "),
426            ));
427            ok &= real == 0;
428            let (out, stats) =
429                greedy(&mut self.model, &self.store, &mut self.timing, &ids, image.as_ref(), toks.len(), |_| {})?;
430            let same = out.iter().zip(&toks).take_while(|(a, b)| a == b).count();
431            say(&format!(
432                "  free run: first {same} of {} tokens as the reference's, {:.1} tokens/s: {:?}",
433                toks.len(),
434                stats.decode_tok_s(),
435                self.tok.decode(&out, true)
436            ));
437        }
438        say(if ok { "PASS" } else { "FAIL" });
439        Ok(ok)
440    }
441}