Skip to main content

qualia_core_db/gguf_bridge/
output.rs

1//! Output projection: chunked argmax + GPU top-k, output RMSNorm
2//! Split from gguf_bridge/mod.rs (structural refactor; no behaviour change).
3use super::*;
4
5impl QTensorEngine {
6    /// Chunked vocabulary projection with streaming argmax (zero heap, stack chunk buffer only).
7    /// `max_chunks`: `0` sweeps the full vocabulary; otherwise caps chunk iterations (tests).
8    pub fn dispatch_output_argmax_chunked(
9        &self,
10        index: &crate::gguf_sharder::GgufTensorIndex,
11        hidden: &[f32],
12        emb_dim: usize,
13        chunk_logits: &mut [f32],
14        max_chunks: u32,
15        sieve_mask: Option<&crate::neuro_symbolic_sieve::SieveStateMask>,
16    ) -> Option<StreamingArgmaxResult> {
17        let info = index.logits_projection_info()?;
18        let (n_in, vocab_size) = Self::matmul_dims(info);
19        if n_in == 0 || vocab_size == 0 || n_in > emb_dim || n_in > hidden.len() {
20            return None;
21        }
22        if chunk_logits.len() < VOCAB_CHUNK_ROWS {
23            return None;
24        }
25        let mmap = self.gguf_mmap.as_deref()?;
26        let full_chunks = vocab_size.div_ceil(VOCAB_CHUNK_ROWS);
27        let n_chunks = if max_chunks == 0 {
28            full_chunks
29        } else {
30            (max_chunks as usize).min(full_chunks)
31        };
32        let mut best_token_id = 0u32;
33        let mut max_logit = f32::NEG_INFINITY;
34
35        for chunk_idx in 0..n_chunks {
36            let row_start = chunk_idx * VOCAB_CHUNK_ROWS;
37            let chunk_rows = VOCAB_CHUNK_ROWS.min(vocab_size - row_start);
38            let raw = crate::ggml_quants::fetch_tensor_row_range_bytes(
39                mmap,
40                index.tensor_data_start,
41                info,
42                row_start,
43                chunk_rows,
44            )
45            .ok()?;
46            if !self.dispatch_gemm_raw_into(
47                info,
48                raw,
49                &hidden[..n_in],
50                &mut chunk_logits[..chunk_rows],
51                n_in,
52                chunk_rows,
53            ) {
54                return None;
55            }
56            if let Some(mask) = sieve_mask {
57                update_streaming_argmax_sieved(
58                    &chunk_logits[..chunk_rows],
59                    chunk_rows,
60                    chunk_idx,
61                    Some(mask),
62                    &mut best_token_id,
63                    &mut max_logit,
64                );
65            } else {
66                update_streaming_argmax(
67                    &chunk_logits[..chunk_rows],
68                    chunk_rows,
69                    chunk_idx,
70                    &mut best_token_id,
71                    &mut max_logit,
72                );
73            }
74            scrub_f32_volatile(&mut chunk_logits[..chunk_rows], chunk_rows);
75        }
76
77        if max_logit == f32::NEG_INFINITY {
78            return None;
79        }
80        Some(StreamingArgmaxResult {
81            best_token_id,
82            max_logit,
83        })
84    }
85
86    /// A1a: create the persistent GPU top-k pipeline + small candidate/staging buffers (once).
87    pub(crate) fn init_output_topk(&mut self) {
88        let shader = self
89            .gpu_device()
90            .create_shader_module(wgpu::ShaderModuleDescriptor {
91                label: Some("output_topk"),
92                source: wgpu::ShaderSource::Wgsl(crate::topk::TOPK_REDUCTION_WGSL.into()),
93            });
94        let pipeline =
95            self.gpu_device()
96                .create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
97                    label: Some("output_topk_pipeline"),
98                    layout: None,
99                    module: &shader,
100                    entry_point: Some("topk_block"),
101                    compilation_options: Default::default(),
102                    cache: {
103                        #[cfg(not(target_arch = "wasm32"))]
104                        {
105                            self.native_pipeline_cache_ref()
106                        }
107                        #[cfg(target_arch = "wasm32")]
108                        {
109                            None
110                        }
111                    },
112                });
113        let pipeline_layout = pipeline.get_bind_group_layout(0);
114        // Multi-chunk mega-pass: hold candidates for a full large-vocab sweep (≤256k ids),
115        // k=1 → one cand per block; oversize for TOPK_MAX_K headroom on smaller vocabs.
116        const MAX_VOCAB_RESIDENT: usize = 262_144;
117        let max_blocks = (MAX_VOCAB_RESIDENT / crate::topk::TOPK_BLOCK_SIZE).max(1);
118        let cand_bytes = ((max_blocks * crate::topk::TOPK_MAX_K).max(1) * 4) as wgpu::BufferAddress;
119        self.topk_cand_val_buf = Some(self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
120            label: Some("TopkCandVal"),
121            size: cand_bytes,
122            usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
123            mapped_at_creation: false,
124        }));
125        self.topk_cand_idx_buf = Some(self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
126            label: Some("TopkCandIdx"),
127            size: cand_bytes,
128            usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
129            mapped_at_creation: false,
130        }));
131        self.topk_cand_staging = Some(self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
132            label: Some("TopkCandStaging"),
133            size: cand_bytes * 2, // packed: [val .. | idx ..]
134            usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
135            mapped_at_creation: false,
136        }));
137        self.topk_params_buf = Some(self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
138            label: Some("TopkParams"),
139            // Browser top-1 uses one 256-byte-aligned slot per vocabulary
140            // chunk; native paths continue to bind the first 16 bytes.
141            size: 32 * 256,
142            usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
143            mapped_at_creation: false,
144        }));
145        self.output_topk_bind_layout = Some(pipeline_layout);
146        self.output_topk_pipeline = Some(pipeline);
147    }
148
149    /// Native decode fast path: output projection plus GPU block argmax (`k=1`) with one tiny
150    /// candidate readback after all vocab chunks have been submitted. This avoids the full
151    /// chunk-logit readback in [`Self::dispatch_output_argmax_chunked`] and avoids heap allocation
152    /// in the decode loop.
153    #[cfg(not(target_arch = "wasm32"))]
154    pub fn dispatch_output_top1_chunked(
155        &self,
156        index: &crate::gguf_sharder::GgufTensorIndex,
157        hidden: &[f32],
158        emb_dim: usize,
159    ) -> Option<StreamingArgmaxResult> {
160        let info = index.logits_projection_info()?;
161        let (n_in, vocab_size) = Self::matmul_dims(info);
162        if n_in == 0 || vocab_size == 0 || n_in > emb_dim || n_in > hidden.len() {
163            return None;
164        }
165        let topk_pipeline = self.output_topk_pipeline.as_ref()?;
166        let input_buf = self.gemm_input_buf.as_ref()?;
167        let weight_buf = self.gemm_weight_buf.as_ref()?;
168        let output_buf = self.gemm_output_buf.as_ref()?;
169        let params_buf = self.gemm_params_buf.as_ref()?;
170        let topk_params_buf = self.topk_params_buf.as_ref()?;
171        let cand_val = self.topk_cand_val_buf.as_ref()?;
172        let cand_idx = self.topk_cand_idx_buf.as_ref()?;
173        let staging = self.topk_cand_staging.as_ref()?;
174        let mmap = self.gguf_mmap.as_deref()?;
175
176        let use_coop = crate::llm_bench::coop_gemv_enabled();
177        let use_mr = use_coop && info.ggml_type == crate::ggml_quants::GGML_TYPE_Q4_K_SOA;
178        let gemm_pipeline: &wgpu::ComputePipeline = if use_mr {
179            &self.coop_gemv_mr_pipeline
180        } else if use_coop {
181            &self.coop_gemv_pipeline
182        } else {
183            &self.pipeline
184        };
185        let gemm_layout = self.native_gemm_bind_layout(use_coop).clone();
186        let topk_layout = self.output_topk_bind_layout.as_ref()?;
187        let topk_bind = self
188            .gpu_device()
189            .create_bind_group(&wgpu::BindGroupDescriptor {
190                label: Some("Top1Bind"),
191                layout: topk_layout,
192                entries: &[
193                    wgpu::BindGroupEntry {
194                        binding: 0,
195                        resource: output_buf.as_entire_binding(),
196                    },
197                    wgpu::BindGroupEntry {
198                        binding: 1,
199                        resource: topk_params_buf.as_entire_binding(),
200                    },
201                    wgpu::BindGroupEntry {
202                        binding: 2,
203                        resource: cand_val.as_entire_binding(),
204                    },
205                    wgpu::BindGroupEntry {
206                        binding: 3,
207                        resource: cand_idx.as_entire_binding(),
208                    },
209                ],
210            });
211        let block_size = crate::topk::TOPK_BLOCK_SIZE;
212        let full_chunks = vocab_size.div_ceil(VOCAB_CHUNK_ROWS);
213        let total_cands = vocab_size.div_ceil(block_size);
214        let cand_capacity = VOCAB_CHUNK_ROWS
215            .div_ceil(block_size)
216            .max(1)
217            .saturating_mul(crate::topk::TOPK_MAX_K);
218        if total_cands == 0 || total_cands > cand_capacity {
219            return None;
220        }
221        let resident_logits = self.mc8_logits_resident_buf.as_ref();
222        let resident_row_bytes = self.mc8_logits_row_bytes as u64;
223
224        self.gpu_queue()
225            .write_buffer(input_buf, 0, bytemuck::cast_slice(&hidden[..n_in]));
226
227        // Fast path: resident logits → ONE submit for all vocab chunks (no per-chunk fence).
228        // Shared uniform buffers force multi-submit when weights must be re-uploaded each chunk.
229        if let Some(res_buf) = resident_logits {
230            if !(ggml_gpu_gemm_supported(info.ggml_type) && n_in <= MAX_STACK_GEMM_IN) {
231                return None;
232            }
233            // 256-byte aligned slots (wgpu min uniform offset) for per-chunk params.
234            const SLOT: usize = 256;
235            let gemm_slot = std::mem::size_of::<GemmGpuParams>().max(32);
236            let mut gemm_slab = vec![0u8; full_chunks * SLOT];
237            let mut topk_slab = vec![0u8; full_chunks * SLOT];
238            let mut chunk_meta: Vec<(u32, u32, u32)> = Vec::with_capacity(full_chunks); // rows, weight_bytes, n_blocks
239            for chunk_idx in 0..full_chunks {
240                let row_start = chunk_idx * VOCAB_CHUNK_ROWS;
241                let chunk_rows = VOCAB_CHUNK_ROWS.min(vocab_size - row_start);
242                if chunk_rows > self.gemm_max_out_dim as usize {
243                    return None;
244                }
245                let weight_byte_len = (chunk_rows as u64 * resident_row_bytes) as u32;
246                let gparams = GemmGpuParams {
247                    n_in: n_in as u32,
248                    n_out: chunk_rows as u32,
249                    weight_ggml_type: info.ggml_type,
250                    weight_row_elems: info.dims[0] as u32,
251                    weight_byte_len,
252                    n_batch: 1,
253                    in_row_stride: 0,
254                    out_row_stride: 0,
255                };
256                let gp = bytemuck::bytes_of(&gparams);
257                let go = chunk_idx * SLOT;
258                gemm_slab[go..go + gemm_slot.min(gp.len())]
259                    .copy_from_slice(&gp[..gemm_slot.min(gp.len())]);
260                let tparams =
261                    crate::topk::topk_params_bytes(chunk_rows as u32, 1, block_size as u32);
262                topk_slab[go..go + tparams.len()].copy_from_slice(&tparams);
263                let num_blocks = chunk_rows.div_ceil(block_size) as u32;
264                chunk_meta.push((chunk_rows as u32, weight_byte_len, num_blocks));
265            }
266            // Upload slabs once; bind with per-chunk offsets via dedicated per-chunk param buffers
267            // (auto layouts rarely allow dynamic offsets). Reuse gemm_params / topk_params only for
268            // the first slot write pattern: create a single multi-slot buffer pair for the fused pass.
269            let gemm_multi = self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
270                label: Some("Top1GemmParamsMulti"),
271                size: gemm_slab.len() as u64,
272                usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
273                mapped_at_creation: false,
274            });
275            let topk_multi = self.gpu_device().create_buffer(&wgpu::BufferDescriptor {
276                label: Some("Top1TopkParamsMulti"),
277                size: topk_slab.len() as u64,
278                usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
279                mapped_at_creation: false,
280            });
281            self.gpu_queue().write_buffer(&gemm_multi, 0, &gemm_slab);
282            self.gpu_queue().write_buffer(&topk_multi, 0, &topk_slab);
283
284            let mut encoder =
285                self.device()
286                    .create_command_encoder(&wgpu::CommandEncoderDescriptor {
287                        label: Some("Top1FusedEncoder"),
288                    });
289            let mut cand_offset = 0usize;
290            for (chunk_idx, &(chunk_rows, _wlen, num_blocks)) in chunk_meta.iter().enumerate() {
291                let row_start = chunk_idx * VOCAB_CHUNK_ROWS;
292                let byte_len = chunk_rows as u64 * resident_row_bytes;
293                let weight_resource = wgpu::BindingResource::Buffer(wgpu::BufferBinding {
294                    buffer: res_buf,
295                    offset: row_start as u64 * resident_row_bytes,
296                    size: std::num::NonZeroU64::new(byte_len),
297                });
298                let go = (chunk_idx * SLOT) as u64;
299                let gsize = std::num::NonZeroU64::new(std::mem::size_of::<GemmGpuParams>() as u64);
300                let tsize = std::num::NonZeroU64::new(16);
301                let mut gemm_entries = vec![
302                    wgpu::BindGroupEntry {
303                        binding: 0,
304                        resource: input_buf.as_entire_binding(),
305                    },
306                    wgpu::BindGroupEntry {
307                        binding: 1,
308                        resource: weight_resource,
309                    },
310                    wgpu::BindGroupEntry {
311                        binding: 2,
312                        resource: wgpu::BindingResource::Buffer(wgpu::BufferBinding {
313                            buffer: &gemm_multi,
314                            offset: go,
315                            size: gsize,
316                        }),
317                    },
318                    wgpu::BindGroupEntry {
319                        binding: 3,
320                        resource: output_buf.as_entire_binding(),
321                    },
322                ];
323                if use_coop {
324                    gemm_entries.push(wgpu::BindGroupEntry {
325                        binding: 4,
326                        resource: input_buf.as_entire_binding(),
327                    });
328                }
329                let gemm_bind = self
330                    .gpu_device()
331                    .create_bind_group(&wgpu::BindGroupDescriptor {
332                        label: Some("Top1GemmBindFused"),
333                        layout: &gemm_layout,
334                        entries: &gemm_entries,
335                    });
336                let topk_bind_c = self
337                    .gpu_device()
338                    .create_bind_group(&wgpu::BindGroupDescriptor {
339                        label: Some("Top1TopkBindFused"),
340                        layout: topk_layout,
341                        entries: &[
342                            wgpu::BindGroupEntry {
343                                binding: 0,
344                                resource: output_buf.as_entire_binding(),
345                            },
346                            wgpu::BindGroupEntry {
347                                binding: 1,
348                                resource: wgpu::BindingResource::Buffer(wgpu::BufferBinding {
349                                    buffer: &topk_multi,
350                                    offset: go,
351                                    size: tsize,
352                                }),
353                            },
354                            wgpu::BindGroupEntry {
355                                binding: 2,
356                                resource: cand_val.as_entire_binding(),
357                            },
358                            wgpu::BindGroupEntry {
359                                binding: 3,
360                                resource: cand_idx.as_entire_binding(),
361                            },
362                        ],
363                    });
364                {
365                    let mut cpass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
366                        label: Some("Top1GemmPass"),
367                        timestamp_writes: None,
368                    });
369                    cpass.set_pipeline(gemm_pipeline);
370                    cpass.set_bind_group(0, &gemm_bind, &[]);
371                    if use_mr {
372                        cpass.dispatch_workgroups(
373                            crate::llm_bench::coop_gemv_workgroups(chunk_rows),
374                            1,
375                            1,
376                        );
377                    } else if use_coop {
378                        cpass.dispatch_workgroups(chunk_rows, 1, 1);
379                    } else {
380                        cpass.dispatch_workgroups((chunk_rows + 63) / 64, 1, 1);
381                    }
382                }
383                {
384                    let mut tpass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
385                        label: Some("Top1ReducePass"),
386                        timestamp_writes: None,
387                    });
388                    tpass.set_pipeline(topk_pipeline);
389                    tpass.set_bind_group(0, &topk_bind_c, &[]);
390                    tpass.dispatch_workgroups(num_blocks, 1, 1);
391                }
392                let cand_count = num_blocks as usize;
393                let cand_bytes = (cand_count * 4) as wgpu::BufferAddress;
394                let val_dst = (cand_offset * 4) as wgpu::BufferAddress;
395                let idx_dst = ((total_cands + cand_offset) * 4) as wgpu::BufferAddress;
396                encoder.copy_buffer_to_buffer(cand_val, 0, staging, val_dst, cand_bytes);
397                encoder.copy_buffer_to_buffer(cand_idx, 0, staging, idx_dst, cand_bytes);
398                cand_offset += cand_count;
399            }
400            crate::llm_gpu_profiler::resolve(&mut encoder);
401            self.gpu_queue().submit(Some(encoder.finish()));
402            crate::llm_gpu_profiler::accumulate(crate::llm_gpu_profiler::Phase::OutputTopk);
403        } else {
404            // Legacy: per-chunk upload + submit (shared weight staging buffer).
405            let mut cand_offset = 0usize;
406            for chunk_idx in 0..full_chunks {
407                let row_start = chunk_idx * VOCAB_CHUNK_ROWS;
408                let chunk_rows = VOCAB_CHUNK_ROWS.min(vocab_size - row_start);
409                let raw = crate::ggml_quants::fetch_tensor_row_range_bytes(
410                    mmap,
411                    index.tensor_data_start,
412                    info,
413                    row_start,
414                    chunk_rows,
415                )
416                .ok()?;
417                if !(ggml_gpu_gemm_supported(info.ggml_type)
418                    && n_in <= MAX_STACK_GEMM_IN
419                    && chunk_rows <= self.gemm_max_out_dim as usize
420                    && raw.len() <= self.max_tensor_bytes)
421                {
422                    return None;
423                }
424                let byte_len = raw.len() as u32;
425                let resident = if crate::llm_bench::resident_weights_enabled() {
426                    self.resident_weight_buffer(raw.as_ptr() as u64, raw)
427                } else {
428                    None
429                };
430                let weight_binding: &wgpu::Buffer = match resident.as_ref() {
431                    Some(r) => r,
432                    None => {
433                        self.write_weight_words(raw, self.max_tensor_bytes);
434                        weight_buf
435                    }
436                };
437
438                let gparams = GemmGpuParams {
439                    n_in: n_in as u32,
440                    n_out: chunk_rows as u32,
441                    weight_ggml_type: info.ggml_type,
442                    weight_row_elems: info.dims[0] as u32,
443                    weight_byte_len: byte_len,
444                    n_batch: 1,
445                    in_row_stride: 0,
446                    out_row_stride: 0,
447                };
448                self.gpu_queue()
449                    .write_buffer(params_buf, 0, bytemuck::bytes_of(&gparams));
450                let tparams =
451                    crate::topk::topk_params_bytes(chunk_rows as u32, 1, block_size as u32);
452                self.gpu_queue().write_buffer(topk_params_buf, 0, &tparams);
453
454                let num_blocks = chunk_rows.div_ceil(block_size);
455                let cand_count = num_blocks;
456
457                let mut gemm_entries = vec![
458                    wgpu::BindGroupEntry {
459                        binding: 0,
460                        resource: input_buf.as_entire_binding(),
461                    },
462                    wgpu::BindGroupEntry {
463                        binding: 1,
464                        resource: weight_binding.as_entire_binding(),
465                    },
466                    wgpu::BindGroupEntry {
467                        binding: 2,
468                        resource: params_buf.as_entire_binding(),
469                    },
470                    wgpu::BindGroupEntry {
471                        binding: 3,
472                        resource: output_buf.as_entire_binding(),
473                    },
474                ];
475                if use_coop {
476                    gemm_entries.push(wgpu::BindGroupEntry {
477                        binding: 4,
478                        resource: input_buf.as_entire_binding(),
479                    });
480                }
481                let gemm_bind = self
482                    .gpu_device()
483                    .create_bind_group(&wgpu::BindGroupDescriptor {
484                        label: Some("Top1GemmBind"),
485                        layout: &gemm_layout,
486                        entries: &gemm_entries,
487                    });
488                let mut encoder =
489                    self.device()
490                        .create_command_encoder(&wgpu::CommandEncoderDescriptor {
491                            label: Some("Top1Encoder"),
492                        });
493                {
494                    let mut cpass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
495                        label: Some("Top1GemmPass"),
496                        timestamp_writes: crate::llm_gpu_profiler::pass_writes_begin(),
497                    });
498                    cpass.set_pipeline(gemm_pipeline);
499                    cpass.set_bind_group(0, &gemm_bind, &[]);
500                    if use_mr {
501                        cpass.dispatch_workgroups(
502                            crate::llm_bench::coop_gemv_workgroups(chunk_rows as u32),
503                            1,
504                            1,
505                        );
506                    } else if use_coop {
507                        cpass.dispatch_workgroups(chunk_rows as u32, 1, 1);
508                    } else {
509                        cpass.dispatch_workgroups((chunk_rows as u32 + 63) / 64, 1, 1);
510                    }
511                }
512                {
513                    let mut tpass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
514                        label: Some("Top1ReducePass"),
515                        timestamp_writes: crate::llm_gpu_profiler::pass_writes_end(),
516                    });
517                    tpass.set_pipeline(topk_pipeline);
518                    tpass.set_bind_group(0, &topk_bind, &[]);
519                    tpass.dispatch_workgroups(num_blocks as u32, 1, 1);
520                }
521                let cand_bytes = (cand_count * 4) as wgpu::BufferAddress;
522                let val_dst = (cand_offset * 4) as wgpu::BufferAddress;
523                let idx_dst = ((total_cands + cand_offset) * 4) as wgpu::BufferAddress;
524                encoder.copy_buffer_to_buffer(cand_val, 0, staging, val_dst, cand_bytes);
525                encoder.copy_buffer_to_buffer(cand_idx, 0, staging, idx_dst, cand_bytes);
526                crate::llm_gpu_profiler::resolve(&mut encoder);
527                self.gpu_queue().submit(Some(encoder.finish()));
528                crate::llm_gpu_profiler::accumulate(crate::llm_gpu_profiler::Phase::OutputTopk);
529                cand_offset += cand_count;
530            }
531        }
532
533        let map_bytes = (total_cands * 8) as wgpu::BufferAddress;
534        let slice = staging.slice(..map_bytes);
535        let (tx, rx) = futures_channel::oneshot::channel();
536        slice.map_async(wgpu::MapMode::Read, move |r| {
537            let _ = tx.send(r);
538        });
539        self.poll_wait();
540        let mapped_ok = if let Ok(handle) = tokio::runtime::Handle::try_current() {
541            handle.block_on(rx).ok().map(|m| m.is_ok()).unwrap_or(false)
542        } else {
543            false
544        };
545        if !mapped_ok {
546            let _ = staging.unmap();
547            return None;
548        }
549
550        let mut best_token_id = 0u32;
551        let mut max_logit = f32::NEG_INFINITY;
552        {
553            let data = slice
554                .get_mapped_range()
555                .expect("wgpu buffer map_range failed");
556            let val_bytes = total_cands * 4;
557            let vals: &[f32] = bytemuck::cast_slice(&data[..val_bytes]);
558            let idxs: &[u32] = bytemuck::cast_slice(&data[val_bytes..val_bytes * 2]);
559            let mut offset = 0usize;
560            for chunk_idx in 0..full_chunks {
561                let row_start = chunk_idx * VOCAB_CHUNK_ROWS;
562                let chunk_rows = VOCAB_CHUNK_ROWS.min(vocab_size - row_start);
563                let cand_count = chunk_rows.div_ceil(block_size);
564                for i in 0..cand_count {
565                    let pos = offset + i;
566                    let v = vals[pos];
567                    let token_id = row_start as u32 + idxs[pos];
568                    if v > f32::NEG_INFINITY
569                        && (v > max_logit || (v == max_logit && token_id < best_token_id))
570                    {
571                        max_logit = v;
572                        best_token_id = token_id;
573                    }
574                }
575                offset += cand_count;
576            }
577        }
578        staging.unmap();
579
580        if max_logit == f32::NEG_INFINITY {
581            None
582        } else {
583            Some(StreamingArgmaxResult {
584                best_token_id,
585                max_logit,
586            })
587        }
588    }
589
590    /// A1a: GPU top-k over the output projection — the logits stay on-GPU (`gemm_output_buf`), the
591    /// top-k kernel reduces them per chunk, and only K `(id, logit)` candidates are read back (vs the
592    /// 196 KB/token full-logit readback + CPU argmax in `dispatch_output_argmax_chunked`). Returns the
593    /// merged global top-K, or `None` to signal the caller to fall back to the argmax path. v1: no
594    /// sieve coupling (caller routes here only when no sieve mask is active).
595    #[cfg(not(target_arch = "wasm32"))]
596    pub fn dispatch_output_topk_chunked(
597        &self,
598        index: &crate::gguf_sharder::GgufTensorIndex,
599        hidden: &[f32],
600        emb_dim: usize,
601        k: usize,
602    ) -> Option<Vec<crate::topk::TopKItem>> {
603        let info = index.logits_projection_info()?;
604        let (n_in, vocab_size) = Self::matmul_dims(info);
605        if n_in == 0 || vocab_size == 0 || n_in > emb_dim || n_in > hidden.len() {
606            return None;
607        }
608        let pipeline = self.output_topk_pipeline.as_ref()?;
609        let topk_layout = self.output_topk_bind_layout.as_ref()?;
610        let input_buf = self.gemm_input_buf.as_ref()?;
611        let weight_buf = self.gemm_weight_buf.as_ref()?;
612        let output_buf = self.gemm_output_buf.as_ref()?;
613        let params_buf = self.gemm_params_buf.as_ref()?;
614        let topk_params_buf = self.topk_params_buf.as_ref()?;
615        let cand_val = self.topk_cand_val_buf.as_ref()?;
616        let cand_idx = self.topk_cand_idx_buf.as_ref()?;
617        let staging = self.topk_cand_staging.as_ref()?;
618        let mmap = self.gguf_mmap.as_deref()?;
619
620        let k = k.clamp(1, crate::topk::TOPK_MAX_K);
621        let block_size = crate::topk::TOPK_BLOCK_SIZE;
622        let full_chunks = vocab_size.div_ceil(VOCAB_CHUNK_ROWS);
623        let use_coop = crate::llm_bench::coop_gemv_enabled();
624        let use_mr = use_coop && info.ggml_type == crate::ggml_quants::GGML_TYPE_Q4_K_SOA;
625        let gemm_pipeline: &wgpu::ComputePipeline = if use_mr {
626            &self.coop_gemv_mr_pipeline
627        } else if use_coop {
628            &self.coop_gemv_pipeline
629        } else {
630            &self.pipeline
631        };
632        let gemm_layout = self.native_gemm_bind_layout(use_coop).clone();
633        let topk_bind = self
634            .gpu_device()
635            .create_bind_group(&wgpu::BindGroupDescriptor {
636                label: Some("TopkBind"),
637                layout: topk_layout,
638                entries: &[
639                    wgpu::BindGroupEntry {
640                        binding: 0,
641                        resource: output_buf.as_entire_binding(),
642                    },
643                    wgpu::BindGroupEntry {
644                        binding: 1,
645                        resource: topk_params_buf.as_entire_binding(),
646                    },
647                    wgpu::BindGroupEntry {
648                        binding: 2,
649                        resource: cand_val.as_entire_binding(),
650                    },
651                    wgpu::BindGroupEntry {
652                        binding: 3,
653                        resource: cand_idx.as_entire_binding(),
654                    },
655                ],
656            });
657
658        let mut all_val: Vec<f32> = Vec::new();
659        let mut all_idx: Vec<u32> = Vec::new();
660
661        // A1a step-2: when the output projection is resident (uploaded once at init), bind the
662        // per-chunk sub-range - zero per-token upload. `VOCAB_CHUNK_ROWS` is a multiple of 256, so
663        // every chunk offset is storage-binding aligned. The bound bytes, quant, shader and params
664        // are identical to the per-chunk-upload fallback, so logits are byte-for-byte equal.
665        let resident_logits = self.mc8_logits_resident_buf.as_ref();
666        let resident_row_bytes = self.mc8_logits_row_bytes as u64;
667
668        for chunk_idx in 0..full_chunks {
669            let row_start = chunk_idx * VOCAB_CHUNK_ROWS;
670            let chunk_rows = VOCAB_CHUNK_ROWS.min(vocab_size - row_start);
671
672            let (weight_resource, weight_byte_len) = if let Some(buf) = resident_logits {
673                // Only the GPU-quant fast path is supported; otherwise signal fallback to argmax.
674                if !(ggml_gpu_gemm_supported(info.ggml_type)
675                    && n_in <= MAX_STACK_GEMM_IN
676                    && chunk_rows <= self.gemm_max_out_dim as usize)
677                {
678                    return None;
679                }
680                let byte_len = chunk_rows as u64 * resident_row_bytes;
681                let res = wgpu::BindingResource::Buffer(wgpu::BufferBinding {
682                    buffer: buf,
683                    offset: row_start as u64 * resident_row_bytes,
684                    size: std::num::NonZeroU64::new(byte_len),
685                });
686                (res, byte_len as u32)
687            } else {
688                let raw = crate::ggml_quants::fetch_tensor_row_range_bytes(
689                    mmap,
690                    index.tensor_data_start,
691                    info,
692                    row_start,
693                    chunk_rows,
694                )
695                .ok()?;
696                if !(ggml_gpu_gemm_supported(info.ggml_type)
697                    && n_in <= MAX_STACK_GEMM_IN
698                    && chunk_rows <= self.gemm_max_out_dim as usize
699                    && raw.len() <= self.max_tensor_bytes)
700                {
701                    return None;
702                }
703                let byte_len = raw.len() as u32;
704                self.write_weight_words(raw, self.max_tensor_bytes);
705                (weight_buf.as_entire_binding(), byte_len)
706            };
707
708            let gparams = GemmGpuParams {
709                n_in: n_in as u32,
710                n_out: chunk_rows as u32,
711                weight_ggml_type: info.ggml_type,
712                weight_row_elems: info.dims[0] as u32,
713                weight_byte_len,
714                n_batch: 1,
715                in_row_stride: 0,
716                out_row_stride: 0,
717            };
718            self.gpu_queue()
719                .write_buffer(input_buf, 0, bytemuck::cast_slice(&hidden[..n_in]));
720            self.gpu_queue()
721                .write_buffer(params_buf, 0, bytemuck::bytes_of(&gparams));
722            let tparams =
723                crate::topk::topk_params_bytes(chunk_rows as u32, k as u32, block_size as u32);
724            self.gpu_queue().write_buffer(topk_params_buf, 0, &tparams);
725
726            let num_blocks = chunk_rows.div_ceil(block_size);
727            let cand_count = num_blocks * k;
728
729            let mut gemm_entries = vec![
730                wgpu::BindGroupEntry {
731                    binding: 0,
732                    resource: input_buf.as_entire_binding(),
733                },
734                wgpu::BindGroupEntry {
735                    binding: 1,
736                    resource: weight_resource,
737                },
738                wgpu::BindGroupEntry {
739                    binding: 2,
740                    resource: params_buf.as_entire_binding(),
741                },
742                wgpu::BindGroupEntry {
743                    binding: 3,
744                    resource: output_buf.as_entire_binding(),
745                },
746            ];
747            if use_coop {
748                gemm_entries.push(wgpu::BindGroupEntry {
749                    binding: 4,
750                    resource: input_buf.as_entire_binding(),
751                });
752            }
753            let gemm_bind = self
754                .gpu_device()
755                .create_bind_group(&wgpu::BindGroupDescriptor {
756                    label: Some("TopkGemmBind"),
757                    layout: &gemm_layout,
758                    entries: &gemm_entries,
759                });
760            let mut encoder =
761                self.device()
762                    .create_command_encoder(&wgpu::CommandEncoderDescriptor {
763                        label: Some("TopkEncoder"),
764                    });
765            {
766                let mut cpass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
767                    label: Some("TopkGemmPass"),
768                    timestamp_writes: crate::llm_gpu_profiler::pass_writes_begin(),
769                });
770                cpass.set_pipeline(gemm_pipeline);
771                cpass.set_bind_group(0, &gemm_bind, &[]);
772                if use_mr {
773                    cpass.dispatch_workgroups(
774                        crate::llm_bench::coop_gemv_workgroups(chunk_rows as u32),
775                        1,
776                        1,
777                    );
778                } else if use_coop {
779                    cpass.dispatch_workgroups(chunk_rows as u32, 1, 1);
780                } else {
781                    cpass.dispatch_workgroups((chunk_rows as u32 + 63) / 64, 1, 1);
782                }
783            }
784            {
785                let mut tpass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
786                    label: Some("TopkReducePass"),
787                    timestamp_writes: crate::llm_gpu_profiler::pass_writes_end(),
788                });
789                tpass.set_pipeline(pipeline);
790                tpass.set_bind_group(0, &topk_bind, &[]);
791                tpass.dispatch_workgroups(num_blocks as u32, 1, 1);
792            }
793            let cand_bytes = (cand_count * 4) as wgpu::BufferAddress;
794            encoder.copy_buffer_to_buffer(cand_val, 0, staging, 0, cand_bytes);
795            encoder.copy_buffer_to_buffer(cand_idx, 0, staging, cand_bytes, cand_bytes);
796            crate::llm_gpu_profiler::resolve(&mut encoder);
797            self.gpu_queue().submit(Some(encoder.finish()));
798            crate::llm_gpu_profiler::accumulate(crate::llm_gpu_profiler::Phase::OutputTopk);
799
800            let map_bytes = cand_bytes * 2;
801            let slice = staging.slice(..map_bytes);
802            let (tx, rx) = futures_channel::oneshot::channel();
803            slice.map_async(wgpu::MapMode::Read, move |r| {
804                let _ = tx.send(r);
805            });
806            self.poll_wait();
807            let mapped_ok = if let Ok(handle) = tokio::runtime::Handle::try_current() {
808                handle.block_on(rx).ok().map(|m| m.is_ok()).unwrap_or(false)
809            } else {
810                false
811            };
812            if !mapped_ok {
813                let _ = staging.unmap();
814                return None;
815            }
816            {
817                let data = slice
818                    .get_mapped_range()
819                    .expect("wgpu buffer map_range failed");
820                let vals: &[f32] = bytemuck::cast_slice(&data[..cand_count * 4]);
821                let idxs: &[u32] = bytemuck::cast_slice(&data[cand_count * 4..cand_count * 8]);
822                for i in 0..cand_count {
823                    let v = vals[i];
824                    if v > f32::NEG_INFINITY {
825                        all_val.push(v);
826                        all_idx.push(row_start as u32 + idxs[i]);
827                    }
828                }
829            }
830            staging.unmap();
831        }
832
833        let top = crate::topk::merge_block_candidates(&all_val, &all_idx, k, None);
834        if top.is_empty() {
835            None
836        } else {
837            Some(top)
838        }
839    }
840
841    /// Final `output_norm` RMSNorm in-place before vocabulary projection (Pre-Norm LLM tail).
842    /// REQUIRED on all targets — native previously skipped it → logits from an un-normed hidden.
843    pub fn apply_output_norm_inplace(
844        &self,
845        index: &crate::gguf_sharder::GgufTensorIndex,
846        hidden: &mut [f32],
847        emb_dim: usize,
848    ) -> bool {
849        let info = match index.output_norm_info() {
850            Some(i) => i,
851            None => return true,
852        };
853        let mmap = match self.gguf_mmap.as_deref() {
854            Some(m) => m,
855            None => return false,
856        };
857        let n_embd = index.hyperparams.n_embd as usize;
858        let n = emb_dim.min(n_embd).min(hidden.len());
859        let mut norm_w = [0f32; MAX_HIDDEN_DIM];
860        if dequant_norm_row_into(mmap, index.tensor_data_start, info, &mut norm_w) < n {
861            return false;
862        }
863        rms_norm_inplace(&mut hidden[..n], &norm_w[..n], RMS_NORM_EPS);
864        true
865    }
866}