1use super::*;
4
5impl QTensorEngine {
6 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 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 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, 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 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 #[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 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 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); 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 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 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 #[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 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 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 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}