1use super::*;
5
6impl QTensorEngine {
7 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 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 #[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 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 #[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 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 #[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 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 #[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 #[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 #[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 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 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 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 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 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 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 #[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 #[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 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 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 #[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 #[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 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 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 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 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 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 #[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 self.gpu_queue()
1560 .write_buffer(batch_buf, 0, bytemuck::cast_slice(&hidden[..n_embd]));
1561
1562 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 let output_norm_info = match index.output_norm_info() {
1775 Some(i) => i,
1776 None => {
1777 self.mc8_flush(&mut enc);
1779 if !self.pipeline_read_hidden(emb_dim, hidden).await {
1780 return None;
1781 }
1782 return None; }
1784 };
1785 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 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 enc.encoder
1821 .copy_buffer_to_buffer(prefill_scratch, 0, batch_buf, 0, n_embd_bytes);
1822
1823 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(¶ms));
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 self.gpu_queue().submit(Some(enc.encoder.finish()));
1962 self.read_browser_top1(staging, top1_plan).await
1963 }
1964
1965 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}