1use super::*;
4
5impl QTensorEngine {
6 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(¶ms));
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 #[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(¶ms));
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 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(¶ms));
569 self.gpu_queue()
570 .write_buffer(mask_buf, 0, bytemuck::cast_slice(&mask_words));
571
572 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}