Skip to main content

memra_engine/
mmq_ffi.rs

1//! FFI to the MMQ prefill GEMMs (cu/mmq_fp4.cu + cu/mmq_q45k.cu) — vendored floor kernels.
2//!
3//! NVFP4: the 5150-pp512 kernel from llama.cpp, ggml-decoupled into a static lib with a C-ABI host
4//! launcher. The launcher quantizes the f32 activation to block_fp4_mmq internally (llama's 2-level
5//! FP8-e8m0/UE4M3 scale = the accurate W4A8-via-FP8 path that fixes memra's W4A4 maxdiff 1.46), then
6//! launches the native mxf4nvf4 block-scale tensor-core mma.
7//!
8//! Q4_K/Q5_K: llama's k-quant int8-MMA MMQ (dequant to int8 at tile-load, q8_1 DS4 activation with
9//! the (d, sum) pair that feeds the k-quant min-offset term, shared m16n8k32 s8 mma inner loop).
10//! Replaces the hand-rolled qmatvec_gemm k-quant GEMMs that dominate prefill (32% + 28% busy).
11//!
12//! All dispatched behind MEMRA_MMQ=1. Always built (no external deps) — unlike cutlass_ffi which is
13//! MEMRA_CUTLASS-gated.
14
15use crate::Engine;
16use cudarc::driver::{CudaSlice, CudaView, DevicePtr, DevicePtrMut};
17
18/// Quantize-once seam state (see `Engine::mmq_act_begin`): window epoch + one cached
19/// (epoch, act_ptr, m, in_f, D4 scratch) slot. Slot drops (freeing the scratch) on each new window.
20static 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);
24/// Stream-k fixup scratch (lazy; sized once per process — one slot per SM).
25static MMQ_FIXUP_SLOT: std::sync::Mutex<Option<cudarc::driver::CudaSlice<u8>>> =
26    std::sync::Mutex::new(None);
27
28/// Model-agnostic expert-major CSR used by grouped projection backends.
29///
30/// `ex_pairs` is a permutation of pair-major output rows. `pair_tok[pair]` selects the activation
31/// row consumed by that pair. Model adapters supply route choices; this type builds, validates,
32/// owns, and uploads the expert-major schedule.
33pub 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    /// Build an expert-major schedule from token-major top-k route choices.
44    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    /// Build an expert-major schedule with an explicit activation row for every output pair.
70    ///
71    /// This is used by chained grouped projections: gate/up pairs select token rows, while the
72    /// down projection selects the corresponding pair-major activation row.
73    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    /// Allocate stable device storage for schedules up to the supplied logical maxima.
309    ///
310    /// `refresh` fills prefixes of these buffers. The grouped kernel receives only the active
311    /// lengths, so route changes do not change any device pointer or allocate in the hot path.
312    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
433/// Caller-owned persistent buffers for grouped block-E4M3 projections.
434///
435/// Allocation and routing-plan upload happen outside the hot projection path. `quantize` and
436/// `project` overwrite their complete buffers and therefore introduce no per-call allocations.
437pub 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    /// Quantize an active prefix while retaining the workspace's stable capacity pointers.
483    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    /// Bytes needed for the block_fp4_mmq activation scratch for (in_f, n_tokens).
742    pub fn memra_mmq_nvfp4_act_bytes(in_f: i32, n_tokens: i32) -> usize;
743    /// Run the NVFP4 W4A4 MMQ prefill GEMM. y[n_tokens, out_f] = act[n_tokens, in_f] @ W[out_f, in_f]^T.
744    ///   W_nvfp4_blocks : raw memra NVFP4 weight rows (block_nvfp4 36B blocks, in_f/64 per row).
745    ///   act_f32        : f32 activation [n_tokens, in_f] (contiguous).
746    ///   y              : f32 output [n_tokens, out_f].
747    ///   act_scratch    : pre-alloc'd quant buffer >= memra_mmq_nvfp4_act_bytes(in_f, n_tokens).
748    /// Returns 0 on success, else (1000 + cudaError).
749    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    /// Same as `memra_mmq_nvfp4`, plus the activation-quantizer selector.
761    ///   per_token_scale = 1: two-level scaling (per-token row amax folded into the GEMM epilogue
762    ///     + per-sub-block UE4M3). This is what `memra_mmq_nvfp4` does.
763    ///   per_token_scale = 0: the v1 sub-block-only quantizer, retained as the numeric oracle so
764    ///     kernel-check can measure what the row scale bought, and as the rollback seam.
765    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    /// Same as `memra_mmq_nvfp4_ex`, plus the residual high-precision channel count.
778    ///   residual_k = 0: off.
779    ///   residual_k > 0: the k largest-magnitude activation channels (ranked across the batch) are
780    ///     zeroed before quantization and their exact f32 contribution is added back as a rank-k
781    ///     correction. Requires per_token_scale = 1. Clamped to MMQ_MAX_RESIDUAL_K (64).
782    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    /// Bytes needed for the block_q8_1_mmq activation scratch for the NVFP4 W4A8 path.
796    pub fn memra_mmq_nvfp4_w4a8_act_bytes(in_f: i32, n_tokens: i32) -> usize;
797    /// Run the NVFP4 W4A8 MMQ prefill GEMM (STAGE 2 accuracy-safe rung). Same fast MMQ tile as
798    /// memra_mmq_nvfp4 (W4A4) but the non-Blackwell int8 pair: weight FP4 LUT-dequantized to int8 at
799    /// tile-load, activation stays q8_1 int8 (D4, the same quant class as the default int8 GEMM).
800    /// `rp`: 0 = GGUF 36B-block weight layout, 1 = A6 split-plane repack (the resident decode
801    /// layout). The rp tile loader is a pure address remap of the GGUF loader (same dequant math,
802    /// same FP op order) — output is bit-identical either way.
803    /// Same contract as memra_mmq_nvfp4 otherwise. Returns 0 or (1000 + cudaError).
804    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    /// Bytes for the block_e4m3_mmq activation scratch (footprint-identical to block_q8_1_mmq).
817    pub fn memra_mmq_nvfp4_f8f4_act_bytes(in_f: i32, n_tokens: i32) -> usize;
818    /// R-B W4A8-FP8 MMQ prefill GEMM (research/prefill-mxf8f6f4-design.md): NVFP4 per-16 scales
819    /// fold into e4m3 weight VALUES at tile load; e4m3 activations; ONE kind::f8f6f4 m16n8k32
820    /// MMA (381-TF class) where the int8 path issues two imma k16. NEW NUMERIC CONFIG — own
821    /// battery. Same contract/rp semantics as memra_mmq_nvfp4_w4a8. Returns 0 / 1000+cudaError /
822    /// 2000+cudaError.
823    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    /// Bytes for the per-block FP8 MMQ activation scratch (delegates to the F8F4 sizing — the
836    /// two arms deliberately share ONE activation format, `block_e4m3_mmq`).
837    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    /// Scale-grid dims for an [out_f x in_f] block-128 FP8 tensor (ceil-div by 128).
866    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    /// PER-BLOCK FP8 MMQ prefill GEMM (cu/mmq_fp8_blk.cu, P1 option (b)): consumes the
869    /// Qwen-official e4m3 weight bytes + the per-[128x128] f32 scale grid DIRECTLY. The weight
870    /// side is never re-quantized (the checkpoint bytes are the MMA A operand), so unlike ARM A's
871    /// per-tensor fold there is no precision loss; unlike ARM B' it does not land on the Q8_0
872    /// floor. `blk_scales` is device f32 [ceil(out_f/128) x ceil(in_f/128)], row-major.
873    /// Requires in_f % 16 == 0. Returns 0 / 1 (bad dims) / 1000+cudaError / 2000+cudaError.
874    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    /// Count e4m3 NaN codes (magnitude 0x7F) in a device weight buffer. Those decode to NaN in
887    /// hardware but to 0.0 in the host/ARM B' convention, so a tensor containing any must NOT
888    /// ride `memra_mmq_fp8_blk`. `out_count` is a device u32 (zeroed by the call).
889    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    /// Bytes needed for the block_q8_1_mmq activation scratch (shared by Q4_K and Q5_K).
896    pub fn memra_mmq_q45k_act_bytes(in_f: i32, n_tokens: i32) -> usize;
897    /// Run the Q4_K W4A8 MMQ prefill GEMM. Same contract as memra_mmq_nvfp4 (raw ggml block_q4_K
898    /// weight rows, in_f/256 144B superblocks per row). Returns 0 or (1000 + cudaError).
899    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    /// Run the Q5_K W4A8 MMQ prefill GEMM (176B superblocks). Same contract as memra_mmq_q4_K.
910    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    /// Bytes needed for the block_q8_1_mmq (D4) activation scratch for the Q8_0 MMQ path.
922    pub fn memra_mmq_q8_0_act_bytes(in_f: i32, n_tokens: i32) -> usize;
923    /// Run the Q8_0 int8-MMA MMQ prefill GEMM (MEMRA_PP_Q8MMQ). Conventional xy-tiling only (no fixup
924    /// scratch). Weight = raw ggml block_q8_0 rows (34B blocks, in_f/32 per row); activation is
925    /// quantized internally to q8_1 D4. Requires in_f % 32 == 0. Returns 0 or (1000 + cudaError).
926    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    // ---- Q1 accumulator instrument (cu/mmq_q8_0_f32acc.cu, lane/fp8-v3-gate) ----
938    // The Q8_0 MMQ floor's GEMM with the accumulator as its ONE free variable: arm S32 is the
939    // floor's `mma...s32.s8.s8.s32`, arm F32 is the same m16n8k32 shape and the same A/B/D fragment
940    // ABI with `mma...kind::f8f6f4...f32.e4m3.e4m3.f32` — the op cu/mmq_fp8_blk.cu accumulates in.
941    // Both take a PRE-QUANTIZED block_q8_1_mmq activation buffer, so the measurement is GEMM-only
942    // and cannot differ by a quantizer. Research instrument only: no dispatch seam, and neither arm's
943    // output is a numeric claim (see the TU header).
944    /// Activation-scratch bytes for the accumulator instrument (same padding rule as the floor).
945    pub fn memra_accprobe_act_bytes(in_f: i32, n_tokens: i32) -> usize;
946    /// ARM S32 — the floor's GEMM verbatim, s32 accumulate. Returns 0, 1, or 1000+cudaError.
947    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    /// ARM F32 — byte-identical kernel, f32 accumulate over the e4m3 reading of the same bytes.
957    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    /// Bytes needed for the block_q8_1_mmq (D4) activation scratch for the Q4_0 MMQ path.
968    pub fn memra_mmq_q4_0_act_bytes(in_f: i32, n_tokens: i32) -> usize;
969    /// Run the Q4_0 int8-MMA MMQ prefill GEMM (MEMRA_PP_Q4MMQ). Nibbles dequant to int8 at
970    /// tile-load (the -8 zero-point folds into the quants, D4 epilogue — same accuracy class as
971    /// the Q8_0 MMQ). `rp`: 0 = raw ggml 18B blocks, 1 = MEMRA_Q4RP split-plane repack (qs plane +
972    /// fp16 d plane) — pure address remap, bit-identical output either way. Requires
973    /// in_f % 32 == 0. Returns 0 or (1000 + cudaError).
974    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    /// Quantize-only entry (quantize-once seam): f32 activation -> block_q8_1_mmq scratch.
986    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    /// GEMM-only entry: consumes a pre-quantized scratch (from memra_mmq_q4_0_quant_act).
994    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    /// Stream-k fixup scratch bytes (one [MMQ_X x MMQ_Y] f32 slot per SM).
1005    pub fn memra_mmq_q4_0_fixup_bytes() -> usize;
1006    /// Force the CLC work-stealing arm: 1 = on, 0 = off (static grid), -1 = MEMRA_MMQ_CLC env
1007    /// default. Schedule-only swap of the xy-tiling kernel — bit-identical output by
1008    /// construction (perf-frontier lever #1). Returns 1 when the CLC kernel is compiled in
1009    /// (SM_100+ gencode), 0 on sm_89/90a builds (force is a no-op there; static grid always).
1010    pub fn memra_mmq_q4_0_set_clc(force: i32) -> i32;
1011    /// Stream-k GEMM entry: deterministic form selection, with the SK form itself
1012    /// falling back to tiling when wave efficiency is at least 90%.
1013    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    // ---- IQ3_S / IQ4_XS expert-segmented int8-MMA MMQ (cu/mmq_iq_experts.cu, MEMRA_MOE_MMA) ----
1026    /// Bytes for the token-major block_q8_1_mmq activation scratch (in_f, n_tokens).
1027    pub fn memra_mmq_iq_experts_act_bytes(in_f: i32, n_tokens: i32) -> usize;
1028    /// Quantize token-major f32 activation [n_tokens, in_f] -> block_q8_1_mmq (D4). Returns 0 or 1000+err.
1029    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    /// Fused act-epilogue: silu/gelu(gate)*up + q8_1_mmq (D4) quantize in ONE launch — no f32 act
1037    /// buffer. gate/up pair-major [n_tokens, in_f]; scratch identical to memra_mmq_iq_quantize_act.
1038    /// act_kind: 0=silu*mul, 1=gelu_tanh*mul. Byte-identical to the two-pass path (kernel-check gated).
1039    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    /// Expert-segmented IQ MMA MMQ. Same CSR shape as moe_pairs_matvec_q8_dec: `table` = [3,n_expert]
1049    /// device slab ptrs, CSR ex_ids/ex_off/ex_pairs group pairs by expert, pair_tok gathers the
1050    /// activation row. y = [n_pairs, out_f] pair-major. `act_scratch` pre-quantized over n_tokens.
1051    /// qtype: 5=IQ4_XS, 6=IQ3_S. Returns 0 or 1000+cudaError.
1052    /// Dense-trunk IQ4_XS MMQ (lane/kquant-tile-loaders): the dense analog of the expert
1053    /// kernel for non-expert IQ4_XS 2-D matmuls (the KAT-Coder trunk class). Quantizes the
1054    /// f32 activation to D4 q8_1_mmq internally; `act_scratch` sized by
1055    /// `memra_mmq_iq_experts_act_bytes`. Requires in_f % 256 == 0.
1056    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    // ---- MoE grouped f16 GEMM (cu/moe_f16_grouped.cu, round 46 arc 2) ----
1087    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    // Single-kernel grouped GEMM (MEMRA_MOE_F16G=2, rounds 49+51): on OUR stream, f32 C with
1134    // the act row-scale folded in — no cublas internal-stream race, no sync. Round 51 runs it
1135    // as a persistent problem-visitor over the real tiles with two tile forms (32x64 tail
1136    // / 128x64x64 3-stage): shape_sel < 0 = the round-49 grid-scan kernel (rollback
1137    // arm); else groups with m_e >= cross ride the 128 form. ex_off_host sizes the visitor
1138    // grids host-side (the offsets are already there at the call site — no extra transfer).
1139    // tail != 0 (lane/sk-tail-form): sub-cross groups ride the DEEP tail (32x64x64 3-stage);
1140    // 0 = the round-51 2-stage 32x64x32 (MEMRA_F16G_TAIL=0 rollback). Byte-identical arms.
1141    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    // DIRECT-FROM-QUANT sk visitor grouped GEMM (lane/kquant-tile-loaders + iq-direct-loaders):
1158    // the visitor forms with the B (weight) tiles dequanted in-register from the expert
1159    // superblocks — no f16 dequant workspace pass. Bit-identical to the workspace path by
1160    // construction (kernel-check "f16g-kq-direct"). qtype: QT_Q4_K | QT_Q6_K | QT_IQ4_XS |
1161    // QT_IQ3_S; rc=2 = not admitted here (caller keeps the dequant-workspace path).
1162    // tail: as memra_moe_f16g_gemm_sk.
1163    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
1185/// W4A8-MMQ DEFAULT-FLIP seam (2026-07-05): the vendored MMQ prefill suite is DEFAULT-ON — NVFP4
1186/// takes the W4A8 MMQ tile (same int8 accuracy class as the int8 GEMM it replaces, all exactness
1187/// gates hold, ~1.9x pp512; the rp tile-loader arm coexists with the A6 split-plane repack) and
1188/// Q4_K/Q5_K take the vendored k-quant int8-MMA MMQ (also int8-class; gated with W4A8 in the same
1189/// battery — the predecessor's `MEMRA_MMQ_W4A8=1` arm engaged BOTH, this flip preserves exactly
1190/// that measured config). `MEMRA_MMQ_W4A8=0` = escape hatch back to the int8 GEMM prefill
1191/// everywhere. `MEMRA_MMQ=1` additionally switches GGUF-layout NVFP4 to the W4A4 mxf4nvf4 tile
1192/// (speed/accuracy tradeoff opt-in, unchanged).
1193pub 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
1202/// Residual high-precision activation channels for the W4A4 MMQ prefill path.
1203/// `MEMRA_MMQ_RESIDUAL_K=<k>` keeps the k largest-magnitude activation channels out of the e2m1
1204/// quantized path and adds their exact f32 contribution back as a rank-k correction. k=0 (default)
1205/// is off; the kernel clamps to MMQ_MAX_RESIDUAL_K (64).
1206///
1207/// Read LIVE per call, not OnceLock'd, for the same reason `MEMRA_MMQ` is: the W4A4 exactness gate
1208/// sweeps arms inside ONE process against ONE set of loaded weights, and a cached first read would
1209/// pin every later arm to whatever the first one saw.
1210pub 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
1218/// Q8_0 MMQ prefill seam (lane/ppmmq lever 2, DEFAULT ON since 2026-07-09 — `MEMRA_PP_Q8MMQ=0`
1219/// reverts): routes Q8_0 dense
1220/// projections (m>=16) through the vendored int8-MMA MMQ (cu/mmq_q8_0.cu) instead of the hand-rolled
1221/// `qmatvec_gemm_q8_0` tiling GEMM. Its own numeric config (MMA f32 reduction order != the tiling
1222/// GEMM's) — gated with the full exactness battery. Default OFF until the battery is green.
1223pub fn mmq_q8_enabled() -> bool {
1224    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1225    // Promotion battery (2026-07-09): argmax MATCH on 35B p1/p2/p3 + 9B p2/p3 (p4-16k OOMs
1226    // identically with and without the flag — pre-existing gate capacity limit, not this seam);
1227    // kernel-check ALL GREEN; run-spec K=1..8 PASS on 9B+35B. 35B pp 2456->3069 free-clock.
1228    *ON.get_or_init(|| {
1229        std::env::var("MEMRA_PP_Q8MMQ")
1230            .map(|v| v != "0")
1231            .unwrap_or(true)
1232    })
1233}
1234
1235/// IQ4_XS dense-trunk MMQ prefill seam (lane/kquant-tile-loaders, 2026-08-02): routes
1236/// NON-expert IQ4_XS 2-D projections (m>=16) through the vendored-machinery int8-MMA dense
1237/// MMQ (cu/mmq_iq_experts.cu `mmq_iq4xs_dense_kernel`) instead of the per-column dp4a grid
1238/// — the KAT-Coder prefill wall (0.169x vs llama; zero weight reuse across tokens,
1239/// research/kat-anomaly-20260802 §6). Its own numeric config (MMA reduction order) — gated
1240/// with the full exactness battery. m=1..15 decode/verify keep dp4a (dispatch parity).
1241/// `MEMRA_PP_IQMMQ=0` reverts.
1242pub 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
1251/// Q4_0 MMQ prefill seam (gemma-4-12B lane, 2026-07-22): routes Q4_0 dense projections (m>=16)
1252/// through the vendored int8-MMA MMQ (cu/mmq_q4_0.cu) instead of the hand-rolled
1253/// `qmatvec_gemm_q4_0[_rp]` tiling GEMM (measured 77% of the 12B prime pass). Its own numeric
1254/// config (MMA f32 reduction order != the tiling GEMM's) — gated with the full exactness battery
1255/// before default-flip; `MEMRA_PP_Q4MMQ=0` reverts.
1256pub 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    // Split-plane weights have no W4A4 loader. Otherwise an explicit W4A8 request wins, then the
1267    // default applies only while MEMRA_MMQ is absent (MEMRA_MMQ=1 selects W4A4).
1268    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    /// True if `w` should take a vendored MMQ GEMM under the current env policy (see
1287    /// `mmq_w4a8_enabled`): NVFP4 needs in_f % 64 == 0, Q4_K/Q5_K need in_f % 256 == 0.
1288    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            // A6 split-plane repacked NVFP4: ONLY the W4A8 loader has an rp arm (pure address
1296            // remap, bit-identical output — mmq_nvfp4_w4a8.cu load_tiles_nvfp4_w4a8<is_rp>).
1297            // The W4A4 loader (mmq_fp4.cu load_tiles_nvfp4_nvfp4) reads 36B GGUF blocks only,
1298            // so an rp weight with W4A8 disabled falls through to the rp-ported int8 GEMM.
1299            // The split-plane layout has only a W4A8 loader. Its int8 MMA is native on sm_100a;
1300            // optional F8F4 uses the existing bit-identical plain-E4M3 rollback form there.
1301            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            // GGUF-layout NVFP4: W4A8 stays the accuracy-safe default on both Blackwell families.
1307            // The new sm_100a tcgen05 W4A4 twin stays behind the EXISTING MEMRA_MMQ=1 opt-in until
1308            // real-B200 exactness and serving gates exist; unmeasured hardware behavior never
1309            // defaults on.
1310            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            // Q8_0 dense projections (35B attn/ssm/shexp): opt-in only (MEMRA_PP_Q8MMQ=1), its own
1321            // numeric config vs qmatvec_gemm_q8_0. in_f % 256 == 0: MMQ_ITER_K=256 loads 8-block
1322            // groups, so a non-multiple row would read a garbage weight tail (fp16 d bytes can be
1323            // NaN-pattern, and NaN * 0-padded-activation = NaN — the 26B ffn_down lesson).
1324            GpuTensor::Quant { qtype, .. } if *qtype == crate::QT_Q8_0 => {
1325                mmq_q8_enabled() && w.in_features().is_multiple_of(256)
1326            }
1327            // Q4_0 dense projections (gemma QAT ggufs): MEMRA_PP_Q4MMQ seam. Both weight layouts
1328            // (raw 18B blocks and the MEMRA_Q4RP split-plane repack) have loader arms. Same
1329            // in_f % 256 == 0 tail rule as Q8_0 (26B ffn_down in_f=2112 NaN'd on the %32 gate);
1330            // non-multiples fall back to the hand-rolled qmatvec_gemm_q4_0[_rp].
1331            GpuTensor::Quant { qtype, .. } if *qtype == crate::QT_Q4_0 => {
1332                mmq_q4_enabled() && w.in_features().is_multiple_of(256)
1333            }
1334            // IQ4_XS dense projections (KAT-Coder trunk): m>=16 prefill only — decode and
1335            // spec-verify (m<16) keep the qmatvec_iq4_XS_dp4a per-column program (the
1336            // kat-anomaly dispatch-parity law). Requires the dp4a fast path itself enabled:
1337            // MEMRA_IQ_FAST=0 (the Stage-A oracle rollback) must also kill this arm so the
1338            // rollback stays a full-path seam. in_f % 256: MMQ_ITER_K walks whole superblocks.
1339            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    /// Unified vendored-MMQ dispatch: routes to the NVFP4 or Q4_K/Q5_K launcher by qtype.
1349    /// Caller MUST have checked `mmq_supports(w)`. `x` is the RAW f32 activation.
1350    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        // NVFP4 tile choice: W4A8 (accuracy-safe int8 pair, DEFAULT since the flip) vs W4A4
1369        // (mxf4nvf4 mma, explicit MEMRA_MMQ=1 speed/accuracy tradeoff). An rp weight ALWAYS takes
1370        // W4A8 — only its loader has the split-plane arm (pure address remap, bit-identical).
1371        // Explicit MEMRA_MMQ_W4A8=1 still overrides a simultaneous MEMRA_MMQ=1 (predecessor rule).
1372        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            // STAGE 2: the accuracy-safe int8 W4A8 MMQ tile (weight FP4->int8 dequant + q8_1
1383            // activation) — handles BOTH weight layouts (rp = A6 split-plane vs GGUF blocks).
1384            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                // wgmma arm (sm_90a, task 8): OPT-IN via MEMRA_WGMMA=1 — v0 measured 3845
1397                // vs MMQ 8692 tok/s pp512 (2026-07-26 N=5), so MMQ stays the default until
1398                // the pipelined wgmma wins. Reads the rp4 split-plane mirror + the engine's
1399                // q8_1 activation planes. Same numeric class as MMQ (exact s32 per 32-block,
1400                // one f32 fold per block, ascending K) — kernel-check tolerance-gated.
1401                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    /// Bare IQ4_XS dense MMQ launch (no macro-scale) — also the kernel_check gate entry.
1441    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    /// Bare Q4_K/Q5_K MMQ launch (no macro-scale) — also the kernel_check accuracy-gate entry.
1484    /// Conventional xy-tiling only (the vendored stream-K arm — MEMRA_MMQ_STREAMK — was removed
1485    /// 2026-07-08: 1.11x per-GEMM but its k-split f32 reorder flipped the model argmax gate;
1486    /// rig5090.jsonl 2026-07-03 has the record).
1487    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    /// Bare Q8_0 int8-MMA MMQ launch (no macro-scale) — the kernel_check accuracy-gate entry and
1534    /// the `qmatvec_mmq` dispatch body. Conventional xy-tiling only (no stream-K / fixup scratch).
1535    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    /// Accumulator-instrument bytes for a pre-quantized block_q8_1_mmq activation buffer
1576    /// (cu/mmq_q8_0_f32acc.cu). The caller synthesizes that buffer itself — see `accprobe_gemm`.
1577    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    /// Run one arm of the Q1 accumulator instrument. `f32acc=false` is the Q8_0 MMQ floor's GEMM
1582    /// verbatim (s32 accumulate); `f32acc=true` is the byte-identical kernel with the f8f6f4 f32
1583    /// accumulate. `act_q` is a PRE-QUANTIZED block_q8_1_mmq buffer of at least
1584    /// `accprobe_act_bytes(in_f, m)` bytes — keeping the quantizer out of the timed region is the
1585    /// point, so this wrapper does not build it. Research instrument: the output is not a numeric
1586    /// claim.
1587    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    /// Open a quantize-once sharing window for the NEXT activation (quantize-once seam): sibling
1637    /// Q4_0 MMQ matmuls on the SAME input (q/k/v; gate/up) quantize its D4 scratch once. Safe by
1638    /// construction: a hit requires the same window epoch AND the same (ptr, m, in_f) — the caller
1639    /// opens a window while it holds the shared input alive, so its address can neither change nor
1640    /// be recycled inside the window. Paths that never call this never hit the cache.
1641    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    /// Bare Q4_0 int8-MMA MMQ launch (no macro-scale) — the kernel_check accuracy-gate entry and
1648    /// the `qmatvec_mmq` dispatch body. `rp` selects the weight layout (MEMRA_Q4RP split-plane vs
1649    /// raw ggml 18B blocks) — pure address remap, bit-identical output.
1650    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        // quantize-once: reuse the window's scratch when the SAME activation comes back.
1669        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            // Stream-k arm (DEFAULT since 2026-07-23; MEMRA_MMQ_SK=0 reverts to xy-tiling):
1700            // small-batch tail-wave fix — the sk entry itself falls back to (bit-identical)
1701            // tiling at >=90% wave efficiency. Band-class fold order below that. Gate: 12B
1702            // pp512 +3.3% (1.005x vs llama), pp1736 +1.0%; 31B +0.5%; D512 sentinel MATCH.
1703            //
1704            // SPEC-SERVING FLIP (2026-07-27, the f16pv/wkv acceptance-law pattern): with
1705            // MEMRA_DRAFT set, big dense models force tiling while MoE/small models defer
1706            // to the fail-closed TILE form. The former shape-timing autotune was removed 2026-08-14:
1707            // its per-process timing coin selected different fold orders on independent
1708            // boots. On the measured 82-SM 5090, TILE is both faster and higher-acceptance
1709            // for the 26B depth cell. Every other hardware class requires its own gate
1710            // before selecting SK without an explicit form override.
1711            // MEMRA_MMQ_SK controls entry and MEMRA_MMQ_SK_FORM pins the numerical form.
1712            // HOPPER DEFAULT OFF (2026-07-31, #23): on sm_90a the SK arm computes WRONG
1713            // values for the 26B a4b's non-rp Q4_0 shapes once the prefill width crosses
1714            // 256 (prefill argmax garbage, maxdiff ~10; MEMRA_MMQ_SK=0 -> MATCH,
1715            // one-variable kill x confirmed on-box). The SK split/fixup is SM-count
1716            // dependent (132 vs 170) — until the kernel is
1717            // fixed for that class, Hopper fails CLOSED to the bit-identical xy-tiling
1718            // (cost on the healthy models: g12 -1.4%, g31 -0.6% prefill, N=3 on-box).
1719            // sm_120a keeps the SK entry on (rig-divergence law). MEMRA_MMQ_SK=1 forces
1720            // entry; MEMRA_MMQ_SK_FORM=sk forces the actual SK numerical form.
1721            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    /// Run the vendored NVFP4 MMQ prefill GEMM from raw weight bytes + f32 activation.
1777    /// y[m, out_f] = x[m, in_f] @ W^T. The per-tensor NVFP4 macro-scale is FOLDED into the MMQ
1778    /// write-back epilogue (was a separate scale_inplace launch + full y round-trip per matmul).
1779    /// Same elementwise multiply -> bit-identical to the two-launch form.
1780    /// `x` is the RAW f32 activation (the launcher quantizes it to block_fp4_mmq internally).
1781    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    /// Bare MMQ launch (no macro-scale) — for the kernel_check accuracy gate.
1794    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    /// Bare MMQ launch on the PRE-PORT activation quantizer (per-sub-block UE4M3 scale only, no
1806    /// per-token row amax). The numeric oracle for the two-level quantizer: kernel-check runs both
1807    /// and reports the accuracy delta, so the port's value is measured rather than asserted.
1808    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    /// Bare MMQ launch with an explicit residual-channel count — for the kernel-check k sweep.
1820    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)] // allow: the parameter list mirrors the kernel/FFI/call contract; bundling into a struct is a refactor, not a lint fix
1845    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    /// STAGE 2 W4A8 MMQ NVFP4: same tile as the W4A4 path, but weight FP4 is LUT-dequantized to
1892    /// int8 at tile-load and the activation stays q8_1 int8 — the accuracy-safe rung. Macro-scale
1893    /// folded into the write-back epilogue (bit-identical to a post-matmul scale_inplace).
1894    /// `rp` selects the weight layout (A6 split-plane vs GGUF blocks) — bit-identical output.
1895    #[allow(clippy::too_many_arguments)] // allow: the parameter list mirrors the kernel/FFI/call contract; bundling into a struct is a refactor, not a lint fix
1896    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    /// Bare W4A8 MMQ launch (no macro-scale, GGUF layout) — for the kernel_check accuracy gate.
1910    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    /// Bare W4A8 MMQ launch on an A6 split-plane repacked weight — the rp-loader bit-identity gate
1922    /// compares this against `qmatvec_mmq_nvfp4_w4a8_raw` on the same weight.
1923    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)] // allow: the parameter list mirrors the kernel/FFI/call contract; bundling into a struct is a refactor, not a lint fix
1935    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        // MEMRA_MMQ_F8F4=1: the R-B W4A8-FP8 tile (own numeric config; battery-gated seam).
1946        // SM100 compiles this route with the existing plain-E4M3 rollback form; SM120 keeps the
1947        // faster block-scale identity form. Both consume the same scratch and numeric program.
1948        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            // Scratch layouts are footprint-identical, so only the entry point swaps.
1964            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    /// PER-BLOCK FP8 MMQ prefill GEMM (cu/mmq_fp8_blk.cu). `w_e4m3` is the raw checkpoint e4m3
2001    /// plane [out_f x in_f] and `blk_scales` the device f32 grid [ceil(out_f/128) x
2002    /// ceil(in_f/128)] — no re-quantization of either.
2003    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)] // allow: the parameter list mirrors the kernel/FFI/call contract; bundling into a struct is a refactor, not a lint fix
2016    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        // allow: explicit (n + k - 1) / k is the load-bearing sizing form, kept textually identical to the kernel-side math
2040        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    /// View-backed twin of `qmatvec_mmq_fp8_blk`. Resident expert banks remain in their
2084    /// layer-wide allocations while the selected expert and token rows are passed as views.
2085    /// The CUDA launcher still performs dynamic E4M3 activation quantization; no Q8 activation
2086    /// sidecar is created.
2087    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    /// Count e4m3 NaN codes (magnitude 0x7F) in a device e4m3 plane. 0 is the precondition for
2159    /// routing that tensor through `qmatvec_mmq_fp8_blk` (hardware decodes them to NaN, the
2160    /// host/ARM B' reference to 0.0).
2161    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    /// Quantize token-major f32 activation [n_tokens, in_f] to the block_q8_1_mmq (D4) scratch the
2187    /// IQ expert-MMA kernel consumes. Returns the scratch buffer (one per proj input per layer).
2188    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    /// Fused act-epilogue (research lever #3): silu/gelu(gate)*up + D4 quantize in one launch —
2217    /// replaces moe_pairs_{silu,gelu}_mul + mmq_iq_quantize_act without materializing the f32 act
2218    /// buffer (saves one full write + one full read pass over [n_pairs x n_ff]). Scratch bytes are
2219    /// BYTE-IDENTICAL to the two-pass path (kernel-check `iq fused act+quant` gates it).
2220    /// `act_kind`: 0 = silu*mul (qwen35moe), 1 = gelu_tanh*mul (gemma4).
2221    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    /// Expert-segmented IQ3_S/IQ4_XS int8-MMA MMQ (the m16n8k16.s8 analog of moe_pairs_matvec_q8_dec).
2255    /// Same CSR inputs (table/ex_ids/ex_off/ex_pairs/pair_tok) + a pre-quantized q8_1_mmq activation
2256    /// scratch (from `mmq_iq_quantize_act` over n_tokens). y = [n_pairs, out_f] pair-major.
2257    #[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    /// Gather+convert the activation to f16 pair-major [n_pairs, in_f] for the grouped
2314    /// GEMM, normalized per row by its amax (raw f16 overflows on gemma's activation
2315    /// spikes — round 46 NaN find). Returns (act_f16, row_scales) — the scales fold back
2316    /// into the GEMM output. `pair_tok` = None when the input is already pair-major.
2317    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    /// One projection through the grouped f16 lane: dequant the active experts' rows to an
2357    /// f16 workspace, then ONE grouped GEMM over the CSR groups (variable m per expert).
2358    /// y = f32 [n_pairs, out_f] pair-major — same layout as mmq_iq_experts.
2359    /// MEMRA_MOE_F16G=1: cublasGemmGroupedBatchedEx (+ h2f pass + per-projection sync — the
2360    /// grouped API runs on internal streams unordered with ours, round-47 ledger).
2361    /// MEMRA_MOE_F16G=2: single-kernel grouped GEMM on the engine stream (round 49) — the
2362    /// row scale folds into the kernel epilogue; no f16 C, no h2f, NO sync (ordered by
2363    /// construction). f16-MIRROR numeric class either way (argmax/spec gated, not
2364    /// byte-identity). Errors on unsupported qtype (caller keeps the MMQ arm as fallback).
2365    #[allow(clippy::too_many_arguments)]
2366    /// Bind the RUNTIME API's current device to `ordinal`. Every raw `<<<>>>` launch in the
2367    /// grouped-MoE FFI follows this, not cudarc's pushed driver context — mandatory before
2368    /// calling the FFI on a non-root rank engine (the TP2 grouped prime), a mismatch is
2369    /// cudaErrorInvalidValue.
2370    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: the parameter list mirrors the kernel/FFI/call contract; bundling into a struct is a refactor, not a lint fix
2380    #[allow(clippy::manual_is_multiple_of)] // allow: divisor is runtime-derived; the modulo form keeps a zero divisor loud (a panic), where is_multiple_of would return false silently
2381    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        // DIRECT-FROM-QUANT lane (lane/kquant-tile-loaders + lane/iq-direct-loaders, default
2400        // ON — MEMRA_F16G_DIRECT=0 is the rollback seam): Q4_K/Q6_K/IQ4_XS/IQ3_S expert
2401        // projections skip the dequant-workspace pass entirely; the sk visitor forms dequant
2402        // B tiles in-register from the superblocks. Bit-identical to the workspace path by
2403        // construction (kernel-check "f16g-kq-direct") — this is a pure data-movement change,
2404        // not a numeric-class change. Admission mirrors the C-side guards; the grid-scan
2405        // rollback arm (MEMRA_F16G_SK=0) keeps the workspace.
2406        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                // v2 slot-major banks read through the same direct lane (kq_fetch's v2 branch),
2416                // which is what keeps the grouped prime off the 1.5 GB/projection dequant
2417                // workspace it otherwise falls back to.
2418                || qtype == crate::QT_NVFP4_V2)
2419            // NVFP4 walks 64-value blocks (its 16-value window is one UE4M3 sub-block);
2420            // the kq/IQ classes walk 256-value superblocks. Mirrors the C-side guard.
2421            && 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        // one-time cublas grouped init (algo heuristics + module load cost ~10% of a cold
2468        // g26 prime when paid inside the first projection): a tiny dummy grouped GEMM at
2469        // first use, synced, so the real prime runs warm. The =2 path never touches cublas.
2470        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        // MODE 1 ONLY: cublasGemmGroupedBatchedEx issues through internal streams NOT ordered
2601        // with ours (round 46: NaN race, clean under sync — 205=205 MATCH). Full sync per
2602        // projection. Mode 2 (single kernel, our stream) is ordered by construction — no sync,
2603        // that is the point of this arc.
2604        if !sk {
2605            self.gpu.stream().synchronize()?;
2606        }
2607        if std::env::var("MEMRA_F16G_DEBUG").is_ok() {
2608            // FULL NaN/Inf scan of w, act (through h2f) and y — localizes the corrupt stage.
2609            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    /// Raw sk grouped-GEMM entry for kernel-check ("f16g-sk" section): explicit shape/cross
2656    /// instead of the env policy. shape_sel < 0 = the round-49 grid-scan rollback arm; else
2657    /// the round-51 problem-visitor split at `cross` (1 forces all-128, i32::MAX all-32).
2658    /// tail: 1 = the deep tail (32x64x64 3-stage, lane/sk-tail-form) on sub-cross groups,
2659    /// 0 = the round-51 2-stage 32x64x32 tail.
2660    /// w_f16 = [n_active][out_f][in_f] f16 bytes, act_f16 = [n_pairs][in_f] f16 bytes.
2661    #[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    /// Raw direct-from-quant sk grouped-GEMM entry for kernel-check ("f16g-kq-direct"):
2716    /// explicit cross/tail instead of the env policy. `table` = device u64 pointer table
2717    /// (proj-major, [n_proj][n_expert] — same contract as moe_f16_grouped), `ex_ids` =
2718    /// active-expert ids (device). Visitor forms only (the C side rejects anything else).
2719    #[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    /// Raw dequant-workspace entry for kernel-check: dequant the active experts' rows to a
2783    /// fresh f16 workspace via the same kernel `moe_f16_grouped` uses (the direct loaders'
2784    /// bitwise reference).
2785    #[allow(clippy::too_many_arguments)] // allow: the parameter list mirrors the kernel/FFI/call contract; bundling into a struct is a refactor, not a lint fix
2786    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}