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 % 16 != 0
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
1265impl Engine {
1266 pub fn mmq_supports(&self, w: &crate::model::GpuTensor) -> bool {
1269 use crate::model::GpuTensor;
1270 if crate::portable_mma_gated() {
1271 return false;
1272 }
1273 let mmq_opt_in = std::env::var("MEMRA_MMQ").is_ok();
1274 match w {
1275 GpuTensor::Quant { qtype, rp, .. } if *qtype == crate::QT_NVFP4 && *rp => {
1283 !cfg!(memra_portable_cuda) && mmq_w4a8_enabled() && w.in_features() % 64 == 0
1284 }
1285 GpuTensor::Quant { qtype, .. } if *qtype == crate::QT_NVFP4 => {
1287 !cfg!(memra_portable_cuda)
1288 && (mmq_w4a8_enabled() || mmq_opt_in)
1289 && w.in_features() % 64 == 0
1290 }
1291 GpuTensor::Quant { qtype, .. }
1292 if *qtype == crate::QT_Q4_K || *qtype == crate::QT_Q5_K =>
1293 {
1294 (mmq_w4a8_enabled() || mmq_opt_in) && w.in_features() % 256 == 0
1295 }
1296 GpuTensor::Quant { qtype, .. } if *qtype == crate::QT_Q8_0 => {
1301 mmq_q8_enabled() && w.in_features() % 256 == 0
1302 }
1303 GpuTensor::Quant { qtype, .. } if *qtype == crate::QT_Q4_0 => {
1308 mmq_q4_enabled() && w.in_features() % 256 == 0
1309 }
1310 GpuTensor::Quant { qtype, .. } if *qtype == crate::QT_IQ4_XS => {
1316 mmq_iq4xs_enabled() && Self::iq_fast_enabled() && w.in_features() % 256 == 0
1317 }
1318 _ => false,
1319 }
1320 }
1321
1322 pub fn qmatvec_mmq(
1325 &self,
1326 w: &crate::model::GpuTensor,
1327 x: &CudaSlice<f32>,
1328 m: usize,
1329 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1330 use crate::model::GpuTensor;
1331 let (in_f, out_f) = (w.in_features(), w.out_features());
1332 let GpuTensor::Quant {
1333 bytes,
1334 scale,
1335 qtype,
1336 rp,
1337 ..
1338 } = w
1339 else {
1340 return Err("qmatvec_mmq: not a Quant tensor".into());
1341 };
1342 let w4a8_explicit = std::env::var("MEMRA_MMQ_W4A8")
1347 .map(|v| v != "0")
1348 .unwrap_or(false);
1349 let use_w4a8 =
1350 *rp || w4a8_explicit || (mmq_w4a8_enabled() && std::env::var("MEMRA_MMQ").is_err());
1351 match *qtype {
1352 q if q == crate::QT_NVFP4 && use_w4a8 => {
1355 self.qmatvec_mmq_nvfp4_w4a8(bytes, x, m, in_f, out_f, *scale, *rp)
1356 }
1357 q if q == crate::QT_NVFP4 => self.qmatvec_mmq_nvfp4(bytes, x, m, in_f, out_f, *scale),
1358 q if q == crate::QT_Q4_K || q == crate::QT_Q5_K => {
1359 let mut y = self.qmatvec_mmq_q45k_raw(bytes, x, m, in_f, out_f, q)?;
1360 if *scale != 1.0 {
1361 self.scale_inplace(&mut y, *scale, m * out_f)?;
1362 }
1363 Ok(y)
1364 }
1365 q if q == crate::QT_Q8_0 => {
1366 if cfg!(memra_hopper_mma) && out_f % 64 == 0 && crate::wgmma_gemm_enabled() {
1372 if let GpuTensor::Quant { rp4: Some(m4), .. } = w {
1373 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
1374 let mut y =
1375 self.qmatvec_gemm_q8_0_wgmma_raw(m4, &aq, &ad, m, in_f, out_f)?;
1376 if *scale != 1.0 {
1377 self.scale_inplace(&mut y, *scale, m * out_f)?;
1378 }
1379 return Ok(y);
1380 }
1381 }
1382 let mut y = self.qmatvec_mmq_q8_0_raw(bytes, x, m, in_f, out_f)?;
1383 if *scale != 1.0 {
1384 self.scale_inplace(&mut y, *scale, m * out_f)?;
1385 }
1386 Ok(y)
1387 }
1388 q if q == crate::QT_Q4_0 => {
1389 let mut y = self.qmatvec_mmq_q4_0_raw(bytes, x, m, in_f, out_f, *rp)?;
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_IQ4_XS => {
1396 let GpuTensor::Quant { row_bytes, .. } = w else {
1397 unreachable!()
1398 };
1399 let mut y = self.qmatvec_mmq_iq4xs_raw(bytes, x, m, in_f, out_f, *row_bytes)?;
1400 if *scale != 1.0 {
1401 self.scale_inplace(&mut y, *scale, m * out_f)?;
1402 }
1403 Ok(y)
1404 }
1405 q => Err(format!("qmatvec_mmq: unsupported qtype {q}").into()),
1406 }
1407 }
1408
1409 pub fn qmatvec_mmq_iq4xs_raw(
1411 &self,
1412 bytes: &CudaSlice<u8>,
1413 x: &CudaSlice<f32>,
1414 m: usize,
1415 in_f: usize,
1416 out_f: usize,
1417 row_bytes: usize,
1418 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1419 assert!(
1420 in_f % 256 == 0,
1421 "MMQ IQ4_XS requires in_f % 256 == 0, got {in_f}"
1422 );
1423 let act_bytes = unsafe { memra_mmq_iq_experts_act_bytes(in_f as i32, m as i32) };
1424 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
1425 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
1426 {
1427 let stream = self.gpu.stream();
1428 let (w_p, _gw) = bytes.device_ptr(&stream);
1429 let (x_p, _gx) = x.device_ptr(&stream);
1430 let (y_p, _gy) = y.device_ptr_mut(&stream);
1431 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
1432 let rc = unsafe {
1433 memra_mmq_iq4xs_dense(
1434 w_p as *const core::ffi::c_void,
1435 x_p as *const f32,
1436 y_p as *mut f32,
1437 in_f as i32,
1438 out_f as i32,
1439 m as i32,
1440 row_bytes as i64,
1441 s_p as *mut core::ffi::c_void,
1442 stream.cu_stream() as *mut core::ffi::c_void,
1443 )
1444 };
1445 if rc != 0 {
1446 return Err(format!("memra_mmq_iq4xs_dense rc={rc}").into());
1447 }
1448 }
1449 Ok(y)
1450 }
1451
1452 pub fn qmatvec_mmq_q45k_raw(
1457 &self,
1458 bytes: &CudaSlice<u8>,
1459 x: &CudaSlice<f32>,
1460 m: usize,
1461 in_f: usize,
1462 out_f: usize,
1463 qtype: i32,
1464 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1465 assert!(
1466 in_f % 256 == 0,
1467 "MMQ Q4_K/Q5_K requires in_f % 256 == 0, got {in_f}"
1468 );
1469 let act_bytes = unsafe { memra_mmq_q45k_act_bytes(in_f as i32, m as i32) };
1470 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
1471 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
1472 {
1473 let stream = self.gpu.stream();
1474 let (w_p, _gw) = bytes.device_ptr(&stream);
1475 let (x_p, _gx) = x.device_ptr(&stream);
1476 let (y_p, _gy) = y.device_ptr_mut(&stream);
1477 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
1478 let launcher = if qtype == crate::QT_Q4_K {
1479 memra_mmq_q4_K
1480 } else {
1481 memra_mmq_q5_K
1482 };
1483 let rc = unsafe {
1484 launcher(
1485 w_p as *const core::ffi::c_void,
1486 x_p as *const f32,
1487 y_p as *mut f32,
1488 in_f as i32,
1489 out_f as i32,
1490 m as i32,
1491 s_p as *mut core::ffi::c_void,
1492 stream.cu_stream() as *mut core::ffi::c_void,
1493 )
1494 };
1495 if rc != 0 {
1496 return Err(format!("memra_mmq_q45k(qtype={qtype}) rc={rc}").into());
1497 }
1498 }
1499 Ok(y)
1500 }
1501
1502 pub fn qmatvec_mmq_q8_0_raw(
1505 &self,
1506 bytes: &CudaSlice<u8>,
1507 x: &CudaSlice<f32>,
1508 m: usize,
1509 in_f: usize,
1510 out_f: usize,
1511 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1512 assert!(
1513 in_f % 32 == 0,
1514 "MMQ Q8_0 requires in_f % 32 == 0, got {in_f}"
1515 );
1516 let act_bytes = unsafe { memra_mmq_q8_0_act_bytes(in_f as i32, m as i32) };
1517 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
1518 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
1519 {
1520 let stream = self.gpu.stream();
1521 let (w_p, _gw) = bytes.device_ptr(&stream);
1522 let (x_p, _gx) = x.device_ptr(&stream);
1523 let (y_p, _gy) = y.device_ptr_mut(&stream);
1524 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
1525 let rc = unsafe {
1526 memra_mmq_q8_0(
1527 w_p as *const core::ffi::c_void,
1528 x_p as *const f32,
1529 y_p as *mut f32,
1530 in_f as i32,
1531 out_f as i32,
1532 m as i32,
1533 s_p as *mut core::ffi::c_void,
1534 stream.cu_stream() as *mut core::ffi::c_void,
1535 )
1536 };
1537 if rc != 0 {
1538 return Err(format!("memra_mmq_q8_0 rc={rc}").into());
1539 }
1540 }
1541 Ok(y)
1542 }
1543
1544 pub fn accprobe_act_bytes(&self, in_f: usize, m: usize) -> usize {
1547 unsafe { memra_accprobe_act_bytes(in_f as i32, m as i32) }
1548 }
1549
1550 pub fn accprobe_gemm(
1557 &self,
1558 w_q8_0: &CudaSlice<u8>,
1559 act_q: &CudaSlice<u8>,
1560 m: usize,
1561 in_f: usize,
1562 out_f: usize,
1563 f32acc: bool,
1564 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1565 assert!(
1566 in_f % 32 == 0,
1567 "accprobe requires in_f % 32 == 0, got {in_f}"
1568 );
1569 assert!(
1570 act_q.len() >= self.accprobe_act_bytes(in_f, m),
1571 "accprobe act_q too small: {} < {}",
1572 act_q.len(),
1573 self.accprobe_act_bytes(in_f, m)
1574 );
1575 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
1576 {
1577 let stream = self.gpu.stream();
1578 let (w_p, _gw) = w_q8_0.device_ptr(&stream);
1579 let (a_p, _ga) = act_q.device_ptr(&stream);
1580 let (y_p, _gy) = y.device_ptr_mut(&stream);
1581 let f = if f32acc {
1582 memra_accprobe_gemm_f32
1583 } else {
1584 memra_accprobe_gemm_s32
1585 };
1586 let rc = unsafe {
1587 f(
1588 w_p as *const core::ffi::c_void,
1589 a_p as *const core::ffi::c_void,
1590 y_p as *mut f32,
1591 in_f as i32,
1592 out_f as i32,
1593 m as i32,
1594 stream.cu_stream() as *mut core::ffi::c_void,
1595 )
1596 };
1597 if rc != 0 {
1598 let arm = if f32acc { "f32" } else { "s32" };
1599 return Err(format!("memra_accprobe_gemm_{arm} rc={rc}").into());
1600 }
1601 }
1602 Ok(y)
1603 }
1604
1605 pub fn mmq_act_begin(&self) {
1611 use std::sync::atomic::Ordering;
1612 MMQ_ACT_EPOCH.fetch_add(1, Ordering::Relaxed);
1613 *MMQ_ACT_SLOT.lock().unwrap() = None;
1614 }
1615
1616 pub fn qmatvec_mmq_q4_0_raw(
1620 &self,
1621 bytes: &CudaSlice<u8>,
1622 x: &CudaSlice<f32>,
1623 m: usize,
1624 in_f: usize,
1625 out_f: usize,
1626 rp: bool,
1627 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1628 use std::sync::atomic::Ordering;
1629 assert!(
1630 in_f % 32 == 0,
1631 "MMQ Q4_0 requires in_f % 32 == 0, got {in_f}"
1632 );
1633 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
1634 let stream = self.gpu.stream();
1635 let (x_p, _gx) = x.device_ptr(&stream);
1636 let epoch = MMQ_ACT_EPOCH.load(Ordering::Relaxed);
1637 let mut slot = MMQ_ACT_SLOT.lock().unwrap();
1639 let hit = matches!(&*slot,
1640 Some((e, p, mm, inf, _)) if *e == epoch && *p == x_p as u64 && *mm == m && *inf == in_f);
1641 if !hit {
1642 let act_bytes = unsafe { memra_mmq_q4_0_act_bytes(in_f as i32, m as i32) };
1643 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
1644 {
1645 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
1646 let rc = unsafe {
1647 memra_mmq_q4_0_quant_act(
1648 x_p as *const f32,
1649 s_p as *mut core::ffi::c_void,
1650 in_f as i32,
1651 m as i32,
1652 stream.cu_stream() as *mut core::ffi::c_void,
1653 )
1654 };
1655 if rc != 0 {
1656 return Err(
1657 format!("memra_mmq_q4_0_quant_act(in_f={in_f}, m={m}) rc={rc}").into(),
1658 );
1659 }
1660 }
1661 *slot = Some((epoch, x_p as u64, m, in_f, scratch));
1662 }
1663 let scratch = &slot.as_ref().unwrap().4;
1664 {
1665 let (w_p, _gw) = bytes.device_ptr(&stream);
1666 let (y_p, _gy) = y.device_ptr_mut(&stream);
1667 let (s_p, _gs) = scratch.device_ptr(&stream);
1668 static SK_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1691 let sk = match crate::MMQ_SK_FORCE.load(std::sync::atomic::Ordering::Relaxed) {
1692 0 => false,
1693 1 => true,
1694 _ => *SK_ON.get_or_init(|| {
1695 std::env::var("MEMRA_MMQ_SK")
1696 .map(|v| v != "0")
1697 .unwrap_or(!cfg!(memra_hopper_mma))
1698 }),
1699 };
1700 let rc = if sk {
1701 let mut fx = MMQ_FIXUP_SLOT.lock().unwrap();
1702 if fx.is_none() {
1703 let nb = unsafe { memra_mmq_q4_0_fixup_bytes() };
1704 *fx = Some(self.alloc_uninit::<u8>(nb)?);
1705 }
1706 let (f_p, _gf) = fx.as_mut().unwrap().device_ptr_mut(&stream);
1707 unsafe {
1708 memra_mmq_q4_0_gemm_sk(
1709 w_p as *const core::ffi::c_void,
1710 s_p as *const core::ffi::c_void,
1711 y_p as *mut f32,
1712 f_p as *mut core::ffi::c_void,
1713 in_f as i32,
1714 out_f as i32,
1715 m as i32,
1716 stream.cu_stream() as *mut core::ffi::c_void,
1717 rp as i32,
1718 )
1719 }
1720 } else {
1721 unsafe {
1722 memra_mmq_q4_0_gemm(
1723 w_p as *const core::ffi::c_void,
1724 s_p as *const core::ffi::c_void,
1725 y_p as *mut f32,
1726 in_f as i32,
1727 out_f as i32,
1728 m as i32,
1729 stream.cu_stream() as *mut core::ffi::c_void,
1730 rp as i32,
1731 )
1732 }
1733 };
1734 if rc != 0 {
1735 return Err(format!(
1736 "memra_mmq_q4_0_gemm(rp={rp}, in_f={in_f}, out_f={out_f}, m={m}, wbytes={}) rc={rc}",
1737 bytes.len()
1738 )
1739 .into());
1740 }
1741 }
1742 Ok(y)
1743 }
1744
1745 pub fn qmatvec_mmq_nvfp4(
1751 &self,
1752 bytes: &CudaSlice<u8>,
1753 x: &CudaSlice<f32>,
1754 m: usize,
1755 in_f: usize,
1756 out_f: usize,
1757 scale: f32,
1758 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1759 self.qmatvec_mmq_nvfp4_scaled(bytes, x, m, in_f, out_f, scale)
1760 }
1761
1762 pub fn qmatvec_mmq_nvfp4_raw(
1764 &self,
1765 bytes: &CudaSlice<u8>,
1766 x: &CudaSlice<f32>,
1767 m: usize,
1768 in_f: usize,
1769 out_f: usize,
1770 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1771 self.qmatvec_mmq_nvfp4_scaled(bytes, x, m, in_f, out_f, 1.0)
1772 }
1773
1774 pub fn qmatvec_mmq_nvfp4_raw_v1(
1778 &self,
1779 bytes: &CudaSlice<u8>,
1780 x: &CudaSlice<f32>,
1781 m: usize,
1782 in_f: usize,
1783 out_f: usize,
1784 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1785 self.qmatvec_mmq_nvfp4_inner(bytes, x, m, in_f, out_f, 1.0, false, 0)
1786 }
1787
1788 pub fn qmatvec_mmq_nvfp4_raw_res(
1790 &self,
1791 bytes: &CudaSlice<u8>,
1792 x: &CudaSlice<f32>,
1793 m: usize,
1794 in_f: usize,
1795 out_f: usize,
1796 residual_k: i32,
1797 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1798 self.qmatvec_mmq_nvfp4_inner(bytes, x, m, in_f, out_f, 1.0, true, residual_k)
1799 }
1800
1801 fn qmatvec_mmq_nvfp4_scaled(
1802 &self,
1803 bytes: &CudaSlice<u8>,
1804 x: &CudaSlice<f32>,
1805 m: usize,
1806 in_f: usize,
1807 out_f: usize,
1808 scale: f32,
1809 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1810 self.qmatvec_mmq_nvfp4_inner(bytes, x, m, in_f, out_f, scale, true, mmq_residual_k())
1811 }
1812
1813 fn qmatvec_mmq_nvfp4_inner(
1814 &self,
1815 bytes: &CudaSlice<u8>,
1816 x: &CudaSlice<f32>,
1817 m: usize,
1818 in_f: usize,
1819 out_f: usize,
1820 scale: f32,
1821 per_token_scale: bool,
1822 residual_k: i32,
1823 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1824 assert!(
1825 in_f % 64 == 0,
1826 "MMQ NVFP4 requires in_f % 64 == 0, got {in_f}"
1827 );
1828 let act_bytes = unsafe { memra_mmq_nvfp4_act_bytes(in_f as i32, m as i32) };
1829 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
1830 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
1831 {
1832 let stream = self.gpu.stream();
1833 let (w_p, _gw) = bytes.device_ptr(&stream);
1834 let (x_p, _gx) = x.device_ptr(&stream);
1835 let (y_p, _gy) = y.device_ptr_mut(&stream);
1836 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
1837 let rc = unsafe {
1838 memra_mmq_nvfp4_ex2(
1839 w_p as *const core::ffi::c_void,
1840 x_p as *const f32,
1841 y_p as *mut f32,
1842 in_f as i32,
1843 out_f as i32,
1844 m as i32,
1845 s_p as *mut core::ffi::c_void,
1846 stream.cu_stream() as *mut core::ffi::c_void,
1847 scale,
1848 per_token_scale as i32,
1849 residual_k,
1850 )
1851 };
1852 if rc != 0 {
1853 return Err(format!("memra_mmq_nvfp4_ex2 rc={rc}").into());
1854 }
1855 }
1856 Ok(y)
1857 }
1858
1859 pub fn qmatvec_mmq_nvfp4_w4a8(
1864 &self,
1865 bytes: &CudaSlice<u8>,
1866 x: &CudaSlice<f32>,
1867 m: usize,
1868 in_f: usize,
1869 out_f: usize,
1870 scale: f32,
1871 rp: bool,
1872 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1873 self.qmatvec_mmq_nvfp4_w4a8_scaled(bytes, x, m, in_f, out_f, scale, rp)
1874 }
1875
1876 pub fn qmatvec_mmq_nvfp4_w4a8_raw(
1878 &self,
1879 bytes: &CudaSlice<u8>,
1880 x: &CudaSlice<f32>,
1881 m: usize,
1882 in_f: usize,
1883 out_f: usize,
1884 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1885 self.qmatvec_mmq_nvfp4_w4a8_scaled(bytes, x, m, in_f, out_f, 1.0, false)
1886 }
1887
1888 pub fn qmatvec_mmq_nvfp4_w4a8_raw_rp(
1891 &self,
1892 bytes: &CudaSlice<u8>,
1893 x: &CudaSlice<f32>,
1894 m: usize,
1895 in_f: usize,
1896 out_f: usize,
1897 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1898 self.qmatvec_mmq_nvfp4_w4a8_scaled(bytes, x, m, in_f, out_f, 1.0, true)
1899 }
1900
1901 fn qmatvec_mmq_nvfp4_w4a8_scaled(
1902 &self,
1903 bytes: &CudaSlice<u8>,
1904 x: &CudaSlice<f32>,
1905 m: usize,
1906 in_f: usize,
1907 out_f: usize,
1908 scale: f32,
1909 rp: bool,
1910 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1911 assert!(
1912 in_f % 64 == 0,
1913 "MMQ NVFP4 W4A8 requires in_f % 64 == 0, got {in_f}"
1914 );
1915 let act_bytes = unsafe { memra_mmq_nvfp4_w4a8_act_bytes(in_f as i32, m as i32) };
1916 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
1917 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
1918 {
1919 let stream = self.gpu.stream();
1920 let (w_p, _gw) = bytes.device_ptr(&stream);
1921 let (x_p, _gx) = x.device_ptr(&stream);
1922 let (y_p, _gy) = y.device_ptr_mut(&stream);
1923 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
1924 static F8F4: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1927 let f8f4 = *F8F4.get_or_init(|| std::env::var("MEMRA_MMQ_F8F4").as_deref() == Ok("1"));
1928 let rc = unsafe {
1929 if f8f4 {
1930 memra_mmq_nvfp4_f8f4(
1931 w_p as *const core::ffi::c_void,
1932 x_p as *const f32,
1933 y_p as *mut f32,
1934 in_f as i32,
1935 out_f as i32,
1936 m as i32,
1937 s_p as *mut core::ffi::c_void,
1938 stream.cu_stream() as *mut core::ffi::c_void,
1939 scale,
1940 rp as i32,
1941 )
1942 } else {
1943 memra_mmq_nvfp4_w4a8(
1944 w_p as *const core::ffi::c_void,
1945 x_p as *const f32,
1946 y_p as *mut f32,
1947 in_f as i32,
1948 out_f as i32,
1949 m as i32,
1950 s_p as *mut core::ffi::c_void,
1951 stream.cu_stream() as *mut core::ffi::c_void,
1952 scale,
1953 rp as i32,
1954 )
1955 }
1956 };
1957 if rc != 0 {
1958 return Err(format!("memra_mmq_nvfp4_w4a8(f8f4={f8f4}) rc={rc}").into());
1959 }
1960 }
1961 Ok(y)
1962 }
1963
1964 pub fn qmatvec_mmq_fp8_blk(
1968 &self,
1969 w_e4m3: &CudaSlice<u8>,
1970 blk_scales: &CudaSlice<f32>,
1971 x: &CudaSlice<f32>,
1972 m: usize,
1973 in_f: usize,
1974 out_f: usize,
1975 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1976 self.qmatvec_mmq_fp8_blk_scaled(w_e4m3, blk_scales, x, m, in_f, out_f, 1.0)
1977 }
1978
1979 pub fn qmatvec_mmq_fp8_blk_scaled(
1980 &self,
1981 w_e4m3: &CudaSlice<u8>,
1982 blk_scales: &CudaSlice<f32>,
1983 x: &CudaSlice<f32>,
1984 m: usize,
1985 in_f: usize,
1986 out_f: usize,
1987 scale: f32,
1988 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1989 assert!(
1990 in_f % 16 == 0,
1991 "per-block FP8 MMQ requires in_f % 16 == 0, got {in_f}"
1992 );
1993 let want_scales = ((out_f + 127) / 128) * ((in_f + 127) / 128);
1994 assert!(
1995 blk_scales.len() >= want_scales,
1996 "blk_scales too small: {} < {want_scales}",
1997 blk_scales.len()
1998 );
1999 assert!(
2000 w_e4m3.len() >= out_f * in_f,
2001 "e4m3 plane too small: {} < {}",
2002 w_e4m3.len(),
2003 out_f * in_f
2004 );
2005 let act_bytes = unsafe { memra_mmq_fp8_blk_act_bytes(in_f as i32, m as i32) };
2006 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
2007 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
2008 {
2009 let stream = self.gpu.stream();
2010 let (w_p, _gw) = w_e4m3.device_ptr(&stream);
2011 let (sc_p, _gsc) = blk_scales.device_ptr(&stream);
2012 let (x_p, _gx) = x.device_ptr(&stream);
2013 let (y_p, _gy) = y.device_ptr_mut(&stream);
2014 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
2015 let rc = unsafe {
2016 memra_mmq_fp8_blk(
2017 w_p as *const core::ffi::c_void,
2018 sc_p as *const f32,
2019 x_p as *const f32,
2020 y_p as *mut f32,
2021 in_f as i32,
2022 out_f as i32,
2023 m as i32,
2024 s_p as *mut core::ffi::c_void,
2025 stream.cu_stream() as *mut core::ffi::c_void,
2026 scale,
2027 )
2028 };
2029 if rc != 0 {
2030 return Err(format!("memra_mmq_fp8_blk rc={rc}").into());
2031 }
2032 }
2033 Ok(y)
2034 }
2035
2036 pub fn qmatvec_mmq_fp8_blk_view(
2041 &self,
2042 w_e4m3: &CudaView<'_, u8>,
2043 blk_scales: &CudaView<'_, f32>,
2044 x: &CudaView<'_, f32>,
2045 m: usize,
2046 in_f: usize,
2047 out_f: usize,
2048 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2049 assert!(
2050 in_f % 16 == 0,
2051 "per-block FP8 MMQ requires in_f % 16 == 0, got {in_f}"
2052 );
2053 let want_scales = out_f.div_ceil(128) * in_f.div_ceil(128);
2054 assert!(
2055 blk_scales.len() >= want_scales,
2056 "blk_scales view too small: {} < {want_scales}",
2057 blk_scales.len()
2058 );
2059 assert!(
2060 w_e4m3.len() >= out_f * in_f,
2061 "e4m3 view too small: {} < {}",
2062 w_e4m3.len(),
2063 out_f * in_f
2064 );
2065 assert!(
2066 x.len() >= m * in_f,
2067 "activation view too small: {} < {}",
2068 x.len(),
2069 m * in_f
2070 );
2071
2072 let act_bytes = unsafe { memra_mmq_fp8_blk_act_bytes(in_f as i32, m as i32) };
2073 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
2074 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
2075 {
2076 let stream = self.gpu.stream();
2077 let (w_p, _gw) = w_e4m3.device_ptr(&stream);
2078 let (sc_p, _gsc) = blk_scales.device_ptr(&stream);
2079 let (x_p, _gx) = x.device_ptr(&stream);
2080 let (y_p, _gy) = y.device_ptr_mut(&stream);
2081 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
2082 let rc = unsafe {
2083 memra_mmq_fp8_blk(
2084 w_p as *const core::ffi::c_void,
2085 sc_p as *const f32,
2086 x_p as *const f32,
2087 y_p as *mut f32,
2088 in_f as i32,
2089 out_f as i32,
2090 m as i32,
2091 s_p as *mut core::ffi::c_void,
2092 stream.cu_stream() as *mut core::ffi::c_void,
2093 1.0,
2094 )
2095 };
2096 if rc != 0 {
2097 return Err(format!("memra_mmq_fp8_blk(view) rc={rc}").into());
2098 }
2099 }
2100 Ok(y)
2101 }
2102
2103 pub fn fp8_blk_nan_count(
2107 &self,
2108 w_e4m3: &CudaSlice<u8>,
2109 ) -> Result<u32, Box<dyn std::error::Error>> {
2110 let mut cnt = self.htod_u32_v(&[0u32])?;
2111 let n = w_e4m3.len();
2112 {
2113 let stream = self.gpu.stream();
2114 let (w_p, _gw) = w_e4m3.device_ptr(&stream);
2115 let (c_p, _gc) = cnt.device_ptr_mut(&stream);
2116 let rc = unsafe {
2117 memra_fp8_blk_count_nan(
2118 w_p as *const core::ffi::c_void,
2119 n,
2120 c_p as *mut u32,
2121 stream.cu_stream() as *mut core::ffi::c_void,
2122 )
2123 };
2124 if rc != 0 {
2125 return Err(format!("memra_fp8_blk_count_nan rc={rc}").into());
2126 }
2127 }
2128 Ok(self.dtoh_u32(&cnt)?[0])
2129 }
2130
2131 pub fn mmq_iq_quantize_act(
2134 &self,
2135 x: &CudaSlice<f32>,
2136 in_f: usize,
2137 n_tokens: usize,
2138 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2139 let act_bytes = unsafe { memra_mmq_iq_experts_act_bytes(in_f as i32, n_tokens as i32) };
2140 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
2141 {
2142 let stream = self.gpu.stream();
2143 let (x_p, _gx) = x.device_ptr(&stream);
2144 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
2145 let rc = unsafe {
2146 memra_mmq_iq_quantize_act(
2147 x_p as *const f32,
2148 s_p as *mut core::ffi::c_void,
2149 in_f as i32,
2150 n_tokens as i32,
2151 stream.cu_stream() as *mut core::ffi::c_void,
2152 )
2153 };
2154 if rc != 0 {
2155 return Err(format!("memra_mmq_iq_quantize_act rc={rc}").into());
2156 }
2157 }
2158 Ok(scratch)
2159 }
2160
2161 pub fn mmq_iq_fused_act_quant(
2167 &self,
2168 gate: &CudaSlice<f32>,
2169 up: &CudaSlice<f32>,
2170 in_f: usize,
2171 n_tokens: usize,
2172 act_kind: i32,
2173 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2174 let act_bytes = unsafe { memra_mmq_iq_experts_act_bytes(in_f as i32, n_tokens as i32) };
2175 let mut scratch = self.alloc_uninit::<u8>(act_bytes)?;
2176 {
2177 let stream = self.gpu.stream();
2178 let (g_p, _gg) = gate.device_ptr(&stream);
2179 let (u_p, _gu) = up.device_ptr(&stream);
2180 let (s_p, _gs) = scratch.device_ptr_mut(&stream);
2181 let rc = unsafe {
2182 memra_mmq_iq_fused_act_quant(
2183 g_p as *const f32,
2184 u_p as *const f32,
2185 s_p as *mut core::ffi::c_void,
2186 in_f as i32,
2187 n_tokens as i32,
2188 act_kind,
2189 stream.cu_stream() as *mut core::ffi::c_void,
2190 )
2191 };
2192 if rc != 0 {
2193 return Err(format!("memra_mmq_iq_fused_act_quant rc={rc}").into());
2194 }
2195 }
2196 Ok(scratch)
2197 }
2198
2199 #[allow(clippy::too_many_arguments)]
2203 pub fn mmq_iq_experts(
2204 &self,
2205 table: &CudaSlice<u64>,
2206 proj: i32,
2207 n_expert: usize,
2208 ex_ids: &CudaSlice<i32>,
2209 ex_off: &CudaSlice<i32>,
2210 ex_pairs: &CudaSlice<i32>,
2211 pair_tok: &CudaSlice<i32>,
2212 act_scratch: &CudaSlice<u8>,
2213 in_f: usize,
2214 out_f: usize,
2215 n_active: usize,
2216 n_pairs: usize,
2217 n_tokens: usize,
2218 qtype: i32,
2219 row_bytes: usize,
2220 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2221 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2222 {
2223 let stream = self.gpu.stream();
2224 let (tab_p, _g0) = table.device_ptr(&stream);
2225 let (ei_p, _g1) = ex_ids.device_ptr(&stream);
2226 let (eo_p, _g2) = ex_off.device_ptr(&stream);
2227 let (ep_p, _g3) = ex_pairs.device_ptr(&stream);
2228 let (pt_p, _g4) = pair_tok.device_ptr(&stream);
2229 let (as_p, _g5) = act_scratch.device_ptr(&stream);
2230 let (y_p, _g6) = y.device_ptr_mut(&stream);
2231 let rc = unsafe {
2232 memra_mmq_iq_experts(
2233 tab_p as *const u64,
2234 proj,
2235 n_expert as i32,
2236 ei_p as *const i32,
2237 eo_p as *const i32,
2238 ep_p as *const i32,
2239 pt_p as *const i32,
2240 as_p as *const core::ffi::c_void,
2241 y_p as *mut f32,
2242 in_f as i32,
2243 out_f as i32,
2244 n_active as i32,
2245 n_tokens as i32,
2246 qtype,
2247 row_bytes as i64,
2248 stream.cu_stream() as *mut core::ffi::c_void,
2249 )
2250 };
2251 if rc != 0 {
2252 return Err(format!("memra_mmq_iq_experts rc={rc}").into());
2253 }
2254 }
2255 Ok(y)
2256 }
2257
2258 pub fn moe_f16g_act(
2263 &self,
2264 x: &CudaSlice<f32>,
2265 pair_tok: Option<&CudaSlice<i32>>,
2266 in_f: usize,
2267 n_pairs: usize,
2268 ) -> Result<(CudaSlice<u8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2269 let mut act = self.alloc_uninit::<u8>(n_pairs * in_f * 2)?;
2270 let mut scales = self.alloc_uninit::<f32>(n_pairs)?;
2271 {
2272 let stream = self.gpu.stream();
2273 let (x_p, _gx) = x.device_ptr(&stream);
2274 let pt_p = match pair_tok {
2275 Some(pt) => {
2276 let (p, _g) = pt.device_ptr(&stream);
2277 p as *const i32
2278 }
2279 None => std::ptr::null(),
2280 };
2281 let (a_p, _ga) = act.device_ptr_mut(&stream);
2282 let (s_p, _gs) = scales.device_ptr_mut(&stream);
2283 let rc = unsafe {
2284 memra_moe_f16g_gather_act(
2285 x_p as *const f32,
2286 pt_p,
2287 a_p as *mut core::ffi::c_void,
2288 s_p as *mut f32,
2289 in_f as i32,
2290 n_pairs as i32,
2291 stream.cu_stream() as *mut core::ffi::c_void,
2292 )
2293 };
2294 if rc != 0 {
2295 return Err(format!("memra_moe_f16g_gather_act rc={rc}").into());
2296 }
2297 }
2298 Ok((act, scales))
2299 }
2300
2301 #[allow(clippy::too_many_arguments)]
2311 pub fn bind_runtime_device(&self, ordinal: i32) -> Result<(), Box<dyn std::error::Error>> {
2316 let rc = unsafe { memra_bind_device(ordinal) };
2317 if rc != 0 {
2318 return Err(format!("cudaSetDevice({ordinal}) rc={rc}").into());
2319 }
2320 Ok(())
2321 }
2322
2323 pub fn moe_f16_grouped(
2324 &self,
2325 table: &CudaSlice<u64>,
2326 proj: i32,
2327 n_expert: usize,
2328 ex_ids: &CudaSlice<i32>,
2329 ex_off_host: &[i32],
2330 ex_off_dev: &CudaSlice<i32>,
2331 act_f16: &CudaSlice<u8>,
2332 act_scale: &CudaSlice<f32>,
2333 in_f: usize,
2334 out_f: usize,
2335 n_active: usize,
2336 n_pairs: usize,
2337 qtype: i32,
2338 row_bytes: usize,
2339 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2340 let sk = crate::moe_f16g_mode() >= 2 && in_f % 32 == 0;
2341 let (shape_sel, cross) = crate::moe_f16g_sk_params();
2349 if sk
2350 && shape_sel >= 0
2351 && crate::moe_f16g_direct_on(qtype)
2352 && (qtype == crate::QT_Q4_K
2353 || qtype == crate::QT_Q6_K
2354 || qtype == crate::QT_IQ4_XS
2355 || qtype == crate::QT_IQ3_S
2356 || qtype == crate::QT_NVFP4
2357 || qtype == crate::QT_NVFP4_V2)
2361 && in_f % (if qtype == crate::QT_NVFP4 || qtype == crate::QT_NVFP4_V2 { 64 } else { 256 }) == 0
2364 && n_active <= 512
2365 && n_active > 0
2366 {
2367 let max_m = ex_off_host
2368 .windows(2)
2369 .map(|w| w[1] - w[0])
2370 .max()
2371 .unwrap_or(0);
2372 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2373 {
2374 let stream = self.gpu.stream();
2375 let (tab_p, _g0) = table.device_ptr(&stream);
2376 let (ei_p, _g1) = ex_ids.device_ptr(&stream);
2377 let (a_p, _g2) = act_f16.device_ptr(&stream);
2378 let (s_p, _g3) = act_scale.device_ptr(&stream);
2379 let (off_p, _g4) = ex_off_dev.device_ptr(&stream);
2380 let (y_p, _g5) = y.device_ptr_mut(&stream);
2381 let rc = unsafe {
2382 memra_moe_kq_gemm_sk(
2383 tab_p as *const u64,
2384 proj,
2385 n_expert as i32,
2386 ei_p as *const i32,
2387 a_p as *const core::ffi::c_void,
2388 y_p as *mut f32,
2389 s_p as *const f32,
2390 off_p as *const i32,
2391 ex_off_host.as_ptr(),
2392 n_active as i32,
2393 max_m,
2394 in_f as i32,
2395 out_f as i32,
2396 qtype,
2397 cross,
2398 crate::moe_f16g_tail_on() as i32,
2399 row_bytes as i64,
2400 stream.cu_stream() as *mut core::ffi::c_void,
2401 )
2402 };
2403 if rc != 0 {
2404 return Err(format!("memra_moe_kq_gemm_sk rc={rc}").into());
2405 }
2406 }
2407 return Ok(y);
2408 }
2409 if !sk {
2413 static WARM: std::sync::Once = std::sync::Once::new();
2414 let mut warm_err = None;
2415 WARM.call_once(|| {
2416 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
2417 let w = self.alloc_uninit::<u8>(2 * 32 * 64 * 2)?;
2418 let a = self.alloc_uninit::<u8>(4 * 64 * 2)?;
2419 let mut yw = self.alloc_uninit::<u8>(4 * 32 * 2)?;
2420 let off = [0i32, 2, 4];
2421 let stream = self.gpu.stream();
2422 let (w_p, _a1) = w.device_ptr(&stream);
2423 let (a_p, _a2) = a.device_ptr(&stream);
2424 let (y_p, _a3) = yw.device_ptr_mut(&stream);
2425 let rc = unsafe {
2426 memra_moe_f16g_gemm(
2427 w_p as *const core::ffi::c_void,
2428 a_p as *const core::ffi::c_void,
2429 y_p as *mut core::ffi::c_void,
2430 off.as_ptr(),
2431 2,
2432 64,
2433 32,
2434 stream.cu_stream() as *mut core::ffi::c_void,
2435 )
2436 };
2437 if rc != 0 {
2438 return Err(format!("f16g warmup rc={rc}").into());
2439 }
2440 self.gpu.stream().synchronize()?;
2441 Ok(())
2442 })();
2443 if let Err(e) = r {
2444 warm_err = Some(e.to_string());
2445 }
2446 });
2447 if let Some(we) = warm_err {
2448 return Err(we.into());
2449 }
2450 }
2451 let w_bytes = n_active * out_f * in_f * 2;
2452 let mut w_f16 = self.alloc_uninit::<u8>(w_bytes)?;
2453 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2454 {
2455 let stream = self.gpu.stream();
2456 let (tab_p, _g0) = table.device_ptr(&stream);
2457 let (ei_p, _g1) = ex_ids.device_ptr(&stream);
2458 let (w_p, _g2) = w_f16.device_ptr_mut(&stream);
2459 let rc = unsafe {
2460 memra_moe_f16g_dequant(
2461 tab_p as *const u64,
2462 proj,
2463 n_expert as i32,
2464 ei_p as *const i32,
2465 w_p as *mut core::ffi::c_void,
2466 in_f as i32,
2467 out_f as i32,
2468 n_active as i32,
2469 qtype,
2470 row_bytes as i64,
2471 stream.cu_stream() as *mut core::ffi::c_void,
2472 )
2473 };
2474 if rc != 0 {
2475 return Err(format!("memra_moe_f16g_dequant rc={rc}").into());
2476 }
2477 let (a_p, _g3) = act_f16.device_ptr(&stream);
2478 let (s_p, _g6) = act_scale.device_ptr(&stream);
2479 let (y_p, _g5) = y.device_ptr_mut(&stream);
2480 if sk {
2481 let max_m = ex_off_host
2482 .windows(2)
2483 .map(|w| w[1] - w[0])
2484 .max()
2485 .unwrap_or(0);
2486 let (off_p, _g7) = ex_off_dev.device_ptr(&stream);
2487 let (shape_sel, cross) = crate::moe_f16g_sk_params();
2488 let rc = unsafe {
2489 memra_moe_f16g_gemm_sk(
2490 w_p as *const core::ffi::c_void,
2491 a_p as *const core::ffi::c_void,
2492 y_p as *mut f32,
2493 s_p as *const f32,
2494 off_p as *const i32,
2495 ex_off_host.as_ptr(),
2496 n_active as i32,
2497 max_m,
2498 in_f as i32,
2499 out_f as i32,
2500 shape_sel,
2501 cross,
2502 crate::moe_f16g_tail_on() as i32,
2503 stream.cu_stream() as *mut core::ffi::c_void,
2504 )
2505 };
2506 if rc != 0 {
2507 return Err(format!("memra_moe_f16g_gemm_sk rc={rc}").into());
2508 }
2509 } else {
2510 let mut y16 = self.alloc_uninit::<u8>(n_pairs * out_f * 2)?;
2511 let (y16_p, _g4) = y16.device_ptr_mut(&stream);
2512 let rc = unsafe {
2513 memra_moe_f16g_gemm(
2514 w_p as *const core::ffi::c_void,
2515 a_p as *const core::ffi::c_void,
2516 y16_p as *mut core::ffi::c_void,
2517 ex_off_host.as_ptr(),
2518 n_active as i32,
2519 in_f as i32,
2520 out_f as i32,
2521 stream.cu_stream() as *mut core::ffi::c_void,
2522 )
2523 };
2524 if rc != 0 {
2525 return Err(format!("memra_moe_f16g_gemm rc={rc}").into());
2526 }
2527 let rc = unsafe {
2528 memra_moe_f16g_h2f_scaled(
2529 y16_p as *const core::ffi::c_void,
2530 y_p as *mut f32,
2531 s_p as *const f32,
2532 out_f as i32,
2533 n_pairs as i32,
2534 stream.cu_stream() as *mut core::ffi::c_void,
2535 )
2536 };
2537 if rc != 0 {
2538 return Err(format!("memra_moe_f16g_h2f_scaled rc={rc}").into());
2539 }
2540 }
2541 }
2542 if !sk {
2547 self.gpu.stream().synchronize()?;
2548 }
2549 if std::env::var("MEMRA_F16G_DEBUG").is_ok() {
2550 let wn = n_active * out_f * in_f;
2552 let an = n_pairs * in_f;
2553 let mut wf = self.alloc_uninit::<f32>(wn)?;
2554 let mut af = self.alloc_uninit::<f32>(an)?;
2555 {
2556 let stream = self.gpu.stream();
2557 let (w_p, _a) = w_f16.device_ptr(&stream);
2558 let (a_p, _b) = act_f16.device_ptr(&stream);
2559 let (wf_p, _c) = wf.device_ptr_mut(&stream);
2560 let (af_p, _d) = af.device_ptr_mut(&stream);
2561 unsafe {
2562 memra_moe_f16g_h2f(
2563 w_p as *const core::ffi::c_void,
2564 wf_p as *mut f32,
2565 wn,
2566 stream.cu_stream() as *mut core::ffi::c_void,
2567 );
2568 memra_moe_f16g_h2f(
2569 a_p as *const core::ffi::c_void,
2570 af_p as *mut f32,
2571 an,
2572 stream.cu_stream() as *mut core::ffi::c_void,
2573 );
2574 }
2575 }
2576 let (wh, ah, yh) = (self.dtoh(&wf)?, self.dtoh(&af)?, self.dtoh(&y)?);
2577 let scan = |v: &[f32]| -> (usize, f32) {
2578 let bad = v.iter().filter(|x| !x.is_finite()).count();
2579 let mx = v
2580 .iter()
2581 .filter(|x| x.is_finite())
2582 .fold(0.0f32, |m, x| m.max(x.abs()));
2583 (bad, mx)
2584 };
2585 let (wb, wm) = scan(&wh);
2586 let (ab, am) = scan(&ah);
2587 let (yb, ym) = scan(&yh);
2588 eprintln!(
2589 "[f16g-debug] proj={proj} w: bad={wb} max={wm:.3e} | act: bad={ab} \
2590 max={am:.3e} | y: bad={yb} max={ym:.3e} (na={n_active} np={n_pairs} \
2591 in={in_f} out={out_f})"
2592 );
2593 }
2594 Ok(y)
2595 }
2596
2597 #[allow(clippy::too_many_arguments)]
2604 pub fn moe_f16g_gemm_sk_raw(
2605 &self,
2606 w_f16: &CudaSlice<u8>,
2607 act_f16: &CudaSlice<u8>,
2608 row_scale: &CudaSlice<f32>,
2609 ex_off_host: &[i32],
2610 ex_off_dev: &CudaSlice<i32>,
2611 in_f: usize,
2612 out_f: usize,
2613 n_pairs: usize,
2614 shape_sel: i32,
2615 cross: i32,
2616 tail: i32,
2617 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2618 let n_active = ex_off_host.len() - 1;
2619 let max_m = ex_off_host
2620 .windows(2)
2621 .map(|w| w[1] - w[0])
2622 .max()
2623 .unwrap_or(0);
2624 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2625 {
2626 let stream = self.gpu.stream();
2627 let (w_p, _g0) = w_f16.device_ptr(&stream);
2628 let (a_p, _g1) = act_f16.device_ptr(&stream);
2629 let (s_p, _g2) = row_scale.device_ptr(&stream);
2630 let (off_p, _g3) = ex_off_dev.device_ptr(&stream);
2631 let (y_p, _g4) = y.device_ptr_mut(&stream);
2632 let rc = unsafe {
2633 memra_moe_f16g_gemm_sk(
2634 w_p as *const core::ffi::c_void,
2635 a_p as *const core::ffi::c_void,
2636 y_p as *mut f32,
2637 s_p as *const f32,
2638 off_p as *const i32,
2639 ex_off_host.as_ptr(),
2640 n_active as i32,
2641 max_m,
2642 in_f as i32,
2643 out_f as i32,
2644 shape_sel,
2645 cross,
2646 tail,
2647 stream.cu_stream() as *mut core::ffi::c_void,
2648 )
2649 };
2650 if rc != 0 {
2651 return Err(format!("memra_moe_f16g_gemm_sk rc={rc}").into());
2652 }
2653 }
2654 Ok(y)
2655 }
2656
2657 #[allow(clippy::too_many_arguments)]
2662 pub fn moe_kq_gemm_sk_raw(
2663 &self,
2664 table: &CudaSlice<u64>,
2665 proj: i32,
2666 n_expert: usize,
2667 ex_ids: &CudaSlice<i32>,
2668 act_f16: &CudaSlice<u8>,
2669 row_scale: &CudaSlice<f32>,
2670 ex_off_host: &[i32],
2671 ex_off_dev: &CudaSlice<i32>,
2672 in_f: usize,
2673 out_f: usize,
2674 n_pairs: usize,
2675 qtype: i32,
2676 row_bytes: usize,
2677 cross: i32,
2678 tail: i32,
2679 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2680 let n_active = ex_off_host.len() - 1;
2681 let max_m = ex_off_host
2682 .windows(2)
2683 .map(|w| w[1] - w[0])
2684 .max()
2685 .unwrap_or(0);
2686 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2687 {
2688 let stream = self.gpu.stream();
2689 let (tab_p, _g0) = table.device_ptr(&stream);
2690 let (ei_p, _g1) = ex_ids.device_ptr(&stream);
2691 let (a_p, _g2) = act_f16.device_ptr(&stream);
2692 let (s_p, _g3) = row_scale.device_ptr(&stream);
2693 let (off_p, _g4) = ex_off_dev.device_ptr(&stream);
2694 let (y_p, _g5) = y.device_ptr_mut(&stream);
2695 let rc = unsafe {
2696 memra_moe_kq_gemm_sk(
2697 tab_p as *const u64,
2698 proj,
2699 n_expert as i32,
2700 ei_p as *const i32,
2701 a_p as *const core::ffi::c_void,
2702 y_p as *mut f32,
2703 s_p as *const f32,
2704 off_p as *const i32,
2705 ex_off_host.as_ptr(),
2706 n_active as i32,
2707 max_m,
2708 in_f as i32,
2709 out_f as i32,
2710 qtype,
2711 cross,
2712 tail,
2713 row_bytes as i64,
2714 stream.cu_stream() as *mut core::ffi::c_void,
2715 )
2716 };
2717 if rc != 0 {
2718 return Err(format!("memra_moe_kq_gemm_sk rc={rc}").into());
2719 }
2720 }
2721 Ok(y)
2722 }
2723
2724 pub fn moe_f16g_dequant_raw(
2728 &self,
2729 table: &CudaSlice<u64>,
2730 proj: i32,
2731 n_expert: usize,
2732 ex_ids: &CudaSlice<i32>,
2733 in_f: usize,
2734 out_f: usize,
2735 n_active: usize,
2736 qtype: i32,
2737 row_bytes: usize,
2738 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2739 let mut w_f16 = self.alloc_uninit::<u8>(n_active * out_f * in_f * 2)?;
2740 {
2741 let stream = self.gpu.stream();
2742 let (tab_p, _g0) = table.device_ptr(&stream);
2743 let (ei_p, _g1) = ex_ids.device_ptr(&stream);
2744 let (w_p, _g2) = w_f16.device_ptr_mut(&stream);
2745 let rc = unsafe {
2746 memra_moe_f16g_dequant(
2747 tab_p as *const u64,
2748 proj,
2749 n_expert as i32,
2750 ei_p as *const i32,
2751 w_p as *mut core::ffi::c_void,
2752 in_f as i32,
2753 out_f as i32,
2754 n_active as i32,
2755 qtype,
2756 row_bytes as i64,
2757 stream.cu_stream() as *mut core::ffi::c_void,
2758 )
2759 };
2760 if rc != 0 {
2761 return Err(format!("memra_moe_f16g_dequant rc={rc}").into());
2762 }
2763 }
2764 Ok(w_f16)
2765 }
2766}