Skip to main content

qualia_core_db/gguf_bridge/
forward.rs

1//! Forward-pass orchestration: prefill (batch/chunk, sync + async + MC8-fused), transformer
2//! layer/forward (sync + async), encode helpers, topology-draft verify, and async lifecycle.
3//! Split from gguf_bridge/mod.rs (structural refactor; no behaviour change).
4use super::*;
5
6impl QTensorEngine {
7    /// One transformer layer for a batched prefill chunk: batched K/V then per-token Q+FFN.
8    pub(crate) fn dispatch_prefill_layer_batch(
9        &mut self,
10        index: &crate::gguf_sharder::GgufTensorIndex,
11        layer: u32,
12        batch_hidden: &mut [f32],
13        emb_dim: usize,
14        n_tokens: u32,
15        batch_start_token_idx: u32,
16        scratch_a: &mut [f32],
17        scratch_b: &mut [f32],
18    ) -> bool {
19        if n_tokens == 0 {
20            wlog("[prefill_layer] FAILED n_tokens=0");
21            return false;
22        }
23        let layout = match self.kv_layout {
24            Some(l) => l,
25            None => {
26                wlog("[prefill_layer] FAILED kv_layout is None");
27                return false;
28            }
29        };
30        let tensors = index.get_layer_tensors(layer);
31        let k_info = match tensors.attn_k.as_ref() {
32            Some(i) => i,
33            None => {
34                wlog(&format!(
35                    "[prefill_layer] FAILED missing attn_k layer={layer}"
36                ));
37                return false;
38            }
39        };
40        let v_info = match tensors.attn_v.as_ref() {
41            Some(i) => i,
42            None => {
43                wlog(&format!(
44                    "[prefill_layer] FAILED missing attn_v layer={layer}"
45                ));
46                return false;
47            }
48        };
49        if tensors.attn_q.is_none() {
50            wlog(&format!(
51                "[prefill_layer] FAILED missing attn_q layer={layer}"
52            ));
53            return false;
54        }
55        let h = index.hyperparams;
56        let n_kv = h.effective_n_kv_head();
57        let n_embd = h.n_embd as usize;
58        let batch_elems = n_embd * n_tokens as usize;
59        if batch_elems > batch_hidden.len() {
60            wlog(&format!(
61                "[prefill_layer] FAILED batch_elems OOB elems={batch_elems} hidden={}",
62                batch_hidden.len()
63            ));
64            return false;
65        }
66        let mmap = match self.gguf_mmap.as_deref() {
67            Some(m) => m,
68            None => {
69                wlog("[prefill_layer] FAILED gguf_mmap is None");
70                return false;
71            }
72        };
73        let k_raw =
74            match crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, k_info) {
75                Ok(s) => s,
76                Err(e) => {
77                    wlog(&format!("[prefill_layer] FAILED fetch attn_k bytes: {e:?}"));
78                    return false;
79                }
80            };
81        let v_raw =
82            match crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, v_info) {
83                Ok(s) => s,
84                Err(e) => {
85                    wlog(&format!("[prefill_layer] FAILED fetch attn_v bytes: {e:?}"));
86                    return false;
87                }
88            };
89        let n_kv_wg = n_tokens.saturating_mul(n_kv);
90        // attn_norm MUST be applied to the K/V projection input on ALL targets — native previously
91        // passed None here → prefill wrote K/V from the RAW residual → the KV cache exploded across
92        // layers → the whole forward (and decode reading it) blew up. (#48)
93        let mut norm_w_attn = [0f32; MAX_HIDDEN_DIM];
94        let norm_weight_attn: Option<&[f32]> = tensors.attn_norm.as_ref().and_then(|info| {
95            let n = dequant_norm_row_into(mmap, index.tensor_data_start, info, &mut norm_w_attn);
96            if n >= n_embd {
97                Some(&norm_w_attn[..n_embd])
98            } else {
99                None
100            }
101        });
102        if !self.dispatch_attention_pass(
103            &batch_hidden[..batch_elems],
104            n_embd,
105            n_tokens,
106            batch_start_token_idx,
107            &layout,
108            layer,
109            batch_start_token_idx,
110            &h,
111            k_info,
112            k_raw,
113            1,
114            n_kv_wg,
115            norm_weight_attn,
116            None,
117        ) {
118            wlog(&format!("[prefill_layer] K pass FAILED layer={layer}"));
119            return false;
120        }
121        if !self.dispatch_attention_pass(
122            &batch_hidden[..batch_elems],
123            n_embd,
124            n_tokens,
125            batch_start_token_idx,
126            &layout,
127            layer,
128            batch_start_token_idx,
129            &h,
130            v_info,
131            v_raw,
132            2,
133            n_kv_wg,
134            norm_weight_attn,
135            None,
136        ) {
137            wlog(&format!("[prefill_layer] V pass FAILED layer={layer}"));
138            return false;
139        }
140        for t in 0..n_tokens {
141            let abs = batch_start_token_idx + t;
142            let off = t as usize * emb_dim;
143            if !self.dispatch_attention_q_ffn_token(
144                index,
145                layer,
146                abs,
147                &mut batch_hidden[off..off + emb_dim],
148                emb_dim,
149                &tensors,
150                scratch_a,
151                scratch_b,
152            ) {
153                wlog(&format!(
154                    "[prefill_layer] q_ffn FAILED layer={layer} t={t} abs={abs}"
155                ));
156                return false;
157            }
158        }
159        true
160    }
161
162    /// Phase 2B: batched prefill layer via async GPU attention (K/V GPU; Q+FFN per token).
163    #[cfg(all(target_arch = "wasm32", feature = "wasm-llm-diagnostics"))]
164    pub(crate) async fn dispatch_prefill_layer_batch_async(
165        &mut self,
166        index: &crate::gguf_sharder::GgufTensorIndex,
167        layer: u32,
168        batch_hidden: &mut [f32],
169        emb_dim: usize,
170        n_tokens: u32,
171        batch_start_token_idx: u32,
172        scratch_a: &mut [f32],
173        scratch_b: &mut [f32],
174    ) -> bool {
175        if n_tokens == 0 {
176            wlog("[prefill_layer] FAILED n_tokens=0");
177            return false;
178        }
179        let layout = match self.kv_layout {
180            Some(l) => l,
181            None => {
182                wlog("[prefill_layer] FAILED kv_layout is None");
183                return false;
184            }
185        };
186        let tensors = index.get_layer_tensors(layer);
187        let k_info = match tensors.attn_k.as_ref() {
188            Some(i) => i,
189            None => {
190                wlog(&format!(
191                    "[prefill_layer] FAILED missing attn_k layer={layer}"
192                ));
193                return false;
194            }
195        };
196        let v_info = match tensors.attn_v.as_ref() {
197            Some(i) => i,
198            None => {
199                wlog(&format!(
200                    "[prefill_layer] FAILED missing attn_v layer={layer}"
201                ));
202                return false;
203            }
204        };
205        if tensors.attn_q.is_none() {
206            wlog(&format!(
207                "[prefill_layer] FAILED missing attn_q layer={layer}"
208            ));
209            return false;
210        }
211        let h = index.hyperparams;
212        let n_kv = h.effective_n_kv_head();
213        let n_embd = h.n_embd as usize;
214        let batch_elems = n_embd * n_tokens as usize;
215        if batch_elems > batch_hidden.len() {
216            wlog(&format!(
217                "[prefill_layer] FAILED batch_elems OOB elems={batch_elems} hidden={}",
218                batch_hidden.len()
219            ));
220            return false;
221        }
222        let mmap = match self.gguf_mmap.as_deref() {
223            Some(m) => m,
224            None => {
225                wlog("[prefill_layer] FAILED gguf_mmap is None");
226                return false;
227            }
228        };
229        let k_raw =
230            match crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, k_info) {
231                Ok(s) => s,
232                Err(e) => {
233                    wlog(&format!("[prefill_layer] FAILED fetch attn_k bytes: {e:?}"));
234                    return false;
235                }
236            };
237        let v_raw =
238            match crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, v_info) {
239                Ok(s) => s,
240                Err(e) => {
241                    wlog(&format!("[prefill_layer] FAILED fetch attn_v bytes: {e:?}"));
242                    return false;
243                }
244            };
245        let mut norm_w_attn = [0f32; MAX_HIDDEN_DIM];
246        let mut norm_scratch = [0f32; PREFILL_CHUNK_STACK_FLOATS];
247        let attn_input: &mut [f32] = if let Some(norm_info) = tensors.attn_norm.as_ref() {
248            let n =
249                dequant_norm_row_into(mmap, index.tensor_data_start, norm_info, &mut norm_w_attn);
250            if n >= n_embd {
251                for t in 0..n_tokens as usize {
252                    let off = t * n_embd;
253                    norm_scratch[off..off + n_embd]
254                        .copy_from_slice(&batch_hidden[off..off + n_embd]);
255                    rms_norm_inplace(
256                        &mut norm_scratch[off..off + n_embd],
257                        &norm_w_attn[..n_embd],
258                        RMS_NORM_EPS,
259                    );
260                }
261                &mut norm_scratch[..batch_elems]
262            } else {
263                batch_hidden
264            }
265        } else {
266            batch_hidden
267        };
268        let n_kv_wg = n_tokens.saturating_mul(n_kv);
269        if !self
270            .dispatch_attention_pass_async(
271                attn_input,
272                n_embd,
273                n_tokens,
274                batch_start_token_idx,
275                &layout,
276                layer,
277                batch_start_token_idx,
278                &h,
279                k_info,
280                k_raw,
281                1,
282                n_kv_wg,
283                None,
284            )
285            .await
286        {
287            wlog(&format!("[prefill_layer] K pass FAILED layer={layer}"));
288            return false;
289        }
290        if !self
291            .dispatch_attention_pass_async(
292                attn_input,
293                n_embd,
294                n_tokens,
295                batch_start_token_idx,
296                &layout,
297                layer,
298                batch_start_token_idx,
299                &h,
300                v_info,
301                v_raw,
302                2,
303                n_kv_wg,
304                None,
305            )
306            .await
307        {
308            wlog(&format!("[prefill_layer] V pass FAILED layer={layer}"));
309            return false;
310        }
311        for t in 0..n_tokens {
312            let abs = batch_start_token_idx + t;
313            let off = t as usize * emb_dim;
314            if !self
315                .dispatch_attention_q_ffn_token_async(
316                    index,
317                    layer,
318                    abs,
319                    &mut batch_hidden[off..off + emb_dim],
320                    emb_dim,
321                    &tensors,
322                    scratch_a,
323                    scratch_b,
324                )
325                .await
326            {
327                wlog(&format!(
328                    "[prefill_layer] q_ffn FAILED layer={layer} t={t} abs={abs}"
329                ));
330                return false;
331            }
332        }
333        true
334    }
335
336    /// Chunked prefill: populate KV arena for `n_tokens` prompt positions starting at `batch_start`.
337    pub fn dispatch_prefill_chunk(
338        &mut self,
339        index: &crate::gguf_sharder::GgufTensorIndex,
340        batch_hidden: &mut [f32],
341        emb_dim: usize,
342        n_tokens: u32,
343        batch_start_token_idx: u32,
344        scratch_a: &mut [f32],
345        scratch_b: &mut [f32],
346        max_layers: u32,
347    ) -> bool {
348        let n_layer = index.hyperparams.n_layer;
349        if n_layer == 0 || n_tokens == 0 {
350            return false;
351        }
352        // W3: resident single-fence-per-chunk arena (toggle-gated, default OFF). Populates the KV
353        // cache for the whole chunk in ONE submit; any ineligibility falls back to the legacy loop.
354        #[cfg(not(target_arch = "wasm32"))]
355        if crate::llm_bench::resident_prefill_enabled() {
356            if self
357                .dispatch_prefill_chunk_resident(
358                    index,
359                    &batch_hidden[..],
360                    emb_dim,
361                    n_tokens,
362                    batch_start_token_idx,
363                    max_layers,
364                )
365                .is_some()
366            {
367                crate::llm_bench::record_resident_prefill_hit();
368                return true;
369            }
370            crate::llm_bench::record_resident_prefill_fallback();
371        }
372        let limit = if max_layers == 0 {
373            n_layer
374        } else {
375            max_layers.min(n_layer)
376        };
377        for layer in 0..limit {
378            if !self.dispatch_prefill_layer_batch(
379                index,
380                layer,
381                batch_hidden,
382                emb_dim,
383                n_tokens,
384                batch_start_token_idx,
385                scratch_a,
386                scratch_b,
387            ) {
388                return false;
389            }
390        }
391        true
392    }
393
394    /// One transformer block using real mmap tensor offsets (stack buffers only).
395    pub fn dispatch_transformer_layer(
396        &mut self,
397        index: &crate::gguf_sharder::GgufTensorIndex,
398        layer: u32,
399        token_idx: u32,
400        hidden: &mut [f32],
401        emb_dim: usize,
402        scratch_a: &mut [f32],
403        scratch_b: &mut [f32],
404    ) -> bool {
405        let tensors = index.get_layer_tensors(layer);
406        let mut attn_ok = false;
407        // Decode-profiler: intra-layer split — attention (this block) vs FFN (below).
408        // Native-only: `llm_bench` is `#[cfg(not(wasm32))]`; the wasm-full LLM bundle skips profiling.
409        #[cfg(not(target_arch = "wasm32"))]
410        let t_attn = std::time::Instant::now();
411
412        if tensors.attn_q.is_some() && tensors.attn_k.is_some() && tensors.attn_v.is_some() {
413            if let Some(n) = self.dispatch_attention_layer(
414                index,
415                layer,
416                token_idx,
417                &hidden[..emb_dim],
418                emb_dim,
419                &tensors,
420                scratch_a,
421                scratch_b,
422            ) {
423                add_residual_inplace(&mut hidden[..emb_dim], &scratch_a[..n], n);
424                attn_ok = true;
425            }
426        } else if let Some(info) = tensors.attn_output {
427            let (n_in, n_out) = Self::matmul_dims(&info);
428            if n_in <= emb_dim
429                && self.dispatch_gemm_into(index, &info, &hidden[..n_in], scratch_a, n_in, n_out)
430            {
431                add_residual_inplace(
432                    &mut hidden[..emb_dim],
433                    &scratch_a[..n_out],
434                    emb_dim.min(n_out),
435                );
436                attn_ok = true;
437            }
438        }
439
440        #[cfg(not(target_arch = "wasm32"))]
441        crate::llm_bench::add_decode_attn_ns(t_attn.elapsed().as_nanos() as u64);
442
443        if !attn_ok && tensors.attn_output.is_none() && tensors.ffn_gate.is_none() {
444            return false;
445        }
446
447        // #48 diagnostic: localize the residual explosion — attention output magnitude per layer.
448        if layer < 3 && std::env::var("QUALIA_LLM_DEBUG_DECODE").is_ok() {
449            let max_attn = scratch_a[..emb_dim]
450                .iter()
451                .fold(0f32, |m, &v| m.max(v.abs()));
452            let max_hid = hidden[..emb_dim].iter().fold(0f32, |m, &v| m.max(v.abs()));
453            eprintln!(
454                "[layer-dbg] L{} attn_ok={} attn_norm={} ffn_norm={} max|attn_out|={:.4} max|hidden_postattn|={:.4}",
455                layer,
456                attn_ok,
457                tensors.attn_norm.is_some(),
458                tensors.ffn_norm.is_some(),
459                max_attn,
460                max_hid,
461            );
462        }
463
464        #[cfg(not(target_arch = "wasm32"))]
465        let t_ffn = std::time::Instant::now();
466        let ffn_ok = self
467            .dispatch_ffn_block_pre_norm(index, hidden, emb_dim, &tensors, scratch_a, scratch_b);
468        #[cfg(not(target_arch = "wasm32"))]
469        crate::llm_bench::add_decode_ffn_ns(t_ffn.elapsed().as_nanos() as u64);
470        ffn_ok
471    }
472
473    /// Attempt the CUDA mega-pass: all layers in one fenced CUDA stream.
474    /// Returns `Some(token_id)` on success, `None` to fall back to per-layer path.
475    #[cfg(all(not(target_arch = "wasm32"), feature = "cuda"))]
476    pub fn try_cuda_mega_pass_decode(
477        &mut self,
478        index: &crate::gguf_sharder::GgufTensorIndex,
479        hidden: &mut [f32],
480        emb_dim: usize,
481        token_idx: u32,
482    ) -> Option<u32> {
483        self.try_prepared_cuda_decode(index, hidden, emb_dim, token_idx)
484    }
485
486    /// Decode directly from a token id using the resident Q8 embedding table.
487    ///
488    /// This path avoids CPU embedding dequantization and the full hidden-state H2D upload.
489    #[cfg(all(not(target_arch = "wasm32"), feature = "cuda"))]
490    pub fn try_cuda_mega_pass_decode_token(
491        &mut self,
492        index: &crate::gguf_sharder::GgufTensorIndex,
493        token_id: u32,
494        hidden: &mut [f32],
495        emb_dim: usize,
496        token_idx: u32,
497    ) -> Option<u32> {
498        self.try_prepared_cuda_decode_token(index, token_id, hidden, emb_dim, token_idx)
499    }
500
501    /// Superseded unprepared implementation retained temporarily for line-by-line parity review
502    /// during R9.4 decomposition. It is excluded from compilation and will be removed after the
503    /// prepared-plan differential test is certified.
504    #[cfg(any())]
505    pub fn try_cuda_mega_pass_decode_unprepared_reference(
506        &self,
507        index: &crate::gguf_sharder::GgufTensorIndex,
508        hidden: &mut [f32],
509        emb_dim: usize,
510        token_idx: u32,
511    ) -> Option<u32> {
512        use super::RMS_NORM_EPS;
513        use crate::ggml_quants::{fetch_tensor_bytes, GGML_TYPE_Q4_K_SOA};
514        use crate::gguf_bridge::cpu_ops::dequant_norm_row_into;
515        use crate::inference::cuda_lane::{MegaPassLayerDims, MegaPassLayerWeights};
516
517        let h = index.hyperparams;
518        let n_layer = h.n_layer;
519        let n_embd = h.n_embd as usize;
520        let n_head = h.n_head as usize;
521        let n_kv = h.effective_n_kv_head() as usize;
522        let head_dim = h.head_dim() as usize;
523        if n_layer == 0 || n_embd == 0 || n_embd > 4096 || n_embd != emb_dim {
524            return None;
525        }
526        let layout = self.kv_layout?;
527        if layout.int8 || layout.dict_k > 0 {
528            return None;
529        }
530        let mmap = self.gguf_mmap.as_deref()?;
531        let tds = index.tensor_data_start;
532
533        // Phase 1: Collect raw weights, dims, and norm weights.
534        struct LayerRaw<'a> {
535            q_raw: &'a [u8],
536            k_raw: &'a [u8],
537            v_raw: &'a [u8],
538            o_raw: &'a [u8],
539            g_raw: &'a [u8],
540            u_raw: &'a [u8],
541            d_raw: &'a [u8],
542        }
543        let mut layer_raws: Vec<LayerRaw<'_>> = Vec::with_capacity(n_layer as usize);
544        let mut layer_dims: Vec<MegaPassLayerDims> = Vec::with_capacity(n_layer as usize);
545        let mut all_attn_norms: Vec<Vec<f32>> = Vec::with_capacity(n_layer as usize);
546        let mut all_ffn_norms: Vec<Vec<f32>> = Vec::with_capacity(n_layer as usize);
547
548        for l in 0..n_layer {
549            let t = index.get_layer_tensors(l);
550
551            let q_info = t.attn_q.as_ref()?;
552            let k_info = t.attn_k.as_ref()?;
553            let v_info = t.attn_v.as_ref()?;
554            let o_info = t.attn_output.as_ref()?;
555            let g_info = t.ffn_gate.as_ref()?;
556            let u_info = t.ffn_up.as_ref()?;
557            let d_info = t.ffn_down.as_ref()?;
558
559            if q_info.ggml_type != GGML_TYPE_Q4_K_SOA
560                || k_info.ggml_type != GGML_TYPE_Q4_K_SOA
561                || v_info.ggml_type != GGML_TYPE_Q4_K_SOA
562                || o_info.ggml_type != GGML_TYPE_Q4_K_SOA
563                || g_info.ggml_type != GGML_TYPE_Q4_K_SOA
564                || u_info.ggml_type != GGML_TYPE_Q4_K_SOA
565                || d_info.ggml_type != GGML_TYPE_Q4_K_SOA
566            {
567                log::debug!(
568                    "mega_pass_decode|skip|layer{l}|not_soa|q={} k={} v={} o={} g={} u={} d={}",
569                    q_info.ggml_type,
570                    k_info.ggml_type,
571                    v_info.ggml_type,
572                    o_info.ggml_type,
573                    g_info.ggml_type,
574                    u_info.ggml_type,
575                    d_info.ggml_type
576                );
577                return None;
578            }
579
580            let (q_in, q_out) = Self::matmul_dims(q_info);
581            let (k_in, k_out) = Self::matmul_dims(k_info);
582            let (_v_in, _v_out) = Self::matmul_dims(v_info);
583            let (o_in, o_out) = Self::matmul_dims(o_info);
584            let (g_in, g_out) = Self::matmul_dims(g_info);
585            let (u_in, u_out) = Self::matmul_dims(u_info);
586            let (d_in, d_out) = Self::matmul_dims(d_info);
587
588            let q_raw = fetch_tensor_bytes(mmap, tds, q_info).ok()?;
589            let k_raw = fetch_tensor_bytes(mmap, tds, k_info).ok()?;
590            let v_raw = fetch_tensor_bytes(mmap, tds, v_info).ok()?;
591            let o_raw = fetch_tensor_bytes(mmap, tds, o_info).ok()?;
592            let g_raw = fetch_tensor_bytes(mmap, tds, g_info).ok()?;
593            let u_raw = fetch_tensor_bytes(mmap, tds, u_info).ok()?;
594            let d_raw = fetch_tensor_bytes(mmap, tds, d_info).ok()?;
595
596            let mut attn_norm = vec![0.0f32; n_embd];
597            let mut ffn_norm = vec![0.0f32; n_embd];
598            if let Some(an_info) = t.attn_norm.as_ref() {
599                if dequant_norm_row_into(mmap, tds, an_info, &mut attn_norm) < n_embd {
600                    return None;
601                }
602            }
603            if let Some(fn_info) = t.ffn_norm.as_ref() {
604                if dequant_norm_row_into(mmap, tds, fn_info, &mut ffn_norm) < n_embd {
605                    return None;
606                }
607            }
608            all_attn_norms.push(attn_norm);
609            all_ffn_norms.push(ffn_norm);
610            layer_raws.push(LayerRaw {
611                q_raw,
612                k_raw,
613                v_raw,
614                o_raw,
615                g_raw,
616                u_raw,
617                d_raw,
618            });
619            layer_dims.push(MegaPassLayerDims {
620                q_in,
621                q_out,
622                kv_in: k_in,
623                kv_out: k_out,
624                o_in,
625                o_out,
626                gate_in: g_in,
627                gate_out: g_out,
628                up_in: u_in,
629                up_out: u_out,
630                down_in: d_in,
631                down_out: d_out,
632            });
633        }
634
635        // Phase 2: Build layer_weights with references into the now-stable norm vectors.
636        let layer_weights: Vec<MegaPassLayerWeights<'_>> = layer_raws
637            .iter()
638            .zip(all_attn_norms.iter().zip(all_ffn_norms.iter()))
639            .map(|(r, (an, fn_))| MegaPassLayerWeights {
640                attn_norm: &an[..n_embd],
641                q_raw: r.q_raw,
642                k_raw: r.k_raw,
643                v_raw: r.v_raw,
644                o_raw: r.o_raw,
645                ffn_norm: &fn_[..n_embd],
646                gate_raw: r.g_raw,
647                up_raw: r.u_raw,
648                down_raw: r.d_raw,
649            })
650            .collect();
651
652        // Output norm.
653        let mut output_norm_buf = vec![0.0f32; n_embd];
654        let output_norm: Option<&[f32]> = if let Some(on_info) = index.output_norm_info() {
655            if dequant_norm_row_into(mmap, tds, on_info, &mut output_norm_buf) >= n_embd {
656                Some(&output_norm_buf[..n_embd])
657            } else {
658                None
659            }
660        } else {
661            None
662        };
663
664        // LM head.
665        let lm_info = index.logits_projection_info()?;
666        let (lm_in, lm_out) = Self::matmul_dims(lm_info);
667        let lm_raw = fetch_tensor_bytes(mmap, tds, lm_info).ok()?;
668        let lm_head_raw = if lm_info.ggml_type == GGML_TYPE_Q4_K_SOA {
669            Some(lm_raw)
670        } else {
671            None
672        };
673
674        let rope_base = h.effective_rope_freq_base();
675        let rope_scale = h.effective_rope_scale();
676        let rms_eps = RMS_NORM_EPS;
677
678        crate::try_cuda_mega_pass(
679            n_embd,
680            n_head,
681            n_kv,
682            head_dim,
683            n_layer,
684            token_idx,
685            layout.max_context,
686            layout.layer_stride,
687            layout.slot_kv_elems,
688            rope_base,
689            rope_scale,
690            rms_eps,
691            &mut hidden[..n_embd],
692            &layer_weights,
693            &layer_dims,
694            output_norm,
695            lm_head_raw,
696            lm_in,
697            lm_out,
698        )
699    }
700
701    /// Sequential layer-by-layer forward (one tensor payload in VRAM at a time).
702    /// `max_layers`: `0` runs all blocks; otherwise caps how many layers execute.
703    pub fn dispatch_transformer_forward(
704        &mut self,
705        index: &crate::gguf_sharder::GgufTensorIndex,
706        hidden: &mut [f32],
707        emb_dim: usize,
708        scratch_a: &mut [f32],
709        scratch_b: &mut [f32],
710        token_idx: u32,
711        max_layers: u32,
712    ) -> u32 {
713        let n_layer = index.hyperparams.n_layer;
714        if n_layer == 0 {
715            return 0;
716        }
717        let limit = if max_layers == 0 {
718            n_layer
719        } else {
720            max_layers.min(n_layer)
721        };
722        // #48 diagnostic: localize where the hidden state turns non-finite (gated; runs once).
723        use std::sync::atomic::{AtomicBool, Ordering as DbgOrdering};
724        static FWD_DBG_DONE: AtomicBool = AtomicBool::new(false);
725        let dbg = std::env::var("QUALIA_LLM_DEBUG_DECODE").is_ok()
726            && !FWD_DBG_DONE.swap(true, DbgOrdering::Relaxed);
727        if dbg {
728            let nf = hidden[..emb_dim].iter().filter(|v| !v.is_finite()).count();
729            eprintln!(
730                "[fwd-dbg] post-embed nonfinite={}/{} sample={:?}",
731                nf,
732                emb_dim,
733                &hidden[..emb_dim.min(4)]
734            );
735        }
736        let mut ran = 0u32;
737        for layer in 0..limit {
738            if self.dispatch_transformer_layer(
739                index, layer, token_idx, hidden, emb_dim, scratch_a, scratch_b,
740            ) {
741                ran += 1;
742            }
743            if dbg {
744                let nf = hidden[..emb_dim].iter().filter(|v| !v.is_finite()).count();
745                eprintln!(
746                    "[fwd-dbg] after layer {} nonfinite={}/{} ran={} sample={:?}",
747                    layer,
748                    nf,
749                    emb_dim,
750                    ran,
751                    &hidden[..emb_dim.min(4)]
752                );
753                if nf > 0 {
754                    break;
755                }
756            }
757        }
758        ran
759    }
760
761    /// Phase 2B: async single-layer forward (GPU `map_async`; CPU path unchanged in sync API).
762    #[cfg(target_arch = "wasm32")]
763    pub async fn dispatch_transformer_layer_async(
764        &mut self,
765        index: &crate::gguf_sharder::GgufTensorIndex,
766        layer: u32,
767        token_idx: u32,
768        hidden: &mut [f32],
769        emb_dim: usize,
770        scratch_a: &mut [f32],
771        scratch_b: &mut [f32],
772    ) -> bool {
773        let tensors = index.get_layer_tensors(layer);
774        let mut attn_ok = false;
775
776        if tensors.attn_q.is_some() && tensors.attn_k.is_some() && tensors.attn_v.is_some() {
777            if let Some(n) = self
778                .dispatch_attention_layer_async(
779                    index,
780                    layer,
781                    token_idx,
782                    &hidden[..emb_dim],
783                    emb_dim,
784                    &tensors,
785                    scratch_a,
786                    scratch_b,
787                )
788                .await
789            {
790                add_residual_inplace(&mut hidden[..emb_dim], &scratch_a[..n], n);
791                attn_ok = true;
792            }
793        } else if let Some(info) = tensors.attn_output {
794            let (n_in, n_out) = Self::matmul_dims(&info);
795            if n_in <= emb_dim
796                && self
797                    .dispatch_gemm_into_async(index, &info, &hidden[..n_in], scratch_a, n_in, n_out)
798                    .await
799            {
800                add_residual_inplace(
801                    &mut hidden[..emb_dim],
802                    &scratch_a[..n_out],
803                    emb_dim.min(n_out),
804                );
805                attn_ok = true;
806            }
807        }
808
809        if !attn_ok && tensors.attn_output.is_none() && tensors.ffn_gate.is_none() {
810            return false;
811        }
812
813        self.dispatch_ffn_block_pre_norm_async(
814            index, hidden, emb_dim, &tensors, scratch_a, scratch_b,
815        )
816        .await
817    }
818
819    /// MC8: Q + o_proj + FFN tail (K/V already written for this token).
820    #[cfg(target_arch = "wasm32")]
821    pub(crate) fn encode_attn_ffn_tail_gpu(
822        &self,
823        pipeline: &mut WasmGpuPipeline,
824        index: &crate::gguf_sharder::GgufTensorIndex,
825        layer: u32,
826        token_idx: u32,
827        emb_dim: usize,
828        tensors: &crate::gguf_sharder::LayerTensors,
829        token_hidden: &wgpu::Buffer,
830        attn_input: Option<&wgpu::Buffer>,
831        work_aliases_hidden: bool,
832    ) -> bool {
833        let mmap = match self.gguf_mmap.as_deref() {
834            Some(m) => m,
835            None => return false,
836        };
837        let h = index.hyperparams;
838        let n_embd = h.n_embd as usize;
839        let layout = match self.kv_layout {
840            Some(l) => l,
841            None => return false,
842        };
843        let work_buf = self.gemm_output_buf.as_ref().unwrap();
844        let aux_buf = self.gemm_aux_buf.as_ref().unwrap();
845        let norm_buf = self.norm_weight_buf.as_ref().unwrap();
846        let q_info = match tensors.attn_q.as_ref() {
847            Some(i) => i,
848            None => return false,
849        };
850        let q_raw =
851            match crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, q_info) {
852                Ok(s) => s,
853                Err(_) => return false,
854            };
855        let q_in_buf = if let Some(pre) = attn_input {
856            pre
857        } else if let Some(norm) = tensors.attn_norm.as_ref() {
858            if !self.upload_norm_weights(mmap, index.tensor_data_start, norm, n_embd) {
859                return false;
860            }
861            self.encode_elem(
862                pipeline,
863                ELEM_OP_RMS_NORM,
864                n_embd as u32,
865                1,
866                token_hidden,
867                norm_buf,
868                aux_buf,
869            );
870            aux_buf
871        } else {
872            token_hidden
873        };
874        let ffn_buf = self.gemm_ffn_buf.as_ref().unwrap();
875        let emb_bytes = (emb_dim * 4) as wgpu::BufferAddress;
876        let q_dim = (h.n_head * h.head_dim()) as usize;
877        let (mask_words, mask_active, mask_word_count) =
878            Self::attention_kv_mask_for_dispatch(&layout, token_idx, 0);
879        let q_params = Self::attention_gpu_params(
880            &h,
881            &layout,
882            layer,
883            token_idx,
884            q_info,
885            q_raw.len(),
886            0,
887            1,
888            token_idx,
889            mask_active,
890            mask_word_count,
891            0,
892        );
893        let q_off = self.mc8_upload_attn_param(&q_params);
894        if mask_active != 0 {
895            self.gpu_queue().write_buffer(
896                self.attention_mask_buf.as_ref().unwrap(),
897                0,
898                bytemuck::cast_slice(&mask_words),
899            );
900        }
901        if !self.encode_attention_pass_gpu(
902            pipeline,
903            q_in_buf,
904            ffn_buf,
905            n_embd,
906            1,
907            token_idx,
908            &layout,
909            layer,
910            token_idx,
911            &h,
912            q_info,
913            q_raw,
914            0,
915            h.n_head,
916            q_off,
917            Mc8WeightRole::AttnQ,
918        ) {
919            return false;
920        }
921        self.mc8_flush(pipeline);
922        if let Some(out_info) = tensors.attn_output.as_ref() {
923            let (o_in, o_out) = Self::matmul_dims(out_info);
924            let o_raw = match crate::ggml_quants::fetch_tensor_bytes(
925                mmap,
926                index.tensor_data_start,
927                out_info,
928            ) {
929                Ok(s) => s,
930                Err(_) => return false,
931            };
932            if work_aliases_hidden {
933                pipeline
934                    .encoder
935                    .copy_buffer_to_buffer(token_hidden, 0, aux_buf, 0, emb_bytes);
936                self.mc8_flush(pipeline);
937            }
938            if o_in > q_dim
939                || !self.encode_gemm_bufs(pipeline, out_info, o_raw, o_in, o_out, ffn_buf, work_buf)
940            {
941                return false;
942            }
943            self.mc8_flush(pipeline);
944            let attn_residual_base: &wgpu::Buffer = if work_aliases_hidden {
945                aux_buf
946            } else {
947                token_hidden
948            };
949            // Never use `prefill_scratch_buf` here — it holds batched attn RMSNorm rows.
950            self.encode_residual_add_gpu(
951                pipeline,
952                attn_residual_base,
953                work_buf,
954                token_hidden,
955                ffn_buf,
956                emb_dim as u32,
957            );
958        } else {
959            self.encode_residual_add_gpu(
960                pipeline,
961                token_hidden,
962                ffn_buf,
963                token_hidden,
964                aux_buf,
965                emb_dim as u32,
966            );
967        }
968        self.mc8_flush(pipeline);
969        let gate_info = match tensors.ffn_gate.as_ref() {
970            Some(i) => i,
971            None => return false,
972        };
973        let up_info = match tensors.ffn_up.as_ref() {
974            Some(i) => i,
975            None => return false,
976        };
977        let down_info = match tensors.ffn_down.as_ref() {
978            Some(i) => i,
979            None => return false,
980        };
981        let (gate_in, n_ffn) = Self::matmul_dims(gate_info);
982        let (up_in, up_out) = Self::matmul_dims(up_info);
983        let (dn_in, dn_out) = Self::matmul_dims(down_info);
984        if gate_in > n_embd
985            || up_in != gate_in
986            || up_out != n_ffn
987            || dn_in != n_ffn
988            || dn_out < n_embd
989        {
990            return false;
991        }
992        let gate_raw = match crate::ggml_quants::fetch_tensor_bytes(
993            mmap,
994            index.tensor_data_start,
995            gate_info,
996        ) {
997            Ok(s) => s,
998            Err(_) => return false,
999        };
1000        let up_raw =
1001            match crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, up_info) {
1002                Ok(s) => s,
1003                Err(_) => return false,
1004            };
1005        let down_raw = match crate::ggml_quants::fetch_tensor_bytes(
1006            mmap,
1007            index.tensor_data_start,
1008            down_info,
1009        ) {
1010            Ok(s) => s,
1011            Err(_) => return false,
1012        };
1013        let base_save = match self.prefill_scratch_buf.as_ref() {
1014            Some(b) => b,
1015            None => return false,
1016        };
1017        pipeline
1018            .encoder
1019            .copy_buffer_to_buffer(token_hidden, 0, base_save, 0, emb_bytes);
1020        self.mc8_flush(pipeline);
1021        if let Some(norm) = tensors.ffn_norm.as_ref() {
1022            if !self.upload_norm_weights(mmap, index.tensor_data_start, norm, n_embd) {
1023                return false;
1024            }
1025            self.encode_elem(
1026                pipeline,
1027                ELEM_OP_RMS_NORM,
1028                n_embd as u32,
1029                1,
1030                token_hidden,
1031                norm_buf,
1032                aux_buf,
1033            );
1034        } else {
1035            pipeline
1036                .encoder
1037                .copy_buffer_to_buffer(token_hidden, 0, aux_buf, 0, emb_bytes);
1038        }
1039        self.mc8_flush(pipeline);
1040        if !self.encode_gemm_bufs(
1041            pipeline, gate_info, gate_raw, gate_in, n_ffn, aux_buf, work_buf,
1042        ) {
1043            return false;
1044        }
1045        self.mc8_flush(pipeline);
1046        if !self.encode_gemm_bufs(pipeline, up_info, up_raw, up_in, n_ffn, aux_buf, ffn_buf) {
1047            return false;
1048        }
1049        self.mc8_flush(pipeline);
1050        self.encode_elem(
1051            pipeline,
1052            ELEM_OP_SILU_MUL,
1053            n_ffn as u32,
1054            1,
1055            work_buf,
1056            ffn_buf,
1057            aux_buf,
1058        );
1059        self.mc8_flush(pipeline);
1060        if !self.encode_gemm_bufs(
1061            pipeline, down_info, down_raw, dn_in, dn_out, aux_buf, work_buf,
1062        ) {
1063            return false;
1064        }
1065        self.mc8_flush(pipeline);
1066        // FFN residual: down output is in work_buf; pre-FFN skip is in base_save.
1067        // Use aux_buf as scratch (SiLU output consumed; down GEMM flushed above).
1068        self.encode_residual_add_gpu(
1069            pipeline,
1070            base_save,
1071            work_buf,
1072            token_hidden,
1073            aux_buf,
1074            emb_dim as u32,
1075        );
1076        self.mc8_flush(pipeline);
1077        true
1078    }
1079
1080    /// MC8: encode one decode layer entirely on GPU (no map_async).
1081    /// Superseded by the Part 3w super-arena decode forward (kept for reference/fallback).
1082    #[cfg(target_arch = "wasm32")]
1083    #[allow(dead_code)]
1084    pub(crate) fn encode_transformer_layer_gpu(
1085        &self,
1086        pipeline: &mut WasmGpuPipeline,
1087        index: &crate::gguf_sharder::GgufTensorIndex,
1088        layer: u32,
1089        token_idx: u32,
1090        emb_dim: usize,
1091    ) -> bool {
1092        let tensors = index.get_layer_tensors(layer);
1093        let layout = match self.kv_layout {
1094            Some(l) => l,
1095            None => return false,
1096        };
1097        let mmap = match self.gguf_mmap.as_deref() {
1098            Some(m) => m,
1099            None => return false,
1100        };
1101        let h = index.hyperparams;
1102        let n_embd = h.n_embd as usize;
1103        if emb_dim < n_embd {
1104            return false;
1105        }
1106        let hidden_buf = self.gemm_input_buf.as_ref().unwrap();
1107        let work_buf = self.gemm_output_buf.as_ref().unwrap();
1108        let aux_buf = self.gemm_aux_buf.as_ref().unwrap();
1109        let norm_buf = self.norm_weight_buf.as_ref().unwrap();
1110
1111        let (k_info, v_info) = match (tensors.attn_k.as_ref(), tensors.attn_v.as_ref()) {
1112            (Some(k), Some(v)) => (k, v),
1113            _ => return false,
1114        };
1115        if tensors.attn_q.is_none() {
1116            return false;
1117        }
1118        let k_raw =
1119            match crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, k_info) {
1120                Ok(s) => s,
1121                Err(_) => return false,
1122            };
1123        let v_raw =
1124            match crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, v_info) {
1125                Ok(s) => s,
1126                Err(_) => return false,
1127            };
1128
1129        let attn_input = if let Some(norm) = tensors.attn_norm.as_ref() {
1130            if !self.upload_norm_weights(mmap, index.tensor_data_start, norm, n_embd) {
1131                return false;
1132            }
1133            self.encode_elem(
1134                pipeline,
1135                ELEM_OP_RMS_NORM,
1136                n_embd as u32,
1137                1,
1138                hidden_buf,
1139                norm_buf,
1140                aux_buf,
1141            );
1142            self.mc8_flush(pipeline);
1143            aux_buf
1144        } else {
1145            hidden_buf
1146        };
1147
1148        let n_kv = h.effective_n_kv_head();
1149        let mut attn_arena = Mc8UniformArena {
1150            bytes: [0u8; MC8_MAX_GEMM_UNIFORM_SLOTS * MC8_UNIFORM_ALIGN],
1151            slots: 0,
1152        };
1153        let k_params = Self::attention_gpu_params(
1154            &h,
1155            &layout,
1156            layer,
1157            token_idx,
1158            k_info,
1159            k_raw.len(),
1160            1,
1161            1,
1162            token_idx,
1163            0,
1164            0,
1165            0,
1166        );
1167        let v_params = Self::attention_gpu_params(
1168            &h,
1169            &layout,
1170            layer,
1171            token_idx,
1172            v_info,
1173            v_raw.len(),
1174            2,
1175            1,
1176            token_idx,
1177            0,
1178            0,
1179            0,
1180        );
1181        let k_off = attn_arena.push(&k_params);
1182        let v_off = attn_arena.push(&v_params);
1183        attn_arena.upload(
1184            self.gpu_queue(),
1185            self.attention_params_buf.as_ref().unwrap(),
1186        );
1187        if !self.encode_attention_pass_gpu(
1188            pipeline,
1189            attn_input,
1190            work_buf,
1191            n_embd,
1192            1,
1193            token_idx,
1194            &layout,
1195            layer,
1196            token_idx,
1197            &h,
1198            k_info,
1199            k_raw,
1200            1,
1201            n_kv,
1202            k_off,
1203            Mc8WeightRole::AttnK,
1204        ) {
1205            return false;
1206        }
1207        if !self.encode_attention_pass_gpu(
1208            pipeline,
1209            attn_input,
1210            work_buf,
1211            n_embd,
1212            1,
1213            token_idx,
1214            &layout,
1215            layer,
1216            token_idx,
1217            &h,
1218            v_info,
1219            v_raw,
1220            2,
1221            n_kv,
1222            v_off,
1223            Mc8WeightRole::AttnV,
1224        ) {
1225            return false;
1226        }
1227        self.mc8_flush(pipeline);
1228        self.encode_attn_ffn_tail_gpu(
1229            pipeline,
1230            index,
1231            layer,
1232            token_idx,
1233            emb_dim,
1234            &tensors,
1235            hidden_buf,
1236            Some(attn_input),
1237            false,
1238        )
1239    }
1240
1241    /// MC8 Part 3w: decode forward via the prefill super-arena (n_tokens=1).
1242    /// Reuses `mc8_stage_prefill_layer_super_arena` + `encode_prefill_q_ffn_tail_fused`
1243    /// (dynamic-offset uniforms + 7 disjoint weight buffers) → **2 submits/layer**
1244    /// (KV-visibility flush + layer-end), down from the legacy 13 flushes/layer.
1245    /// A single decode token at absolute position `token_idx` is a 1-row prefill chunk:
1246    /// dense causal Q (`mask_active=0`, `logical <= abs_pos`) is correct for decode.
1247    #[cfg(target_arch = "wasm32")]
1248    pub async fn dispatch_transformer_forward_async(
1249        &mut self,
1250        index: &crate::gguf_sharder::GgufTensorIndex,
1251        hidden: &mut [f32],
1252        emb_dim: usize,
1253        _scratch_a: &mut [f32],
1254        _scratch_b: &mut [f32],
1255        token_idx: u32,
1256        max_layers: u32,
1257    ) -> u32 {
1258        let n_layer = index.hyperparams.n_layer;
1259        if n_layer == 0 || !self.mc8_buffers_ready() {
1260            return 0;
1261        }
1262        if self.prefill_work_buf_a.is_none() || self.prefill_work_buf_b.is_none() {
1263            wlog("[MC8] decode forward: prefill work buffers missing — cannot run super-arena");
1264            return 0;
1265        }
1266        // Part 3x: upload all layer weights to GPU once (idempotent; falls back if it fails).
1267        if !self.mc8_weights_resident {
1268            let _ = self.mc8_upload_all_resident_weights(index);
1269        }
1270        let limit = if max_layers == 0 {
1271            n_layer
1272        } else {
1273            max_layers.min(n_layer)
1274        };
1275        let n_embd = index.hyperparams.n_embd as usize;
1276        if emb_dim < n_embd || n_embd > hidden.len() || n_embd > self.gemm_max_input_floats {
1277            return 0;
1278        }
1279        let prefill_scratch = match self.prefill_scratch_buf.as_ref() {
1280            Some(b) => b,
1281            None => return 0,
1282        };
1283        let batch_buf = self.gemm_input_buf.as_ref().unwrap();
1284        let token_buf = self.gemm_output_buf.as_ref().unwrap();
1285        if self.norm_weight_buf.is_none() {
1286            return 0;
1287        }
1288        self.gpu_queue()
1289            .write_buffer(batch_buf, 0, bytemuck::cast_slice(&hidden[..n_embd]));
1290        let mmap = match self.gguf_mmap.as_deref() {
1291            Some(m) => m,
1292            None => return 0,
1293        };
1294        let layout = match self.kv_layout {
1295            Some(l) => l,
1296            None => return 0,
1297        };
1298        let n_tokens = 1u32;
1299        let mut ran = 0u32;
1300        // Phase 5.4 single-submit: one encoder + monotonic uniform cursors across the whole forward;
1301        // flush only at MC8_LAYERS_PER_ENCODER chunk boundaries → 1 submit for ≤64-layer models (vs
1302        // the old 2/layer). Per-layer write_buffer races are gone (resident norms + accumulating
1303        // cursors), so KV/work-buffer visibility relies on WebGPU intra-encoder barriers.
1304        let mut layer_uniform_cursors = Mc8ChunkUniformCursors {
1305            attn: 0,
1306            elem: 0,
1307            gemm: 0,
1308        };
1309        let mut enc = WasmGpuPipeline::begin(self);
1310        for layer in 0..limit {
1311            if layer > 0 && (layer % MC8_LAYERS_PER_ENCODER) == 0 {
1312                self.mc8_flush(&mut enc);
1313                layer_uniform_cursors.reset();
1314            }
1315            let tensors = index.get_layer_tensors(layer);
1316            let k_info = match tensors.attn_k.as_ref() {
1317                Some(i) => i,
1318                None => break,
1319            };
1320            let v_info = match tensors.attn_v.as_ref() {
1321                Some(i) => i,
1322                None => break,
1323            };
1324            let k_raw =
1325                match crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, k_info)
1326                {
1327                    Ok(s) => s,
1328                    Err(_) => break,
1329                };
1330            let v_raw =
1331                match crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, v_info)
1332                {
1333                    Ok(s) => s,
1334                    Err(_) => break,
1335                };
1336            let h = index.hyperparams;
1337            let n_kv = h.effective_n_kv_head();
1338            let used_attn_norm = tensors.attn_norm.is_some();
1339            let (uniforms, geom) = match self.mc8_stage_prefill_layer_super_arena(
1340                index,
1341                layer,
1342                &tensors,
1343                token_idx,
1344                n_tokens,
1345                emb_dim,
1346                used_attn_norm,
1347                k_info,
1348                &k_raw,
1349                v_info,
1350                &v_raw,
1351                &mut layer_uniform_cursors,
1352            ) {
1353                Some(v) => v,
1354                None => break,
1355            };
1356            let attn_src = if used_attn_norm {
1357                if let (Some(norm), Some(off)) =
1358                    (tensors.attn_norm.as_ref(), uniforms.attn_norm_elem_off)
1359                {
1360                    let (norm_b, norm_b_off) = match self.mc8_norm_source(
1361                        mmap,
1362                        index.tensor_data_start,
1363                        norm,
1364                        n_embd,
1365                        layer,
1366                        false,
1367                    ) {
1368                        Some(v) => v,
1369                        None => break,
1370                    };
1371                    self.encode_elem_offset(
1372                        &mut enc,
1373                        ELEM_OP_RMS_NORM,
1374                        n_embd as u32,
1375                        n_tokens,
1376                        batch_buf,
1377                        0,
1378                        geom.batch_in_bytes,
1379                        norm_b,
1380                        norm_b_off,
1381                        geom.n_embd_bytes,
1382                        prefill_scratch,
1383                        0,
1384                        geom.batch_in_bytes,
1385                        0,
1386                        0,
1387                        0,
1388                        0,
1389                        0,
1390                        0,
1391                        off,
1392                    );
1393                }
1394                prefill_scratch
1395            } else {
1396                batch_buf
1397            };
1398            let n_kv_wg = n_tokens.saturating_mul(n_kv);
1399            // Phase 5.5: K/V projection on the parallel GEMM → proj buffers; the (now lightweight)
1400            // attention shader reads the pre-computed projection instead of matmul-ing on 1 thread.
1401            let kv_dim = (n_kv * h.head_dim()) as usize;
1402            let kv_proj_bytes = (kv_dim * n_tokens as usize * 4) as wgpu::BufferAddress;
1403            let k_proj = self.mc8_k_proj_buf.as_ref().unwrap();
1404            let v_proj = self.mc8_v_proj_buf.as_ref().unwrap();
1405            if !self.encode_gemm_bufs_offset(
1406                &mut enc,
1407                k_info,
1408                k_raw,
1409                n_embd,
1410                kv_dim,
1411                attn_src,
1412                0,
1413                geom.batch_in_bytes,
1414                k_proj,
1415                0,
1416                kv_proj_bytes,
1417                n_tokens,
1418                n_embd as u32,
1419                kv_dim as u32,
1420                uniforms.off_k_gemm,
1421                layer,
1422                Mc8WeightRole::AttnK,
1423            ) {
1424                break;
1425            }
1426            if !self.encode_gemm_bufs_offset(
1427                &mut enc,
1428                v_info,
1429                v_raw,
1430                n_embd,
1431                kv_dim,
1432                attn_src,
1433                0,
1434                geom.batch_in_bytes,
1435                v_proj,
1436                0,
1437                kv_proj_bytes,
1438                n_tokens,
1439                n_embd as u32,
1440                kv_dim as u32,
1441                uniforms.off_v_gemm,
1442                layer,
1443                Mc8WeightRole::AttnV,
1444            ) {
1445                break;
1446            }
1447            if !self.encode_attention_pass_gpu(
1448                &mut enc,
1449                k_proj,
1450                token_buf,
1451                n_embd,
1452                n_tokens,
1453                token_idx,
1454                &layout,
1455                layer,
1456                token_idx,
1457                &h,
1458                k_info,
1459                k_raw,
1460                1,
1461                n_kv_wg,
1462                uniforms.k_off,
1463                Mc8WeightRole::AttnK,
1464            ) {
1465                break;
1466            }
1467            if !self.encode_attention_pass_gpu(
1468                &mut enc,
1469                v_proj,
1470                token_buf,
1471                n_embd,
1472                n_tokens,
1473                token_idx,
1474                &layout,
1475                layer,
1476                token_idx,
1477                &h,
1478                v_info,
1479                v_raw,
1480                2,
1481                n_kv_wg,
1482                uniforms.v_off,
1483                Mc8WeightRole::AttnV,
1484            ) {
1485                break;
1486            }
1487            // Phase 5.4: NO per-layer flush. KV-cache + work-buffer visibility now relies on WebGPU's
1488            // automatic intra-encoder barriers between compute passes (the per-layer write_buffer
1489            // races that previously forced a flush are gone: resident norms + accumulating uniforms).
1490            let work_a = self.prefill_work_buf_a.as_ref().unwrap();
1491            let work_b = self.prefill_work_buf_b.as_ref().unwrap();
1492            if !self.encode_prefill_q_ffn_tail_fused(
1493                &mut enc,
1494                index,
1495                layer,
1496                &tensors,
1497                batch_buf,
1498                attn_src,
1499                work_a,
1500                work_b,
1501                n_tokens,
1502                token_idx,
1503                emb_dim,
1504                used_attn_norm,
1505                &uniforms,
1506                &geom,
1507            ) {
1508                break;
1509            }
1510            ran += 1;
1511        }
1512        // Phase 5.4: ONE submit for the whole forward (or per chunk if n_layer > MC8_LAYERS_PER_ENCODER).
1513        self.mc8_flush(&mut enc);
1514        if ran > 0 && !self.pipeline_read_hidden(emb_dim, hidden).await {
1515            return 0;
1516        }
1517        ran
1518    }
1519
1520    /// Fused forward + output norm + argmax in a single encoder/submit/readback.
1521    /// Eliminates 1 of 2 GPU round-trips per token vs the separate forward+argmax path.
1522    #[cfg(target_arch = "wasm32")]
1523    pub async fn dispatch_forward_and_argmax_fused_async(
1524        &mut self,
1525        index: &crate::gguf_sharder::GgufTensorIndex,
1526        hidden: &mut [f32],
1527        emb_dim: usize,
1528        token_idx: u32,
1529        max_layers: u32,
1530        _chunk_logits: &mut [f32],
1531        max_chunks: u32,
1532    ) -> Option<StreamingArgmaxResult> {
1533        let n_layer = index.hyperparams.n_layer;
1534        if n_layer == 0 || !self.mc8_buffers_ready() {
1535            return None;
1536        }
1537        if self.prefill_work_buf_a.is_none() || self.prefill_work_buf_b.is_none() {
1538            return None;
1539        }
1540        if !self.mc8_weights_resident {
1541            let _ = self.mc8_upload_all_resident_weights(index);
1542        }
1543        let limit = if max_layers == 0 {
1544            n_layer
1545        } else {
1546            max_layers.min(n_layer)
1547        };
1548        let n_embd = index.hyperparams.n_embd as usize;
1549        if emb_dim < n_embd || n_embd > hidden.len() || n_embd > self.gemm_max_input_floats {
1550            return None;
1551        }
1552        let prefill_scratch = self.prefill_scratch_buf.as_ref()?;
1553        let batch_buf = self.gemm_input_buf.as_ref().unwrap();
1554        let norm_buf = self.norm_weight_buf.as_ref().unwrap();
1555        let mmap = self.gguf_mmap.as_deref()?;
1556        let layout = self.kv_layout?;
1557
1558        // Upload hidden to batch_buf
1559        self.gpu_queue()
1560            .write_buffer(batch_buf, 0, bytemuck::cast_slice(&hidden[..n_embd]));
1561
1562        // === Forward pass (same as dispatch_transformer_forward_async but without flush) ===
1563        let n_tokens = 1u32;
1564        let mut ran = 0u32;
1565        let mut layer_uniform_cursors = Mc8ChunkUniformCursors {
1566            attn: 0,
1567            elem: 0,
1568            gemm: 0,
1569        };
1570        let mut enc = WasmGpuPipeline::begin(self);
1571        for layer in 0..limit {
1572            if layer > 0 && (layer % MC8_LAYERS_PER_ENCODER) == 0 {
1573                self.mc8_flush(&mut enc);
1574                layer_uniform_cursors.reset();
1575            }
1576            let tensors = index.get_layer_tensors(layer);
1577            let k_info = match tensors.attn_k.as_ref() {
1578                Some(i) => i,
1579                None => break,
1580            };
1581            let v_info = match tensors.attn_v.as_ref() {
1582                Some(i) => i,
1583                None => break,
1584            };
1585            let k_raw =
1586                match crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, k_info)
1587                {
1588                    Ok(s) => s,
1589                    Err(_) => break,
1590                };
1591            let v_raw =
1592                match crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, v_info)
1593                {
1594                    Ok(s) => s,
1595                    Err(_) => break,
1596                };
1597            let h = index.hyperparams;
1598            let n_kv = h.effective_n_kv_head();
1599            let used_attn_norm = tensors.attn_norm.is_some();
1600            let (uniforms, geom) = match self.mc8_stage_prefill_layer_super_arena(
1601                index,
1602                layer,
1603                &tensors,
1604                token_idx,
1605                n_tokens,
1606                emb_dim,
1607                used_attn_norm,
1608                k_info,
1609                &k_raw,
1610                v_info,
1611                &v_raw,
1612                &mut layer_uniform_cursors,
1613            ) {
1614                Some(v) => v,
1615                None => break,
1616            };
1617            let attn_src = if used_attn_norm {
1618                if let (Some(norm), Some(off)) =
1619                    (tensors.attn_norm.as_ref(), uniforms.attn_norm_elem_off)
1620                {
1621                    let (norm_b, norm_b_off) = match self.mc8_norm_source(
1622                        mmap,
1623                        index.tensor_data_start,
1624                        norm,
1625                        n_embd,
1626                        layer,
1627                        false,
1628                    ) {
1629                        Some(v) => v,
1630                        None => break,
1631                    };
1632                    self.encode_elem_offset(
1633                        &mut enc,
1634                        ELEM_OP_RMS_NORM,
1635                        n_embd as u32,
1636                        n_tokens,
1637                        batch_buf,
1638                        0,
1639                        geom.batch_in_bytes,
1640                        norm_b,
1641                        norm_b_off,
1642                        geom.n_embd_bytes,
1643                        prefill_scratch,
1644                        0,
1645                        geom.batch_in_bytes,
1646                        0,
1647                        0,
1648                        0,
1649                        0,
1650                        0,
1651                        0,
1652                        off,
1653                    );
1654                }
1655                prefill_scratch
1656            } else {
1657                batch_buf
1658            };
1659            let n_kv_wg = n_tokens.saturating_mul(n_kv);
1660            let kv_dim = (n_kv * h.head_dim()) as usize;
1661            let kv_proj_bytes = (kv_dim * n_tokens as usize * 4) as wgpu::BufferAddress;
1662            let k_proj = self.mc8_k_proj_buf.as_ref().unwrap();
1663            let v_proj = self.mc8_v_proj_buf.as_ref().unwrap();
1664            if !self.encode_gemm_bufs_offset(
1665                &mut enc,
1666                k_info,
1667                k_raw,
1668                n_embd,
1669                kv_dim,
1670                attn_src,
1671                0,
1672                geom.batch_in_bytes,
1673                k_proj,
1674                0,
1675                kv_proj_bytes,
1676                n_tokens,
1677                n_embd as u32,
1678                kv_dim as u32,
1679                uniforms.off_k_gemm,
1680                layer,
1681                Mc8WeightRole::AttnK,
1682            ) {
1683                break;
1684            }
1685            if !self.encode_gemm_bufs_offset(
1686                &mut enc,
1687                v_info,
1688                v_raw,
1689                n_embd,
1690                kv_dim,
1691                attn_src,
1692                0,
1693                geom.batch_in_bytes,
1694                v_proj,
1695                0,
1696                kv_proj_bytes,
1697                n_tokens,
1698                n_embd as u32,
1699                kv_dim as u32,
1700                uniforms.off_v_gemm,
1701                layer,
1702                Mc8WeightRole::AttnV,
1703            ) {
1704                break;
1705            }
1706            if !self.encode_attention_pass_gpu(
1707                &mut enc,
1708                k_proj,
1709                self.gemm_output_buf.as_ref().unwrap(),
1710                n_embd,
1711                n_tokens,
1712                token_idx,
1713                &layout,
1714                layer,
1715                token_idx,
1716                &h,
1717                k_info,
1718                k_raw,
1719                1,
1720                n_kv_wg,
1721                uniforms.k_off,
1722                Mc8WeightRole::AttnK,
1723            ) {
1724                break;
1725            }
1726            if !self.encode_attention_pass_gpu(
1727                &mut enc,
1728                v_proj,
1729                self.gemm_output_buf.as_ref().unwrap(),
1730                n_embd,
1731                n_tokens,
1732                token_idx,
1733                &layout,
1734                layer,
1735                token_idx,
1736                &h,
1737                v_info,
1738                v_raw,
1739                2,
1740                n_kv_wg,
1741                uniforms.v_off,
1742                Mc8WeightRole::AttnV,
1743            ) {
1744                break;
1745            }
1746            let work_a = self.prefill_work_buf_a.as_ref().unwrap();
1747            let work_b = self.prefill_work_buf_b.as_ref().unwrap();
1748            if !self.encode_prefill_q_ffn_tail_fused(
1749                &mut enc,
1750                index,
1751                layer,
1752                &tensors,
1753                batch_buf,
1754                attn_src,
1755                work_a,
1756                work_b,
1757                n_tokens,
1758                token_idx,
1759                emb_dim,
1760                used_attn_norm,
1761                &uniforms,
1762                &geom,
1763            ) {
1764                break;
1765            }
1766            ran += 1;
1767        }
1768        if ran == 0 {
1769            self.mc8_flush(&mut enc);
1770            return None;
1771        }
1772
1773        // === Output norm on GPU ===
1774        let output_norm_info = match index.output_norm_info() {
1775            Some(i) => i,
1776            None => {
1777                // No output norm — just read hidden and fall back to CPU argmax
1778                self.mc8_flush(&mut enc);
1779                if !self.pipeline_read_hidden(emb_dim, hidden).await {
1780                    return None;
1781                }
1782                return None; // caller will use separate argmax
1783            }
1784        };
1785        // Upload output norm weights to norm_weight_buf
1786        let mut norm_w = [0f32; MAX_HIDDEN_DIM];
1787        if dequant_norm_row_into(mmap, index.tensor_data_start, output_norm_info, &mut norm_w)
1788            < n_embd
1789        {
1790            self.mc8_flush(&mut enc);
1791            return None;
1792        }
1793        self.gpu_queue()
1794            .write_buffer(norm_buf, 0, bytemuck::cast_slice(&norm_w[..n_embd]));
1795        // RMS norm: read from batch_buf, write to prefill_scratch
1796        let n_embd_bytes = (n_embd * 4) as wgpu::BufferAddress;
1797        self.encode_elem_offset(
1798            &mut enc,
1799            ELEM_OP_RMS_NORM,
1800            n_embd as u32,
1801            1,
1802            batch_buf,
1803            0,
1804            n_embd_bytes,
1805            norm_buf,
1806            0,
1807            n_embd_bytes,
1808            prefill_scratch,
1809            0,
1810            n_embd_bytes,
1811            0,
1812            0,
1813            0,
1814            0,
1815            0,
1816            0,
1817            0,
1818        );
1819        // Copy normed hidden back to batch_buf for argmax GEMM input
1820        enc.encoder
1821            .copy_buffer_to_buffer(prefill_scratch, 0, batch_buf, 0, n_embd_bytes);
1822
1823        // === Batched argmax in same encoder ===
1824        let resident_buf = self.mc8_logits_resident_buf.as_ref()?;
1825        let row_bytes = self.mc8_logits_row_bytes as u64;
1826        let output_buf = self.gemm_output_buf.as_ref().unwrap();
1827        let params_buf = self.gemm_params_buf.as_ref().unwrap();
1828        let staging = self.gemm_output_staging.as_ref().unwrap();
1829        let logits_info = index.logits_projection_info()?;
1830        let (n_in, vocab_size) = Self::matmul_dims(logits_info);
1831        if n_in == 0 || vocab_size == 0 || n_in > emb_dim {
1832            self.mc8_flush(&mut enc);
1833            return None;
1834        }
1835        let full_chunks = vocab_size.div_ceil(VOCAB_CHUNK_ROWS);
1836        let n_chunks = if max_chunks == 0 {
1837            full_chunks
1838        } else {
1839            (max_chunks as usize).min(full_chunks)
1840        };
1841        let top1_plan =
1842            crate::gguf_bridge::browser::webgpu::BrowserTop1Plan::new(vocab_size, n_chunks as u32)?;
1843        if !self.prepare_browser_top1(top1_plan) {
1844            self.mc8_flush(&mut enc);
1845            return None;
1846        }
1847        let params = GemmGpuParams {
1848            n_in: n_in as u32,
1849            n_out: VOCAB_CHUNK_ROWS as u32,
1850            weight_ggml_type: logits_info.ggml_type,
1851            weight_row_elems: logits_info.dims[0] as u32,
1852            weight_byte_len: (VOCAB_CHUNK_ROWS as u64 * row_bytes) as u32,
1853            n_batch: 1,
1854            in_row_stride: 0,
1855            out_row_stride: 0,
1856        };
1857        self.gpu_queue()
1858            .write_buffer(params_buf, 0, bytemuck::bytes_of(&params));
1859        #[cfg(target_arch = "wasm32")]
1860        let use_mmv_q8_0 =
1861            logits_info.ggml_type == crate::ggml_quants::GGML_TYPE_Q8_0 && (n_in % 32 == 0);
1862        #[cfg(target_arch = "wasm32")]
1863        let logits_pipeline: &wgpu::ComputePipeline = if use_mmv_q8_0 {
1864            &self.mmv_q8_0_pipeline
1865        } else {
1866            &self.pipeline
1867        };
1868        #[cfg(not(target_arch = "wasm32"))]
1869        let logits_pipeline: &wgpu::ComputePipeline = &self.pipeline;
1870        let bind_layout = logits_pipeline.get_bind_group_layout(0);
1871        for chunk_idx in 0..n_chunks {
1872            let row_start = chunk_idx * VOCAB_CHUNK_ROWS;
1873            let chunk_rows = VOCAB_CHUNK_ROWS.min(vocab_size - row_start);
1874            let weight = wgpu::BindingResource::Buffer(wgpu::BufferBinding {
1875                buffer: resident_buf,
1876                offset: row_start as u64 * row_bytes,
1877                size: std::num::NonZeroU64::new(chunk_rows as u64 * row_bytes),
1878            });
1879            let key = mc8_bg_hash(&[
1880                7,
1881                chunk_idx as u64,
1882                row_start as u64,
1883                row_bytes as u64,
1884                chunk_rows as u64,
1885                resident_buf as *const _ as u64,
1886                batch_buf as *const _ as u64,
1887                output_buf as *const _ as u64,
1888            ]);
1889            let bind_group = {
1890                let cached = self
1891                    .mc8_bg_cache
1892                    .lock()
1893                    .ok()
1894                    .and_then(|c| c.get(&key).cloned());
1895                match cached {
1896                    Some(bg) => bg,
1897                    None => {
1898                        let bg = self
1899                            .gpu_device()
1900                            .create_bind_group(&wgpu::BindGroupDescriptor {
1901                                label: Some("FusedArgmaxBind"),
1902                                layout: &bind_layout,
1903                                entries: &[
1904                                    wgpu::BindGroupEntry {
1905                                        binding: 0,
1906                                        resource: batch_buf.as_entire_binding(),
1907                                    },
1908                                    wgpu::BindGroupEntry {
1909                                        binding: 1,
1910                                        resource: weight,
1911                                    },
1912                                    wgpu::BindGroupEntry {
1913                                        binding: 2,
1914                                        resource: Self::mc8_dynamic_uniform_binding(params_buf),
1915                                    },
1916                                    wgpu::BindGroupEntry {
1917                                        binding: 3,
1918                                        resource: output_buf.as_entire_binding(),
1919                                    },
1920                                ],
1921                            });
1922                        if let Ok(mut c) = self.mc8_bg_cache.lock() {
1923                            c.insert(key, bg.clone());
1924                        }
1925                        bg
1926                    }
1927                }
1928            };
1929            {
1930                let mut cpass = enc
1931                    .encoder
1932                    .begin_compute_pass(&wgpu::ComputePassDescriptor {
1933                        label: None,
1934                        timestamp_writes: None,
1935                    });
1936                cpass.set_pipeline(logits_pipeline);
1937                cpass.set_bind_group(0, &bind_group, &[0]);
1938                #[cfg(target_arch = "wasm32")]
1939                if use_mmv_q8_0 {
1940                    cpass.dispatch_workgroups((chunk_rows as u32 + 3) / 4, 1, 1);
1941                } else {
1942                    cpass.dispatch_workgroups((chunk_rows as u32 + 63) / 64, 1, 1);
1943                }
1944                #[cfg(not(target_arch = "wasm32"))]
1945                cpass.dispatch_workgroups((chunk_rows as u32 + 63) / 64, 1, 1);
1946            }
1947            if !self.encode_browser_top1_chunk(
1948                &mut enc.encoder,
1949                output_buf,
1950                staging,
1951                top1_plan,
1952                chunk_idx,
1953            ) {
1954                self.mc8_flush(&mut enc);
1955                return None;
1956            }
1957        }
1958
1959        // One submit/fence. Logits remain on-device; only bounded block
1960        // candidates are mapped for the deterministic global merge.
1961        self.gpu_queue().submit(Some(enc.encoder.finish()));
1962        self.read_browser_top1(staging, top1_plan).await
1963    }
1964
1965    /// Topological speculative verify — accept longest draft prefix (B3.1d).
1966    pub fn verify_topology_draft_batch(
1967        &mut self,
1968        index: &crate::gguf_sharder::GgufTensorIndex,
1969        ctx: &mut Vec<u32>,
1970        draft: &crate::compute_universe::TopologyDraftBatch,
1971        emb_dim: usize,
1972        emb_buf: &mut [f32],
1973        scratch_a: &mut [f32],
1974        scratch_b: &mut [f32],
1975        max_layers: u32,
1976        max_vocab_chunks: u32,
1977    ) -> u32 {
1978        let mmap = match self.gguf_mmap.clone() {
1979            Some(m) => m,
1980            None => return 0,
1981        };
1982        let gamma = draft.draft_len as usize;
1983        if gamma == 0 || ctx.is_empty() {
1984            return 0;
1985        }
1986        let mut accepted = 0u32;
1987        for i in 0..gamma {
1988            let cur = *ctx.last().unwrap();
1989            let token_idx = ctx.len().saturating_sub(1) as u32;
1990            let hidden_ok =
1991                index.dequantize_token_embedding_into(mmap.as_ref(), cur, &mut emb_buf[..emb_dim]);
1992            if hidden_ok == 0 {
1993                break;
1994            }
1995            let _ = self.dispatch_transformer_forward(
1996                index,
1997                &mut emb_buf[..emb_dim],
1998                emb_dim,
1999                scratch_a,
2000                scratch_b,
2001                token_idx,
2002                max_layers,
2003            );
2004            let pred = if let Some(argmax) = self.dispatch_output_argmax_chunked(
2005                index,
2006                &emb_buf[..emb_dim],
2007                emb_dim,
2008                scratch_a,
2009                max_vocab_chunks,
2010                None,
2011            ) {
2012                if argmax.max_logit > f32::NEG_INFINITY {
2013                    argmax.best_token_id
2014                } else {
2015                    break;
2016                }
2017            } else {
2018                break;
2019            };
2020            if pred != draft.draft_ids[i] {
2021                break;
2022            }
2023            ctx.push(pred);
2024            accepted += 1;
2025        }
2026        accepted
2027    }
2028}