1use crate::Engine;
16use cudarc::driver::{CudaSlice, CudaView, DevicePtr, DevicePtrMut};
17
18static MMQ_ACT_EPOCH: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
21#[allow(clippy::type_complexity)]
22static MMQ_ACT_SLOT: std::sync::Mutex<Option<(u64, u64, usize, usize, CudaSlice<u8>)>> =
23 std::sync::Mutex::new(None);
24static MMQ_FIXUP_SLOT: std::sync::Mutex<Option<cudarc::driver::CudaSlice<u8>>> =
26 std::sync::Mutex::new(None);
27
28pub struct ExpertCsr {
34 ex_ids: Vec<i32>,
35 ex_off: Vec<i32>,
36 ex_pairs: Vec<i32>,
37 pair_tok: Vec<i32>,
38 n_expert: usize,
39 n_tokens: usize,
40}
41
42impl ExpertCsr {
43 pub fn from_token_routes(
45 n_expert: usize,
46 n_tokens: usize,
47 experts_per_token: usize,
48 selected: &[usize],
49 ) -> Result<Self, String> {
50 if experts_per_token == 0 {
51 return Err("grouped expert CSR experts/token must be nonzero".into());
52 }
53 let n_pairs = n_tokens
54 .checked_mul(experts_per_token)
55 .ok_or("grouped expert CSR route count overflow")?;
56 if selected.len() != n_pairs {
57 return Err(format!(
58 "grouped expert CSR selected routes {} != {n_tokens}x{experts_per_token} \
59 ({n_pairs})",
60 selected.len()
61 ));
62 }
63 let pair_tok = (0..n_pairs)
64 .map(|pair| pair / experts_per_token)
65 .collect::<Vec<_>>();
66 Self::from_pair_rows(n_expert, n_tokens, selected, &pair_tok)
67 }
68
69 pub fn from_pair_rows(
74 n_expert: usize,
75 n_tokens: usize,
76 selected: &[usize],
77 pair_tok: &[usize],
78 ) -> Result<Self, String> {
79 if n_expert == 0 || n_tokens == 0 || selected.is_empty() {
80 return Err("grouped expert CSR requires non-empty experts, tokens, and pairs".into());
81 }
82 if n_expert > i32::MAX as usize
83 || n_tokens > i32::MAX as usize
84 || selected.len() > i32::MAX as usize
85 {
86 return Err("grouped expert CSR dimensions exceed the i32 kernel ABI".into());
87 }
88 if pair_tok.len() != selected.len() {
89 return Err(format!(
90 "grouped expert CSR pair rows {} != selected routes {}",
91 pair_tok.len(),
92 selected.len()
93 ));
94 }
95
96 let mut counts = vec![0usize; n_expert];
97 for &expert in selected {
98 let count = counts.get_mut(expert).ok_or_else(|| {
99 format!("grouped expert CSR expert {expert} outside 0..{n_expert}")
100 })?;
101 *count += 1;
102 }
103 if let Some(&token) = pair_tok.iter().find(|&&token| token >= n_tokens) {
104 return Err(format!(
105 "grouped expert CSR token {token} outside 0..{n_tokens}"
106 ));
107 }
108
109 let mut prefix = vec![0usize; n_expert + 1];
110 for expert in 0..n_expert {
111 prefix[expert + 1] = prefix[expert] + counts[expert];
112 }
113 let mut ex_ids = Vec::with_capacity(n_expert.min(selected.len()));
114 let mut ex_off = Vec::with_capacity(ex_ids.capacity() + 1);
115 for expert in 0..n_expert {
116 if counts[expert] != 0 {
117 ex_ids.push(expert as i32);
118 ex_off.push(prefix[expert] as i32);
119 }
120 }
121 ex_off.push(selected.len() as i32);
122
123 let mut cursor = prefix[..n_expert].to_vec();
124 let mut ex_pairs = vec![0i32; selected.len()];
125 for (pair, &expert) in selected.iter().enumerate() {
126 ex_pairs[cursor[expert]] = pair as i32;
127 cursor[expert] += 1;
128 }
129 let pair_tok = pair_tok.iter().map(|&token| token as i32).collect();
130 Self::from_parts(n_expert, n_tokens, ex_ids, ex_off, ex_pairs, pair_tok)
131 }
132
133 fn from_parts(
134 n_expert: usize,
135 n_tokens: usize,
136 ex_ids: Vec<i32>,
137 ex_off: Vec<i32>,
138 ex_pairs: Vec<i32>,
139 pair_tok: Vec<i32>,
140 ) -> Result<Self, String> {
141 if n_expert == 0
142 || n_tokens == 0
143 || ex_ids.is_empty()
144 || ex_pairs.is_empty()
145 || n_expert > i32::MAX as usize
146 || n_tokens > i32::MAX as usize
147 || ex_pairs.len() > i32::MAX as usize
148 {
149 return Err("grouped expert CSR requires non-empty experts, tokens, and pairs".into());
150 }
151 if ex_off.len() != ex_ids.len() + 1 || ex_off.first() != Some(&0) {
152 return Err(format!(
153 "grouped expert CSR offsets {} != active experts {} + 1 or do not start at zero",
154 ex_off.len(),
155 ex_ids.len()
156 ));
157 }
158 let n_pairs = i32::try_from(ex_pairs.len())
159 .map_err(|_| "grouped expert CSR pair count exceeds i32")?;
160 if pair_tok.len() != ex_pairs.len() || ex_off.last().copied() != Some(n_pairs) {
161 return Err(format!(
162 "grouped expert CSR pair lengths offsets_end={:?} pairs={} pair_tok={}",
163 ex_off.last(),
164 ex_pairs.len(),
165 pair_tok.len()
166 ));
167 }
168 for pair in ex_ids.windows(2) {
169 if pair[0] >= pair[1] {
170 return Err("grouped expert CSR expert ids must be strictly increasing".into());
171 }
172 }
173 if ex_ids
174 .iter()
175 .any(|&expert| expert < 0 || expert as usize >= n_expert)
176 {
177 return Err(format!(
178 "grouped expert CSR expert id outside 0..{n_expert}: {ex_ids:?}"
179 ));
180 }
181 let mut seen = vec![false; ex_pairs.len()];
182 for &pair in &ex_pairs {
183 if pair < 0 || pair as usize >= ex_pairs.len() {
184 return Err(format!(
185 "grouped expert CSR pair {pair} outside 0..{}",
186 ex_pairs.len()
187 ));
188 }
189 if std::mem::replace(&mut seen[pair as usize], true) {
190 return Err(format!(
191 "grouped expert CSR pair {pair} appears more than once"
192 ));
193 }
194 let token = pair_tok[pair as usize];
195 if token < 0 || token as usize >= n_tokens {
196 return Err(format!(
197 "grouped expert CSR token {token} outside 0..{n_tokens}"
198 ));
199 }
200 }
201 for offsets in ex_off.windows(2) {
202 if offsets[0] >= offsets[1] {
203 return Err("grouped expert CSR segments must be non-empty and increasing".into());
204 }
205 }
206 Ok(Self {
207 ex_ids,
208 ex_off,
209 ex_pairs,
210 pair_tok,
211 n_expert,
212 n_tokens,
213 })
214 }
215
216 pub fn upload(&self, engine: &Engine) -> Result<DeviceExpertCsr, Box<dyn std::error::Error>> {
217 Ok(DeviceExpertCsr {
218 ex_ids: engine.htod_i32(&self.ex_ids)?,
219 ex_off: engine.htod_i32(&self.ex_off)?,
220 ex_pairs: engine.htod_i32(&self.ex_pairs)?,
221 pair_tok: engine.htod_i32(&self.pair_tok)?,
222 n_expert: self.n_expert,
223 active_experts: self.ex_ids.len(),
224 n_tokens: self.n_tokens,
225 n_pairs: self.ex_pairs.len(),
226 max_tokens: self.n_tokens,
227 max_pairs: self.ex_pairs.len(),
228 })
229 }
230}
231
232pub struct DeviceExpertCsr {
233 ex_ids: CudaSlice<i32>,
234 ex_off: CudaSlice<i32>,
235 ex_pairs: CudaSlice<i32>,
236 pair_tok: CudaSlice<i32>,
237 n_expert: usize,
238 active_experts: usize,
239 n_tokens: usize,
240 n_pairs: usize,
241 max_tokens: usize,
242 max_pairs: usize,
243}
244
245#[derive(Debug, Clone, Copy, PartialEq, Eq)]
246struct DeviceExpertCsrCapacity {
247 n_expert: usize,
248 max_active_experts: usize,
249 max_tokens: usize,
250 max_pairs: usize,
251}
252
253fn validate_device_expert_csr_capacity(
254 n_expert: usize,
255 max_tokens: usize,
256 max_pairs: usize,
257) -> Result<DeviceExpertCsrCapacity, String> {
258 if n_expert == 0
259 || max_tokens == 0
260 || max_pairs == 0
261 || n_expert > i32::MAX as usize
262 || max_tokens > i32::MAX as usize
263 || max_pairs > i32::MAX as usize
264 {
265 return Err(format!(
266 "invalid device expert CSR capacity experts={n_expert} tokens={max_tokens} \
267 pairs={max_pairs}"
268 ));
269 }
270 Ok(DeviceExpertCsrCapacity {
271 n_expert,
272 max_active_experts: n_expert.min(max_pairs),
273 max_tokens,
274 max_pairs,
275 })
276}
277
278fn validate_device_expert_csr_refresh(
279 capacity: DeviceExpertCsrCapacity,
280 n_expert: usize,
281 active_experts: usize,
282 n_tokens: usize,
283 n_pairs: usize,
284) -> Result<(), String> {
285 if n_expert != capacity.n_expert {
286 return Err(format!(
287 "device expert CSR expert count changed {n_expert} != {}",
288 capacity.n_expert
289 ));
290 }
291 if active_experts == 0
292 || n_tokens == 0
293 || n_pairs == 0
294 || active_experts > capacity.max_active_experts
295 || n_tokens > capacity.max_tokens
296 || n_pairs > capacity.max_pairs
297 {
298 return Err(format!(
299 "device expert CSR active shape experts={active_experts} tokens={n_tokens} \
300 pairs={n_pairs} exceeds capacity experts={} tokens={} pairs={}",
301 capacity.max_active_experts, capacity.max_tokens, capacity.max_pairs
302 ));
303 }
304 Ok(())
305}
306
307impl DeviceExpertCsr {
308 pub fn with_capacity(
313 engine: &Engine,
314 n_expert: usize,
315 max_tokens: usize,
316 max_pairs: usize,
317 ) -> Result<Self, Box<dyn std::error::Error>> {
318 let capacity = validate_device_expert_csr_capacity(n_expert, max_tokens, max_pairs)?;
319 Ok(Self {
320 ex_ids: engine.htod_i32(&vec![0; capacity.max_active_experts])?,
321 ex_off: engine.htod_i32(&vec![0; capacity.max_active_experts + 1])?,
322 ex_pairs: engine.htod_i32(&vec![0; capacity.max_pairs])?,
323 pair_tok: engine.htod_i32(&vec![0; capacity.max_pairs])?,
324 n_expert,
325 active_experts: 0,
326 n_tokens: 0,
327 n_pairs: 0,
328 max_tokens,
329 max_pairs,
330 })
331 }
332
333 pub fn refresh(
334 &mut self,
335 engine: &Engine,
336 csr: &ExpertCsr,
337 ) -> Result<(), Box<dyn std::error::Error>> {
338 let capacity =
339 validate_device_expert_csr_capacity(self.n_expert, self.max_tokens, self.max_pairs)?;
340 validate_device_expert_csr_refresh(
341 capacity,
342 csr.n_expert,
343 csr.ex_ids.len(),
344 csr.n_tokens,
345 csr.ex_pairs.len(),
346 )?;
347 let device = engine.ctx().ordinal();
348 if self.ex_ids.ordinal() != device
349 || self.ex_off.ordinal() != device
350 || self.ex_pairs.ordinal() != device
351 || self.pair_tok.ordinal() != device
352 {
353 return Err(
354 format!("device expert CSR capacity is not resident on device {device}").into(),
355 );
356 }
357 engine.htod_i32_into(&mut self.ex_ids, &csr.ex_ids)?;
358 engine.htod_i32_into(&mut self.ex_off, &csr.ex_off)?;
359 engine.htod_i32_into(&mut self.ex_pairs, &csr.ex_pairs)?;
360 engine.htod_i32_into(&mut self.pair_tok, &csr.pair_tok)?;
361 self.active_experts = csr.ex_ids.len();
362 self.n_tokens = csr.n_tokens;
363 self.n_pairs = csr.ex_pairs.len();
364 Ok(())
365 }
366
367 pub fn clear(&mut self) {
368 self.active_experts = 0;
369 self.n_tokens = 0;
370 self.n_pairs = 0;
371 }
372}
373
374#[derive(Debug, Clone, Copy, PartialEq, Eq)]
375struct GroupedFp8WorkspaceShape {
376 activation_len: usize,
377 output_len: usize,
378}
379
380fn validate_grouped_fp8_workspace_shape(
381 in_features: usize,
382 out_features: usize,
383 n_tokens: usize,
384 n_pairs: usize,
385) -> Result<GroupedFp8WorkspaceShape, String> {
386 if in_features == 0
387 || out_features == 0
388 || n_tokens == 0
389 || n_pairs == 0
390 || !in_features.is_multiple_of(16)
391 || in_features > i32::MAX as usize
392 || out_features > i32::MAX as usize
393 || n_tokens > i32::MAX as usize
394 || n_pairs > i32::MAX as usize
395 {
396 return Err(format!(
397 "invalid grouped FP8 workspace in={in_features} out={out_features} \
398 tokens={n_tokens} pairs={n_pairs}"
399 ));
400 }
401 let activation_len = n_tokens
402 .checked_mul(in_features)
403 .ok_or("grouped FP8 activation length overflow")?;
404 let output_len = n_pairs
405 .checked_mul(out_features)
406 .ok_or("grouped FP8 output length overflow")?;
407 Ok(GroupedFp8WorkspaceShape {
408 activation_len,
409 output_len,
410 })
411}
412
413fn validate_grouped_fp8_workspace_active_shape(
414 in_features: usize,
415 out_features: usize,
416 max_tokens: usize,
417 max_pairs: usize,
418 n_tokens: usize,
419 n_pairs: usize,
420) -> Result<GroupedFp8WorkspaceShape, String> {
421 validate_grouped_fp8_workspace_shape(in_features, out_features, max_tokens, max_pairs)?;
422 let active =
423 validate_grouped_fp8_workspace_shape(in_features, out_features, n_tokens, n_pairs)?;
424 if n_tokens > max_tokens || n_pairs > max_pairs {
425 return Err(format!(
426 "grouped FP8 active shape tokens={n_tokens} pairs={n_pairs} exceeds capacity \
427 tokens={max_tokens} pairs={max_pairs}"
428 ));
429 }
430 Ok(active)
431}
432
433pub struct Fp8GroupedWorkspace {
438 act_scratch: CudaSlice<u8>,
439 output: CudaSlice<f32>,
440 in_features: usize,
441 out_features: usize,
442 n_tokens: usize,
443 n_pairs: usize,
444 max_tokens: usize,
445 max_pairs: usize,
446}
447
448impl Fp8GroupedWorkspace {
449 pub fn new(
450 engine: &Engine,
451 in_features: usize,
452 out_features: usize,
453 n_tokens: usize,
454 n_pairs: usize,
455 ) -> Result<Self, Box<dyn std::error::Error>> {
456 let shape =
457 validate_grouped_fp8_workspace_shape(in_features, out_features, n_tokens, n_pairs)?;
458 let act_bytes = unsafe { memra_mmq_fp8_blk_act_bytes(in_features as i32, n_tokens as i32) };
459 if act_bytes == 0 {
460 return Err("grouped FP8 activation scratch size is zero".into());
461 }
462 Ok(Self {
463 act_scratch: engine.alloc_u8_uninit(act_bytes)?,
464 output: engine.uninit(shape.output_len)?,
465 in_features,
466 out_features,
467 n_tokens,
468 n_pairs,
469 max_tokens: n_tokens,
470 max_pairs: n_pairs,
471 })
472 }
473
474 pub fn quantize(
475 &mut self,
476 engine: &Engine,
477 activations: &CudaSlice<f32>,
478 ) -> Result<(), Box<dyn std::error::Error>> {
479 self.quantize_for_shape(engine, activations, self.n_tokens, self.n_pairs)
480 }
481
482 pub fn quantize_for_shape(
484 &mut self,
485 engine: &Engine,
486 activations: &CudaSlice<f32>,
487 n_tokens: usize,
488 n_pairs: usize,
489 ) -> Result<(), Box<dyn std::error::Error>> {
490 let shape = validate_grouped_fp8_workspace_active_shape(
491 self.in_features,
492 self.out_features,
493 self.max_tokens,
494 self.max_pairs,
495 n_tokens,
496 n_pairs,
497 )?;
498 let device = engine.ctx().ordinal();
499 if activations.len() < shape.activation_len
500 || activations.ordinal() != device
501 || self.act_scratch.ordinal() != device
502 {
503 return Err(format!(
504 "grouped FP8 activation len/device {}/{} does not cover {}x{} on device {}",
505 activations.len(),
506 activations.ordinal(),
507 n_tokens,
508 self.in_features,
509 device,
510 )
511 .into());
512 }
513 let stream = engine.gpu.stream();
514 let (x_p, _gx) = activations.device_ptr(&stream);
515 let (scratch_p, _gs) = self.act_scratch.device_ptr_mut(&stream);
516 let rc = unsafe {
517 memra_mmq_fp8_blk_quantize_act(
518 x_p as *const f32,
519 scratch_p as *mut core::ffi::c_void,
520 self.in_features as i32,
521 n_tokens as i32,
522 stream.cu_stream() as *mut core::ffi::c_void,
523 )
524 };
525 if rc != 0 {
526 return Err(format!("memra_mmq_fp8_blk_quantize_act rc={rc}").into());
527 }
528 self.n_tokens = n_tokens;
529 self.n_pairs = n_pairs;
530 Ok(())
531 }
532
533 #[allow(clippy::too_many_arguments)]
534 pub fn project(
535 &mut self,
536 engine: &Engine,
537 bank_codes: &CudaSlice<u8>,
538 bank_scales: &CudaSlice<f32>,
539 csr: &DeviceExpertCsr,
540 code_stride: usize,
541 scale_stride: usize,
542 out_scale: f32,
543 ) -> Result<(), Box<dyn std::error::Error>> {
544 if csr.n_tokens != self.n_tokens || csr.n_pairs != self.n_pairs {
545 return Err(format!(
546 "grouped FP8 CSR/workspace mismatch tokens {} != {}, pairs {} != {}",
547 csr.n_tokens, self.n_tokens, csr.n_pairs, self.n_pairs
548 )
549 .into());
550 }
551 let want_code_stride = self
552 .in_features
553 .checked_mul(self.out_features)
554 .ok_or("grouped FP8 code stride overflow")?;
555 let want_scale_stride = self.in_features.div_ceil(128) * self.out_features.div_ceil(128);
556 if code_stride < want_code_stride || scale_stride < want_scale_stride {
557 return Err(format!(
558 "grouped FP8 expert strides codes {code_stride} < {want_code_stride}, \
559 scales {scale_stride} < {want_scale_stride}"
560 )
561 .into());
562 }
563 let code_count = csr
564 .n_expert
565 .checked_mul(code_stride)
566 .ok_or("grouped FP8 expert code count overflow")?;
567 let scale_count = csr
568 .n_expert
569 .checked_mul(scale_stride)
570 .ok_or("grouped FP8 expert scale count overflow")?;
571 if bank_codes.len() < code_count || bank_scales.len() < scale_count {
572 return Err(format!(
573 "grouped FP8 expert bank too small codes {} < {}, scales {} < {}",
574 bank_codes.len(),
575 code_count,
576 bank_scales.len(),
577 scale_count,
578 )
579 .into());
580 }
581 if !out_scale.is_finite() {
582 return Err(format!("grouped FP8 output scale is not finite: {out_scale}").into());
583 }
584 let device = engine.ctx().ordinal();
585 if bank_codes.ordinal() != device
586 || bank_scales.ordinal() != device
587 || csr.ex_ids.ordinal() != device
588 || csr.ex_off.ordinal() != device
589 || csr.ex_pairs.ordinal() != device
590 || csr.pair_tok.ordinal() != device
591 || self.act_scratch.ordinal() != device
592 || self.output.ordinal() != device
593 {
594 return Err(format!(
595 "grouped FP8 bank, CSR, and workspace must all reside on device {device}"
596 )
597 .into());
598 }
599 let stream = engine.gpu.stream();
600 let (codes_p, _gc) = bank_codes.device_ptr(&stream);
601 let (scales_p, _gs) = bank_scales.device_ptr(&stream);
602 let (ids_p, _gi) = csr.ex_ids.device_ptr(&stream);
603 let (off_p, _go) = csr.ex_off.device_ptr(&stream);
604 let (pairs_p, _gp) = csr.ex_pairs.device_ptr(&stream);
605 let (tok_p, _gt) = csr.pair_tok.device_ptr(&stream);
606 let (act_p, _ga) = self.act_scratch.device_ptr(&stream);
607 let (output_p, _gy) = self.output.device_ptr_mut(&stream);
608 let rc = unsafe {
609 memra_mmq_fp8_blk_grouped(
610 codes_p as *const core::ffi::c_void,
611 scales_p as *const f32,
612 ids_p as *const i32,
613 off_p as *const i32,
614 pairs_p as *const i32,
615 tok_p as *const i32,
616 act_p as *const core::ffi::c_void,
617 output_p as *mut f32,
618 self.in_features as i32,
619 self.out_features as i32,
620 csr.n_expert as i32,
621 csr.active_experts as i32,
622 csr.n_pairs as i32,
623 csr.n_tokens as i32,
624 code_stride,
625 scale_stride,
626 stream.cu_stream() as *mut core::ffi::c_void,
627 out_scale,
628 )
629 };
630 if rc != 0 {
631 return Err(format!("memra_mmq_fp8_blk_grouped rc={rc}").into());
632 }
633 Ok(())
634 }
635
636 pub fn output(&self) -> &CudaSlice<f32> {
637 &self.output
638 }
639
640 pub fn output_len(&self) -> usize {
641 self.n_pairs * self.out_features
642 }
643}
644
645#[cfg(test)]
646mod grouped_fp8_tests {
647 use super::{
648 DeviceExpertCsrCapacity, ExpertCsr, GroupedFp8WorkspaceShape,
649 validate_device_expert_csr_capacity, validate_device_expert_csr_refresh,
650 validate_grouped_fp8_workspace_active_shape, validate_grouped_fp8_workspace_shape,
651 };
652
653 #[test]
654 fn token_routes_build_stable_expert_major_csr() {
655 let csr = ExpertCsr::from_token_routes(4, 2, 3, &[2, 0, 2, 1, 0, 3]).unwrap();
656 assert_eq!(csr.ex_ids, vec![0, 1, 2, 3]);
657 assert_eq!(csr.ex_off, vec![0, 2, 3, 5, 6]);
658 assert_eq!(csr.ex_pairs, vec![1, 4, 3, 0, 2, 5]);
659 assert_eq!(csr.pair_tok, vec![0, 0, 0, 1, 1, 1]);
660 }
661
662 #[test]
663 fn explicit_pair_rows_remain_indexed_by_pair_id() {
664 let csr = ExpertCsr::from_pair_rows(2, 3, &[1, 0, 1], &[2, 0, 1]).unwrap();
665 assert_eq!(csr.ex_ids, vec![0, 1]);
666 assert_eq!(csr.ex_off, vec![0, 1, 3]);
667 assert_eq!(csr.ex_pairs, vec![1, 0, 2]);
668 assert_eq!(csr.pair_tok, vec![2, 0, 1]);
669 }
670
671 #[test]
672 fn csr_validation_rejects_bad_routes_and_parts() {
673 assert!(ExpertCsr::from_token_routes(4, 2, 3, &[0, 1]).is_err());
674 assert!(ExpertCsr::from_pair_rows(2, 1, &[2], &[0]).is_err());
675 assert!(ExpertCsr::from_pair_rows(2, 1, &[0], &[1]).is_err());
676 assert!(
677 ExpertCsr::from_parts(2, 2, vec![0, 1], vec![0, 1, 2], vec![0, 0], vec![0, 1]).is_err()
678 );
679 assert!(
680 ExpertCsr::from_parts(2, 2, vec![1, 0], vec![0, 1, 2], vec![0, 1], vec![0, 1]).is_err()
681 );
682 }
683
684 #[test]
685 fn csr_segments_are_not_limited_to_one_kernel_tile() {
686 let selected = vec![0usize; 17];
687 let rows = (0..17).collect::<Vec<_>>();
688 let csr = ExpertCsr::from_pair_rows(1, 17, &selected, &rows).unwrap();
689 assert_eq!(csr.ex_off, vec![0, 17]);
690 assert_eq!(csr.ex_pairs, (0..17).collect::<Vec<i32>>());
691 }
692
693 #[test]
694 fn workspace_shape_validation_is_pure_and_checked() {
695 assert_eq!(
696 validate_grouped_fp8_workspace_shape(4096, 1280, 2, 16).unwrap(),
697 GroupedFp8WorkspaceShape {
698 activation_len: 8192,
699 output_len: 20480,
700 }
701 );
702 assert!(validate_grouped_fp8_workspace_shape(15, 128, 1, 1).is_err());
703 assert!(validate_grouped_fp8_workspace_shape(i32::MAX as usize + 1, 128, 1, 1,).is_err());
704 }
705
706 #[test]
707 fn device_csr_capacity_admits_smaller_dynamic_schedules() {
708 let capacity = validate_device_expert_csr_capacity(72, 8, 64).unwrap();
709 assert_eq!(
710 capacity,
711 DeviceExpertCsrCapacity {
712 n_expert: 72,
713 max_active_experts: 64,
714 max_tokens: 8,
715 max_pairs: 64,
716 }
717 );
718 validate_device_expert_csr_refresh(capacity, 72, 5, 3, 17).unwrap();
719 assert!(validate_device_expert_csr_refresh(capacity, 72, 5, 9, 17).is_err());
720 assert!(validate_device_expert_csr_refresh(capacity, 72, 5, 3, 65).is_err());
721 assert!(validate_device_expert_csr_refresh(capacity, 71, 5, 3, 17).is_err());
722 assert!(validate_device_expert_csr_refresh(capacity, 72, 0, 3, 17).is_err());
723 }
724
725 #[test]
726 fn grouped_workspace_capacity_accepts_only_bounded_active_shapes() {
727 assert_eq!(
728 validate_grouped_fp8_workspace_active_shape(4096, 1280, 8, 64, 3, 17).unwrap(),
729 GroupedFp8WorkspaceShape {
730 activation_len: 3 * 4096,
731 output_len: 17 * 1280,
732 }
733 );
734 assert!(validate_grouped_fp8_workspace_active_shape(4096, 1280, 8, 64, 9, 17).is_err());
735 assert!(validate_grouped_fp8_workspace_active_shape(4096, 1280, 8, 64, 3, 65).is_err());
736 }
737}
738
739unsafe extern "C" {
740 fn memra_bind_device(dev: i32) -> i32;
741 pub fn memra_mmq_nvfp4_act_bytes(in_f: i32, n_tokens: i32) -> usize;
743 pub fn memra_mmq_nvfp4(
750 w_nvfp4_blocks: *const core::ffi::c_void,
751 act_f32: *const f32,
752 y: *mut f32,
753 in_f: i32,
754 out_f: i32,
755 n_tokens: i32,
756 act_scratch: *mut core::ffi::c_void,
757 stream: *mut core::ffi::c_void,
758 out_scale: f32,
759 ) -> i32;
760 pub fn memra_mmq_nvfp4_ex(
766 w_nvfp4_blocks: *const core::ffi::c_void,
767 act_f32: *const f32,
768 y: *mut f32,
769 in_f: i32,
770 out_f: i32,
771 n_tokens: i32,
772 act_scratch: *mut core::ffi::c_void,
773 stream: *mut core::ffi::c_void,
774 out_scale: f32,
775 per_token_scale: i32,
776 ) -> i32;
777 pub fn memra_mmq_nvfp4_ex2(
783 w_nvfp4_blocks: *const core::ffi::c_void,
784 act_f32: *const f32,
785 y: *mut f32,
786 in_f: i32,
787 out_f: i32,
788 n_tokens: i32,
789 act_scratch: *mut core::ffi::c_void,
790 stream: *mut core::ffi::c_void,
791 out_scale: f32,
792 per_token_scale: i32,
793 residual_k: i32,
794 ) -> i32;
795 pub fn memra_mmq_nvfp4_w4a8_act_bytes(in_f: i32, n_tokens: i32) -> usize;
797 pub fn memra_mmq_nvfp4_w4a8(
805 w_nvfp4_blocks: *const core::ffi::c_void,
806 act_f32: *const f32,
807 y: *mut f32,
808 in_f: i32,
809 out_f: i32,
810 n_tokens: i32,
811 act_scratch: *mut core::ffi::c_void,
812 stream: *mut core::ffi::c_void,
813 out_scale: f32,
814 rp: i32,
815 ) -> i32;
816 pub fn memra_mmq_nvfp4_f8f4_act_bytes(in_f: i32, n_tokens: i32) -> usize;
818 pub fn memra_mmq_nvfp4_f8f4(
824 w_nvfp4_blocks: *const core::ffi::c_void,
825 act_f32: *const f32,
826 y: *mut f32,
827 in_f: i32,
828 out_f: i32,
829 n_tokens: i32,
830 act_scratch: *mut core::ffi::c_void,
831 stream: *mut core::ffi::c_void,
832 out_scale: f32,
833 rp: i32,
834 ) -> i32;
835 pub fn memra_mmq_fp8_blk_act_bytes(in_f: i32, n_tokens: i32) -> usize;
838 pub fn memra_mmq_fp8_blk_quantize_act(
839 act_f32: *const f32,
840 act_scratch: *mut core::ffi::c_void,
841 in_f: i32,
842 n_tokens: i32,
843 stream: *mut core::ffi::c_void,
844 ) -> i32;
845 pub fn memra_mmq_fp8_blk_grouped(
846 bank_codes: *const core::ffi::c_void,
847 bank_scales: *const f32,
848 ex_ids: *const i32,
849 ex_off: *const i32,
850 ex_pairs: *const i32,
851 pair_tok: *const i32,
852 act_scratch: *const core::ffi::c_void,
853 y: *mut f32,
854 in_f: i32,
855 out_f: i32,
856 n_expert: i32,
857 n_active: i32,
858 n_pairs: i32,
859 n_tokens: i32,
860 code_stride: usize,
861 scale_stride: usize,
862 stream: *mut core::ffi::c_void,
863 out_scale: f32,
864 ) -> i32;
865 pub fn memra_mmq_fp8_blk_scale_rows(out_f: i32) -> i32;
867 pub fn memra_mmq_fp8_blk_scale_cols(in_f: i32) -> i32;
868 pub fn memra_mmq_fp8_blk(
875 w_e4m3: *const core::ffi::c_void,
876 blk_scales: *const f32,
877 act_f32: *const f32,
878 y: *mut f32,
879 in_f: i32,
880 out_f: i32,
881 n_tokens: i32,
882 act_scratch: *mut core::ffi::c_void,
883 stream: *mut core::ffi::c_void,
884 out_scale: f32,
885 ) -> i32;
886 pub fn memra_fp8_blk_count_nan(
890 w_e4m3: *const core::ffi::c_void,
891 nbytes: usize,
892 out_count: *mut u32,
893 stream: *mut core::ffi::c_void,
894 ) -> i32;
895 pub fn memra_mmq_q45k_act_bytes(in_f: i32, n_tokens: i32) -> usize;
897 pub fn memra_mmq_q4_K(
900 w_q4k_blocks: *const core::ffi::c_void,
901 act_f32: *const f32,
902 y: *mut f32,
903 in_f: i32,
904 out_f: i32,
905 n_tokens: i32,
906 act_scratch: *mut core::ffi::c_void,
907 stream: *mut core::ffi::c_void,
908 ) -> i32;
909 pub fn memra_mmq_q5_K(
911 w_q5k_blocks: *const core::ffi::c_void,
912 act_f32: *const f32,
913 y: *mut f32,
914 in_f: i32,
915 out_f: i32,
916 n_tokens: i32,
917 act_scratch: *mut core::ffi::c_void,
918 stream: *mut core::ffi::c_void,
919 ) -> i32;
920
921 pub fn memra_mmq_q8_0_act_bytes(in_f: i32, n_tokens: i32) -> usize;
923 pub fn memra_mmq_q8_0(
927 w_q8_0_blocks: *const core::ffi::c_void,
928 act_f32: *const f32,
929 y: *mut f32,
930 in_f: i32,
931 out_f: i32,
932 n_tokens: i32,
933 act_scratch: *mut core::ffi::c_void,
934 stream: *mut core::ffi::c_void,
935 ) -> i32;
936
937 pub fn memra_accprobe_act_bytes(in_f: i32, n_tokens: i32) -> usize;
946 pub fn memra_accprobe_gemm_s32(
948 w_q8_0_blocks: *const core::ffi::c_void,
949 act_q: *const core::ffi::c_void,
950 y: *mut f32,
951 in_f: i32,
952 out_f: i32,
953 n_tokens: i32,
954 stream: *mut core::ffi::c_void,
955 ) -> i32;
956 pub fn memra_accprobe_gemm_f32(
958 w_q8_0_blocks: *const core::ffi::c_void,
959 act_q: *const core::ffi::c_void,
960 y: *mut f32,
961 in_f: i32,
962 out_f: i32,
963 n_tokens: i32,
964 stream: *mut core::ffi::c_void,
965 ) -> i32;
966
967 pub fn memra_mmq_q4_0_act_bytes(in_f: i32, n_tokens: i32) -> usize;
969 pub fn memra_mmq_q4_0(
975 w_q4_0: *const core::ffi::c_void,
976 act_f32: *const f32,
977 y: *mut f32,
978 in_f: i32,
979 out_f: i32,
980 n_tokens: i32,
981 act_scratch: *mut core::ffi::c_void,
982 stream: *mut core::ffi::c_void,
983 rp: i32,
984 ) -> i32;
985 pub fn memra_mmq_q4_0_quant_act(
987 act_f32: *const f32,
988 act_scratch: *mut core::ffi::c_void,
989 in_f: i32,
990 n_tokens: i32,
991 stream: *mut core::ffi::c_void,
992 ) -> i32;
993 pub fn memra_mmq_q4_0_gemm(
995 w_q4_0: *const core::ffi::c_void,
996 act_scratch: *const core::ffi::c_void,
997 y: *mut f32,
998 in_f: i32,
999 out_f: i32,
1000 n_tokens: i32,
1001 stream: *mut core::ffi::c_void,
1002 rp: i32,
1003 ) -> i32;
1004 pub fn memra_mmq_q4_0_fixup_bytes() -> usize;
1006 pub fn memra_mmq_q4_0_set_clc(force: i32) -> i32;
1011 pub fn memra_mmq_q4_0_gemm_sk(
1014 w_q4_0: *const core::ffi::c_void,
1015 act_scratch: *const core::ffi::c_void,
1016 y: *mut f32,
1017 fixup_scratch: *mut core::ffi::c_void,
1018 in_f: i32,
1019 out_f: i32,
1020 n_tokens: i32,
1021 stream: *mut core::ffi::c_void,
1022 rp: i32,
1023 ) -> i32;
1024
1025 pub fn memra_mmq_iq_experts_act_bytes(in_f: i32, n_tokens: i32) -> usize;
1028 pub fn memra_mmq_iq_quantize_act(
1030 act_f32: *const f32,
1031 act_scratch: *mut core::ffi::c_void,
1032 in_f: i32,
1033 n_tokens: i32,
1034 stream: *mut core::ffi::c_void,
1035 ) -> i32;
1036 pub fn memra_mmq_iq_fused_act_quant(
1040 gate: *const f32,
1041 up: *const f32,
1042 act_scratch: *mut core::ffi::c_void,
1043 in_f: i32,
1044 n_tokens: i32,
1045 act_kind: i32,
1046 stream: *mut core::ffi::c_void,
1047 ) -> i32;
1048 pub fn memra_mmq_iq4xs_dense(
1057 w_blocks: *const core::ffi::c_void,
1058 act_f32: *const f32,
1059 y: *mut f32,
1060 in_f: i32,
1061 out_f: i32,
1062 n_tokens: i32,
1063 row_bytes: i64,
1064 act_scratch: *mut core::ffi::c_void,
1065 stream: *mut core::ffi::c_void,
1066 ) -> i32;
1067 pub fn memra_mmq_iq_experts(
1068 table: *const u64,
1069 proj: i32,
1070 n_expert: i32,
1071 ex_ids: *const i32,
1072 ex_off: *const i32,
1073 ex_pairs: *const i32,
1074 pair_tok: *const i32,
1075 act_scratch: *const core::ffi::c_void,
1076 y: *mut f32,
1077 in_f: i32,
1078 out_f: i32,
1079 n_active: i32,
1080 n_tokens: i32,
1081 qtype: i32,
1082 row_bytes: i64,
1083 stream: *mut core::ffi::c_void,
1084 ) -> i32;
1085
1086 pub fn memra_moe_f16g_dequant(
1088 table: *const u64,
1089 proj: i32,
1090 n_expert: i32,
1091 ex_ids: *const i32,
1092 w_f16: *mut core::ffi::c_void,
1093 in_f: i32,
1094 out_f: i32,
1095 n_active: i32,
1096 qtype: i32,
1097 row_bytes: i64,
1098 stream: *mut core::ffi::c_void,
1099 ) -> i32;
1100 pub fn memra_moe_f16g_gather_act(
1101 x: *const f32,
1102 pair_tok_or_null: *const i32,
1103 act_f16: *mut core::ffi::c_void,
1104 row_scale: *mut f32,
1105 in_f: i32,
1106 n_pairs: i32,
1107 stream: *mut core::ffi::c_void,
1108 ) -> i32;
1109 pub fn memra_moe_f16g_h2f_scaled(
1110 src_f16: *const core::ffi::c_void,
1111 dst: *mut f32,
1112 row_scale: *const f32,
1113 ncols: i32,
1114 nrows: i32,
1115 stream: *mut core::ffi::c_void,
1116 ) -> i32;
1117 pub fn memra_moe_f16g_gemm(
1118 w_f16: *const core::ffi::c_void,
1119 act_f16: *const core::ffi::c_void,
1120 y_f16: *mut core::ffi::c_void,
1121 ex_off_host: *const i32,
1122 n_active: i32,
1123 in_f: i32,
1124 out_f: i32,
1125 stream: *mut core::ffi::c_void,
1126 ) -> i32;
1127 pub fn memra_moe_f16g_h2f(
1128 src_f16: *const core::ffi::c_void,
1129 dst: *mut f32,
1130 n: usize,
1131 stream: *mut core::ffi::c_void,
1132 ) -> i32;
1133 pub fn memra_moe_f16g_gemm_sk(
1142 w_f16: *const core::ffi::c_void,
1143 act_f16: *const core::ffi::c_void,
1144 y_f32: *mut f32,
1145 row_scale: *const f32,
1146 ex_off_dev: *const i32,
1147 ex_off_host: *const i32,
1148 n_active: i32,
1149 max_m: i32,
1150 in_f: i32,
1151 out_f: i32,
1152 shape_sel: i32,
1153 cross: i32,
1154 tail: i32,
1155 stream: *mut core::ffi::c_void,
1156 ) -> i32;
1157 pub fn memra_moe_kq_gemm_sk(
1164 table: *const u64,
1165 proj: i32,
1166 n_expert: i32,
1167 ex_ids: *const i32,
1168 act_f16: *const core::ffi::c_void,
1169 y_f32: *mut f32,
1170 row_scale: *const f32,
1171 ex_off_dev: *const i32,
1172 ex_off_host: *const i32,
1173 n_active: i32,
1174 max_m: i32,
1175 in_f: i32,
1176 out_f: i32,
1177 qtype: i32,
1178 cross: i32,
1179 tail: i32,
1180 row_bytes: i64,
1181 stream: *mut core::ffi::c_void,
1182 ) -> i32;
1183}
1184
1185pub fn mmq_w4a8_enabled() -> bool {
1194 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1195 *ON.get_or_init(|| {
1196 std::env::var("MEMRA_MMQ_W4A8")
1197 .map(|v| v != "0")
1198 .unwrap_or(true)
1199 })
1200}
1201
1202pub fn mmq_residual_k() -> i32 {
1211 std::env::var("MEMRA_MMQ_RESIDUAL_K")
1212 .ok()
1213 .and_then(|v| v.parse::<i32>().ok())
1214 .unwrap_or(0)
1215 .clamp(0, 64)
1216}
1217
1218pub fn mmq_q8_enabled() -> bool {
1224 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1225 *ON.get_or_init(|| {
1229 std::env::var("MEMRA_PP_Q8MMQ")
1230 .map(|v| v != "0")
1231 .unwrap_or(true)
1232 })
1233}
1234
1235pub fn mmq_iq4xs_enabled() -> bool {
1243 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1244 *ON.get_or_init(|| {
1245 std::env::var("MEMRA_PP_IQMMQ")
1246 .map(|v| v != "0")
1247 .unwrap_or(true)
1248 })
1249}
1250
1251pub fn mmq_q4_enabled() -> bool {
1257 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1258 *ON.get_or_init(|| {
1259 std::env::var("MEMRA_PP_Q4MMQ")
1260 .map(|v| v != "0")
1261 .unwrap_or(true)
1262 })
1263}
1264
1265fn nvfp4_use_w4a8(rp: bool, w4a8_explicit: bool, w4a8_default: bool, mmq_explicit: bool) -> bool {
1266 rp || w4a8_explicit || (w4a8_default && !mmq_explicit)
1269}
1270
1271#[cfg(test)]
1272mod b200_dry_policy_tests {
1273 use super::nvfp4_use_w4a8;
1274
1275 #[test]
1276 fn nvfp4_default_and_explicit_routes_do_not_reach_sm100_stubs() {
1277 assert!(nvfp4_use_w4a8(false, false, true, false));
1278 assert!(!nvfp4_use_w4a8(false, false, true, true));
1279 assert!(nvfp4_use_w4a8(false, true, true, true));
1280 assert!(nvfp4_use_w4a8(true, false, false, true));
1281 assert!(!nvfp4_use_w4a8(false, false, false, false));
1282 }
1283}
1284
1285impl Engine {
1286 pub fn mmq_supports(&self, w: &crate::model::GpuTensor) -> bool {
1289 use crate::model::GpuTensor;
1290 if crate::portable_mma_gated() {
1291 return false;
1292 }
1293 let mmq_opt_in = std::env::var("MEMRA_MMQ").is_ok();
1294 match w {
1295 GpuTensor::Quant { qtype, rp, .. } if *qtype == crate::QT_NVFP4 && *rp => {
1302 !cfg!(memra_portable_cuda)
1303 && mmq_w4a8_enabled()
1304 && w.in_features().is_multiple_of(64)
1305 }
1306 GpuTensor::Quant { qtype, .. } if *qtype == crate::QT_NVFP4 => {
1311 !cfg!(memra_portable_cuda)
1312 && (mmq_w4a8_enabled() || mmq_opt_in)
1313 && w.in_features().is_multiple_of(64)
1314 }
1315 GpuTensor::Quant { qtype, .. }
1316 if *qtype == crate::QT_Q4_K || *qtype == crate::QT_Q5_K =>
1317 {
1318 (mmq_w4a8_enabled() || mmq_opt_in) && w.in_features().is_multiple_of(256)
1319 }
1320 GpuTensor::Quant { qtype, .. } if *qtype == crate::QT_Q8_0 => {
1325 mmq_q8_enabled() && w.in_features().is_multiple_of(256)
1326 }
1327 GpuTensor::Quant { qtype, .. } if *qtype == crate::QT_Q4_0 => {
1332 mmq_q4_enabled() && w.in_features().is_multiple_of(256)
1333 }
1334 GpuTensor::Quant { qtype, .. } if *qtype == crate::QT_IQ4_XS => {
1340 mmq_iq4xs_enabled()
1341 && Self::iq_fast_enabled()
1342 && w.in_features().is_multiple_of(256)
1343 }
1344 _ => false,
1345 }
1346 }
1347
1348 pub fn qmatvec_mmq(
1351 &self,
1352 w: &crate::model::GpuTensor,
1353 x: &CudaSlice<f32>,
1354 m: usize,
1355 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1356 use crate::model::GpuTensor;
1357 let (in_f, out_f) = (w.in_features(), w.out_features());
1358 let GpuTensor::Quant {
1359 bytes,
1360 scale,
1361 qtype,
1362 rp,
1363 ..
1364 } = w
1365 else {
1366 return Err("qmatvec_mmq: not a Quant tensor".into());
1367 };
1368 let w4a8_explicit = std::env::var("MEMRA_MMQ_W4A8")
1373 .map(|v| v != "0")
1374 .unwrap_or(false);
1375 let use_w4a8 = nvfp4_use_w4a8(
1376 *rp,
1377 w4a8_explicit,
1378 mmq_w4a8_enabled(),
1379 std::env::var("MEMRA_MMQ").is_ok(),
1380 );
1381 match *qtype {
1382 q if q == crate::QT_NVFP4 && use_w4a8 => {
1385 self.qmatvec_mmq_nvfp4_w4a8(bytes, x, m, in_f, out_f, *scale, *rp)
1386 }
1387 q if q == crate::QT_NVFP4 => self.qmatvec_mmq_nvfp4(bytes, x, m, in_f, out_f, *scale),
1388 q if q == crate::QT_Q4_K || q == crate::QT_Q5_K => {
1389 let mut y = self.qmatvec_mmq_q45k_raw(bytes, x, m, in_f, out_f, q)?;
1390 if *scale != 1.0 {
1391 self.scale_inplace(&mut y, *scale, m * out_f)?;
1392 }
1393 Ok(y)
1394 }
1395 q if q == crate::QT_Q8_0 => {
1396 if cfg!(memra_hopper_mma)
1402 && out_f % 64 == 0
1403 && crate::wgmma_gemm_enabled()
1404 && let GpuTensor::Quant { rp4: Some(m4), .. } = w
1405 {
1406 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
1407 let mut y = self.qmatvec_gemm_q8_0_wgmma_raw(m4, &aq, &ad, m, in_f, out_f)?;
1408 if *scale != 1.0 {
1409 self.scale_inplace(&mut y, *scale, m * out_f)?;
1410 }
1411 return Ok(y);
1412 }
1413 let mut y = self.qmatvec_mmq_q8_0_raw(bytes, x, m, in_f, out_f)?;
1414 if *scale != 1.0 {
1415 self.scale_inplace(&mut y, *scale, m * out_f)?;
1416 }
1417 Ok(y)
1418 }
1419 q if q == crate::QT_Q4_0 => {
1420 let mut y = self.qmatvec_mmq_q4_0_raw(bytes, x, m, in_f, out_f, *rp)?;
1421 if *scale != 1.0 {
1422 self.scale_inplace(&mut y, *scale, m * out_f)?;
1423 }
1424 Ok(y)
1425 }
1426 q if q == crate::QT_IQ4_XS => {
1427 let GpuTensor::Quant { row_bytes, .. } = w else {
1428 unreachable!()
1429 };
1430 let mut y = self.qmatvec_mmq_iq4xs_raw(bytes, x, m, in_f, out_f, *row_bytes)?;
1431 if *scale != 1.0 {
1432 self.scale_inplace(&mut y, *scale, m * out_f)?;
1433 }
1434 Ok(y)
1435 }
1436 q => Err(format!("qmatvec_mmq: unsupported qtype {q}").into()),
1437 }
1438 }
1439
1440 pub fn qmatvec_mmq_iq4xs_raw(
1442 &self,
1443 bytes: &CudaSlice<u8>,
1444 x: &CudaSlice<f32>,
1445 m: usize,
1446 in_f: usize,
1447 out_f: usize,
1448 row_bytes: usize,
1449 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1450 assert!(
1451 in_f.is_multiple_of(256),
1452 "MMQ IQ4_XS requires in_f % 256 == 0, got {in_f}"
1453 );
1454 let act_bytes = unsafe { memra_mmq_iq_experts_act_bytes(in_f as i32, m as i32) };
1455 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
1456 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
1457 {
1458 let stream = self.gpu.stream();
1459 let (w_p, _gw) = bytes.device_ptr(&stream);
1460 let (x_p, _gx) = x.device_ptr(&stream);
1461 let (y_p, _gy) = y.device_ptr_mut(&stream);
1462 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
1463 let rc = unsafe {
1464 memra_mmq_iq4xs_dense(
1465 w_p as *const core::ffi::c_void,
1466 x_p as *const f32,
1467 y_p as *mut f32,
1468 in_f as i32,
1469 out_f as i32,
1470 m as i32,
1471 row_bytes as i64,
1472 s_p as *mut core::ffi::c_void,
1473 stream.cu_stream() as *mut core::ffi::c_void,
1474 )
1475 };
1476 if rc != 0 {
1477 return Err(format!("memra_mmq_iq4xs_dense rc={rc}").into());
1478 }
1479 }
1480 Ok(y)
1481 }
1482
1483 pub fn qmatvec_mmq_q45k_raw(
1488 &self,
1489 bytes: &CudaSlice<u8>,
1490 x: &CudaSlice<f32>,
1491 m: usize,
1492 in_f: usize,
1493 out_f: usize,
1494 qtype: i32,
1495 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1496 assert!(
1497 in_f.is_multiple_of(256),
1498 "MMQ Q4_K/Q5_K requires in_f % 256 == 0, got {in_f}"
1499 );
1500 let act_bytes = unsafe { memra_mmq_q45k_act_bytes(in_f as i32, m as i32) };
1501 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
1502 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
1503 {
1504 let stream = self.gpu.stream();
1505 let (w_p, _gw) = bytes.device_ptr(&stream);
1506 let (x_p, _gx) = x.device_ptr(&stream);
1507 let (y_p, _gy) = y.device_ptr_mut(&stream);
1508 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
1509 let launcher = if qtype == crate::QT_Q4_K {
1510 memra_mmq_q4_K
1511 } else {
1512 memra_mmq_q5_K
1513 };
1514 let rc = unsafe {
1515 launcher(
1516 w_p as *const core::ffi::c_void,
1517 x_p as *const f32,
1518 y_p as *mut f32,
1519 in_f as i32,
1520 out_f as i32,
1521 m as i32,
1522 s_p as *mut core::ffi::c_void,
1523 stream.cu_stream() as *mut core::ffi::c_void,
1524 )
1525 };
1526 if rc != 0 {
1527 return Err(format!("memra_mmq_q45k(qtype={qtype}) rc={rc}").into());
1528 }
1529 }
1530 Ok(y)
1531 }
1532
1533 pub fn qmatvec_mmq_q8_0_raw(
1536 &self,
1537 bytes: &CudaSlice<u8>,
1538 x: &CudaSlice<f32>,
1539 m: usize,
1540 in_f: usize,
1541 out_f: usize,
1542 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1543 assert!(
1544 in_f.is_multiple_of(32),
1545 "MMQ Q8_0 requires in_f % 32 == 0, got {in_f}"
1546 );
1547 let act_bytes = unsafe { memra_mmq_q8_0_act_bytes(in_f as i32, m as i32) };
1548 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
1549 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
1550 {
1551 let stream = self.gpu.stream();
1552 let (w_p, _gw) = bytes.device_ptr(&stream);
1553 let (x_p, _gx) = x.device_ptr(&stream);
1554 let (y_p, _gy) = y.device_ptr_mut(&stream);
1555 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
1556 let rc = unsafe {
1557 memra_mmq_q8_0(
1558 w_p as *const core::ffi::c_void,
1559 x_p as *const f32,
1560 y_p as *mut f32,
1561 in_f as i32,
1562 out_f as i32,
1563 m as i32,
1564 s_p as *mut core::ffi::c_void,
1565 stream.cu_stream() as *mut core::ffi::c_void,
1566 )
1567 };
1568 if rc != 0 {
1569 return Err(format!("memra_mmq_q8_0 rc={rc}").into());
1570 }
1571 }
1572 Ok(y)
1573 }
1574
1575 pub fn accprobe_act_bytes(&self, in_f: usize, m: usize) -> usize {
1578 unsafe { memra_accprobe_act_bytes(in_f as i32, m as i32) }
1579 }
1580
1581 pub fn accprobe_gemm(
1588 &self,
1589 w_q8_0: &CudaSlice<u8>,
1590 act_q: &CudaSlice<u8>,
1591 m: usize,
1592 in_f: usize,
1593 out_f: usize,
1594 f32acc: bool,
1595 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1596 assert!(
1597 in_f.is_multiple_of(32),
1598 "accprobe requires in_f % 32 == 0, got {in_f}"
1599 );
1600 assert!(
1601 act_q.len() >= self.accprobe_act_bytes(in_f, m),
1602 "accprobe act_q too small: {} < {}",
1603 act_q.len(),
1604 self.accprobe_act_bytes(in_f, m)
1605 );
1606 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
1607 {
1608 let stream = self.gpu.stream();
1609 let (w_p, _gw) = w_q8_0.device_ptr(&stream);
1610 let (a_p, _ga) = act_q.device_ptr(&stream);
1611 let (y_p, _gy) = y.device_ptr_mut(&stream);
1612 let f = if f32acc {
1613 memra_accprobe_gemm_f32
1614 } else {
1615 memra_accprobe_gemm_s32
1616 };
1617 let rc = unsafe {
1618 f(
1619 w_p as *const core::ffi::c_void,
1620 a_p as *const core::ffi::c_void,
1621 y_p as *mut f32,
1622 in_f as i32,
1623 out_f as i32,
1624 m as i32,
1625 stream.cu_stream() as *mut core::ffi::c_void,
1626 )
1627 };
1628 if rc != 0 {
1629 let arm = if f32acc { "f32" } else { "s32" };
1630 return Err(format!("memra_accprobe_gemm_{arm} rc={rc}").into());
1631 }
1632 }
1633 Ok(y)
1634 }
1635
1636 pub fn mmq_act_begin(&self) {
1642 use std::sync::atomic::Ordering;
1643 MMQ_ACT_EPOCH.fetch_add(1, Ordering::Relaxed);
1644 *MMQ_ACT_SLOT.lock().unwrap() = None;
1645 }
1646
1647 pub fn qmatvec_mmq_q4_0_raw(
1651 &self,
1652 bytes: &CudaSlice<u8>,
1653 x: &CudaSlice<f32>,
1654 m: usize,
1655 in_f: usize,
1656 out_f: usize,
1657 rp: bool,
1658 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1659 use std::sync::atomic::Ordering;
1660 assert!(
1661 in_f.is_multiple_of(32),
1662 "MMQ Q4_0 requires in_f % 32 == 0, got {in_f}"
1663 );
1664 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
1665 let stream = self.gpu.stream();
1666 let (x_p, _gx) = x.device_ptr(&stream);
1667 let epoch = MMQ_ACT_EPOCH.load(Ordering::Relaxed);
1668 let mut slot = MMQ_ACT_SLOT.lock().unwrap();
1670 let hit = matches!(&*slot,
1671 Some((e, p, mm, inf, _)) if *e == epoch && *p == x_p && *mm == m && *inf == in_f);
1672 if !hit {
1673 let act_bytes = unsafe { memra_mmq_q4_0_act_bytes(in_f as i32, m as i32) };
1674 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
1675 {
1676 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
1677 let rc = unsafe {
1678 memra_mmq_q4_0_quant_act(
1679 x_p as *const f32,
1680 s_p as *mut core::ffi::c_void,
1681 in_f as i32,
1682 m as i32,
1683 stream.cu_stream() as *mut core::ffi::c_void,
1684 )
1685 };
1686 if rc != 0 {
1687 return Err(
1688 format!("memra_mmq_q4_0_quant_act(in_f={in_f}, m={m}) rc={rc}").into(),
1689 );
1690 }
1691 }
1692 *slot = Some((epoch, x_p, m, in_f, scratch));
1693 }
1694 let scratch = &slot.as_ref().unwrap().4;
1695 {
1696 let (w_p, _gw) = bytes.device_ptr(&stream);
1697 let (y_p, _gy) = y.device_ptr_mut(&stream);
1698 let (s_p, _gs) = scratch.device_ptr(&stream);
1699 static SK_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1722 let sk = match crate::MMQ_SK_FORCE.load(std::sync::atomic::Ordering::Relaxed) {
1723 0 => false,
1724 1 => true,
1725 _ => *SK_ON.get_or_init(|| {
1726 std::env::var("MEMRA_MMQ_SK")
1727 .map(|v| v != "0")
1728 .unwrap_or(!cfg!(memra_hopper_mma))
1729 }),
1730 };
1731 let rc = if sk {
1732 let mut fx = MMQ_FIXUP_SLOT.lock().unwrap();
1733 if fx.is_none() {
1734 let nb = unsafe { memra_mmq_q4_0_fixup_bytes() };
1735 *fx = Some(self.alloc_uninit::<u8>(nb)?);
1736 }
1737 let (f_p, _gf) = fx.as_mut().unwrap().device_ptr_mut(&stream);
1738 unsafe {
1739 memra_mmq_q4_0_gemm_sk(
1740 w_p as *const core::ffi::c_void,
1741 s_p as *const core::ffi::c_void,
1742 y_p as *mut f32,
1743 f_p as *mut core::ffi::c_void,
1744 in_f as i32,
1745 out_f as i32,
1746 m as i32,
1747 stream.cu_stream() as *mut core::ffi::c_void,
1748 rp as i32,
1749 )
1750 }
1751 } else {
1752 unsafe {
1753 memra_mmq_q4_0_gemm(
1754 w_p as *const core::ffi::c_void,
1755 s_p as *const core::ffi::c_void,
1756 y_p as *mut f32,
1757 in_f as i32,
1758 out_f as i32,
1759 m as i32,
1760 stream.cu_stream() as *mut core::ffi::c_void,
1761 rp as i32,
1762 )
1763 }
1764 };
1765 if rc != 0 {
1766 return Err(format!(
1767 "memra_mmq_q4_0_gemm(rp={rp}, in_f={in_f}, out_f={out_f}, m={m}, wbytes={}) rc={rc}",
1768 bytes.len()
1769 )
1770 .into());
1771 }
1772 }
1773 Ok(y)
1774 }
1775
1776 pub fn qmatvec_mmq_nvfp4(
1782 &self,
1783 bytes: &CudaSlice<u8>,
1784 x: &CudaSlice<f32>,
1785 m: usize,
1786 in_f: usize,
1787 out_f: usize,
1788 scale: f32,
1789 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1790 self.qmatvec_mmq_nvfp4_scaled(bytes, x, m, in_f, out_f, scale)
1791 }
1792
1793 pub fn qmatvec_mmq_nvfp4_raw(
1795 &self,
1796 bytes: &CudaSlice<u8>,
1797 x: &CudaSlice<f32>,
1798 m: usize,
1799 in_f: usize,
1800 out_f: usize,
1801 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1802 self.qmatvec_mmq_nvfp4_scaled(bytes, x, m, in_f, out_f, 1.0)
1803 }
1804
1805 pub fn qmatvec_mmq_nvfp4_raw_v1(
1809 &self,
1810 bytes: &CudaSlice<u8>,
1811 x: &CudaSlice<f32>,
1812 m: usize,
1813 in_f: usize,
1814 out_f: usize,
1815 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1816 self.qmatvec_mmq_nvfp4_inner(bytes, x, m, in_f, out_f, 1.0, false, 0)
1817 }
1818
1819 pub fn qmatvec_mmq_nvfp4_raw_res(
1821 &self,
1822 bytes: &CudaSlice<u8>,
1823 x: &CudaSlice<f32>,
1824 m: usize,
1825 in_f: usize,
1826 out_f: usize,
1827 residual_k: i32,
1828 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1829 self.qmatvec_mmq_nvfp4_inner(bytes, x, m, in_f, out_f, 1.0, true, residual_k)
1830 }
1831
1832 fn qmatvec_mmq_nvfp4_scaled(
1833 &self,
1834 bytes: &CudaSlice<u8>,
1835 x: &CudaSlice<f32>,
1836 m: usize,
1837 in_f: usize,
1838 out_f: usize,
1839 scale: f32,
1840 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1841 self.qmatvec_mmq_nvfp4_inner(bytes, x, m, in_f, out_f, scale, true, mmq_residual_k())
1842 }
1843
1844 #[allow(clippy::too_many_arguments)] fn qmatvec_mmq_nvfp4_inner(
1846 &self,
1847 bytes: &CudaSlice<u8>,
1848 x: &CudaSlice<f32>,
1849 m: usize,
1850 in_f: usize,
1851 out_f: usize,
1852 scale: f32,
1853 per_token_scale: bool,
1854 residual_k: i32,
1855 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1856 assert!(
1857 in_f.is_multiple_of(64),
1858 "MMQ NVFP4 requires in_f % 64 == 0, got {in_f}"
1859 );
1860 let act_bytes = unsafe { memra_mmq_nvfp4_act_bytes(in_f as i32, m as i32) };
1861 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
1862 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
1863 {
1864 let stream = self.gpu.stream();
1865 let (w_p, _gw) = bytes.device_ptr(&stream);
1866 let (x_p, _gx) = x.device_ptr(&stream);
1867 let (y_p, _gy) = y.device_ptr_mut(&stream);
1868 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
1869 let rc = unsafe {
1870 memra_mmq_nvfp4_ex2(
1871 w_p as *const core::ffi::c_void,
1872 x_p as *const f32,
1873 y_p as *mut f32,
1874 in_f as i32,
1875 out_f as i32,
1876 m as i32,
1877 s_p as *mut core::ffi::c_void,
1878 stream.cu_stream() as *mut core::ffi::c_void,
1879 scale,
1880 per_token_scale as i32,
1881 residual_k,
1882 )
1883 };
1884 if rc != 0 {
1885 return Err(format!("memra_mmq_nvfp4_ex2 rc={rc}").into());
1886 }
1887 }
1888 Ok(y)
1889 }
1890
1891 #[allow(clippy::too_many_arguments)] pub fn qmatvec_mmq_nvfp4_w4a8(
1897 &self,
1898 bytes: &CudaSlice<u8>,
1899 x: &CudaSlice<f32>,
1900 m: usize,
1901 in_f: usize,
1902 out_f: usize,
1903 scale: f32,
1904 rp: bool,
1905 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1906 self.qmatvec_mmq_nvfp4_w4a8_scaled(bytes, x, m, in_f, out_f, scale, rp)
1907 }
1908
1909 pub fn qmatvec_mmq_nvfp4_w4a8_raw(
1911 &self,
1912 bytes: &CudaSlice<u8>,
1913 x: &CudaSlice<f32>,
1914 m: usize,
1915 in_f: usize,
1916 out_f: usize,
1917 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1918 self.qmatvec_mmq_nvfp4_w4a8_scaled(bytes, x, m, in_f, out_f, 1.0, false)
1919 }
1920
1921 pub fn qmatvec_mmq_nvfp4_w4a8_raw_rp(
1924 &self,
1925 bytes: &CudaSlice<u8>,
1926 x: &CudaSlice<f32>,
1927 m: usize,
1928 in_f: usize,
1929 out_f: usize,
1930 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1931 self.qmatvec_mmq_nvfp4_w4a8_scaled(bytes, x, m, in_f, out_f, 1.0, true)
1932 }
1933
1934 #[allow(clippy::too_many_arguments)] fn qmatvec_mmq_nvfp4_w4a8_scaled(
1936 &self,
1937 bytes: &CudaSlice<u8>,
1938 x: &CudaSlice<f32>,
1939 m: usize,
1940 in_f: usize,
1941 out_f: usize,
1942 scale: f32,
1943 rp: bool,
1944 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1945 static F8F4: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1949 let f8f4 = *F8F4.get_or_init(|| std::env::var("MEMRA_MMQ_F8F4").as_deref() == Ok("1"));
1950 assert!(
1951 in_f.is_multiple_of(64),
1952 "MMQ NVFP4 W4A8 requires in_f % 64 == 0, got {in_f}"
1953 );
1954 let act_bytes = unsafe { memra_mmq_nvfp4_w4a8_act_bytes(in_f as i32, m as i32) };
1955 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
1956 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
1957 {
1958 let stream = self.gpu.stream();
1959 let (w_p, _gw) = bytes.device_ptr(&stream);
1960 let (x_p, _gx) = x.device_ptr(&stream);
1961 let (y_p, _gy) = y.device_ptr_mut(&stream);
1962 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
1963 let rc = unsafe {
1965 if f8f4 {
1966 memra_mmq_nvfp4_f8f4(
1967 w_p as *const core::ffi::c_void,
1968 x_p as *const f32,
1969 y_p as *mut f32,
1970 in_f as i32,
1971 out_f as i32,
1972 m as i32,
1973 s_p as *mut core::ffi::c_void,
1974 stream.cu_stream() as *mut core::ffi::c_void,
1975 scale,
1976 rp as i32,
1977 )
1978 } else {
1979 memra_mmq_nvfp4_w4a8(
1980 w_p as *const core::ffi::c_void,
1981 x_p as *const f32,
1982 y_p as *mut f32,
1983 in_f as i32,
1984 out_f as i32,
1985 m as i32,
1986 s_p as *mut core::ffi::c_void,
1987 stream.cu_stream() as *mut core::ffi::c_void,
1988 scale,
1989 rp as i32,
1990 )
1991 }
1992 };
1993 if rc != 0 {
1994 return Err(format!("memra_mmq_nvfp4_w4a8(f8f4={f8f4}) rc={rc}").into());
1995 }
1996 }
1997 Ok(y)
1998 }
1999
2000 pub fn qmatvec_mmq_fp8_blk(
2004 &self,
2005 w_e4m3: &CudaSlice<u8>,
2006 blk_scales: &CudaSlice<f32>,
2007 x: &CudaSlice<f32>,
2008 m: usize,
2009 in_f: usize,
2010 out_f: usize,
2011 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2012 self.qmatvec_mmq_fp8_blk_scaled(w_e4m3, blk_scales, x, m, in_f, out_f, 1.0)
2013 }
2014
2015 #[allow(clippy::too_many_arguments)] pub fn qmatvec_mmq_fp8_blk_scaled(
2017 &self,
2018 w_e4m3: &CudaSlice<u8>,
2019 blk_scales: &CudaSlice<f32>,
2020 x: &CudaSlice<f32>,
2021 m: usize,
2022 in_f: usize,
2023 out_f: usize,
2024 scale: f32,
2025 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2026 if cfg!(memra_sm100_tcgen05) && std::env::var("MEMRA_FP8_MMQ").as_deref() != Ok("1") {
2027 return Err(
2028 "B200 block-FP8 tcgen05 is NativeReference but not tuned; set \
2029 MEMRA_FP8_MMQ=1 only for explicit qualification or research (the pinned \
2030 pp1483 receipt measured 0.173x the established fallback)"
2031 .into(),
2032 );
2033 }
2034 assert!(
2035 in_f.is_multiple_of(16),
2036 "per-block FP8 MMQ requires in_f % 16 == 0, got {in_f}"
2037 );
2038 #[allow(clippy::manual_div_ceil)]
2039 let want_scales = ((out_f + 127) / 128) * ((in_f + 127) / 128);
2041 assert!(
2042 blk_scales.len() >= want_scales,
2043 "blk_scales too small: {} < {want_scales}",
2044 blk_scales.len()
2045 );
2046 assert!(
2047 w_e4m3.len() >= out_f * in_f,
2048 "e4m3 plane too small: {} < {}",
2049 w_e4m3.len(),
2050 out_f * in_f
2051 );
2052 let act_bytes = unsafe { memra_mmq_fp8_blk_act_bytes(in_f as i32, m as i32) };
2053 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
2054 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
2055 {
2056 let stream = self.gpu.stream();
2057 let (w_p, _gw) = w_e4m3.device_ptr(&stream);
2058 let (sc_p, _gsc) = blk_scales.device_ptr(&stream);
2059 let (x_p, _gx) = x.device_ptr(&stream);
2060 let (y_p, _gy) = y.device_ptr_mut(&stream);
2061 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
2062 let rc = unsafe {
2063 memra_mmq_fp8_blk(
2064 w_p as *const core::ffi::c_void,
2065 sc_p as *const f32,
2066 x_p as *const f32,
2067 y_p as *mut f32,
2068 in_f as i32,
2069 out_f as i32,
2070 m as i32,
2071 s_p as *mut core::ffi::c_void,
2072 stream.cu_stream() as *mut core::ffi::c_void,
2073 scale,
2074 )
2075 };
2076 if rc != 0 {
2077 return Err(format!("memra_mmq_fp8_blk rc={rc}").into());
2078 }
2079 }
2080 Ok(y)
2081 }
2082
2083 pub fn qmatvec_mmq_fp8_blk_view(
2088 &self,
2089 w_e4m3: &CudaView<'_, u8>,
2090 blk_scales: &CudaView<'_, f32>,
2091 x: &CudaView<'_, f32>,
2092 m: usize,
2093 in_f: usize,
2094 out_f: usize,
2095 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2096 if cfg!(memra_sm100_tcgen05) && std::env::var("MEMRA_FP8_MMQ").as_deref() != Ok("1") {
2097 return Err(
2098 "B200 block-FP8 tcgen05 is NativeReference but not tuned; set \
2099 MEMRA_FP8_MMQ=1 only for explicit qualification or research (the pinned \
2100 pp1483 receipt measured 0.173x the established fallback)"
2101 .into(),
2102 );
2103 }
2104 assert!(
2105 in_f.is_multiple_of(16),
2106 "per-block FP8 MMQ requires in_f % 16 == 0, got {in_f}"
2107 );
2108 let want_scales = out_f.div_ceil(128) * in_f.div_ceil(128);
2109 assert!(
2110 blk_scales.len() >= want_scales,
2111 "blk_scales view too small: {} < {want_scales}",
2112 blk_scales.len()
2113 );
2114 assert!(
2115 w_e4m3.len() >= out_f * in_f,
2116 "e4m3 view too small: {} < {}",
2117 w_e4m3.len(),
2118 out_f * in_f
2119 );
2120 assert!(
2121 x.len() >= m * in_f,
2122 "activation view too small: {} < {}",
2123 x.len(),
2124 m * in_f
2125 );
2126
2127 let act_bytes = unsafe { memra_mmq_fp8_blk_act_bytes(in_f as i32, m as i32) };
2128 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
2129 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
2130 {
2131 let stream = self.gpu.stream();
2132 let (w_p, _gw) = w_e4m3.device_ptr(&stream);
2133 let (sc_p, _gsc) = blk_scales.device_ptr(&stream);
2134 let (x_p, _gx) = x.device_ptr(&stream);
2135 let (y_p, _gy) = y.device_ptr_mut(&stream);
2136 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
2137 let rc = unsafe {
2138 memra_mmq_fp8_blk(
2139 w_p as *const core::ffi::c_void,
2140 sc_p as *const f32,
2141 x_p as *const f32,
2142 y_p as *mut f32,
2143 in_f as i32,
2144 out_f as i32,
2145 m as i32,
2146 s_p as *mut core::ffi::c_void,
2147 stream.cu_stream() as *mut core::ffi::c_void,
2148 1.0,
2149 )
2150 };
2151 if rc != 0 {
2152 return Err(format!("memra_mmq_fp8_blk(view) rc={rc}").into());
2153 }
2154 }
2155 Ok(y)
2156 }
2157
2158 pub fn fp8_blk_nan_count(
2162 &self,
2163 w_e4m3: &CudaSlice<u8>,
2164 ) -> Result<u32, Box<dyn std::error::Error>> {
2165 let mut cnt = self.htod_u32_v(&[0u32])?;
2166 let n = w_e4m3.len();
2167 {
2168 let stream = self.gpu.stream();
2169 let (w_p, _gw) = w_e4m3.device_ptr(&stream);
2170 let (c_p, _gc) = cnt.device_ptr_mut(&stream);
2171 let rc = unsafe {
2172 memra_fp8_blk_count_nan(
2173 w_p as *const core::ffi::c_void,
2174 n,
2175 c_p as *mut u32,
2176 stream.cu_stream() as *mut core::ffi::c_void,
2177 )
2178 };
2179 if rc != 0 {
2180 return Err(format!("memra_fp8_blk_count_nan rc={rc}").into());
2181 }
2182 }
2183 Ok(self.dtoh_u32(&cnt)?[0])
2184 }
2185
2186 pub fn mmq_iq_quantize_act(
2189 &self,
2190 x: &CudaSlice<f32>,
2191 in_f: usize,
2192 n_tokens: usize,
2193 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2194 let act_bytes = unsafe { memra_mmq_iq_experts_act_bytes(in_f as i32, n_tokens as i32) };
2195 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
2196 {
2197 let stream = self.gpu.stream();
2198 let (x_p, _gx) = x.device_ptr(&stream);
2199 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
2200 let rc = unsafe {
2201 memra_mmq_iq_quantize_act(
2202 x_p as *const f32,
2203 s_p as *mut core::ffi::c_void,
2204 in_f as i32,
2205 n_tokens as i32,
2206 stream.cu_stream() as *mut core::ffi::c_void,
2207 )
2208 };
2209 if rc != 0 {
2210 return Err(format!("memra_mmq_iq_quantize_act rc={rc}").into());
2211 }
2212 }
2213 Ok(scratch)
2214 }
2215
2216 pub fn mmq_iq_fused_act_quant(
2222 &self,
2223 gate: &CudaSlice<f32>,
2224 up: &CudaSlice<f32>,
2225 in_f: usize,
2226 n_tokens: usize,
2227 act_kind: i32,
2228 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2229 let act_bytes = unsafe { memra_mmq_iq_experts_act_bytes(in_f as i32, n_tokens as i32) };
2230 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
2231 {
2232 let stream = self.gpu.stream();
2233 let (g_p, _gg) = gate.device_ptr(&stream);
2234 let (u_p, _gu) = up.device_ptr(&stream);
2235 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
2236 let rc = unsafe {
2237 memra_mmq_iq_fused_act_quant(
2238 g_p as *const f32,
2239 u_p as *const f32,
2240 s_p as *mut core::ffi::c_void,
2241 in_f as i32,
2242 n_tokens as i32,
2243 act_kind,
2244 stream.cu_stream() as *mut core::ffi::c_void,
2245 )
2246 };
2247 if rc != 0 {
2248 return Err(format!("memra_mmq_iq_fused_act_quant rc={rc}").into());
2249 }
2250 }
2251 Ok(scratch)
2252 }
2253
2254 #[allow(clippy::too_many_arguments)]
2258 pub fn mmq_iq_experts(
2259 &self,
2260 table: &CudaSlice<u64>,
2261 proj: i32,
2262 n_expert: usize,
2263 ex_ids: &CudaSlice<i32>,
2264 ex_off: &CudaSlice<i32>,
2265 ex_pairs: &CudaSlice<i32>,
2266 pair_tok: &CudaSlice<i32>,
2267 act_scratch: &CudaSlice<u8>,
2268 in_f: usize,
2269 out_f: usize,
2270 n_active: usize,
2271 n_pairs: usize,
2272 n_tokens: usize,
2273 qtype: i32,
2274 row_bytes: usize,
2275 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2276 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2277 {
2278 let stream = self.gpu.stream();
2279 let (tab_p, _g0) = table.device_ptr(&stream);
2280 let (ei_p, _g1) = ex_ids.device_ptr(&stream);
2281 let (eo_p, _g2) = ex_off.device_ptr(&stream);
2282 let (ep_p, _g3) = ex_pairs.device_ptr(&stream);
2283 let (pt_p, _g4) = pair_tok.device_ptr(&stream);
2284 let (as_p, _g5) = act_scratch.device_ptr(&stream);
2285 let (y_p, _g6) = y.device_ptr_mut(&stream);
2286 let rc = unsafe {
2287 memra_mmq_iq_experts(
2288 tab_p as *const u64,
2289 proj,
2290 n_expert as i32,
2291 ei_p as *const i32,
2292 eo_p as *const i32,
2293 ep_p as *const i32,
2294 pt_p as *const i32,
2295 as_p as *const core::ffi::c_void,
2296 y_p as *mut f32,
2297 in_f as i32,
2298 out_f as i32,
2299 n_active as i32,
2300 n_tokens as i32,
2301 qtype,
2302 row_bytes as i64,
2303 stream.cu_stream() as *mut core::ffi::c_void,
2304 )
2305 };
2306 if rc != 0 {
2307 return Err(format!("memra_mmq_iq_experts rc={rc}").into());
2308 }
2309 }
2310 Ok(y)
2311 }
2312
2313 pub fn moe_f16g_act(
2318 &self,
2319 x: &CudaSlice<f32>,
2320 pair_tok: Option<&CudaSlice<i32>>,
2321 in_f: usize,
2322 n_pairs: usize,
2323 ) -> Result<(CudaSlice<u8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2324 let mut act = self.alloc_uninit::<u8>(n_pairs * in_f * 2)?;
2325 let mut scales = self.alloc_uninit::<f32>(n_pairs)?;
2326 {
2327 let stream = self.gpu.stream();
2328 let (x_p, _gx) = x.device_ptr(&stream);
2329 let pt_p = match pair_tok {
2330 Some(pt) => {
2331 let (p, _g) = pt.device_ptr(&stream);
2332 p as *const i32
2333 }
2334 None => std::ptr::null(),
2335 };
2336 let (a_p, _ga) = act.device_ptr_mut(&stream);
2337 let (s_p, _gs) = scales.device_ptr_mut(&stream);
2338 let rc = unsafe {
2339 memra_moe_f16g_gather_act(
2340 x_p as *const f32,
2341 pt_p,
2342 a_p as *mut core::ffi::c_void,
2343 s_p as *mut f32,
2344 in_f as i32,
2345 n_pairs as i32,
2346 stream.cu_stream() as *mut core::ffi::c_void,
2347 )
2348 };
2349 if rc != 0 {
2350 return Err(format!("memra_moe_f16g_gather_act rc={rc}").into());
2351 }
2352 }
2353 Ok((act, scales))
2354 }
2355
2356 #[allow(clippy::too_many_arguments)]
2366 pub fn bind_runtime_device(&self, ordinal: i32) -> Result<(), Box<dyn std::error::Error>> {
2371 let rc = unsafe { memra_bind_device(ordinal) };
2372 if rc != 0 {
2373 return Err(format!("cudaSetDevice({ordinal}) rc={rc}").into());
2374 }
2375 Ok(())
2376 }
2377
2378 #[allow(clippy::too_many_arguments)]
2379 #[allow(clippy::manual_is_multiple_of)] pub fn moe_f16_grouped(
2382 &self,
2383 table: &CudaSlice<u64>,
2384 proj: i32,
2385 n_expert: usize,
2386 ex_ids: &CudaSlice<i32>,
2387 ex_off_host: &[i32],
2388 ex_off_dev: &CudaSlice<i32>,
2389 act_f16: &CudaSlice<u8>,
2390 act_scale: &CudaSlice<f32>,
2391 in_f: usize,
2392 out_f: usize,
2393 n_active: usize,
2394 n_pairs: usize,
2395 qtype: i32,
2396 row_bytes: usize,
2397 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2398 let sk = crate::moe_f16g_mode() >= 2 && in_f.is_multiple_of(32);
2399 let (shape_sel, cross) = crate::moe_f16g_sk_params();
2407 if sk
2408 && shape_sel >= 0
2409 && crate::moe_f16g_direct_on(qtype)
2410 && (qtype == crate::QT_Q4_K
2411 || qtype == crate::QT_Q6_K
2412 || qtype == crate::QT_IQ4_XS
2413 || qtype == crate::QT_IQ3_S
2414 || qtype == crate::QT_NVFP4
2415 || qtype == crate::QT_NVFP4_V2)
2419 && in_f % (if qtype == crate::QT_NVFP4 || qtype == crate::QT_NVFP4_V2 { 64 } else { 256 }) == 0
2422 && n_active <= 512
2423 && n_active > 0
2424 {
2425 let max_m = ex_off_host
2426 .windows(2)
2427 .map(|w| w[1] - w[0])
2428 .max()
2429 .unwrap_or(0);
2430 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2431 {
2432 let stream = self.gpu.stream();
2433 let (tab_p, _g0) = table.device_ptr(&stream);
2434 let (ei_p, _g1) = ex_ids.device_ptr(&stream);
2435 let (a_p, _g2) = act_f16.device_ptr(&stream);
2436 let (s_p, _g3) = act_scale.device_ptr(&stream);
2437 let (off_p, _g4) = ex_off_dev.device_ptr(&stream);
2438 let (y_p, _g5) = y.device_ptr_mut(&stream);
2439 let rc = unsafe {
2440 memra_moe_kq_gemm_sk(
2441 tab_p as *const u64,
2442 proj,
2443 n_expert as i32,
2444 ei_p as *const i32,
2445 a_p as *const core::ffi::c_void,
2446 y_p as *mut f32,
2447 s_p as *const f32,
2448 off_p as *const i32,
2449 ex_off_host.as_ptr(),
2450 n_active as i32,
2451 max_m,
2452 in_f as i32,
2453 out_f as i32,
2454 qtype,
2455 cross,
2456 crate::moe_f16g_tail_on() as i32,
2457 row_bytes as i64,
2458 stream.cu_stream() as *mut core::ffi::c_void,
2459 )
2460 };
2461 if rc != 0 {
2462 return Err(format!("memra_moe_kq_gemm_sk rc={rc}").into());
2463 }
2464 }
2465 return Ok(y);
2466 }
2467 if !sk {
2471 static WARM: std::sync::Once = std::sync::Once::new();
2472 let mut warm_err = None;
2473 WARM.call_once(|| {
2474 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
2475 let w = self.alloc_uninit::<u8>(2 * 32 * 64 * 2)?;
2476 let a = self.alloc_uninit::<u8>(4 * 64 * 2)?;
2477 let mut yw = self.alloc_uninit::<u8>(4 * 32 * 2)?;
2478 let off = [0i32, 2, 4];
2479 let stream = self.gpu.stream();
2480 let (w_p, _a1) = w.device_ptr(&stream);
2481 let (a_p, _a2) = a.device_ptr(&stream);
2482 let (y_p, _a3) = yw.device_ptr_mut(&stream);
2483 let rc = unsafe {
2484 memra_moe_f16g_gemm(
2485 w_p as *const core::ffi::c_void,
2486 a_p as *const core::ffi::c_void,
2487 y_p as *mut core::ffi::c_void,
2488 off.as_ptr(),
2489 2,
2490 64,
2491 32,
2492 stream.cu_stream() as *mut core::ffi::c_void,
2493 )
2494 };
2495 if rc != 0 {
2496 return Err(format!("f16g warmup rc={rc}").into());
2497 }
2498 self.gpu.stream().synchronize()?;
2499 Ok(())
2500 })();
2501 if let Err(e) = r {
2502 warm_err = Some(e.to_string());
2503 }
2504 });
2505 if let Some(we) = warm_err {
2506 return Err(we.into());
2507 }
2508 }
2509 let w_bytes = n_active * out_f * in_f * 2;
2510 let mut w_f16 = self.alloc_uninit::<u8>(w_bytes)?;
2511 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2512 {
2513 let stream = self.gpu.stream();
2514 let (tab_p, _g0) = table.device_ptr(&stream);
2515 let (ei_p, _g1) = ex_ids.device_ptr(&stream);
2516 let (w_p, _g2) = w_f16.device_ptr_mut(&stream);
2517 let rc = unsafe {
2518 memra_moe_f16g_dequant(
2519 tab_p as *const u64,
2520 proj,
2521 n_expert as i32,
2522 ei_p as *const i32,
2523 w_p as *mut core::ffi::c_void,
2524 in_f as i32,
2525 out_f as i32,
2526 n_active as i32,
2527 qtype,
2528 row_bytes as i64,
2529 stream.cu_stream() as *mut core::ffi::c_void,
2530 )
2531 };
2532 if rc != 0 {
2533 return Err(format!("memra_moe_f16g_dequant rc={rc}").into());
2534 }
2535 let (a_p, _g3) = act_f16.device_ptr(&stream);
2536 let (s_p, _g6) = act_scale.device_ptr(&stream);
2537 let (y_p, _g5) = y.device_ptr_mut(&stream);
2538 if sk {
2539 let max_m = ex_off_host
2540 .windows(2)
2541 .map(|w| w[1] - w[0])
2542 .max()
2543 .unwrap_or(0);
2544 let (off_p, _g7) = ex_off_dev.device_ptr(&stream);
2545 let (shape_sel, cross) = crate::moe_f16g_sk_params();
2546 let rc = unsafe {
2547 memra_moe_f16g_gemm_sk(
2548 w_p as *const core::ffi::c_void,
2549 a_p as *const core::ffi::c_void,
2550 y_p as *mut f32,
2551 s_p as *const f32,
2552 off_p as *const i32,
2553 ex_off_host.as_ptr(),
2554 n_active as i32,
2555 max_m,
2556 in_f as i32,
2557 out_f as i32,
2558 shape_sel,
2559 cross,
2560 crate::moe_f16g_tail_on() as i32,
2561 stream.cu_stream() as *mut core::ffi::c_void,
2562 )
2563 };
2564 if rc != 0 {
2565 return Err(format!("memra_moe_f16g_gemm_sk rc={rc}").into());
2566 }
2567 } else {
2568 let mut y16 = self.alloc_uninit::<u8>(n_pairs * out_f * 2)?;
2569 let (y16_p, _g4) = y16.device_ptr_mut(&stream);
2570 let rc = unsafe {
2571 memra_moe_f16g_gemm(
2572 w_p as *const core::ffi::c_void,
2573 a_p as *const core::ffi::c_void,
2574 y16_p as *mut core::ffi::c_void,
2575 ex_off_host.as_ptr(),
2576 n_active as i32,
2577 in_f as i32,
2578 out_f as i32,
2579 stream.cu_stream() as *mut core::ffi::c_void,
2580 )
2581 };
2582 if rc != 0 {
2583 return Err(format!("memra_moe_f16g_gemm rc={rc}").into());
2584 }
2585 let rc = unsafe {
2586 memra_moe_f16g_h2f_scaled(
2587 y16_p as *const core::ffi::c_void,
2588 y_p as *mut f32,
2589 s_p as *const f32,
2590 out_f as i32,
2591 n_pairs as i32,
2592 stream.cu_stream() as *mut core::ffi::c_void,
2593 )
2594 };
2595 if rc != 0 {
2596 return Err(format!("memra_moe_f16g_h2f_scaled rc={rc}").into());
2597 }
2598 }
2599 }
2600 if !sk {
2605 self.gpu.stream().synchronize()?;
2606 }
2607 if std::env::var("MEMRA_F16G_DEBUG").is_ok() {
2608 let wn = n_active * out_f * in_f;
2610 let an = n_pairs * in_f;
2611 let mut wf = self.alloc_uninit::<f32>(wn)?;
2612 let mut af = self.alloc_uninit::<f32>(an)?;
2613 {
2614 let stream = self.gpu.stream();
2615 let (w_p, _a) = w_f16.device_ptr(&stream);
2616 let (a_p, _b) = act_f16.device_ptr(&stream);
2617 let (wf_p, _c) = wf.device_ptr_mut(&stream);
2618 let (af_p, _d) = af.device_ptr_mut(&stream);
2619 unsafe {
2620 memra_moe_f16g_h2f(
2621 w_p as *const core::ffi::c_void,
2622 wf_p as *mut f32,
2623 wn,
2624 stream.cu_stream() as *mut core::ffi::c_void,
2625 );
2626 memra_moe_f16g_h2f(
2627 a_p as *const core::ffi::c_void,
2628 af_p as *mut f32,
2629 an,
2630 stream.cu_stream() as *mut core::ffi::c_void,
2631 );
2632 }
2633 }
2634 let (wh, ah, yh) = (self.dtoh(&wf)?, self.dtoh(&af)?, self.dtoh(&y)?);
2635 let scan = |v: &[f32]| -> (usize, f32) {
2636 let bad = v.iter().filter(|x| !x.is_finite()).count();
2637 let mx = v
2638 .iter()
2639 .filter(|x| x.is_finite())
2640 .fold(0.0f32, |m, x| m.max(x.abs()));
2641 (bad, mx)
2642 };
2643 let (wb, wm) = scan(&wh);
2644 let (ab, am) = scan(&ah);
2645 let (yb, ym) = scan(&yh);
2646 eprintln!(
2647 "[f16g-debug] proj={proj} w: bad={wb} max={wm:.3e} | act: bad={ab} \
2648 max={am:.3e} | y: bad={yb} max={ym:.3e} (na={n_active} np={n_pairs} \
2649 in={in_f} out={out_f})"
2650 );
2651 }
2652 Ok(y)
2653 }
2654
2655 #[allow(clippy::too_many_arguments)]
2662 pub fn moe_f16g_gemm_sk_raw(
2663 &self,
2664 w_f16: &CudaSlice<u8>,
2665 act_f16: &CudaSlice<u8>,
2666 row_scale: &CudaSlice<f32>,
2667 ex_off_host: &[i32],
2668 ex_off_dev: &CudaSlice<i32>,
2669 in_f: usize,
2670 out_f: usize,
2671 n_pairs: usize,
2672 shape_sel: i32,
2673 cross: i32,
2674 tail: i32,
2675 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2676 let n_active = ex_off_host.len() - 1;
2677 let max_m = ex_off_host
2678 .windows(2)
2679 .map(|w| w[1] - w[0])
2680 .max()
2681 .unwrap_or(0);
2682 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2683 {
2684 let stream = self.gpu.stream();
2685 let (w_p, _g0) = w_f16.device_ptr(&stream);
2686 let (a_p, _g1) = act_f16.device_ptr(&stream);
2687 let (s_p, _g2) = row_scale.device_ptr(&stream);
2688 let (off_p, _g3) = ex_off_dev.device_ptr(&stream);
2689 let (y_p, _g4) = y.device_ptr_mut(&stream);
2690 let rc = unsafe {
2691 memra_moe_f16g_gemm_sk(
2692 w_p as *const core::ffi::c_void,
2693 a_p as *const core::ffi::c_void,
2694 y_p as *mut f32,
2695 s_p as *const f32,
2696 off_p as *const i32,
2697 ex_off_host.as_ptr(),
2698 n_active as i32,
2699 max_m,
2700 in_f as i32,
2701 out_f as i32,
2702 shape_sel,
2703 cross,
2704 tail,
2705 stream.cu_stream() as *mut core::ffi::c_void,
2706 )
2707 };
2708 if rc != 0 {
2709 return Err(format!("memra_moe_f16g_gemm_sk rc={rc}").into());
2710 }
2711 }
2712 Ok(y)
2713 }
2714
2715 #[allow(clippy::too_many_arguments)]
2720 pub fn moe_kq_gemm_sk_raw(
2721 &self,
2722 table: &CudaSlice<u64>,
2723 proj: i32,
2724 n_expert: usize,
2725 ex_ids: &CudaSlice<i32>,
2726 act_f16: &CudaSlice<u8>,
2727 row_scale: &CudaSlice<f32>,
2728 ex_off_host: &[i32],
2729 ex_off_dev: &CudaSlice<i32>,
2730 in_f: usize,
2731 out_f: usize,
2732 n_pairs: usize,
2733 qtype: i32,
2734 row_bytes: usize,
2735 cross: i32,
2736 tail: i32,
2737 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2738 let n_active = ex_off_host.len() - 1;
2739 let max_m = ex_off_host
2740 .windows(2)
2741 .map(|w| w[1] - w[0])
2742 .max()
2743 .unwrap_or(0);
2744 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2745 {
2746 let stream = self.gpu.stream();
2747 let (tab_p, _g0) = table.device_ptr(&stream);
2748 let (ei_p, _g1) = ex_ids.device_ptr(&stream);
2749 let (a_p, _g2) = act_f16.device_ptr(&stream);
2750 let (s_p, _g3) = row_scale.device_ptr(&stream);
2751 let (off_p, _g4) = ex_off_dev.device_ptr(&stream);
2752 let (y_p, _g5) = y.device_ptr_mut(&stream);
2753 let rc = unsafe {
2754 memra_moe_kq_gemm_sk(
2755 tab_p as *const u64,
2756 proj,
2757 n_expert as i32,
2758 ei_p as *const i32,
2759 a_p as *const core::ffi::c_void,
2760 y_p as *mut f32,
2761 s_p as *const f32,
2762 off_p as *const i32,
2763 ex_off_host.as_ptr(),
2764 n_active as i32,
2765 max_m,
2766 in_f as i32,
2767 out_f as i32,
2768 qtype,
2769 cross,
2770 tail,
2771 row_bytes as i64,
2772 stream.cu_stream() as *mut core::ffi::c_void,
2773 )
2774 };
2775 if rc != 0 {
2776 return Err(format!("memra_moe_kq_gemm_sk rc={rc}").into());
2777 }
2778 }
2779 Ok(y)
2780 }
2781
2782 #[allow(clippy::too_many_arguments)] pub fn moe_f16g_dequant_raw(
2787 &self,
2788 table: &CudaSlice<u64>,
2789 proj: i32,
2790 n_expert: usize,
2791 ex_ids: &CudaSlice<i32>,
2792 in_f: usize,
2793 out_f: usize,
2794 n_active: usize,
2795 qtype: i32,
2796 row_bytes: usize,
2797 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2798 let mut w_f16 = self.alloc_uninit::<u8>(n_active * out_f * in_f * 2)?;
2799 {
2800 let stream = self.gpu.stream();
2801 let (tab_p, _g0) = table.device_ptr(&stream);
2802 let (ei_p, _g1) = ex_ids.device_ptr(&stream);
2803 let (w_p, _g2) = w_f16.device_ptr_mut(&stream);
2804 let rc = unsafe {
2805 memra_moe_f16g_dequant(
2806 tab_p as *const u64,
2807 proj,
2808 n_expert as i32,
2809 ei_p as *const i32,
2810 w_p as *mut core::ffi::c_void,
2811 in_f as i32,
2812 out_f as i32,
2813 n_active as i32,
2814 qtype,
2815 row_bytes as i64,
2816 stream.cu_stream() as *mut core::ffi::c_void,
2817 )
2818 };
2819 if rc != 0 {
2820 return Err(format!("memra_moe_f16g_dequant rc={rc}").into());
2821 }
2822 }
2823 Ok(w_f16)
2824 }
2825}