Skip to main content

qualia_core_db/gguf_bridge/
async_dispatch.rs

1//! Decode-tail dispatch: final logits projection + async GEMM / async attention paths.
2//! Split from gguf_bridge/forward.rs (structural; no behaviour change).
3use super::*;
4
5impl QTensorEngine {
6    /// Final logits via chunked projection into `logits_out` (fills min(vocab, buf) rows).
7    pub fn dispatch_output_logits_into(
8        &self,
9        index: &crate::gguf_sharder::GgufTensorIndex,
10        hidden: &[f32],
11        emb_dim: usize,
12        logits_out: &mut [f32],
13    ) -> usize {
14        let Some(info) = index.logits_projection_info() else {
15            let n = emb_dim.min(logits_out.len());
16            logits_out[..n].copy_from_slice(&hidden[..n]);
17            return n;
18        };
19        let (n_in, vocab_size) = Self::matmul_dims(info);
20        let fill = vocab_size.min(logits_out.len());
21        if n_in > emb_dim || fill == 0 {
22            let n = emb_dim.min(logits_out.len());
23            logits_out[..n].copy_from_slice(&hidden[..n]);
24            return n;
25        }
26        let mmap = match self.gguf_mmap.as_deref() {
27            Some(m) => m,
28            None => {
29                let n = emb_dim.min(logits_out.len());
30                logits_out[..n].copy_from_slice(&hidden[..n]);
31                return n;
32            }
33        };
34        let mut written = 0usize;
35        let n_chunks = vocab_size.div_ceil(VOCAB_CHUNK_ROWS);
36        for chunk_idx in 0..n_chunks {
37            if written >= fill {
38                break;
39            }
40            let row_start = chunk_idx * VOCAB_CHUNK_ROWS;
41            let chunk_rows = VOCAB_CHUNK_ROWS.min(vocab_size - row_start);
42            let raw = match crate::ggml_quants::fetch_tensor_row_range_bytes(
43                mmap,
44                index.tensor_data_start,
45                info,
46                row_start,
47                chunk_rows,
48            ) {
49                Ok(s) => s,
50                Err(_) => break,
51            };
52            let out_rows = chunk_rows.min(fill - written);
53            if !self.dispatch_gemm_raw_into(
54                info,
55                raw,
56                &hidden[..n_in],
57                &mut logits_out[written..written + out_rows],
58                n_in,
59                out_rows,
60            ) {
61                break;
62            }
63            written += out_rows;
64        }
65        if written > 0 {
66            written
67        } else {
68            let n = emb_dim.min(logits_out.len());
69            logits_out[..n].copy_from_slice(&hidden[..n]);
70            n
71        }
72    }
73
74    pub fn decode_lexicon_bound(&self, _logits: &[f32], valid_lexicon_ids: &[u64]) -> u64 {
75        if valid_lexicon_ids.is_empty() {
76            0
77        } else {
78            valid_lexicon_ids[0]
79        }
80    }
81
82    #[cfg(target_arch = "wasm32")]
83    pub(crate) async fn dispatch_gemm_raw_into_async(
84        &self,
85        info: &GgufTensorInfo,
86        raw: &[u8],
87        input: &[f32],
88        out: &mut [f32],
89        n_in: usize,
90        n_out: usize,
91    ) -> bool {
92        if n_in > input.len() || n_out > out.len() {
93            return false;
94        }
95
96        let weight_bytes = raw.len();
97        if ggml_gpu_attention_shader_supported(info.ggml_type)
98            && n_in <= MAX_STACK_GEMM_IN
99            && n_out <= self.gemm_max_out_dim as usize
100            && weight_bytes <= self.max_tensor_bytes
101            && self.gemm_input_buf.is_some()
102        {
103            let params = GemmGpuParams {
104                n_in: n_in as u32,
105                n_out: n_out as u32,
106                weight_ggml_type: info.ggml_type,
107                weight_row_elems: info.dims[0] as u32,
108                weight_byte_len: raw.len() as u32,
109                n_batch: 1,
110                in_row_stride: 0,
111                out_row_stride: 0,
112            };
113            let input_buf = self.gemm_input_buf.as_ref().unwrap();
114            let weight_buf = self.gemm_weight_buf.as_ref().unwrap();
115            let output_buf = self.gemm_output_buf.as_ref().unwrap();
116            let params_buf = self.gemm_params_buf.as_ref().unwrap();
117            let staging = self.gemm_output_staging.as_ref().unwrap();
118
119            self.gpu_queue()
120                .write_buffer(input_buf, 0, bytemuck::cast_slice(&input[..n_in]));
121            self.write_weight_words(raw, self.max_tensor_bytes);
122            self.gpu_queue()
123                .write_buffer(params_buf, 0, bytemuck::bytes_of(&params));
124
125            let bind_layout = self.pipeline.get_bind_group_layout(0);
126            let bind_group = self
127                .gpu_device()
128                .create_bind_group(&wgpu::BindGroupDescriptor {
129                    label: Some("LayerGemmBindGroup"),
130                    layout: &bind_layout,
131                    entries: &[
132                        wgpu::BindGroupEntry {
133                            binding: 0,
134                            resource: input_buf.as_entire_binding(),
135                        },
136                        wgpu::BindGroupEntry {
137                            binding: 1,
138                            resource: weight_buf.as_entire_binding(),
139                        },
140                        wgpu::BindGroupEntry {
141                            binding: 2,
142                            resource: Self::mc8_dynamic_uniform_binding(params_buf),
143                        },
144                        wgpu::BindGroupEntry {
145                            binding: 3,
146                            resource: output_buf.as_entire_binding(),
147                        },
148                    ],
149                });
150
151            let mut encoder =
152                self.device()
153                    .create_command_encoder(&wgpu::CommandEncoderDescriptor {
154                        label: Some("LayerGemmEncoder"),
155                    });
156            {
157                let mut cpass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
158                    label: None,
159                    timestamp_writes: crate::llm_gpu_profiler::pass_writes_both(),
160                });
161                cpass.set_pipeline(&self.pipeline);
162                cpass.set_bind_group(0, &bind_group, &[0]);
163                cpass.dispatch_workgroups((n_out as u32 + 63) / 64, 1, 1);
164            }
165            let out_bytes = (n_out * 4) as wgpu::BufferAddress;
166            encoder.copy_buffer_to_buffer(output_buf, 0, staging, 0, out_bytes);
167            crate::llm_gpu_profiler::resolve(&mut encoder);
168            self.gpu_queue().submit(Some(encoder.finish()));
169            crate::llm_gpu_profiler::accumulate(crate::llm_gpu_profiler::Phase::Gemm);
170
171            let slice = staging.slice(..out_bytes);
172            if await_wgpu_map(slice).await {
173                let data = slice
174                    .get_mapped_range()
175                    .expect("wgpu buffer map_range failed");
176                let floats: &[f32] = bytemuck::cast_slice(&data);
177                out[..n_out].copy_from_slice(&floats[..n_out]);
178                drop(data);
179                staging.unmap();
180                return true;
181            }
182            let _ = staging.unmap();
183        }
184
185        stack_gemm_quant(raw, info, input, out, n_in, n_out)
186    }
187
188    /// Phase 5.3: quantized GEMM against a **resident** weight sub-range (no per-call weight upload).
189    /// Same dispatch+readback as `dispatch_gemm_raw_into_async`, but binds the caller-provided
190    /// `weight` BindingResource (a slice of the resident output-projection buffer) instead of
191    /// `write_buffer`-ing the chunk. Used by the decode argmax once the logits projection is resident.
192    #[cfg(target_arch = "wasm32")]
193    pub(crate) async fn dispatch_gemm_resident_chunk_async(
194        &self,
195        weight_ggml_type: u32,
196        weight_row_elems: u32,
197        weight: wgpu::BindingResource<'_>,
198        weight_byte_len: u32,
199        input: &[f32],
200        out: &mut [f32],
201        n_in: usize,
202        n_out: usize,
203    ) -> bool {
204        if n_in > input.len() || n_out > out.len() {
205            return false;
206        }
207        if n_in > MAX_STACK_GEMM_IN || n_out > self.gemm_max_out_dim as usize {
208            return false;
209        }
210        let input_buf = match self.gemm_input_buf.as_ref() {
211            Some(b) => b,
212            None => return false,
213        };
214        let output_buf = match self.gemm_output_buf.as_ref() {
215            Some(b) => b,
216            None => return false,
217        };
218        let params_buf = match self.gemm_params_buf.as_ref() {
219            Some(b) => b,
220            None => return false,
221        };
222        let staging = match self.gemm_output_staging.as_ref() {
223            Some(b) => b,
224            None => return false,
225        };
226        let params = GemmGpuParams {
227            n_in: n_in as u32,
228            n_out: n_out as u32,
229            weight_ggml_type,
230            weight_row_elems,
231            weight_byte_len,
232            n_batch: 1,
233            in_row_stride: 0,
234            out_row_stride: 0,
235        };
236        self.gpu_queue()
237            .write_buffer(input_buf, 0, bytemuck::cast_slice(&input[..n_in]));
238        self.gpu_queue()
239            .write_buffer(params_buf, 0, bytemuck::bytes_of(&params));
240
241        let use_mmv = weight_ggml_type == crate::ggml_quants::GGML_TYPE_Q8_0 && (n_in % 32 == 0);
242        let active_pipeline = if use_mmv {
243            &self.mmv_q8_0_pipeline
244        } else {
245            &self.pipeline
246        };
247        let bind_layout = active_pipeline.get_bind_group_layout(0);
248        let bind_group = self
249            .gpu_device()
250            .create_bind_group(&wgpu::BindGroupDescriptor {
251                label: Some("ResidentLogitsBind"),
252                layout: &bind_layout,
253                entries: &[
254                    wgpu::BindGroupEntry {
255                        binding: 0,
256                        resource: input_buf.as_entire_binding(),
257                    },
258                    wgpu::BindGroupEntry {
259                        binding: 1,
260                        resource: weight,
261                    },
262                    // self.pipeline uses MC8GemmBGL, whose binding 2 is a DYNAMIC uniform — bind it
263                    // as a 256-sized dynamic slice and pass offset 0 below (matches encode_gemm_bufs_offset).
264                    wgpu::BindGroupEntry {
265                        binding: 2,
266                        resource: Self::mc8_dynamic_uniform_binding(params_buf),
267                    },
268                    wgpu::BindGroupEntry {
269                        binding: 3,
270                        resource: output_buf.as_entire_binding(),
271                    },
272                ],
273            });
274
275        let mut encoder = self
276            .device()
277            .create_command_encoder(&wgpu::CommandEncoderDescriptor {
278                label: Some("ResidentLogitsEncoder"),
279            });
280        {
281            let mut cpass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
282                label: None,
283                timestamp_writes: crate::llm_gpu_profiler::pass_writes_both(),
284            });
285            cpass.set_pipeline(active_pipeline);
286            cpass.set_bind_group(0, &bind_group, &[0]);
287            if use_mmv {
288                cpass.dispatch_workgroups((n_out as u32 + 3) / 4, 1, 1);
289            } else {
290                cpass.dispatch_workgroups((n_out as u32 + 63) / 64, 1, 1);
291            }
292        }
293        let out_bytes = (n_out * 4) as wgpu::BufferAddress;
294        encoder.copy_buffer_to_buffer(output_buf, 0, staging, 0, out_bytes);
295        crate::llm_gpu_profiler::resolve(&mut encoder);
296        self.gpu_queue().submit(Some(encoder.finish()));
297        crate::llm_gpu_profiler::accumulate(crate::llm_gpu_profiler::Phase::Gemm);
298
299        let slice = staging.slice(..out_bytes);
300        if await_wgpu_map(slice).await {
301            let data = slice
302                .get_mapped_range()
303                .expect("wgpu buffer map_range failed");
304            let floats: &[f32] = bytemuck::cast_slice(&data);
305            out[..n_out].copy_from_slice(&floats[..n_out]);
306            drop(data);
307            staging.unmap();
308            return true;
309        }
310        let _ = staging.unmap();
311        false
312    }
313    #[cfg(target_arch = "wasm32")]
314    pub async fn dispatch_gemm_into_async(
315        &self,
316        index: &crate::gguf_sharder::GgufTensorIndex,
317        info: &GgufTensorInfo,
318        input: &[f32],
319        out: &mut [f32],
320        n_in: usize,
321        n_out: usize,
322    ) -> bool {
323        if n_in > input.len() || n_out > out.len() {
324            wlog(&format!(
325                "[gemm_into_async] GUARD n_in={n_in} n_out={n_out} input={} out={}",
326                input.len(),
327                out.len()
328            ));
329            return false;
330        }
331        let mmap = match self.gguf_mmap.as_deref() {
332            Some(m) => m,
333            None => return false,
334        };
335        let raw = match crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, info)
336        {
337            Ok(s) => s,
338            Err(_) => return false,
339        };
340        self.dispatch_gemm_raw_into_async(info, raw, input, out, n_in, n_out)
341            .await
342    }
343    #[cfg(target_arch = "wasm32")]
344    pub(crate) async fn dispatch_attention_layer_async(
345        &mut self,
346        index: &crate::gguf_sharder::GgufTensorIndex,
347        layer: u32,
348        token_idx: u32,
349        hidden: &[f32],
350        emb_dim: usize,
351        tensors: &crate::gguf_sharder::LayerTensors,
352        scratch_a: &mut [f32],
353        scratch_b: &mut [f32],
354    ) -> Option<usize> {
355        let layout = self.kv_layout?;
356        let q_info = tensors.attn_q.as_ref()?;
357        let k_info = tensors.attn_k.as_ref()?;
358        let v_info = tensors.attn_v.as_ref()?;
359        let h = index.hyperparams;
360        let n_head = h.n_head as usize;
361        let n_kv = h.effective_n_kv_head() as usize;
362        let head_dim = h.head_dim() as usize;
363        if head_dim == 0 || n_head == 0 || n_kv == 0 {
364            return None;
365        }
366        let q_dim = n_head * head_dim;
367        if q_dim > scratch_a.len() || q_dim > scratch_b.len() || emb_dim < h.n_embd as usize {
368            return None;
369        }
370        if !ggml_gpu_quant_supported(q_info.ggml_type)
371            || !ggml_gpu_quant_supported(k_info.ggml_type)
372            || !ggml_gpu_quant_supported(v_info.ggml_type)
373        {
374            return None;
375        }
376
377        let mmap = self.gguf_mmap.as_deref()?;
378        let k_raw =
379            crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, k_info).ok()?;
380        let v_raw =
381            crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, v_info).ok()?;
382        let q_raw =
383            crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, q_info).ok()?;
384        let n_embd = h.n_embd as usize;
385
386        let mut norm_w_attn = [0f32; MAX_HIDDEN_DIM];
387        let mut h_norm_attn = [0f32; MAX_HIDDEN_DIM];
388        let hidden_input = prepare_pre_norm_input(
389            &hidden[..emb_dim],
390            emb_dim,
391            tensors.attn_norm.as_ref(),
392            Some(&mmap[..]),
393            index.tensor_data_start,
394            &mut h_norm_attn,
395            &mut norm_w_attn,
396        );
397
398        if !self
399            .dispatch_attention_pass_async(
400                hidden_input,
401                n_embd,
402                1,
403                token_idx,
404                &layout,
405                layer,
406                token_idx,
407                &h,
408                k_info,
409                k_raw,
410                1,
411                n_kv as u32,
412                None,
413            )
414            .await
415        {
416            return None;
417        }
418        if !self
419            .dispatch_attention_pass_async(
420                hidden_input,
421                n_embd,
422                1,
423                token_idx,
424                &layout,
425                layer,
426                token_idx,
427                &h,
428                v_info,
429                v_raw,
430                2,
431                n_kv as u32,
432                None,
433            )
434            .await
435        {
436            return None;
437        }
438        if !self
439            .dispatch_attention_pass_async(
440                hidden_input,
441                n_embd,
442                1,
443                token_idx,
444                &layout,
445                layer,
446                token_idx,
447                &h,
448                q_info,
449                q_raw,
450                0,
451                n_head as u32,
452                Some(&mut scratch_b[..q_dim]),
453            )
454            .await
455        {
456            return None;
457        }
458
459        if let Some(out_info) = tensors.attn_output {
460            let (o_in, o_out) = Self::matmul_dims(&out_info);
461            if o_in <= q_dim
462                && self
463                    .dispatch_gemm_into_async(
464                        index,
465                        &out_info,
466                        &scratch_b[..o_in],
467                        &mut scratch_a[..o_out],
468                        o_in,
469                        o_out,
470                    )
471                    .await
472            {
473                return Some(o_out.min(emb_dim));
474            }
475        }
476        let n = q_dim.min(emb_dim);
477        scratch_a[..n].copy_from_slice(&scratch_b[..n]);
478        Some(n)
479    }
480
481    #[cfg(target_arch = "wasm32")]
482    pub(crate) async fn dispatch_attention_pass_async(
483        &self,
484        hidden: &[f32],
485        n_embd: usize,
486        num_tokens_in_batch: u32,
487        batch_start_token_idx: u32,
488        layout: &KvCacheLayout,
489        layer: u32,
490        token_idx: u32,
491        h: &crate::gguf_sharder::GgufHyperparams,
492        info: &GgufTensorInfo,
493        raw_weights: &[u8],
494        proj_kind: u32,
495        n_workgroups: u32,
496        readback_out: Option<&mut [f32]>,
497    ) -> bool {
498        if !ggml_gpu_quant_supported(info.ggml_type) {
499            wlog(&format!(
500                "[attn_pass_async] GUARD unsupported quant kind={proj_kind}"
501            ));
502            return false;
503        }
504        if !ggml_gpu_attention_shader_supported(info.ggml_type) {
505            return self.cpu_attention_pass(
506                hidden,
507                n_embd,
508                num_tokens_in_batch,
509                batch_start_token_idx,
510                layout,
511                layer,
512                h,
513                info,
514                raw_weights,
515                proj_kind,
516                None,
517                readback_out,
518            );
519        }
520        let batch = num_tokens_in_batch.max(1) as usize;
521        let hidden_elems = n_embd.checked_mul(batch).unwrap_or(0);
522        if hidden_elems > hidden.len()
523            || hidden_elems > self.gemm_max_input_floats
524            || raw_weights.len() > self.max_tensor_bytes
525            || self.gemm_input_buf.is_none()
526            || self.kv_cache_gpu.is_none()
527            || self.attention_params_buf.is_none()
528            || self.attention_mask_buf.is_none()
529        {
530            wlog(&format!(
531                "[attn_pass_async] GUARD buffers kind={proj_kind} hidden_elems={hidden_elems} hidden={} gemm_in={} raw_w={} max_w={}",
532                hidden.len(),
533                self.gemm_max_input_floats,
534                raw_weights.len(),
535                self.max_tensor_bytes,
536            ));
537            return false;
538        }
539
540        let (mask_words, mask_active, mask_word_count) =
541            Self::attention_kv_mask_for_dispatch(layout, token_idx, proj_kind);
542        let params = Self::attention_gpu_params(
543            h,
544            layout,
545            layer,
546            token_idx,
547            info,
548            raw_weights.len(),
549            proj_kind,
550            num_tokens_in_batch.max(1),
551            batch_start_token_idx,
552            mask_active,
553            mask_word_count,
554            0,
555        );
556        let input_buf = self.gemm_input_buf.as_ref().unwrap();
557        let weight_buf = self.gemm_weight_buf.as_ref().unwrap();
558        let output_buf = self.gemm_output_buf.as_ref().unwrap();
559        let params_buf = self.attention_params_buf.as_ref().unwrap();
560        let mask_buf = self.attention_mask_buf.as_ref().unwrap();
561        let kv_buf = self.kv_cache_gpu.as_ref().unwrap();
562        let staging = self.gemm_output_staging.as_ref().unwrap();
563
564        self.gpu_queue()
565            .write_buffer(input_buf, 0, bytemuck::cast_slice(&hidden[..hidden_elems]));
566        self.write_weight_words(raw_weights, self.max_tensor_bytes);
567        self.gpu_queue()
568            .write_buffer(params_buf, 0, bytemuck::bytes_of(&params));
569        self.gpu_queue()
570            .write_buffer(mask_buf, 0, bytemuck::cast_slice(&mask_words));
571
572        // Bind one layer slice of the KV arena (full arena exceeds 128 MiB wgpu binding cap).
573        let layer_f32s = layout.layer_stride as usize;
574        let layer_bytes = (layer_f32s * std::mem::size_of::<f32>()) as wgpu::BufferAddress;
575        let layer_offset =
576            (layer as usize * layer_f32s * std::mem::size_of::<f32>()) as wgpu::BufferAddress;
577        let kv_binding = wgpu::BufferBinding {
578            buffer: kv_buf,
579            offset: layer_offset,
580            size: std::num::NonZeroU64::new(layer_bytes.max(4)),
581        };
582        let (wg_x, wg_y) = if proj_kind == 0 && num_tokens_in_batch > 1 {
583            (h.n_head, num_tokens_in_batch)
584        } else {
585            (n_workgroups.max(1), 1)
586        };
587
588        let bind_layout = self.attention_pipeline.get_bind_group_layout(0);
589        let bind_group = self
590            .gpu_device()
591            .create_bind_group(&wgpu::BindGroupDescriptor {
592                label: Some("FusedAttentionBindGroup"),
593                layout: &bind_layout,
594                entries: &[
595                    wgpu::BindGroupEntry {
596                        binding: 0,
597                        resource: input_buf.as_entire_binding(),
598                    },
599                    wgpu::BindGroupEntry {
600                        binding: 1,
601                        resource: weight_buf.as_entire_binding(),
602                    },
603                    wgpu::BindGroupEntry {
604                        binding: 2,
605                        resource: params_buf.as_entire_binding(),
606                    },
607                    wgpu::BindGroupEntry {
608                        binding: 3,
609                        resource: wgpu::BindingResource::Buffer(kv_binding),
610                    },
611                    wgpu::BindGroupEntry {
612                        binding: 4,
613                        resource: output_buf.as_entire_binding(),
614                    },
615                    wgpu::BindGroupEntry {
616                        binding: 5,
617                        resource: mask_buf.as_entire_binding(),
618                    },
619                ],
620            });
621
622        let mut encoder = self
623            .device()
624            .create_command_encoder(&wgpu::CommandEncoderDescriptor {
625                label: Some("FusedAttentionEncoder"),
626            });
627        {
628            let mut cpass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
629                label: Some("FusedAttentionPass"),
630                timestamp_writes: crate::llm_gpu_profiler::pass_writes_both(),
631            });
632            cpass.set_pipeline(&self.attention_pipeline);
633            cpass.set_bind_group(0, &bind_group, &[]);
634            cpass.dispatch_workgroups(wg_x, wg_y, 1);
635        }
636
637        let readback_elems = readback_out.as_ref().map(|o| o.len()).unwrap_or(0);
638        if readback_elems > 0 {
639            let out_bytes = (readback_elems * 4) as wgpu::BufferAddress;
640            encoder.copy_buffer_to_buffer(output_buf, 0, staging, 0, out_bytes);
641        }
642        crate::llm_gpu_profiler::resolve(&mut encoder);
643        self.gpu_queue().submit(Some(encoder.finish()));
644        crate::llm_gpu_profiler::accumulate(crate::llm_gpu_profiler::Phase::Attention);
645
646        if readback_elems == 0 {
647            return true;
648        }
649
650        let out_bytes = (readback_elems * 4) as wgpu::BufferAddress;
651        let slice = staging.slice(..out_bytes);
652        if await_wgpu_map(slice).await {
653            let data = slice
654                .get_mapped_range()
655                .expect("wgpu buffer map_range failed");
656            let floats: &[f32] = bytemuck::cast_slice(&data);
657            if let Some(out) = readback_out {
658                out[..readback_elems].copy_from_slice(&floats[..readback_elems]);
659            }
660            drop(data);
661            staging.unmap();
662            return true;
663        }
664        let _ = staging.unmap();
665        false
666    }
667    #[cfg(all(target_arch = "wasm32", feature = "wasm-llm-diagnostics"))]
668    pub(crate) async fn dispatch_attention_q_ffn_token_async(
669        &mut self,
670        index: &crate::gguf_sharder::GgufTensorIndex,
671        layer: u32,
672        token_idx: u32,
673        hidden: &mut [f32],
674        emb_dim: usize,
675        tensors: &crate::gguf_sharder::LayerTensors,
676        scratch_a: &mut [f32],
677        scratch_b: &mut [f32],
678    ) -> bool {
679        let layout = match self.kv_layout {
680            Some(l) => l,
681            None => return false,
682        };
683        let q_info = match tensors.attn_q.as_ref() {
684            Some(i) => i,
685            None => return false,
686        };
687        let h = index.hyperparams;
688        let n_head = h.n_head as usize;
689        let head_dim = h.head_dim() as usize;
690        let q_dim = n_head * head_dim;
691        if q_dim > scratch_a.len() || q_dim > scratch_b.len() || emb_dim < h.n_embd as usize {
692            return false;
693        }
694        let mmap = match self.gguf_mmap.as_deref() {
695            Some(m) => m,
696            None => return false,
697        };
698        let q_raw =
699            match crate::ggml_quants::fetch_tensor_bytes(mmap, index.tensor_data_start, q_info) {
700                Ok(s) => s,
701                Err(_) => return false,
702            };
703        let n_embd = h.n_embd as usize;
704        let mut norm_w_attn = [0f32; MAX_HIDDEN_DIM];
705        let mut h_norm_attn = [0f32; MAX_HIDDEN_DIM];
706        let hidden_input = prepare_pre_norm_input(
707            &hidden[..emb_dim],
708            emb_dim,
709            tensors.attn_norm.as_ref(),
710            Some(&mmap[..]),
711            index.tensor_data_start,
712            &mut h_norm_attn,
713            &mut norm_w_attn,
714        );
715        if !self
716            .dispatch_attention_pass_async(
717                hidden_input,
718                n_embd,
719                1,
720                token_idx,
721                &layout,
722                layer,
723                token_idx,
724                &h,
725                q_info,
726                q_raw,
727                0,
728                n_head as u32,
729                Some(&mut scratch_b[..q_dim]),
730            )
731            .await
732        {
733            return false;
734        }
735        let mut attn_ok = false;
736        if let Some(out_info) = tensors.attn_output.as_ref() {
737            let (o_in, o_out) = Self::matmul_dims(out_info);
738            if o_in <= q_dim
739                && self
740                    .dispatch_gemm_into_async(
741                        index,
742                        out_info,
743                        &scratch_b[..o_in],
744                        &mut scratch_a[..o_out],
745                        o_in,
746                        o_out,
747                    )
748                    .await
749            {
750                add_residual_inplace(
751                    &mut hidden[..emb_dim],
752                    &scratch_a[..o_out],
753                    emb_dim.min(o_out),
754                );
755                attn_ok = true;
756            }
757        } else {
758            let n = q_dim.min(emb_dim);
759            add_residual_inplace(&mut hidden[..emb_dim], &scratch_b[..n], n);
760            attn_ok = true;
761        }
762        if !attn_ok {
763            return false;
764        }
765        self.dispatch_ffn_block_pre_norm_async(
766            index, hidden, emb_dim, tensors, scratch_a, scratch_b,
767        )
768        .await
769    }
770}