1use crate::Engine;
6use crate::model::{EmbedHost, GpuTensor, HostExps};
7use cudarc::driver::CudaSlice;
8use memra_gguf::config::ModelConfig;
9use memra_gguf::model_plan::{AttentionPlan, MlpPlan};
10use memra_gguf::source::{GgufSource, TensorSource};
11use memra_gguf::{GgmlType, GgufFile};
12use std::collections::HashMap;
13use std::sync::Arc;
14
15fn load_t(
18 e: &Engine,
19 src: &dyn TensorSource,
20 name: &str,
21) -> Result<GpuTensor, Box<dyn std::error::Error>> {
22 GpuTensor::load_from_source(e, src, name)
23}
24fn load_opt(
25 e: &Engine,
26 src: &dyn TensorSource,
27 name: &str,
28) -> Result<Option<GpuTensor>, Box<dyn std::error::Error>> {
29 GpuTensor::load_opt_from_source(e, src, name)
30}
31
32struct ResidencyBytes {
33 experts: HashMap<usize, usize>,
34 rest: usize,
35 saw_experts: bool,
36}
37
38fn block_index(name: &str) -> Option<usize> {
39 name.strip_prefix("blk.")?.split('.').next()?.parse().ok()
40}
41
42fn residency_bytes_by_device<'a>(
43 tensors: impl IntoIterator<Item = (&'a str, usize)>,
44 layer_devices: &[usize],
45 primary_device: usize,
46) -> ResidencyBytes {
47 let mut out = ResidencyBytes {
48 experts: HashMap::new(),
49 rest: 0,
50 saw_experts: false,
51 };
52 for (name, bytes) in tensors {
53 if name.starts_with("blk.") && name.contains("_exps.") {
54 let device = block_index(name)
55 .and_then(|il| layer_devices.get(il).copied())
56 .unwrap_or(primary_device);
57 *out.experts.entry(device).or_default() += bytes;
58 out.saw_experts = true;
59 } else {
60 out.rest += bytes;
61 }
62 }
63 out
64}
65
66pub(crate) struct ResidentPlan {
69 primary_device: usize,
70 layer_devices: Vec<usize>,
71 layer_counts: HashMap<usize, usize>,
72 exact_expert_bytes: Option<HashMap<usize, usize>>,
73 trunk_bytes: usize,
74 decisions: HashMap<usize, bool>,
75 pp: bool,
76}
77
78#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
88enum StepExpertArtifact {
89 #[default]
90 E4m3,
91 Nvfp4,
92}
93
94#[derive(Clone, Debug, Default)]
95struct StepParallelLoadConfig {
96 ep_specs: Vec<crate::tp::StepEpLayerSpec>,
97 tp_specs: Vec<crate::tp::StepTpLayerSpec>,
98 native_p2p: bool,
99 ep_device_arithmetic: bool,
100 f32_mirror: bool,
101 bulk_p2p: bool,
102 expert_artifact: StepExpertArtifact,
103}
104
105#[derive(Default)]
106pub(crate) struct StepParallelRuntimeRegistry {
107 config: StepParallelLoadConfig,
108 runtimes: HashMap<(Vec<usize>, bool, bool, bool), Arc<crate::tp::TpE4m3HostBounce>>,
109}
110
111#[derive(Clone, Copy, Debug, PartialEq, Eq)]
112enum StepExpertLayout {
113 TensorParallel,
114 ExpertParallel,
115}
116
117#[derive(Clone, Debug, PartialEq, Eq)]
118struct StepExpertSelection {
119 spec: crate::tp::StepEpLayerSpec,
120 layout: StepExpertLayout,
121 configured_by_tp: bool,
122}
123
124fn select_step_expert_layout(
125 layer: usize,
126 ep_specs: &[crate::tp::StepEpLayerSpec],
127 tp_specs: &[crate::tp::StepTpLayerSpec],
128) -> Result<Option<StepExpertSelection>, String> {
129 let ep = ep_specs.iter().find(|spec| spec.layer == layer);
130 let tp = tp_specs.iter().find(|spec| spec.layer == layer);
131 if ep.is_some() && tp.is_some() {
132 return Err(format!(
133 "Step layer {layer} cannot enable MEMRA_STEP_EP and MEMRA_STEP_TP together"
134 ));
135 }
136 Ok(match (ep, tp) {
137 (Some(spec), None) => Some(StepExpertSelection {
138 spec: spec.clone(),
139 layout: StepExpertLayout::ExpertParallel,
140 configured_by_tp: false,
141 }),
142 (None, Some(spec)) => Some(StepExpertSelection {
143 spec: spec.clone(),
144 layout: if spec.devices.len() > 2 {
145 StepExpertLayout::ExpertParallel
146 } else {
147 StepExpertLayout::TensorParallel
148 },
149 configured_by_tp: true,
150 }),
151 (None, None) => None,
152 (Some(_), Some(_)) => unreachable!(),
153 })
154}
155
156impl StepParallelRuntimeRegistry {
157 fn with_config(config: StepParallelLoadConfig) -> Self {
158 Self {
159 config,
160 runtimes: HashMap::new(),
161 }
162 }
163
164 fn tp_spec(&self, layer: usize) -> Option<&crate::tp::StepTpLayerSpec> {
165 self.config.tp_specs.iter().find(|spec| spec.layer == layer)
166 }
167
168 fn expert_selection(&self, layer: usize) -> Result<Option<StepExpertSelection>, String> {
169 select_step_expert_layout(layer, &self.config.ep_specs, &self.config.tp_specs)
170 }
171
172 fn runtime(
173 &mut self,
174 devices: &[usize],
175 native_p2p: bool,
176 ep_device_arithmetic: bool,
177 ) -> Result<Arc<crate::tp::TpE4m3HostBounce>, Box<dyn std::error::Error>> {
178 let bulk_p2p = self.config.bulk_p2p && native_p2p;
179 let key = (devices.to_vec(), native_p2p, ep_device_arithmetic, bulk_p2p);
180 if let Some(runtime) = self.runtimes.get(&key) {
181 return Ok(Arc::clone(runtime));
182 }
183 let runtime = Arc::new(crate::tp::TpE4m3HostBounce::new_configured(
184 devices,
185 native_p2p,
186 ep_device_arithmetic,
187 bulk_p2p,
188 )?);
189 let names = runtime.device_names()?;
190 if names
191 .iter()
192 .any(|name| !name.contains("RTX PRO 6000") || !name.contains("Blackwell"))
193 {
194 return Err(format!(
195 "Step distributed execution is qualified only on RTX PRO 6000 Blackwell, \
196 got {names:?}"
197 )
198 .into());
199 }
200 self.runtimes.insert(key, Arc::clone(&runtime));
201 Ok(runtime)
202 }
203}
204
205impl ResidentPlan {
206 fn from_layout(
207 src: &dyn TensorSource,
208 primary_device: usize,
209 layer_devices: Vec<usize>,
210 pp: bool,
211 ) -> Self {
212 let mut layer_counts = HashMap::new();
213 for &device in &layer_devices {
214 *layer_counts.entry(device).or_default() += 1;
215 }
216 let (exact_expert_bytes, trunk_bytes) = match src.gguf() {
217 Some(g) => {
218 let bytes = residency_bytes_by_device(
219 g.tensors
220 .iter()
221 .map(|t| (t.name.as_str(), t.n_bytes as usize)),
222 &layer_devices,
223 primary_device,
224 );
225 if bytes.saw_experts {
226 (Some(bytes.experts), bytes.rest)
227 } else {
228 (None, 0)
229 }
230 }
231 None => (None, 0),
232 };
233 Self {
234 primary_device,
235 layer_devices,
236 layer_counts,
237 exact_expert_bytes,
238 trunk_bytes,
239 decisions: HashMap::new(),
240 pp,
241 }
242 }
243
244 pub(crate) fn unsharded(e: &Engine, src: &dyn TensorSource, cfg: &ModelConfig) -> Self {
245 let device = e.ctx().ordinal();
246 Self::from_layout(src, device, vec![device; cfg.n_layer as usize], false)
247 }
248
249 pub(crate) fn pp(
250 e: &Engine,
251 src: &dyn TensorSource,
252 cfg: &ModelConfig,
253 n_trunk: usize,
254 ) -> Result<Self, Box<dyn std::error::Error>> {
255 let primary = e.ctx().ordinal();
256 let Some(_fence) = crate::pp::pp_cuts(n_trunk) else {
257 return Ok(Self::unsharded(e, src, cfg));
258 };
259 let mut layer_devices = vec![primary; cfg.n_layer as usize];
260 for (il, device) in layer_devices.iter_mut().take(n_trunk).enumerate() {
261 *device = crate::pp::layer_engine(e, n_trunk, il)?.ctx().ordinal();
262 }
263 Ok(Self::from_layout(src, primary, layer_devices, true))
264 }
265
266 fn should_reside(&mut self, e: &Engine, il: usize, per_layer: usize) -> bool {
267 let device = self
268 .layer_devices
269 .get(il)
270 .copied()
271 .unwrap_or(self.primary_device);
272 debug_assert_eq!(e.ctx().ordinal(), device);
273 if let Some(&decision) = self.decisions.get(&device) {
274 return decision;
275 }
276 if std::env::var("MEMRA_MOE_RESIDENT").as_deref() == Ok("0") {
277 self.decisions.insert(device, false);
278 return false;
279 }
280 let (free, _total) = match e.ctx().mem_get_info() {
281 Ok(v) => v,
282 Err(_) => {
283 self.decisions.insert(device, false);
284 return false;
285 }
286 };
287 let projected = self
288 .exact_expert_bytes
289 .as_ref()
290 .map(|bytes| bytes.get(&device).copied().unwrap_or(0))
291 .unwrap_or(per_layer * self.layer_counts.get(&device).copied().unwrap_or(1));
292 let budget = std::env::var("MEMRA_MOE_RESIDENT_GB")
293 .ok()
294 .and_then(|v| v.parse::<f64>().ok())
295 .map(|gb| (gb * 1e9) as usize)
296 .unwrap_or_else(|| {
297 let reserve = std::env::var("MEMRA_MOE_RESIDENT_HEADROOM_GB")
298 .ok()
299 .and_then(|v| v.parse::<f64>().ok())
300 .map(|gb| (gb * 1e9) as usize)
301 .unwrap_or(2_000_000_000);
302 (free as usize).saturating_sub(self.trunk_bytes + reserve)
303 });
304 let ok = projected <= budget;
305 eprintln!(
306 "[moe] resident-experts decision ({}dev{}): experts {:.2}GB + trunk {:.2}GB vs free {:.2}GB (expert budget {:.2}GB) -> {}",
307 if self.pp { "PP " } else { "" },
308 device,
309 projected as f64 / 1e9,
310 self.trunk_bytes as f64 / 1e9,
311 free as f64 / 1e9,
312 budget as f64 / 1e9,
313 if ok { "RESIDENT" } else { "SLRU cache" }
314 );
315 self.decisions.insert(device, ok);
316 ok
317 }
318}
319
320fn load_mixer_kind(
322 e: &Engine,
323 src: &dyn TensorSource,
324 cfg: &ModelConfig,
325 il: u32,
326 attention: &AttentionPlan,
327 step_runtimes: &mut StepParallelRuntimeRegistry,
328) -> Result<Mixer, Box<dyn std::error::Error>> {
329 let p = |s: &str| format!("blk.{il}.{s}");
330 Ok(match attention {
331 AttentionPlan::Mla(mla) => Mixer::Mla(MlaAttnLayer::load(e, src, il, mla)?),
332 AttentionPlan::Full(full)
333 | AttentionPlan::SlidingWindow {
334 attention: full, ..
335 } => {
336 Mixer::Full(FullAttnLayer {
337 wq: load_t(e, src, &p("attn_q.weight"))?,
338 wk: load_t(e, src, &p("attn_k.weight"))?,
339 wv: match load_opt(e, src, &p("attn_v.weight"))? {
344 Some(v) => v,
345 None => load_t(e, src, &p("attn_k.weight"))?,
346 },
347 wo: load_t(e, src, &p("attn_output.weight"))?,
348 q_norm: load_t(e, src, &p("attn_q_norm.weight"))?,
349 k_norm: load_t(e, src, &p("attn_k_norm.weight"))?,
350 attn_gate: if full.output_gate
354 == memra_gguf::config::AttentionGateKind::SeparateHead
355 {
356 Some(load_t(e, src, &p("attn_gate.weight"))?)
357 } else {
358 None
359 },
360 step_tp_qkv: build_step_tp_qkv(e, src, cfg, il as usize, step_runtimes)?,
361 })
362 }
363 AttentionPlan::GatedDeltaNet(geometry) => Mixer::Linear(LinearAttnLayer {
364 geometry: *geometry,
365 wqkv: load_t(e, src, &p("attn_qkv.weight"))?,
366 wqkv_gate: load_t(e, src, &p("attn_gate.weight"))?,
367 ssm_beta: load_t(e, src, &p("ssm_beta.weight"))?,
368 ssm_alpha: load_t(e, src, &p("ssm_alpha.weight"))?,
369 ssm_a: load_t(e, src, &p("ssm_a"))?,
370 ssm_dt: load_t(e, src, &p("ssm_dt.bias"))?,
371 ssm_conv1d: load_t(e, src, &p("ssm_conv1d.weight"))?,
372 ssm_norm: load_t(e, src, &p("ssm_norm.weight"))?,
373 ssm_out: load_t(e, src, &p("ssm_out.weight"))?,
374 }),
375 })
376}
377
378pub(crate) fn load_ffn(
385 e: &Engine,
386 src: &dyn TensorSource,
387 cfg: &ModelConfig,
388 mlp: &MlpPlan,
389 il: u32,
390 spill: Option<(&GgufFile, &mut crate::spill::SpillCtx)>,
391 resident: &mut ResidentPlan,
392 step_runtimes: &mut StepParallelRuntimeRegistry,
393) -> Result<Ffn, Box<dyn std::error::Error>> {
394 let p = |s: &str| format!("blk.{il}.{s}");
395 let artifact_dense = matches!(mlp, MlpPlan::Moe(_))
402 && !src.has(&p("ffn_gate_exps.weight"))
403 && !src.has(&p("ffn_gate_up_exps.weight"))
404 && src.has(&p("ffn_gate.weight"));
405 Ok(if artifact_dense {
406 Ffn::Dense {
407 ffn_gate: GpuTensor::load_from_source(e, src, &p("ffn_gate.weight"))?,
408 ffn_up: GpuTensor::load_from_source(e, src, &p("ffn_up.weight"))?,
409 ffn_down: GpuTensor::load_from_source(e, src, &p("ffn_down.weight"))?,
410 }
411 } else if let MlpPlan::Moe(moe) = mlp {
412 let n_expert = moe.expert_count as usize;
413 let (gate_exps, up_exps, down_exps) = match spill {
419 Some((g, ctx)) => (
420 HostExps::load_tiered(e, g, &p("ffn_gate_exps.weight"), ctx)?,
421 HostExps::load_tiered(e, g, &p("ffn_up_exps.weight"), ctx)?,
422 HostExps::load_tiered(e, g, &p("ffn_down_exps.weight"), ctx)?,
423 ),
424 None => {
425 let exps = |e: &Engine, n: &str| -> Result<HostExps, Box<dyn std::error::Error>> {
426 if src.has(n) {
427 HostExps::load_stacked_from_source(e, src, n)
428 } else {
429 HostExps::load_from_source(e, src, n, n_expert)
430 }
431 };
432 let fused = p("ffn_gate_up_exps.weight");
434 if !src.has(&p("ffn_gate_exps.weight")) && src.has(&fused) {
435 let ff = moe.expert_intermediate_size as usize;
436 (
437 HostExps::load_stacked_split_from_source(e, src, &fused, 0, ff)?,
438 HostExps::load_stacked_split_from_source(e, src, &fused, ff, 2 * ff)?,
439 exps(e, &p("ffn_down_exps.weight"))?,
440 )
441 } else {
442 (
443 exps(e, &p("ffn_gate_exps.weight"))?,
444 exps(e, &p("ffn_up_exps.weight"))?,
445 exps(e, &p("ffn_down_exps.weight"))?,
446 )
447 }
448 }
449 };
450 let (step_ep, step_tp) = build_step_distributed_exps(
451 e,
452 cfg,
453 src,
454 il as usize,
455 &gate_exps,
456 &up_exps,
457 &down_exps,
458 step_runtimes,
459 )?;
460 let dev_exps = if step_ep.is_some() || step_tp.is_some() {
466 None
467 } else {
468 build_dev_exps(e, resident, il as usize, &gate_exps, &up_exps, &down_exps)?
469 };
470 let mut macro_row = vec![1.0f32; 3 * n_expert];
472 for (slot, exps) in [(0usize, &gate_exps), (1, &up_exps), (2, &down_exps)] {
473 if let Some(ms) = exps.macros.as_ref() {
474 macro_row[slot * n_expert..(slot + 1) * n_expert].copy_from_slice(ms);
475 }
476 }
477 let has_macros = macro_row.iter().any(|&m| m != 1.0);
478 let dev_macros = e.htod(¯o_row)?;
479 let exp_probs_b = src
482 .find(&p("exp_probs_b.bias"))
483 .map(|v| memra_gguf::dequant::dequantize(v.ggml_type, &v.bytes, n_expert));
484 let active_experts = src.active_experts(il).map(<[bool]>::to_vec);
485 let route_bias = exp_probs_b.clone().unwrap_or_else(|| vec![0.0; n_expert]);
486 let active_row: Vec<u8> = active_experts
487 .as_ref()
488 .map(|mask| mask.iter().map(|&is_active| u8::from(is_active)).collect())
489 .unwrap_or_else(|| vec![1; n_expert]);
490 let exp_probs_b_dev = e.htod(&route_bias)?;
491 let active_experts_dev = e.htod_bytes(&active_row)?;
492 Ffn::Moe(MoeWeights {
493 gate_inp: load_t(e, src, &p("ffn_gate_inp.weight"))?,
494 gate_inp_shexp: load_opt(e, src, &p("ffn_gate_inp_shexp.weight"))?,
495 exp_probs_b,
496 exp_probs_b_dev,
497 active_experts,
498 active_experts_dev,
499 gate_exps,
500 up_exps,
501 down_exps,
502 gate_shexp: load_opt(e, src, &p("ffn_gate_shexp.weight"))?,
503 up_shexp: load_opt(e, src, &p("ffn_up_shexp.weight"))?,
504 down_shexp: load_opt(e, src, &p("ffn_down_shexp.weight"))?,
505 dev_exps,
506 step_ep,
507 step_tp,
508 dev_macros,
509 has_macros,
510 })
511 } else {
512 Ffn::Dense {
513 ffn_gate: load_t(e, src, &p("ffn_gate.weight"))?,
514 ffn_up: load_t(e, src, &p("ffn_up.weight"))?,
515 ffn_down: load_t(e, src, &p("ffn_down.weight"))?,
516 }
517 })
518}
519
520fn host_e4m3_bank(
521 exps: &HostExps,
522) -> Result<crate::tp::E4m3ExpertBank<'_>, Box<dyn std::error::Error>> {
523 if exps.qtype != crate::QT_F8_E4M3_BLK {
524 return Err(format!(
525 "Step EP requires native block-E4M3 expert banks, got qtype {}",
526 exps.qtype
527 )
528 .into());
529 }
530 let scales = exps
531 .fp8_blk
532 .as_ref()
533 .ok_or("Step EP native expert bank has no block-E4M3 scale plane")?;
534 Ok(crate::tp::E4m3ExpertBank {
535 codes: exps.bytes.as_bytes(),
536 scales: &scales.scales,
537 expert_count: exps.n_expert,
538 out_features: exps.out_f,
539 in_features: exps.in_f,
540 })
541}
542
543fn validate_step_expert_specs(
544 contract: &crate::parallel::ModelParallelContract,
545 flag: &str,
546 specs: &[crate::tp::StepEpLayerSpec],
547 allow_dense_attention_only: bool,
548) -> Result<(), Box<dyn std::error::Error>> {
549 for candidate in specs {
550 if candidate.layer >= contract.trunk_layers {
551 return Err(format!(
552 "{flag} layer {} is outside Step trunk layers 0..{}",
553 candidate.layer, contract.trunk_layers
554 )
555 .into());
556 }
557 if candidate.layer < contract.dense_prefix_layers {
558 if allow_dense_attention_only {
559 continue;
560 }
561 return Err(format!(
562 "{flag} layer {} is outside Step routed-expert layers {}..{}",
563 candidate.layer, contract.dense_prefix_layers, contract.trunk_layers
564 )
565 .into());
566 }
567 }
568 Ok(())
569}
570
571fn validate_step_expert_activation_layout(
572 cfg: &ModelConfig,
573 flag: &str,
574 selection: &StepExpertSelection,
575) -> Result<(), Box<dyn std::error::Error>> {
576 let _ = (cfg, flag, selection);
582 Ok(())
583}
584
585fn prepare_step_parallel_load(
586 e: &Engine,
587 src: &dyn TensorSource,
588 cfg: &ModelConfig,
589 trunk_layers: usize,
590) -> Result<StepParallelLoadConfig, Box<dyn std::error::Error>> {
591 let tp_specs = crate::tp::step_tp_layer_specs()?;
592 let ep_specs = crate::tp::step_ep_layer_specs()?;
593 let device_arithmetic = crate::tp::step_ep_device_arithmetic_enabled()?;
594 let f32_mirror = crate::tp::step_tp_f32_mirror_enabled()?;
595 let bulk_p2p = crate::tp::step_tp_bulk_p2p_enabled()?;
596 if tp_specs.is_empty() {
597 if device_arithmetic || f32_mirror || bulk_p2p {
598 return Err(
599 "MEMRA_STEP_EP_DEVICE_ARITHMETIC=1, MEMRA_STEP_TP_F32_MIRROR=1, or \
600 MEMRA_STEP_TP_BULK_P2P=1 requires MEMRA_STEP_TP; device arithmetic and bulk \
601 transport also require MEMRA_STEP_TP_NATIVE_P2P=1"
602 .into(),
603 );
604 }
605 let expert_artifact = if ep_specs.is_empty() {
608 StepExpertArtifact::default()
609 } else {
610 let contract = crate::parallel::ModelParallelContract::from_model(cfg)?;
611 match crate::parallel::validate_step_fp8_checkpoint(src, &contract) {
612 Ok(_) => StepExpertArtifact::E4m3,
613 Err(fp8_error) => {
614 match crate::parallel::validate_step_nvfp4_checkpoint(src, &contract) {
615 Ok(_) => StepExpertArtifact::Nvfp4,
616 Err(nvfp4_error) => {
617 return Err(format!(
618 "Step checkpoint qualifies as neither native expert artifact \
619 class: [E4M3] {fp8_error} [NVFP4] {nvfp4_error}"
620 )
621 .into());
622 }
623 }
624 }
625 }
626 };
627 return Ok(StepParallelLoadConfig {
628 ep_specs,
629 expert_artifact,
630 ..StepParallelLoadConfig::default()
631 });
632 }
633 let contract = crate::parallel::ModelParallelContract::from_model(cfg)?;
634 validate_step_expert_specs(&contract, "MEMRA_STEP_EP", &ep_specs, false)?;
635 validate_step_expert_specs(&contract, "MEMRA_STEP_TP", &tp_specs, true)?;
636 for spec in &tp_specs {
637 let selection = select_step_expert_layout(spec.layer, &ep_specs, &tp_specs)?
638 .ok_or("Step TP expert selection disappeared during preflight")?;
639 validate_step_expert_activation_layout(cfg, "MEMRA_STEP_TP", &selection)?;
640 }
641
642 let layer_owners = (0..trunk_layers)
643 .map(|layer| {
644 crate::pp::layer_engine(e, trunk_layers, layer).map(|engine| engine.ctx().ordinal())
645 })
646 .collect::<Result<Vec<_>, _>>()?;
647 let plan = contract.preflight_step_tp_specs(
648 tp_specs
649 .iter()
650 .map(|spec| (spec.layer, spec.devices.as_slice())),
651 &layer_owners,
652 )?;
653
654 for devices in &plan.runtime_groups {
655 let hardware = crate::parallel::detect_uniform_hardware(devices)?;
656 if !contract.hardware_targets.contains(&hardware) {
657 return Err(format!(
658 "{} has no qualified {hardware:?} TP contract for devices {devices:?}",
659 contract.variant
660 )
661 .into());
662 }
663 }
664
665 let native_p2p = crate::tp::step_tp_native_p2p_enabled()?;
666 if bulk_p2p && !native_p2p {
667 return Err("MEMRA_STEP_TP_BULK_P2P=1 requires MEMRA_STEP_TP_NATIVE_P2P=1".into());
668 }
669 if device_arithmetic
670 && (!ep_specs.is_empty()
671 || !native_p2p
672 || plan.expert_parallel_layers() == 0
673 || plan.tensor_parallel_expert_layers() != 0)
674 {
675 return Err(
676 "MEMRA_STEP_EP_DEVICE_ARITHMETIC=1 requires native-P2P TP4/TP8 \
677 expert ownership for every selected routed-expert layer"
678 .into(),
679 );
680 }
681 let (qualified_experts, expert_artifact) =
685 match crate::parallel::validate_step_fp8_checkpoint(src, &contract) {
686 Ok(qualified) => (qualified, StepExpertArtifact::E4m3),
687 Err(fp8_error) => match crate::parallel::validate_step_nvfp4_checkpoint(src, &contract)
688 {
689 Ok(qualified) => (qualified, StepExpertArtifact::Nvfp4),
690 Err(nvfp4_error) => {
691 return Err(format!(
692 "Step checkpoint qualifies as neither native expert artifact class: \
693 [E4M3] {fp8_error} [NVFP4] {nvfp4_error}"
694 )
695 .into());
696 }
697 },
698 };
699 if expert_artifact == StepExpertArtifact::Nvfp4 {
700 if device_arithmetic {
701 return Err(
702 "MEMRA_STEP_EP_DEVICE_ARITHMETIC=1 is qualified for the E4M3 expert artifact \
703 only; the NVFP4 expert program is host-canonical in this increment"
704 .into(),
705 );
706 }
707 if bulk_p2p {
712 return Err(
713 "MEMRA_STEP_TP_BULK_P2P=1 is qualified for the E4M3 expert artifact only; the \
714 NVFP4 bank transport increment has not landed"
715 .into(),
716 );
717 }
718 }
719
720 if f32_mirror {
721 eprintln!(
722 "[step-tp-preflight] layers={} full_trunk={} runtime_groups={} \
723 dense_attention_layers={} tensor_expert_layers={} expert_owner_layers={} \
724 qualified_fp8_expert_projection_slices={} owner_first=true \
725 hardware=rtx-pro-6000-blackwell \
726 native_p2p={} bulk_p2p={} device_arithmetic={} bf16_residency=f32-mirror \
727 weights_loaded=false performance_claim=false",
728 plan.layers.len(),
729 plan.full_trunk,
730 plan.runtime_groups.len(),
731 plan.dense_attention_layers(),
732 plan.tensor_parallel_expert_layers(),
733 plan.expert_parallel_layers(),
734 qualified_experts,
735 native_p2p,
736 bulk_p2p,
737 device_arithmetic,
738 );
739 } else {
740 eprintln!(
741 "[step-tp-preflight] layers={} full_trunk={} runtime_groups={} \
742 dense_attention_layers={} tensor_expert_layers={} expert_owner_layers={} \
743 qualified_fp8_expert_projection_slices={} owner_first=true \
744 hardware=rtx-pro-6000-blackwell \
745 native_p2p={} bulk_p2p={} device_arithmetic={} \
746 weights_loaded=false performance_claim=false",
747 plan.layers.len(),
748 plan.full_trunk,
749 plan.runtime_groups.len(),
750 plan.dense_attention_layers(),
751 plan.tensor_parallel_expert_layers(),
752 plan.expert_parallel_layers(),
753 qualified_experts,
754 native_p2p,
755 bulk_p2p,
756 device_arithmetic,
757 );
758 }
759 Ok(StepParallelLoadConfig {
760 ep_specs,
761 tp_specs,
762 native_p2p,
763 ep_device_arithmetic: device_arithmetic,
764 f32_mirror,
765 bulk_p2p,
766 expert_artifact,
767 })
768}
769
770fn step_nvfp4_native<'a>(
772 src: &'a dyn TensorSource,
773 layer: usize,
774 proj: &str,
775) -> Result<memra_gguf::source::Nvfp4StackedNative<'a>, Box<dyn std::error::Error>> {
776 let name = format!("blk.{layer}.ffn_{proj}_exps.weight");
777 src.find_nvfp4_stacked_native(&name)
778 .ok_or_else(|| format!("Step NVFP4 expert program is missing native bank {name}").into())
779}
780
781fn step_nvfp4_bank<'a>(
783 native: &'a memra_gguf::source::Nvfp4StackedNative<'a>,
784) -> crate::tp::Nvfp4ExpertBank<'a> {
785 crate::tp::Nvfp4ExpertBank {
786 codes: native.codes,
787 scales: native.scales,
788 macros: &native.macros,
789 expert_count: native.n_expert,
790 out_features: native.out_f,
791 in_features: native.in_f,
792 }
793}
794
795fn build_step_distributed_exps(
796 e: &Engine,
797 cfg: &ModelConfig,
798 src: &dyn TensorSource,
799 layer: usize,
800 gate: &HostExps,
801 up: &HostExps,
802 down: &HostExps,
803 step_runtimes: &mut StepParallelRuntimeRegistry,
804) -> Result<(Option<StepEpExps>, Option<StepTpExps>), Box<dyn std::error::Error>> {
805 let ep_device_arithmetic = step_runtimes.config.ep_device_arithmetic;
806 if step_runtimes.config.ep_specs.is_empty() && step_runtimes.config.tp_specs.is_empty() {
807 if ep_device_arithmetic {
808 return Err(
809 "MEMRA_STEP_EP_DEVICE_ARITHMETIC=1 requires MEMRA_STEP_TP and \
810 MEMRA_STEP_TP_NATIVE_P2P=1"
811 .into(),
812 );
813 }
814 return Ok((None, None));
815 }
816 let contract = crate::parallel::ModelParallelContract::from_model(cfg)?;
817 validate_step_expert_specs(
818 &contract,
819 "MEMRA_STEP_EP",
820 &step_runtimes.config.ep_specs,
821 false,
822 )?;
823 validate_step_expert_specs(
824 &contract,
825 "MEMRA_STEP_TP",
826 &step_runtimes.config.tp_specs,
827 true,
828 )?;
829 let Some(selection) = step_runtimes.expert_selection(layer)? else {
830 return Ok((None, None));
831 };
832 validate_step_expert_activation_layout(
833 cfg,
834 if selection.configured_by_tp {
835 "MEMRA_STEP_TP"
836 } else {
837 "MEMRA_STEP_EP"
838 },
839 &selection,
840 )?;
841 let activation_limit = cfg.clamp_exp_at(layer as u32);
842 let owner = e.ctx().ordinal();
843 if !selection.spec.devices.contains(&owner) {
844 let flag = if selection.configured_by_tp {
845 "MEMRA_STEP_TP"
846 } else {
847 "MEMRA_STEP_EP"
848 };
849 return Err(format!(
850 "{flag} layer {layer} owning PP device {owner} is absent from rank devices {:?}",
851 selection.spec.devices
852 )
853 .into());
854 }
855 let expert_parallel = selection.layout == StepExpertLayout::ExpertParallel;
856 if selection.configured_by_tp {
857 contract.plan(crate::parallel::TopologyRequest {
858 pipeline: 1,
859 tensor: selection.spec.devices.len(),
860 expert_parallel,
861 available_devices: selection.spec.devices.len(),
862 hardware: crate::parallel::HardwareTarget::RtxPro6000Blackwell,
863 })?;
864 }
865 let native_p2p = selection.configured_by_tp && step_runtimes.config.native_p2p;
866 if ep_device_arithmetic
867 && (!selection.configured_by_tp
868 || selection.layout != StepExpertLayout::ExpertParallel
869 || !native_p2p)
870 {
871 return Err(
872 "MEMRA_STEP_EP_DEVICE_ARITHMETIC=1 requires a MEMRA_STEP_TP TP4/TP8 \
873 expert-owner layer and MEMRA_STEP_TP_NATIVE_P2P=1"
874 .into(),
875 );
876 }
877 let runtime =
878 step_runtimes.runtime(&selection.spec.devices, native_p2p, ep_device_arithmetic)?;
879 let expert_artifact = step_runtimes.config.expert_artifact;
880 match selection.layout {
881 StepExpertLayout::ExpertParallel => {
882 if expert_artifact == StepExpertArtifact::Nvfp4 {
883 let runtime = step_runtimes.runtime(
890 &selection.spec.devices,
891 step_runtimes.config.native_p2p,
892 false,
893 )?;
894 let gate_native = step_nvfp4_native(src, layer, "gate")?;
895 let up_native = step_nvfp4_native(src, layer, "up")?;
896 let down_native = step_nvfp4_native(src, layer, "down")?;
897 let experts = runtime.upload_expert_parallel_nvfp4(
898 step_nvfp4_bank(&gate_native),
899 step_nvfp4_bank(&up_native),
900 step_nvfp4_bank(&down_native),
901 )?;
902 eprintln!(
903 "[step-ep] layer={layer} devices={:?} experts={} artifact=nvfp4 \
904 expert_layout=expert-parallel expert_transport=host-bounce \
905 macro_fold=post-kernel-once native_p2p=false performance_claim=false",
906 selection.spec.devices, contract.expert_count
907 );
908 if let Some(limit) = activation_limit {
909 eprintln!(
910 "[step-ep-clamp] load layer={layer} routed_clamp={limit} \
911 formula=min-silu-times-clamped-up performance_claim=false"
912 );
913 }
914 return Ok((
915 Some(StepEpExps {
916 runtime,
917 experts: StepEpExpertBank::Nvfp4(experts),
918 devices: selection.spec.devices,
919 configured_by_tp: selection.configured_by_tp,
920 activation_limit,
921 grouped_decode: None,
922 }),
923 None,
924 ));
925 }
926 let experts = runtime.upload_expert_parallel(
927 host_e4m3_bank(gate)?,
928 host_e4m3_bank(up)?,
929 host_e4m3_bank(down)?,
930 )?;
931 let grouped_decode = if ep_device_arithmetic {
932 let tokens = 1;
933 let selected = (0..contract.experts_per_token).collect::<Vec<_>>();
934 let input = vec![0.0f32; contract.hidden_size];
935 let route_weights = vec![1.0f32; contract.experts_per_token];
936 let projection = runtime.prepare_step_grouped_expert_parallel_gate_with_capacity(
937 &experts,
938 &input,
939 tokens,
940 &selected,
941 activation_limit,
942 tokens,
943 )?;
944 let combine = runtime
945 .prepare_step_grouped_expert_parallel_combine(&projection, &route_weights)?;
946 Some(std::sync::Mutex::new(StepEpGroupedDecode {
947 projection,
948 combine,
949 }))
950 } else {
951 None
952 };
953 if selection.configured_by_tp {
954 eprintln!(
955 "[step-tp-ep] layer={layer} devices={:?} experts={} tp={} \
956 attention_layout=tensor-parallel expert_layout=expert-parallel \
957 expert_transport={} tp_transport={} native_p2p={} \
958 activation={} accumulation={} output={} \
959 grouped_decode_prepared={} grouped_decode_capacity=1 \
960 performance_claim=false",
961 selection.spec.devices,
962 contract.expert_count,
963 selection.spec.devices.len(),
964 runtime.transport_label(),
965 runtime.transport_label(),
966 runtime.native_p2p(),
967 runtime.expert_activation_label(),
968 runtime.expert_accumulation_label(),
969 runtime.expert_output_label(),
970 grouped_decode.is_some(),
971 );
972 } else {
973 eprintln!(
974 "[step-ep] layer={layer} devices={:?} experts={} \
975 expert_layout=expert-parallel expert_transport=host-bounce \
976 native_p2p=false performance_claim=false",
977 selection.spec.devices, contract.expert_count
978 );
979 }
980 if let Some(limit) = activation_limit {
981 eprintln!(
982 "[step-ep-clamp] load layer={layer} routed_clamp={limit} \
983 formula=min-silu-times-clamped-up performance_claim=false"
984 );
985 }
986 Ok((
987 Some(StepEpExps {
988 runtime,
989 experts: StepEpExpertBank::E4m3(experts),
990 devices: selection.spec.devices,
991 configured_by_tp: selection.configured_by_tp,
992 activation_limit,
993 grouped_decode,
994 }),
995 None,
996 ))
997 }
998 StepExpertLayout::TensorParallel => {
999 if activation_limit.is_some() && expert_artifact == StepExpertArtifact::E4m3 {
1000 return Err(format!(
1001 "layer {layer} uses the routed SwiGLU clamp and the E4M3 TP expert \
1002 program has no clamp arm; select EP for this layer (the NVFP4 TP \
1003 program carries the clamp)"
1004 )
1005 .into());
1006 }
1007 let experts = if expert_artifact == StepExpertArtifact::Nvfp4 {
1008 let gate_native = step_nvfp4_native(src, layer, "gate")?;
1009 let up_native = step_nvfp4_native(src, layer, "up")?;
1010 let down_native = step_nvfp4_native(src, layer, "down")?;
1011 StepTpExpertBank::Nvfp4(runtime.upload_tensor_parallel_nvfp4(
1012 step_nvfp4_bank(&gate_native),
1013 step_nvfp4_bank(&up_native),
1014 step_nvfp4_bank(&down_native),
1015 )?)
1016 } else {
1017 StepTpExpertBank::E4m3(runtime.upload_tensor_parallel(
1018 host_e4m3_bank(gate)?,
1019 host_e4m3_bank(up)?,
1020 host_e4m3_bank(down)?,
1021 )?)
1022 };
1023 eprintln!(
1024 "[step-tp] layer={layer} devices={:?} experts={} tp={} artifact={} \
1025 expert_layout=tensor-parallel transport={} native_p2p={} \
1026 performance_claim=false",
1027 selection.spec.devices,
1028 contract.expert_count,
1029 selection.spec.devices.len(),
1030 match expert_artifact {
1031 StepExpertArtifact::E4m3 => "e4m3",
1032 StepExpertArtifact::Nvfp4 => "nvfp4",
1033 },
1034 runtime.transport_label(),
1035 runtime.native_p2p(),
1036 );
1037 if let Some(limit) = activation_limit {
1038 eprintln!(
1039 "[step-tp-clamp] load layer={layer} routed_clamp={limit} \
1040 formula=min-silu-times-clamped-up performance_claim=false"
1041 );
1042 }
1043 Ok((
1044 None,
1045 Some(StepTpExps {
1046 runtime,
1047 experts,
1048 devices: selection.spec.devices,
1049 activation_limit,
1050 }),
1051 ))
1052 }
1053 }
1054}
1055
1056fn upload_step_bf16_column(
1057 runtime: &crate::tp::TpE4m3HostBounce,
1058 src: &dyn TensorSource,
1059 name: &str,
1060 expected_in: usize,
1061 expected_out: usize,
1062 f32_mirror: bool,
1063) -> Result<crate::tp::ResidentBf16ColumnParallel, Box<dyn std::error::Error>> {
1064 let tensor = src
1065 .find(name)
1066 .ok_or_else(|| format!("Step TP projection is missing {name}"))?;
1067 if tensor.ggml_type != GgmlType::BF16 {
1068 return Err(format!(
1069 "Step TP projection {name} must preserve checkpoint BF16 bytes, got {:?}",
1070 tensor.ggml_type
1071 )
1072 .into());
1073 }
1074 if tensor.ne.len() != 2 {
1075 return Err(format!(
1076 "Step TP projection {name} must be a 2-D matrix, got shape {:?}",
1077 tensor.ne
1078 )
1079 .into());
1080 }
1081 let matrix = crate::tp::Bf16Matrix {
1082 bytes: tensor.bytes.as_ref(),
1083 in_features: tensor.ne[0] as usize,
1084 out_features: tensor.ne[1] as usize,
1085 };
1086 matrix.validate()?;
1087 if matrix.in_features != expected_in || matrix.out_features != expected_out {
1088 return Err(format!(
1089 "Step TP projection {name} shape {}x{} != registered {expected_out}x{expected_in}",
1090 matrix.out_features, matrix.in_features
1091 )
1092 .into());
1093 }
1094 Ok(if f32_mirror {
1095 runtime.upload_step_bf16_column_parallel_f32_mirror(matrix)?
1096 } else {
1097 runtime.upload_step_bf16_column_parallel(matrix)?
1098 })
1099}
1100
1101fn upload_step_bf16_row(
1102 runtime: &crate::tp::TpE4m3HostBounce,
1103 src: &dyn TensorSource,
1104 name: &str,
1105 expected_in: usize,
1106 expected_out: usize,
1107 f32_mirror: bool,
1108) -> Result<crate::tp::ResidentStepBf16RowParallel, Box<dyn std::error::Error>> {
1109 let tensor = src
1110 .find(name)
1111 .ok_or_else(|| format!("Step TP projection is missing {name}"))?;
1112 if tensor.ggml_type != GgmlType::BF16 {
1113 return Err(format!(
1114 "Step TP projection {name} must preserve checkpoint BF16 bytes, got {:?}",
1115 tensor.ggml_type
1116 )
1117 .into());
1118 }
1119 if tensor.ne.len() != 2 {
1120 return Err(format!(
1121 "Step TP projection {name} must be a 2-D matrix, got shape {:?}",
1122 tensor.ne
1123 )
1124 .into());
1125 }
1126 let matrix = crate::tp::Bf16Matrix {
1127 bytes: tensor.bytes.as_ref(),
1128 in_features: tensor.ne[0] as usize,
1129 out_features: tensor.ne[1] as usize,
1130 };
1131 matrix.validate()?;
1132 if matrix.in_features != expected_in || matrix.out_features != expected_out {
1133 return Err(format!(
1134 "Step TP projection {name} shape {}x{} != registered {expected_out}x{expected_in}",
1135 matrix.out_features, matrix.in_features
1136 )
1137 .into());
1138 }
1139 Ok(if f32_mirror {
1140 runtime.upload_step_bf16_row_parallel_f32_mirror(matrix)?
1141 } else {
1142 runtime.upload_step_bf16_row_parallel(matrix)?
1143 })
1144}
1145
1146fn upload_step_tp_f32_copies(
1147 runtime: &crate::tp::TpE4m3HostBounce,
1148 src: &dyn TensorSource,
1149 name: &str,
1150 expected: usize,
1151) -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
1152 let tensor = src
1153 .find(name)
1154 .ok_or_else(|| format!("Step TP attention is missing {name}"))?;
1155 let values = memra_gguf::dequant::dequantize(
1156 tensor.ggml_type,
1157 &tensor.bytes,
1158 tensor.ne.iter().product::<u64>() as usize,
1159 );
1160 if values.len() != expected || values.iter().any(|value| !value.is_finite()) {
1161 return Err(format!(
1162 "Step TP attention {name} has {} finite values, expected {expected}",
1163 values.len()
1164 )
1165 .into());
1166 }
1167 let mut copies = Vec::with_capacity(runtime.devices().len());
1168 for rank in 0..runtime.devices().len() {
1169 let engine = runtime
1170 .rank_engine(rank)
1171 .ok_or_else(|| format!("Step TP attention has no engine for rank {rank}"))?;
1172 let _main = engine.gpu.enter_main()?;
1173 copies.push(engine.htod(&values)?);
1174 }
1175 Ok(copies)
1176}
1177
1178fn upload_step_tp_f32_row_shards(
1183 runtime: &crate::tp::TpE4m3HostBounce,
1184 src: &dyn TensorSource,
1185 name: &str,
1186 rows: usize,
1187 cols: usize,
1188) -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
1189 let tensor = src
1190 .find(name)
1191 .ok_or_else(|| format!("Step TP attention is missing {name}"))?;
1192 let values = memra_gguf::dequant::dequantize(
1193 tensor.ggml_type,
1194 &tensor.bytes,
1195 tensor.ne.iter().product::<u64>() as usize,
1196 );
1197 let world = runtime.devices().len();
1198 if values.len() != rows * cols || rows % world != 0 || values.iter().any(|v| !v.is_finite()) {
1199 return Err(format!(
1200 "Step TP attention {name} has {} finite values, expected {rows}x{cols} \
1201 (rows divisible by world {world})",
1202 values.len()
1203 )
1204 .into());
1205 }
1206 let local_rows = rows / world;
1207 let mut shards = Vec::with_capacity(world);
1208 for rank in 0..world {
1209 let engine = runtime
1210 .rank_engine(rank)
1211 .ok_or_else(|| format!("Step TP attention has no engine for rank {rank}"))?;
1212 let _main = engine.gpu.enter_main()?;
1213 shards
1214 .push(engine.htod(&values[rank * local_rows * cols..(rank + 1) * local_rows * cols])?);
1215 }
1216 Ok(shards)
1217}
1218
1219fn upload_step_tp_bf16_row_shards(
1221 runtime: &crate::tp::TpE4m3HostBounce,
1222 src: &dyn TensorSource,
1223 name: &str,
1224 rows: usize,
1225 cols: usize,
1226) -> Result<Vec<CudaSlice<u8>>, Box<dyn std::error::Error>> {
1227 let tensor = src
1228 .find(name)
1229 .ok_or_else(|| format!("Step TP attention is missing {name}"))?;
1230 if tensor.ggml_type != memra_gguf::GgmlType::BF16 || tensor.bytes.len() != rows * cols * 2 {
1231 return Err(format!(
1232 "Step TP attention {name} is not a bf16 [{rows}, {cols}] tensor ({} bytes, {:?})",
1233 tensor.bytes.len(),
1234 tensor.ggml_type
1235 )
1236 .into());
1237 }
1238 let world = runtime.devices().len();
1239 if rows % world != 0 {
1240 return Err(format!("{name} rows {rows} not divisible by world {world}").into());
1241 }
1242 let local = rows / world * cols * 2;
1243 let mut shards = Vec::with_capacity(world);
1244 for rank in 0..world {
1245 let engine = runtime
1246 .rank_engine(rank)
1247 .ok_or_else(|| format!("Step TP attention has no engine for rank {rank}"))?;
1248 let _main = engine.gpu.enter_main()?;
1249 shards.push(engine.htod_bytes(&tensor.bytes[rank * local..(rank + 1) * local])?);
1250 }
1251 Ok(shards)
1252}
1253
1254#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1255enum StepTpAttentionPlacement {
1256 RankLocalGlobal,
1257 RankLocalSwa,
1258 OwnerSwa,
1259 OwnerTransportFallback,
1260}
1261
1262impl StepTpAttentionPlacement {
1263 fn resolve(native_p2p: bool, window: Option<u32>) -> Self {
1264 match (native_p2p, window.is_some()) {
1265 (true, true) => Self::RankLocalSwa,
1266 (false, true) => Self::OwnerSwa,
1267 (true, false) => Self::RankLocalGlobal,
1268 (false, false) => Self::OwnerTransportFallback,
1269 }
1270 }
1271
1272 fn is_rank_local(self) -> bool {
1273 matches!(self, Self::RankLocalGlobal | Self::RankLocalSwa)
1274 }
1275
1276 fn label(self) -> &'static str {
1277 match self {
1278 Self::RankLocalGlobal => "rank-local-global",
1279 Self::RankLocalSwa => "rank-local-swa-ring",
1280 Self::OwnerSwa => "owner-swa",
1281 Self::OwnerTransportFallback => "owner-transport-fallback",
1282 }
1283 }
1284}
1285
1286fn build_step_tp_qkv(
1287 e: &Engine,
1288 src: &dyn TensorSource,
1289 cfg: &ModelConfig,
1290 layer: usize,
1291 step_runtimes: &mut StepParallelRuntimeRegistry,
1292) -> Result<Option<StepTpQkv>, Box<dyn std::error::Error>> {
1293 let Some(spec) = step_runtimes.tp_spec(layer).cloned() else {
1294 return Ok(None);
1295 };
1296 let contract = crate::parallel::ModelParallelContract::from_model(cfg)?;
1297 if layer >= contract.trunk_layers {
1298 return Err(format!(
1299 "MEMRA_STEP_TP layer {layer} is outside Step trunk layers 0..{}",
1300 contract.trunk_layers
1301 )
1302 .into());
1303 }
1304 let owner = e.ctx().ordinal();
1305 if spec.devices.first().copied() != Some(owner) {
1306 return Err(format!(
1307 "MEMRA_STEP_TP layer {layer} owning PP device {owner} must be the first QKV rank, \
1308 got {:?}",
1309 spec.devices
1310 )
1311 .into());
1312 }
1313 let plan = contract.plan(crate::parallel::TopologyRequest {
1314 pipeline: 1,
1315 tensor: spec.devices.len(),
1316 expert_parallel: spec.devices.len() > 2,
1317 available_devices: spec.devices.len(),
1318 hardware: crate::parallel::HardwareTarget::RtxPro6000Blackwell,
1319 })?;
1320 for rank in 0..spec.devices.len() {
1321 plan.query_head_range(layer, rank).ok_or_else(|| {
1322 format!("Step TP layer {layer} has no query-head range for rank {rank}")
1323 })?;
1324 plan.kv_head_range(layer, rank)
1325 .ok_or_else(|| format!("Step TP layer {layer} has no KV-head range for rank {rank}"))?;
1326 }
1327 let native_p2p = step_runtimes.config.native_p2p;
1328 let ep_device_arithmetic = step_runtimes.config.ep_device_arithmetic;
1329 let f32_mirror = step_runtimes.config.f32_mirror;
1330 if ep_device_arithmetic && (!native_p2p || !matches!(spec.devices.len(), 4 | 8)) {
1331 return Err(
1332 "MEMRA_STEP_EP_DEVICE_ARITHMETIC=1 requires a MEMRA_STEP_TP TP4/TP8 \
1333 expert-owner layer and MEMRA_STEP_TP_NATIVE_P2P=1"
1334 .into(),
1335 );
1336 }
1337 let runtime = step_runtimes.runtime(&spec.devices, native_p2p, ep_device_arithmetic)?;
1338 let p = |suffix: &str| format!("blk.{layer}.{suffix}");
1339 let q = upload_step_bf16_column(
1340 &runtime,
1341 src,
1342 &p("attn_q.weight"),
1343 contract.hidden_size,
1344 contract.query_heads[layer] * contract.head_dim,
1345 f32_mirror,
1346 )?;
1347 let k = upload_step_bf16_column(
1348 &runtime,
1349 src,
1350 &p("attn_k.weight"),
1351 contract.hidden_size,
1352 contract.kv_heads[layer] * contract.head_dim,
1353 f32_mirror,
1354 )?;
1355 let v = upload_step_bf16_column(
1356 &runtime,
1357 src,
1358 &p("attn_v.weight"),
1359 contract.hidden_size,
1360 contract.kv_heads[layer] * contract.head_dim,
1361 f32_mirror,
1362 )?;
1363 let o = upload_step_bf16_row(
1364 &runtime,
1365 src,
1366 &p("attn_output.weight"),
1367 contract.query_heads[layer] * contract.head_dim,
1368 contract.hidden_size,
1369 f32_mirror,
1370 )?;
1371 let geometry = cfg.full_attention_geometry_at(layer as u32);
1372 let attention_placement =
1373 StepTpAttentionPlacement::resolve(runtime.native_p2p(), geometry.window);
1374 let attention = if attention_placement.is_rank_local() {
1375 let decode_input = if ep_device_arithmetic || crate::tp::step_tp_decode_v2_enabled()? {
1380 Some(std::sync::Mutex::new(
1381 runtime.allocate_replicated_device_rows(1, contract.hidden_size)?,
1382 ))
1383 } else {
1384 None
1385 };
1386 let gate_fused =
1389 crate::tp::step_tp_qkv_fused_enabled()? && src.find(&p("attn_gate.weight")).is_some();
1390 let gate_shards = if gate_fused && f32_mirror {
1391 Some(upload_step_tp_f32_row_shards(
1392 &runtime,
1393 src,
1394 &p("attn_gate.weight"),
1395 contract.query_heads[layer],
1396 contract.hidden_size,
1397 )?)
1398 } else {
1399 None
1400 };
1401 let gate_shards_bf16 = if gate_fused && !f32_mirror {
1402 Some(upload_step_tp_bf16_row_shards(
1403 &runtime,
1404 src,
1405 &p("attn_gate.weight"),
1406 contract.query_heads[layer],
1407 contract.hidden_size,
1408 )?)
1409 } else {
1410 None
1411 };
1412 Some(StepTpAttention {
1413 q_norm: upload_step_tp_f32_copies(
1414 &runtime,
1415 src,
1416 &p("attn_q_norm.weight"),
1417 contract.head_dim,
1418 )?,
1419 k_norm: upload_step_tp_f32_copies(
1420 &runtime,
1421 src,
1422 &p("attn_k_norm.weight"),
1423 contract.head_dim,
1424 )?,
1425 decode_input,
1426 gate_shards,
1427 gate_shards_bf16,
1428 })
1429 } else {
1430 None
1431 };
1432 if f32_mirror {
1433 eprintln!(
1434 "[step-tp-qkv] load layer={layer} devices={:?} projections=qkv \
1435 qkv_tensor_parallel=true attention_local=true kv_local=true output_local=true \
1436 transport={} native_p2p={} bf16_residency=f32-mirror \
1437 output=root-readback performance_claim=false",
1438 spec.devices,
1439 runtime.transport_label(),
1440 runtime.native_p2p(),
1441 );
1442 } else {
1443 eprintln!(
1444 "[step-tp-qkv] load layer={layer} devices={:?} projections=qkv \
1445 qkv_tensor_parallel=true attention_local=true kv_local=true output_local=true \
1446 transport={} native_p2p={} output=root-readback performance_claim=false",
1447 spec.devices,
1448 runtime.transport_label(),
1449 runtime.native_p2p(),
1450 );
1451 }
1452 eprintln!(
1453 "[step-tp-attn-plan] load layer={layer} devices={:?} \
1454 qkv_tensor_parallel=true attention_tensor_parallel={} kv_cache_distributed={} \
1455 attention_scope={} transport={} native_p2p={} replicated_decode_input_prepared={} \
1456 performance_claim=false",
1457 spec.devices,
1458 attention_placement.is_rank_local(),
1459 attention_placement.is_rank_local(),
1460 attention_placement.label(),
1461 runtime.transport_label(),
1462 runtime.native_p2p(),
1463 attention
1464 .as_ref()
1465 .is_some_and(|attention| attention.decode_input.is_some()),
1466 );
1467 if f32_mirror {
1468 eprintln!(
1469 "[step-tp-o] load layer={layer} devices={:?} projection=o \
1470 o_tensor_parallel=true attention_local=true kv_local=true \
1471 transport={} native_p2p={} reduction=global-tp8-block-order \
1472 bf16_residency=f32-mirror output=root-readback performance_claim=false",
1473 spec.devices,
1474 runtime.transport_label(),
1475 runtime.native_p2p(),
1476 );
1477 } else {
1478 eprintln!(
1479 "[step-tp-o] load layer={layer} devices={:?} projection=o \
1480 o_tensor_parallel=true attention_local=true kv_local=true \
1481 transport={} native_p2p={} reduction=global-tp8-block-order \
1482 output=root-readback performance_claim=false",
1483 spec.devices,
1484 runtime.transport_label(),
1485 runtime.native_p2p(),
1486 );
1487 }
1488 Ok(Some(StepTpQkv {
1489 runtime,
1490 q,
1491 k,
1492 v,
1493 o,
1494 attention,
1495 devices: spec.devices,
1496 layer,
1497 }))
1498}
1499
1500fn build_dev_exps(
1513 e: &Engine,
1514 resident: &mut ResidentPlan,
1515 il: usize,
1516 gate: &HostExps,
1517 up: &HostExps,
1518 down: &HostExps,
1519) -> Result<Option<crate::hybrid::DevExps>, Box<dyn std::error::Error>> {
1520 if !gate.is_uniform_layout() || !up.is_uniform_layout() || !down.is_uniform_layout() {
1523 return Ok(None);
1524 }
1525 let fp8_host = match (&gate.fp8_blk, &up.fp8_blk, &down.fp8_blk) {
1526 (None, None, None) => None,
1527 (Some(g), Some(u), Some(d)) => Some((g, u, d)),
1528 _ => {
1529 return Err("resident expert projections disagree on block-E4M3 scale carriage".into());
1530 }
1531 };
1532 let scale_bytes = fp8_host
1533 .map(|(g, u, d)| (g.scales.len() + u.scales.len() + d.scales.len()) * size_of::<f32>())
1534 .unwrap_or(0);
1535 let per_layer = gate.bytes.as_bytes().len()
1536 + up.bytes.as_bytes().len()
1537 + down.bytes.as_bytes().len()
1538 + scale_bytes;
1539 if gate.tiers.is_some() {
1540 return Ok(None); }
1542 let fits = resident.should_reside(e, il, per_layer);
1543 if !fits {
1544 return Ok(None);
1545 }
1546 use cudarc::driver::DevicePtr;
1547 let gu_il = std::env::var("MEMRA_MOE_GU_IL").as_deref() == Ok("1")
1548 && gate.out_f == up.out_f
1549 && gate.in_f == up.in_f
1550 && fp8_host.is_none();
1551 let n_expert = gate.n_expert;
1552 let (g, u) = if gu_il {
1553 let (rbg, rbu) = (gate.row_bytes, up.row_bytes);
1555 let n_rows = gate.out_f;
1556 let gb = gate.bytes.as_bytes();
1557 let ub = up.bytes.as_bytes();
1558 let mut il = vec![0u8; n_expert * n_rows * (rbg + rbu)];
1559 for ex in 0..n_expert {
1560 for o in 0..n_rows {
1561 let dst = (ex * n_rows + o) * (rbg + rbu);
1562 let sg = ex * gate.expert_stride + o * rbg;
1563 let su = ex * up.expert_stride + o * rbu;
1564 il[dst..dst + rbg].copy_from_slice(&gb[sg..sg + rbg]);
1565 il[dst + rbg..dst + rbg + rbu].copy_from_slice(&ub[su..su + rbu]);
1566 }
1567 }
1568 let ild = e.htod_bytes_padded(&il, 8)?;
1569 (ild, e.htod_bytes(&[0u8; 16])?)
1572 } else {
1573 (
1574 e.htod_bytes_padded(gate.bytes.as_bytes(), 8)?,
1575 e.htod_bytes_padded(up.bytes.as_bytes(), 8)?,
1576 )
1577 };
1578 let d = e.htod_bytes_padded(down.bytes.as_bytes(), 144)?;
1583 let fp8_blk = match fp8_host {
1584 Some((gate, up, down)) => {
1585 if e.fp8_blk_nan_count(&g)? != 0
1586 || e.fp8_blk_nan_count(&u)? != 0
1587 || e.fp8_blk_nan_count(&d)? != 0
1588 {
1589 return Err("native stacked block-E4M3 expert bank contains NaN codes".into());
1590 }
1591 Some(DevExpertFp8BlockScales {
1592 gate: DevExpertFp8ProjectionScales::upload(e, gate, n_expert)?,
1593 up: DevExpertFp8ProjectionScales::upload(e, up, n_expert)?,
1594 down: DevExpertFp8ProjectionScales::upload(e, down, n_expert)?,
1595 })
1596 }
1597 None => None,
1598 };
1599 let mut host = vec![0u64; 3 * n_expert];
1600 let (pg, pu, pd) = {
1601 let __s_e0 = e.stream();
1602 let (pg, _e0) = g.device_ptr(&__s_e0);
1603 let __s_e1 = e.stream();
1604 let (pu, _e1) = u.device_ptr(&__s_e1);
1605 let __s_e2 = e.stream();
1606 let (pd, _e2) = d.device_ptr(&__s_e2);
1607 (pg as u64, pu as u64, pd as u64)
1608 };
1609 for ex in 0..n_expert {
1610 if gu_il {
1611 let stride = gate.out_f * (gate.row_bytes + up.row_bytes);
1612 host[ex] = pg + (ex * stride) as u64;
1613 host[n_expert + ex] = pg + (ex * stride + gate.row_bytes) as u64;
1614 } else {
1615 host[ex] = pg + (ex * gate.expert_stride) as u64;
1616 host[n_expert + ex] = pu + (ex * up.expert_stride) as u64;
1617 }
1618 host[2 * n_expert + ex] = pd + (ex * down.expert_stride) as u64;
1619 }
1620 if gu_il {
1621 eprintln!("[moe] gate/up dev slab INTERLEAVED (MEMRA_MOE_GU_IL)");
1622 }
1623 let ptr_row = e.htod_u64(&host)?;
1624 Ok(Some(crate::hybrid::DevExps {
1625 gate: g,
1626 up: u,
1627 down: d,
1628 ptr_row,
1629 gu_il,
1630 dev: e.ctx().ordinal(),
1631 fp8_blk,
1632 }))
1633}
1634
1635pub struct FullAttnLayer {
1636 pub wq: GpuTensor,
1637 pub wk: GpuTensor,
1638 pub wv: GpuTensor,
1639 pub wo: GpuTensor,
1640 pub q_norm: GpuTensor,
1641 pub k_norm: GpuTensor,
1642 pub attn_gate: Option<GpuTensor>,
1653 pub step_tp_qkv: Option<StepTpQkv>,
1657}
1658
1659pub struct StepTpQkv {
1660 pub runtime: Arc<crate::tp::TpE4m3HostBounce>,
1661 pub q: crate::tp::ResidentBf16ColumnParallel,
1662 pub k: crate::tp::ResidentBf16ColumnParallel,
1663 pub v: crate::tp::ResidentBf16ColumnParallel,
1664 pub o: crate::tp::ResidentStepBf16RowParallel,
1665 pub attention: Option<StepTpAttention>,
1666 pub devices: Vec<usize>,
1667 pub layer: usize,
1668}
1669
1670pub struct StepTpAttention {
1671 pub q_norm: Vec<CudaSlice<f32>>,
1672 pub k_norm: Vec<CudaSlice<f32>>,
1673 pub decode_input: Option<std::sync::Mutex<crate::tp::ResidentReplicatedDeviceRows>>,
1674 pub gate_shards: Option<Vec<CudaSlice<f32>>>,
1677 pub gate_shards_bf16: Option<Vec<CudaSlice<u8>>>,
1679}
1680
1681#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1682pub struct StepTpKvDeviceAdmission {
1683 pub device: usize,
1684 pub bytes: usize,
1685}
1686
1687#[derive(Clone, Copy, Debug)]
1691pub struct MlaGeom {
1692 pub n_head: usize, pub d_nope: usize, pub d_rope: usize, pub d_v: usize, pub kv_rank: usize, pub latent_dim: usize, pub scale: f32, }
1700
1701pub struct MlaAttnLayer {
1705 pub wq_a: GpuTensor, pub q_a_norm: GpuTensor, pub wq_b: GpuTensor, pub wkv_a: GpuTensor, pub kv_a_norm: GpuTensor, pub wk_b: GpuTensor, pub wv_b: GpuTensor, pub wo: GpuTensor, pub geom: MlaGeom,
1715}
1716
1717impl MlaAttnLayer {
1718 pub fn load(
1727 e: &Engine,
1728 src: &dyn TensorSource,
1729 il: u32,
1730 plan: &memra_gguf::model_plan::MlaAttentionPlan,
1731 ) -> Result<Self, Box<dyn std::error::Error>> {
1732 let memra_gguf::model_plan::MlaAttentionPlan::LatentKv {
1733 query_heads,
1734 q_lora_rank,
1735 kv_lora_rank,
1736 qk_head_dim,
1737 rope_head_dim,
1738 value_head_dim,
1739 ..
1740 } = plan
1741 else {
1742 return Err(format!(
1743 "native MLA loader has no compressed-KV implementation for block {il}"
1744 )
1745 .into());
1746 };
1747 let d_nope = qk_head_dim
1748 .checked_sub(*rope_head_dim)
1749 .ok_or("MLA rope head width exceeds total QK head width")?;
1750 let p = |s: &str| format!("blk.{il}.{s}");
1751 let geom = MlaGeom {
1752 n_head: *query_heads as usize,
1753 d_nope: d_nope as usize,
1754 d_rope: *rope_head_dim as usize,
1755 d_v: *value_head_dim as usize,
1756 kv_rank: *kv_lora_rank as usize,
1757 latent_dim: (*kv_lora_rank + *rope_head_dim) as usize,
1758 scale: 1.0 / (*qk_head_dim as f32).sqrt(),
1759 };
1760 let wq_a = load_t(e, src, &p("attn_q_a.weight"))?;
1761 let wq_b = load_t(e, src, &p("attn_q_b.weight"))?;
1762 let wkv_a = load_t(e, src, &p("attn_kv_a_mqa.weight"))?;
1763 let wk_b = load_t(e, src, &p("attn_k_b.weight"))?;
1764 let wv_b = load_t(e, src, &p("attn_v_b.weight"))?;
1765 let wo = load_t(e, src, &p("attn_output.weight"))?;
1766 let n_head = wq_b.out_features() / (geom.d_nope + geom.d_rope);
1768 assert_eq!(
1769 wq_b.out_features(),
1770 n_head * (geom.d_nope + geom.d_rope),
1771 "wq_b out {} not a multiple of qk_head_dim {}",
1772 wq_b.out_features(),
1773 geom.d_nope + geom.d_rope
1774 );
1775 assert_eq!(
1776 wq_a.in_features(),
1777 wkv_a.in_features(),
1778 "q_a/kv_a hidden mismatch"
1779 );
1780 assert_eq!(
1781 wq_b.in_features(),
1782 *q_lora_rank as usize,
1783 "wq_b in != q_lora_rank"
1784 );
1785 assert_eq!(
1786 n_head, geom.n_head,
1787 "MLA checkpoint head count != ModelPlan"
1788 );
1789 assert_eq!(
1790 wkv_a.out_features(),
1791 geom.latent_dim,
1792 "wkv_a out != kv_lora_rank + rope"
1793 );
1794 assert_eq!(
1795 wk_b.ne(),
1796 &[geom.d_nope as u64, geom.kv_rank as u64, n_head as u64],
1797 "attn_k_b must be the TRANSPOSED (nope, kv_rank, head) conversion split"
1798 );
1799 assert_eq!(
1800 wv_b.ne(),
1801 &[geom.kv_rank as u64, geom.d_v as u64, n_head as u64],
1802 "attn_v_b must be the (kv_rank, v, head) conversion split"
1803 );
1804 assert_eq!(
1805 wo.in_features(),
1806 n_head * geom.d_v,
1807 "wo in != n_head * v_head_dim"
1808 );
1809 Ok(MlaAttnLayer {
1810 wq_a,
1811 q_a_norm: load_t(e, src, &p("attn_q_a_norm.weight"))?,
1812 wq_b,
1813 wkv_a,
1814 kv_a_norm: load_t(e, src, &p("attn_kv_a_norm.weight"))?,
1815 wk_b,
1816 wv_b,
1817 wo,
1818 geom,
1819 })
1820 }
1821}
1822
1823#[track_caller]
1827pub(crate) fn mla_forward_unimplemented() -> ! {
1828 panic!(
1829 "Mixer::Mla has no forward arm yet — glm-dsa is loader-only in increment 2; \
1830 the CUDA forward lands in increment 4 (research/mla-bringup-20260801/DESIGN.md §4)"
1831 )
1832}
1833
1834pub struct LinearAttnLayer {
1835 pub geometry: memra_gguf::model_plan::GatedDeltaNetPlan,
1836 pub wqkv: GpuTensor, pub wqkv_gate: GpuTensor, pub ssm_beta: GpuTensor, pub ssm_alpha: GpuTensor, pub ssm_a: GpuTensor, pub ssm_dt: GpuTensor, pub ssm_conv1d: GpuTensor, pub ssm_norm: GpuTensor, pub ssm_out: GpuTensor, }
1846
1847pub enum Mixer {
1848 Full(FullAttnLayer),
1849 Linear(LinearAttnLayer),
1850 Mla(MlaAttnLayer),
1852}
1853
1854pub struct MoeWeights {
1861 pub gate_inp: GpuTensor, pub gate_inp_shexp: Option<GpuTensor>, pub exp_probs_b: Option<Vec<f32>>,
1867 pub exp_probs_b_dev: CudaSlice<f32>,
1868 pub active_experts: Option<Vec<bool>>,
1872 pub active_experts_dev: CudaSlice<u8>,
1873 pub gate_exps: HostExps, pub up_exps: HostExps, pub down_exps: HostExps, pub gate_shexp: Option<GpuTensor>,
1877 pub up_shexp: Option<GpuTensor>,
1878 pub down_shexp: Option<GpuTensor>,
1879 pub dev_exps: Option<DevExps>,
1886 pub step_ep: Option<StepEpExps>,
1890 pub step_tp: Option<StepTpExps>,
1894 pub dev_macros: cudarc::driver::CudaSlice<f32>,
1900 pub has_macros: bool,
1901}
1902
1903pub enum StepEpExpertBank {
1905 E4m3(crate::tp::ResidentExpertParallel),
1906 Nvfp4(crate::tp::ResidentNvfp4ExpertParallel),
1907}
1908
1909impl StepEpExpertBank {
1910 pub fn e4m3(&self) -> Result<&crate::tp::ResidentExpertParallel, String> {
1914 match self {
1915 Self::E4m3(bank) => Ok(bank),
1916 Self::Nvfp4(_) => Err(
1917 "Step grouped expert program reached an NVFP4 bank; this path is qualified \
1918 for the E4M3 artifact only"
1919 .to_string(),
1920 ),
1921 }
1922 }
1923}
1924
1925pub struct StepEpExps {
1926 pub runtime: Arc<crate::tp::TpE4m3HostBounce>,
1927 pub experts: StepEpExpertBank,
1928 pub devices: Vec<usize>,
1929 pub configured_by_tp: bool,
1930 pub activation_limit: Option<f32>,
1931 pub grouped_decode: Option<std::sync::Mutex<StepEpGroupedDecode>>,
1934}
1935
1936pub struct StepEpGroupedDecode {
1937 pub(crate) projection: crate::tp::PreparedStepGroupedExpertParallelGate,
1938 pub(crate) combine: crate::tp::PreparedPeerWeightedRouteCombine,
1939}
1940
1941#[derive(Default)]
1942pub(crate) struct StepEpGroupedPrefill {
1943 pub(crate) state: Option<StepEpGroupedPrefillState>,
1944}
1945
1946pub(crate) struct StepEpGroupedPrefillState {
1947 pub(crate) devices: Vec<usize>,
1948 pub(crate) grouped: StepEpGroupedDecode,
1949}
1950
1951pub enum StepTpExpertBank {
1953 E4m3(crate::tp::ResidentTensorParallel),
1954 Nvfp4(crate::tp::ResidentNvfp4TensorParallel),
1955}
1956
1957pub struct StepTpExps {
1958 pub runtime: Arc<crate::tp::TpE4m3HostBounce>,
1959 pub experts: StepTpExpertBank,
1960 pub devices: Vec<usize>,
1961 pub activation_limit: Option<f32>,
1964}
1965
1966impl MoeWeights {
1967 #[inline]
1968 pub fn has_uniform_expert_layout(&self) -> bool {
1969 self.gate_exps.is_uniform_layout()
1970 && self.up_exps.is_uniform_layout()
1971 && self.down_exps.is_uniform_layout()
1972 }
1973
1974 #[inline]
1975 pub fn active_count(&self) -> usize {
1976 self.active_experts
1977 .as_ref()
1978 .map(|mask| mask.iter().filter(|&&active| active).count())
1979 .unwrap_or(self.gate_exps.n_expert)
1980 }
1981}
1982
1983pub struct DevExps {
1986 pub gate: CudaSlice<u8>,
1987 pub up: CudaSlice<u8>,
1988 pub down: CudaSlice<u8>,
1989 pub ptr_row: CudaSlice<u64>,
1991 pub dev: usize,
1999 pub gu_il: bool,
2005 pub fp8_blk: Option<DevExpertFp8BlockScales>,
2009}
2010
2011pub struct DevExpertFp8BlockScales {
2012 pub gate: DevExpertFp8ProjectionScales,
2013 pub up: DevExpertFp8ProjectionScales,
2014 pub down: DevExpertFp8ProjectionScales,
2015}
2016
2017pub struct DevExpertFp8ProjectionScales {
2018 pub scales: CudaSlice<f32>,
2019 pub rows: usize,
2020 pub cols: usize,
2021 pub expert_stride: usize,
2022}
2023
2024impl DevExpertFp8ProjectionScales {
2025 fn validate(
2026 host: &crate::model::HostExpertFp8BlockScales,
2027 n_expert: usize,
2028 ) -> Result<(), String> {
2029 if host.expert_stride == 0 {
2030 return Err("block-E4M3 expert scale stride must be nonzero".into());
2031 }
2032 if host.rows * host.cols != host.expert_stride {
2033 return Err(format!(
2034 "block-E4M3 expert scale stride mismatch: {}x{} != {}",
2035 host.rows, host.cols, host.expert_stride
2036 ));
2037 }
2038 let want = n_expert
2039 .checked_mul(host.expert_stride)
2040 .ok_or("block-E4M3 expert scale slab length overflow")?;
2041 if host.scales.len() != want {
2042 return Err(format!(
2043 "block-E4M3 scale slab length mismatch: got {}, want {n_expert}x{}={want}",
2044 host.scales.len(),
2045 host.expert_stride
2046 ));
2047 }
2048 Ok(())
2049 }
2050
2051 fn upload(
2052 e: &Engine,
2053 host: &crate::model::HostExpertFp8BlockScales,
2054 n_expert: usize,
2055 ) -> Result<Self, Box<dyn std::error::Error>> {
2056 Self::validate(host, n_expert)?;
2057 Ok(Self {
2058 scales: e.htod(&host.scales)?,
2059 rows: host.rows,
2060 cols: host.cols,
2061 expert_stride: host.expert_stride,
2062 })
2063 }
2064}
2065
2066pub enum Ffn {
2068 Dense {
2069 ffn_gate: GpuTensor,
2070 ffn_up: GpuTensor,
2071 ffn_down: GpuTensor,
2072 },
2073 Moe(MoeWeights),
2074}
2075
2076pub struct HybridLayer {
2077 pub attn_norm: GpuTensor,
2078 pub post_attn_norm: GpuTensor, pub mixer: Mixer,
2080 pub ffn: Ffn,
2081 pub gemma4: Option<Gemma4LayerBits>,
2082}
2083
2084pub struct Gemma4LayerBits {
2088 pub ffn_norm: GpuTensor, pub post_ffw_norm: GpuTensor, pub moe_bits: Option<Gemma4MoeBits>,
2093 pub layer_scale: f32, pub e4b: Option<Gemma4E4bLayer>,
2096}
2097
2098pub struct Gemma4E4bLayer {
2103 pub inp_gate: GpuTensor, pub proj: GpuTensor, pub post_norm: GpuTensor, pub qkv_cat: Option<GpuTensor>,
2110 pub kv_share: Option<u32>,
2114}
2115
2116pub struct Gemma4E4bModel {
2120 pub tok_tbl_gpu: std::sync::OnceLock<CudaSlice<u8>>,
2123 pub tok_embd_bytes: Vec<u8>,
2124 pub tok_embd_qt: i32,
2125 pub tok_embd_row_bytes: usize,
2126 pub model_proj: GpuTensor, pub proj_norm: GpuTensor, pub n_epl: usize,
2129}
2130
2131pub struct Gemma4MoeBits {
2132 pub post_ffw_norm_1: GpuTensor, pub pre_ffw_norm_2: GpuTensor, pub post_ffw_norm_2: GpuTensor, pub shared_gate: GpuTensor,
2136 pub shared_up: GpuTensor,
2137 pub shared_down: GpuTensor,
2138 pub router_scale_pre: CudaSlice<f32>,
2143 pub per_expert_scale: Vec<f32>, pub per_expert_scale_d: CudaSlice<f32>, }
2146
2147pub struct MtpHead {
2152 pub enorm: GpuTensor, pub hnorm: GpuTensor, pub eh_proj: GpuTensor, pub attn_norm: GpuTensor, pub post_attn_norm: GpuTensor, pub mixer: Mixer, pub ffn: Ffn, pub shared_head_norm: Option<GpuTensor>, pub shared_head_head: Option<GpuTensor>, pub d2t: Option<Vec<u32>>,
2166 pub d2t_from_target_head: bool,
2170 pub geom: Option<DraftGeom>,
2176 pub step35: Option<Step35MtpGeom>,
2181}
2182
2183#[derive(Debug, Clone)]
2197pub struct Step35MtpGeom {
2198 pub il: u32,
2200 pub n_head: usize, pub n_head_kv: usize, pub n_rot: usize, pub rope_base: f32, pub swa: bool, pub window: usize, pub clamp_shexp: Option<f32>,
2211}
2212
2213impl Step35MtpGeom {
2214 pub fn from_plan(layer: &memra_gguf::model_plan::LayerPlan) -> Result<Self, String> {
2216 use memra_gguf::model_plan::{ActivationPlan, AttentionPlan};
2217
2218 let (attention, window) = match &layer.attention {
2219 AttentionPlan::Full(attention) => (attention, None),
2220 AttentionPlan::SlidingWindow { attention, window } => (attention, Some(*window)),
2221 other => {
2222 return Err(format!(
2223 "MTP block {} has unsupported tuned attention {other:?}",
2224 layer.index
2225 ));
2226 }
2227 };
2228 if attention.output_gate != memra_gguf::config::AttentionGateKind::SeparateHead {
2229 return Err(format!(
2230 "MTP block {} does not declare a separate attention gate",
2231 layer.index
2232 ));
2233 }
2234 let activation = match &layer.mlp {
2235 MlpPlan::Dense(dense) => &dense.activation,
2236 MlpPlan::Moe(moe) => &moe.activation,
2237 };
2238 let clamp_shexp = match activation {
2239 ActivationPlan::SwiGluClamped { limit } if *limit > 0.0 => Some(*limit),
2240 _ => None,
2241 };
2242 Ok(Step35MtpGeom {
2243 il: layer.index,
2244 n_head: attention.query_heads as usize,
2245 n_head_kv: attention.kv_heads as usize,
2246 n_rot: attention.rope.dimensions as usize,
2247 rope_base: attention.rope.base,
2248 swa: window.is_some(),
2249 window: window.unwrap_or(0) as usize,
2250 clamp_shexp,
2251 })
2252 }
2253}
2254
2255pub struct DraftGeom {
2257 pub d_inner: usize, pub n_head: usize, pub n_head_kv: usize,
2260 pub out_up: GpuTensor, }
2262
2263pub fn draft_head_tensor(has: impl Fn(&str) -> bool, n: u32) -> String {
2272 let own = format!("blk.{n}.nextn.shared_head_head.weight");
2273 if has(&own) {
2274 return own;
2275 }
2276 let legacy = format!("blk.{n}.nextn.shared_head.weight");
2279 if has(&legacy) {
2280 return legacy;
2281 }
2282 "output.weight".to_string()
2284}
2285
2286impl MtpHead {
2287 pub fn load_draft(
2294 e: &Engine,
2295 g: &GgufFile,
2296 main_cfg: &ModelConfig,
2297 ) -> Result<Self, Box<dyn std::error::Error>> {
2298 let src = GgufSource(g);
2299 let dcfg = src.config();
2300 let draft_plan = match memra_gguf::model_packs::for_config(&dcfg) {
2301 Some(pack) => pack.compile_plan(&dcfg)?,
2302 None => memra_gguf::model_plan::ModelPlan::compile(&dcfg)?,
2303 };
2304 let main_plan = match memra_gguf::model_packs::for_config(main_cfg) {
2305 Some(pack) => pack.compile_plan(main_cfg)?,
2306 None => memra_gguf::model_plan::ModelPlan::compile(main_cfg)?,
2307 };
2308 if dcfg.nextn_predict_layers == 0 {
2313 return Err(format!(
2314 "draft GGUF has no nextn_predict_layers (arch {:?}) — not a NextN/MTP regime \
2315 draft; gemma assistant drafters attach via MEMRA_DRAFT, not '+draft'",
2316 g.arch()
2317 )
2318 .into());
2319 }
2320 let n = dcfg.n_layer - dcfg.nextn_predict_layers;
2321 let draft_block = draft_plan
2322 .mtp_blocks
2323 .iter()
2324 .find(|block| block.layer.index == n)
2325 .ok_or_else(|| format!("draft ModelPlan has no MTP block {n}"))?;
2326 let p = |s: &str| format!("blk.{n}.{s}");
2327
2328 let student = src.has(&p("nextn.out_up.weight"));
2332 assert_eq!(dcfg.n_embd, main_cfg.n_embd, "draft n_embd != model n_embd");
2333 assert_eq!(
2334 dcfg.head_dim_k, main_cfg.head_dim_k,
2335 "draft head_dim != model head_dim"
2336 );
2337 let main_sliding_gated = crate::plan_backend::decode_batch_program(&main_plan)
2344 == crate::plan_backend::DecodeBatchProgram::SlidingGatedMoe;
2345 let draft_sliding_gated = crate::plan_backend::decode_batch_program(&draft_plan)
2346 == crate::plan_backend::DecodeBatchProgram::SlidingGatedMoe;
2347 let step35 = match (main_sliding_gated, draft_sliding_gated) {
2348 (true, true) => {
2349 let g = Step35MtpGeom::from_plan(&draft_block.layer)?;
2350 let out_f = |t: &str| -> Option<usize> {
2352 src.find(&p(t))
2353 .and_then(|v| v.ne.get(1).copied())
2354 .map(|x| x as usize)
2355 };
2356 let hd = dcfg.head_dim_k as usize;
2357 let wq_out =
2358 out_f("attn_q.weight").ok_or("step35 draft block has no attn_q.weight")?;
2359 assert_eq!(
2360 wq_out,
2361 g.n_head * hd,
2362 "step35 draft blk.{n}: attn_q out {wq_out} != n_head({}) * head_dim({hd}) — \
2363 the draft file's head_count array disagrees with its own tensors",
2364 g.n_head
2365 );
2366 let wg_out = out_f("attn_gate.weight")
2369 .ok_or("step35 draft block has no attn_gate.weight (head-wise gate)")?;
2370 assert_eq!(
2371 wg_out, g.n_head,
2372 "step35 draft blk.{n}: attn_gate out {wg_out} != n_head({})",
2373 g.n_head
2374 );
2375 assert_eq!(
2379 g.n_head_kv, main_cfg.n_head_kv as usize,
2380 "step35 draft blk.{n} KV heads {} != trunk n_head_kv {} — the MTP scratch \
2381 rows are sized from the trunk cfg, so a differing draft KV width would \
2382 write past the row",
2383 g.n_head_kv, main_cfg.n_head_kv
2384 );
2385 eprintln!(
2386 "[mtp-draft] step35 MTP geometry blk.{n}: n_head={} n_head_kv={} n_rot={} \
2387 rope_base={:.0} swa={} window={}",
2388 g.n_head, g.n_head_kv, g.n_rot, g.rope_base, g.swa, g.window
2389 );
2390 Some(g)
2391 }
2392 (true, false) => {
2393 return Err(format!(
2394 "MEMRA_MTP_DRAFT operations are incompatible with the model's \
2395 sliding-gated-MoE program (draft arch {:?})",
2396 g.arch()
2397 )
2398 .into());
2399 }
2400 (false, true) => {
2401 return Err(
2402 "MEMRA_MTP_DRAFT requires sliding-gated-MoE operations but the model does not"
2403 .into(),
2404 );
2405 }
2406 (false, false) => None,
2407 };
2408 if step35.is_none() && !student {
2409 assert_eq!(dcfg.n_head, main_cfg.n_head, "draft n_head != model n_head");
2412 assert_eq!(
2413 dcfg.n_head_kv, main_cfg.n_head_kv,
2414 "draft n_head_kv != model n_head_kv"
2415 );
2416 }
2417
2418 let head_name = draft_head_tensor(|t| src.has(t), n);
2445 let head = load_t(e, &src, &head_name)?;
2446 let head_norm = match load_opt(e, &src, &p("nextn.shared_head_norm.weight"))? {
2447 Some(t) => Some(t),
2448 None => load_opt(e, &src, "output_norm.weight")?,
2449 };
2450
2451 let d2t: Option<Vec<u32>> = g.find("d2t").map(|t| {
2453 let bytes = g.tensor_data(t);
2454 match t.ggml_type {
2455 GgmlType::I32 => bytes
2456 .chunks_exact(4)
2457 .map(|c| i32::from_le_bytes(c.try_into().unwrap()) as u32)
2458 .collect(),
2459 GgmlType::I64 => bytes
2460 .chunks_exact(8)
2461 .map(|c| i64::from_le_bytes(c.try_into().unwrap()) as u32)
2462 .collect(),
2463 other => panic!("d2t must be I32/I64, got {other:?}"),
2464 }
2465 });
2466 if let Some(map) = &d2t {
2467 assert_eq!(
2468 map.len(),
2469 head.out_features(),
2470 "d2t len {} != draft head rows {}",
2471 map.len(),
2472 head.out_features()
2473 );
2474 let n_vocab = main_cfg.n_vocab as u64;
2475 assert!(
2476 map.iter().all(|&t| (t as u64) < n_vocab),
2477 "d2t contains token id >= model n_vocab {n_vocab}"
2478 );
2479 }
2480 let eh_proj = load_t(e, &src, &p("nextn.eh_proj.weight"))?;
2481 assert_eq!(
2484 eh_proj.in_features(),
2485 2 * main_cfg.n_embd as usize,
2486 "eh_proj in dim != 2*n_embd"
2487 );
2488 let geom = if student {
2489 let out_up = load_t(e, &src, &p("nextn.out_up.weight"))?;
2490 let d_inner = eh_proj.out_features();
2491 assert_eq!(
2492 out_up.out_features(),
2493 main_cfg.n_embd as usize,
2494 "out_up out dim != n_embd"
2495 );
2496 assert_eq!(
2497 out_up.in_features(),
2498 d_inner,
2499 "out_up in dim != eh_proj out dim (d_inner)"
2500 );
2501 assert!(
2502 dcfg.n_head >= 1 && dcfg.n_head_kv >= 1 && dcfg.n_head % dcfg.n_head_kv == 0,
2503 "student head counts malformed ({}/{})",
2504 dcfg.n_head,
2505 dcfg.n_head_kv
2506 );
2507 Some(DraftGeom {
2508 d_inner,
2509 n_head: dcfg.n_head as usize,
2510 n_head_kv: dcfg.n_head_kv as usize,
2511 out_up,
2512 })
2513 } else {
2514 None
2515 };
2516 let blk_prefix = format!("blk.{n}.");
2520 let head_src = head_name.strip_prefix(&blk_prefix).unwrap_or(&head_name);
2521 eprintln!(
2522 "[mtp-draft] external draft head: blk.{n}, source={}, head_vocab={}{}{}",
2523 head_src,
2524 head.out_features(),
2525 if d2t.is_some() {
2526 " (trimmed, d2t map)"
2527 } else {
2528 " (full)"
2529 },
2530 match &geom {
2531 Some(g) => format!(
2532 " (student d_inner={} heads={}/{})",
2533 g.d_inner, g.n_head, g.n_head_kv
2534 ),
2535 None => String::new(),
2536 }
2537 );
2538
2539 let mut resident = ResidentPlan::unsharded(e, &src, &dcfg);
2540 let mut step_runtimes = StepParallelRuntimeRegistry::default();
2541 Ok(MtpHead {
2542 enorm: load_t(e, &src, &p("nextn.enorm.weight"))?,
2543 hnorm: load_t(e, &src, &p("nextn.hnorm.weight"))?,
2544 eh_proj,
2545 attn_norm: load_t(e, &src, &p("attn_norm.weight"))?,
2546 post_attn_norm: load_opt(e, &src, &p("post_attention_norm.weight"))?
2547 .or(load_opt(e, &src, &p("ffn_norm.weight"))?)
2548 .expect("draft NextN block needs post_attention_norm or ffn_norm"),
2549 mixer: load_mixer_kind(
2550 e,
2551 &src,
2552 &dcfg,
2553 n,
2554 &draft_block.layer.attention,
2555 &mut step_runtimes,
2556 )?,
2557 ffn: load_ffn(
2558 e,
2559 &src,
2560 &dcfg,
2561 &draft_block.layer.mlp,
2562 n,
2563 None,
2564 &mut resident,
2565 &mut step_runtimes,
2566 )?,
2567 shared_head_norm: head_norm,
2568 shared_head_head: Some(head),
2569 d2t,
2570 d2t_from_target_head: false,
2571 geom,
2572 step35,
2573 })
2574 }
2575}
2576
2577pub struct GemmaAux {
2579 pub rope_freqs: Option<Vec<(usize, CudaSlice<f32>)>>,
2582 pub ones: Vec<(usize, CudaSlice<f32>)>,
2585 pub suppress_d: Option<(CudaSlice<i32>, usize)>,
2588 pub e4b: Option<Gemma4E4bModel>,
2590}
2591
2592impl GemmaAux {
2593 pub fn rope_freqs(&self, e: &Engine) -> Option<&CudaSlice<f32>> {
2594 self.rope_freqs.as_ref().map(|copies| {
2595 let dev = e.ctx().ordinal();
2596 &copies
2597 .iter()
2598 .find(|(d, _)| *d == dev)
2599 .unwrap_or_else(|| panic!("gemma4 rope_freqs has no local copy for device {dev}"))
2600 .1
2601 })
2602 }
2603
2604 pub fn ones(&self, e: &Engine) -> &CudaSlice<f32> {
2605 let dev = e.ctx().ordinal();
2606 &self
2607 .ones
2608 .iter()
2609 .find(|(d, _)| *d == dev)
2610 .unwrap_or_else(|| panic!("gemma4 ones has no local copy for device {dev}"))
2611 .1
2612 }
2613}
2614
2615pub struct Step35Aux {
2618 pub rope_freqs: Option<Vec<(usize, CudaSlice<f32>)>>,
2624}
2625
2626impl Step35Aux {
2627 pub fn rope_freqs(&self, e: &Engine) -> Option<&CudaSlice<f32>> {
2628 self.rope_freqs.as_ref().map(|copies| {
2629 let dev = e.ctx().ordinal();
2630 &copies
2631 .iter()
2632 .find(|(d, _)| *d == dev)
2633 .unwrap_or_else(|| panic!("step35 rope_freqs has no local copy for device {dev}"))
2634 .1
2635 })
2636 }
2637}
2638
2639pub struct HybridModel {
2640 pub cfg: ModelConfig,
2641 pub plan: memra_gguf::model_plan::ModelPlan,
2642 pub rewrite_qualifications: Option<memra_gguf::execution_manifest::RewriteQualifications>,
2643 pub embd: EmbedHost,
2644 pub output_norm: GpuTensor,
2645 pub output: GpuTensor,
2646 pub layers: Vec<HybridLayer>,
2647 pub mtp: Option<MtpHead>, pub mtp_extra: Vec<MtpHead>,
2651 pub embd_gpu: std::sync::OnceLock<cudarc::driver::CudaSlice<u8>>,
2654 pub gemma4_aux: Option<GemmaAux>,
2655 pub step35_aux: Option<Step35Aux>,
2657 pub prime_slabs: std::sync::Mutex<
2665 std::collections::HashMap<
2666 usize,
2667 std::sync::Arc<std::sync::Mutex<crate::hybrid_forward::PrimeSlabs>>,
2668 >,
2669 >,
2670 pub(crate) dspark_vgraphs: std::sync::Mutex<Option<crate::spec::DsparkVerifyGraphs>>,
2683 pub(crate) step_grouped_prefill: std::sync::Mutex<StepEpGroupedPrefill>,
2689 pub(crate) step35_token_graph:
2692 std::sync::Mutex<Option<crate::hybrid_forward::Step35TokenGraphState>>,
2693}
2694
2695impl HybridModel {
2696 pub fn install_rewrite_bundle(
2697 &mut self,
2698 bundle: &std::path::Path,
2699 ) -> Result<(), Box<dyn std::error::Error>> {
2700 self.rewrite_qualifications = Some(
2701 memra_gguf::execution_manifest::RewriteQualifications::load(bundle, &self.plan)
2702 .map_err(|error| format!("rewrite qualification: {error}"))?,
2703 );
2704 Ok(())
2705 }
2706
2707 pub fn rewrite_allowed(&self, surface: memra_gguf::execution_manifest::RewriteSurface) -> bool {
2708 self.rewrite_qualifications
2709 .as_ref()
2710 .is_none_or(|qualifications| qualifications.allows(surface))
2711 }
2712
2713 pub fn step_tp_unmaterialized_kv_bytes(
2719 &self,
2720 cache: Option<&crate::cache::Cache>,
2721 capacity: usize,
2722 ) -> Result<Vec<StepTpKvDeviceAdmission>, String> {
2723 if let Some(cache) = cache
2724 && cache.tp_kv.len() < self.layers.len()
2725 {
2726 return Err(format!(
2727 "Step TP admission cache has {} layers, model trunk has {}",
2728 cache.tp_kv.len(),
2729 self.layers.len()
2730 ));
2731 }
2732
2733 let mut by_device: HashMap<usize, usize> = HashMap::new();
2734 for (layer, weights) in self.layers.iter().enumerate() {
2735 let Mixer::Full(attention) = &weights.mixer else {
2736 continue;
2737 };
2738 let Some(tp) = attention
2739 .step_tp_qkv
2740 .as_ref()
2741 .filter(|tp| tp.attention.is_some())
2742 else {
2743 continue;
2744 };
2745 if cache.is_some_and(|cache| cache.tp_kv[layer].is_some()) {
2746 continue;
2747 }
2748 let geometry = self.cfg.full_attention_geometry_at(layer as u32);
2749 let shape = crate::cache::tp_kv_rank_allocation_shape(
2750 geometry.n_head_kv as usize * geometry.head_dim_k as usize,
2751 geometry.n_head_kv as usize * geometry.head_dim_v as usize,
2752 tp.devices.len(),
2753 )?;
2754 let physical_rows = geometry
2755 .window
2756 .map(|window| crate::cache::swa_ring_rows(window as usize, capacity))
2757 .unwrap_or(capacity);
2758 let bytes = shape.allocation_bytes(physical_rows);
2759 for &device in &tp.devices {
2760 let total = by_device.entry(device).or_default();
2761 *total = total.saturating_add(bytes);
2762 }
2763 }
2764
2765 let mut out: Vec<_> = by_device
2766 .into_iter()
2767 .map(|(device, bytes)| StepTpKvDeviceAdmission { device, bytes })
2768 .collect();
2769 out.sort_unstable_by_key(|charge| charge.device);
2770 Ok(out)
2771 }
2772
2773 pub fn step_tp_rank_engine(&self, device: usize) -> Option<&Engine> {
2775 self.layers.iter().find_map(|weights| {
2776 let Mixer::Full(attention) = &weights.mixer else {
2777 return None;
2778 };
2779 let tp = attention.step_tp_qkv.as_ref()?;
2780 let rank = tp
2781 .runtime
2782 .devices()
2783 .iter()
2784 .position(|&rank| rank == device)?;
2785 tp.runtime.rank_engine(rank)
2786 })
2787 }
2788
2789 pub(crate) fn step_tp_runtime_for_layer(
2790 &self,
2791 layer: usize,
2792 ) -> Option<&crate::tp::TpE4m3HostBounce> {
2793 let Mixer::Full(attention) = &self.layers.get(layer)?.mixer else {
2794 return None;
2795 };
2796 let tp = attention.step_tp_qkv.as_ref()?;
2797 tp.attention.as_ref()?;
2798 Some(tp.runtime.as_ref())
2799 }
2800
2801 pub fn decode_batch_program(&self) -> crate::plan_backend::DecodeBatchProgram {
2802 crate::plan_backend::decode_batch_program(&self.plan)
2803 }
2804
2805 pub fn uses_gemma_program(&self) -> bool {
2806 self.decode_batch_program() == crate::plan_backend::DecodeBatchProgram::Gemma
2807 }
2808
2809 pub fn uses_sliding_gated_moe_program(&self) -> bool {
2810 self.decode_batch_program() == crate::plan_backend::DecodeBatchProgram::SlidingGatedMoe
2811 }
2812
2813 pub fn has_plan_operation(&self, operation: memra_gguf::model_plan::OperationKind) -> bool {
2814 self.plan.trunk_operations().contains(&operation)
2815 }
2816
2817 pub fn load(e: &Engine, g: &GgufFile) -> Result<Self, Box<dyn std::error::Error>> {
2819 Self::load_from_source(e, &GgufSource(g))
2820 }
2821
2822 pub fn load_without_mtp(e: &Engine, g: &GgufFile) -> Result<Self, Box<dyn std::error::Error>> {
2825 Self::load_from_source_impl(e, &GgufSource(g), false)
2826 }
2827
2828 pub fn load_from_source(
2832 e: &Engine,
2833 src: &dyn TensorSource,
2834 ) -> Result<Self, Box<dyn std::error::Error>> {
2835 Self::load_from_source_impl(e, src, true)
2836 }
2837
2838 pub fn load_from_source_without_mtp(
2840 e: &Engine,
2841 src: &dyn TensorSource,
2842 ) -> Result<Self, Box<dyn std::error::Error>> {
2843 Self::load_from_source_impl(e, src, false)
2844 }
2845
2846 fn load_from_source_impl(
2847 e: &Engine,
2848 src: &dyn TensorSource,
2849 load_mtp: bool,
2850 ) -> Result<Self, Box<dyn std::error::Error>> {
2851 let cfg = src.config();
2852 let plan = match memra_gguf::model_packs::for_config(&cfg) {
2853 Some(pack) => pack.compile_plan(&cfg)?,
2854 None => memra_gguf::model_plan::ModelPlan::compile(&cfg)?,
2855 };
2856 let batch_program = crate::plan_backend::decode_batch_program(&plan);
2857 let gemma_program = batch_program == crate::plan_backend::DecodeBatchProgram::Gemma;
2858 let sliding_gated_moe_program =
2859 batch_program == crate::plan_backend::DecodeBatchProgram::SlidingGatedMoe;
2860 cfg.validate_attention_gate_layout()?;
2865 if cfg.sigmoid_router().is_some() {
2872 let host_oracle = std::env::var("MEMRA_SIG_ROUTER").as_deref() == Ok("0");
2873 match crate::sigrouter_contract::verify_host_expf() {
2874 Ok(()) => {}
2875 Err(e) if host_oracle => return Err(e.into()),
2876 Err(e) => eprintln!(
2877 "[sigrouter] WARN: host expf probe mismatch ({e}); device routing is \
2878 unaffected, but host-oracle replay/comparison cells are invalid on this host"
2879 ),
2880 }
2881 }
2882 if std::env::var("MEMRA_DRAFT").is_ok() && std::env::var("MEMRA_MMQ_SK").is_err() {
2891 let force = if cfg.n_embd >= 3500 { 0i8 } else { -1i8 };
2892 crate::MMQ_SK_FORCE.store(force, std::sync::atomic::Ordering::Relaxed);
2893 }
2894 crate::KV_FP8_FORCE.store(0, std::sync::atomic::Ordering::Relaxed);
2898
2899 let n_trunk = (cfg.n_layer - cfg.nextn_predict_layers) as usize;
2904 crate::pp::init_model_transport(e, &cfg, n_trunk)?;
2905 let step_parallel = prepare_step_parallel_load(e, src, &cfg, n_trunk)?;
2906 let embd = EmbedHost::from_source(src, "token_embd.weight");
2907 let e_head = crate::pp::layer_engine(e, n_trunk, n_trunk - 1)?;
2911 let output_norm = load_t(e_head, src, "output_norm.weight")?;
2912 let mut output = if src.has("output.weight") {
2914 load_t(e_head, src, "output.weight")?
2915 } else {
2916 load_t(e_head, src, "token_embd.weight")?
2917 };
2918 let mut resident = ResidentPlan::pp(e, src, &cfg, n_trunk)?;
2919 let mut step_runtimes = StepParallelRuntimeRegistry::with_config(step_parallel);
2920
2921 let gguf: Option<&GgufFile> = src.gguf();
2928 let mut spill: Option<crate::spill::SpillCtx> = if cfg
2931 .moe
2932 .as_ref()
2933 .is_some_and(|m| m.expert_count > 0)
2934 && crate::spill::disk_tier_enabled()
2935 && gguf.is_some()
2936 {
2937 let budget = crate::spill::MemBudget::probe(e)?;
2938 let ctx = crate::spill::SpillCtx::open(gguf.unwrap(), &budget)?;
2939 eprintln!(
2940 "[spill] disk tier ON: free_vram={} MiB free_pinnable_ram={} MiB (MemAvailable*resolved_frac)",
2941 budget.free_vram >> 20,
2942 budget.free_pinnable_ram >> 20
2943 );
2944 Some(ctx)
2945 } else {
2946 None
2947 };
2948
2949 let mut layers = Vec::with_capacity(n_trunk);
2952 for il in 0..n_trunk as u32 {
2953 let p = |s: &str| format!("blk.{il}.{s}");
2954 let layer_plan = plan
2955 .layers
2956 .get(il as usize)
2957 .ok_or_else(|| format!("ModelPlan has no trunk layer {il}"))?;
2958 let e = crate::pp::layer_engine(e, n_trunk, il as usize)?;
2962 layers.push(HybridLayer {
2964 attn_norm: load_t(e, src, &p("attn_norm.weight"))?,
2965 post_attn_norm: load_opt(e, src, &p("post_attention_norm.weight"))?
2966 .or(load_opt(e, src, &p("ffn_norm.weight"))?)
2967 .expect("need post_attention_norm or ffn_norm"),
2968 mixer: {
2969 let g4_shared = cfg.gemma4.as_ref().map(|g| g.shared_kv_layers).unwrap_or(0);
2973 let kv_from = n_trunk as u32 - g4_shared;
2974 if g4_shared > 0
2975 && il >= kv_from
2976 && !src.has(&format!("blk.{il}.attn_k.weight"))
2977 {
2978 let g4 = cfg.gemma4.as_ref().unwrap();
2979 let swa = g4.swa_pattern.get(il as usize).copied().unwrap_or(true);
2980 let tgt = kv_from - if swa { 2 } else { 1 };
2981 let tp = |s: &str| format!("blk.{tgt}.{s}");
2982 Mixer::Full(FullAttnLayer {
2983 wq: load_t(e, src, &p("attn_q.weight"))?,
2984 wk: load_t(e, src, &tp("attn_k.weight"))?,
2985 wv: load_t(e, src, &tp("attn_v.weight"))?,
2986 wo: load_t(e, src, &p("attn_output.weight"))?,
2987 q_norm: load_t(e, src, &p("attn_q_norm.weight"))?,
2988 k_norm: load_t(e, src, &tp("attn_k_norm.weight"))?,
2989 attn_gate: None, step_tp_qkv: None,
2991 })
2992 } else {
2993 load_mixer_kind(
2994 e,
2995 src,
2996 &cfg,
2997 il,
2998 &layer_plan.attention,
2999 &mut step_runtimes,
3000 )?
3001 }
3002 },
3003 ffn: load_ffn(
3004 e,
3005 src,
3006 &cfg,
3007 &layer_plan.mlp,
3008 il,
3009 spill.as_mut().map(|c| (gguf.unwrap(), c)),
3010 &mut resident,
3011 &mut step_runtimes,
3012 )?,
3013 gemma4: if gemma_program {
3014 let scalar = |n: &str| -> f32 {
3015 let t = src.find(&p(n)).unwrap_or_else(|| panic!("missing {n}"));
3016 memra_gguf::dequant::dequantize(t.ggml_type, &t.bytes, 1)[0]
3017 };
3018 let vecf = |n: &str| -> Vec<f32> {
3019 let t = src.find(&p(n)).unwrap_or_else(|| panic!("missing {n}"));
3020 memra_gguf::dequant::dequantize(
3021 t.ggml_type,
3022 &t.bytes,
3023 t.ne.iter().product::<u64>() as usize,
3024 )
3025 };
3026 let moe_bits = if src.find(&p("ffn_gate_inp.scale")).is_some() {
3027 Some(crate::hybrid::Gemma4MoeBits {
3028 post_ffw_norm_1: load_t(e, src, &p("post_ffw_norm_1.weight"))?,
3029 pre_ffw_norm_2: load_t(e, src, &p("pre_ffw_norm_2.weight"))?,
3030 post_ffw_norm_2: load_t(e, src, &p("post_ffw_norm_2.weight"))?,
3031 shared_gate: load_t(e, src, &p("ffn_gate.weight"))?,
3032 shared_up: load_t(e, src, &p("ffn_up.weight"))?,
3033 shared_down: load_t(e, src, &p("ffn_down.weight"))?,
3034 router_scale_pre: {
3035 let inv = 1.0 / (cfg.n_embd as f32).sqrt();
3036 let v: Vec<f32> =
3037 vecf("ffn_gate_inp.scale").iter().map(|x| x * inv).collect();
3038 e.htod(&v)?
3039 },
3040 per_expert_scale: vecf("ffn_down_exps.scale"),
3041 per_expert_scale_d: e.htod(&vecf("ffn_down_exps.scale"))?,
3042 })
3043 } else {
3044 None
3045 };
3046 let e4b = if src.has(&p("inp_gate.weight")) {
3048 let g4 = cfg.gemma4.as_ref().unwrap();
3049 let kv_from = n_trunk as u32 - g4.shared_kv_layers;
3050 let kv_share = if g4.shared_kv_layers > 0 && il >= kv_from {
3051 let swa = g4.swa_pattern.get(il as usize).copied().unwrap_or(true);
3052 Some(kv_from - if swa { 2 } else { 1 })
3053 } else {
3054 None
3055 };
3056 Some(crate::hybrid::Gemma4E4bLayer {
3057 inp_gate: load_t(e, src, &p("inp_gate.weight"))?,
3058 proj: load_t(e, src, &p("proj.weight"))?,
3059 post_norm: load_t(e, src, &p("post_norm.weight"))?,
3060 kv_share,
3061 qkv_cat: None, })
3063 } else {
3064 None
3065 };
3066 Some(Gemma4LayerBits {
3067 ffn_norm: load_t(e, src, &p("ffn_norm.weight"))?,
3068 post_ffw_norm: load_t(e, src, &p("post_ffw_norm.weight"))?,
3069 moe_bits,
3070 layer_scale: scalar("layer_output_scale.weight"),
3071 e4b,
3072 })
3073 } else {
3074 None
3075 },
3076 });
3077 }
3078
3079 let external_mtp_requested =
3083 load_mtp && std::env::var("MEMRA_MTP_DRAFT").is_ok_and(|path| !path.is_empty());
3084 let trim_mtp_requested = load_mtp
3085 && !crate::model::full_prec_enabled()
3086 && std::env::var("MEMRA_FRSPEC_TRIM").is_ok_and(|path| !path.is_empty());
3087 let embedded_head_count = if external_mtp_requested {
3088 0
3089 } else if trim_mtp_requested {
3090 1
3091 } else {
3092 cfg.nextn_predict_layers
3093 };
3094 let mut embedded_mtp = Vec::new();
3095 if load_mtp && embedded_head_count > 0 {
3096 for offset in 0..embedded_head_count {
3097 let n = n_trunk as u32 + offset;
3098 let p = |s: &str| format!("blk.{n}.{s}");
3099 let mtp_plan = plan
3100 .mtp_blocks
3101 .iter()
3102 .find(|block| block.layer.index == n)
3103 .ok_or_else(|| format!("ModelPlan has no embedded MTP block {n}"))?;
3104 if !src.has(&p("nextn.eh_proj.weight")) {
3105 if offset == 0 {
3106 break;
3107 }
3108 return Err(format!(
3109 "embedded MTP chain declares {} heads but blk.{n} has no \
3110 nextn.eh_proj.weight",
3111 cfg.nextn_predict_layers
3112 )
3113 .into());
3114 }
3115 embedded_mtp.push(MtpHead {
3116 enorm: load_t(e, src, &p("nextn.enorm.weight"))?,
3117 hnorm: load_t(e, src, &p("nextn.hnorm.weight"))?,
3118 eh_proj: load_t(e, src, &p("nextn.eh_proj.weight"))?,
3119 attn_norm: load_t(e, src, &p("attn_norm.weight"))?,
3120 post_attn_norm: load_opt(e, src, &p("post_attention_norm.weight"))?
3121 .or(load_opt(e, src, &p("ffn_norm.weight"))?)
3122 .expect("MTP block needs post_attention_norm or ffn_norm"),
3123 mixer: load_mixer_kind(
3124 e,
3125 src,
3126 &cfg,
3127 n,
3128 &mtp_plan.layer.attention,
3129 &mut step_runtimes,
3130 )?,
3131 ffn: load_ffn(
3132 e,
3133 src,
3134 &cfg,
3135 &mtp_plan.layer.mlp,
3136 n,
3137 spill.as_mut().map(|c| (gguf.unwrap(), c)),
3138 &mut resident,
3139 &mut step_runtimes,
3140 )?,
3141 shared_head_norm: load_opt(e, src, &p("nextn.shared_head_norm.weight"))?,
3142 shared_head_head: load_opt(e, src, &p("nextn.shared_head_head.weight"))?
3151 .or(load_opt(e, src, &p("nextn.shared_head.weight"))?),
3152 d2t: None,
3153 d2t_from_target_head: false,
3154 geom: None,
3155 step35: if sliding_gated_moe_program {
3156 Some(Step35MtpGeom::from_plan(&mtp_plan.layer)?)
3157 } else {
3158 None
3159 },
3160 });
3161 }
3162 }
3163 let mut embedded_mtp = embedded_mtp.into_iter();
3164 let mut mtp = embedded_mtp.next();
3165 let mut mtp_extra: Vec<MtpHead> = embedded_mtp.collect();
3166
3167 mtp = if load_mtp {
3171 match std::env::var("MEMRA_MTP_DRAFT") {
3172 Ok(path) if !path.is_empty() => {
3173 eprintln!("[mtp-draft] loading external MTP draft: {path}");
3174 let dg = GgufFile::open(&path)?;
3175 mtp_extra.clear();
3176 Some(MtpHead::load_draft(e, &dg, &cfg)?)
3177 }
3178 _ => mtp,
3179 }
3180 } else {
3181 None
3182 };
3183
3184 let trim_env = if load_mtp {
3195 std::env::var("MEMRA_FRSPEC_TRIM")
3196 } else {
3197 Err(std::env::VarError::NotPresent)
3198 };
3199 if crate::model::full_prec_enabled()
3200 && trim_env.as_deref().map(|p| !p.is_empty()).unwrap_or(false)
3201 {
3202 eprintln!(
3203 "[frspec-trim] DISABLED under MEMRA_FULL_PREC — using the natural full MTP head"
3204 );
3205 }
3206 mtp = match (
3207 if crate::model::full_prec_enabled() {
3208 Err(std::env::VarError::NotPresent)
3209 } else {
3210 trim_env
3211 },
3212 mtp,
3213 ) {
3214 (Ok(path), Some(mut head)) if !path.is_empty() => {
3215 let path = memra_gguf::hf::resolve_arg(&path)
3219 .map_err(|err| format!("MEMRA_FRSPEC_TRIM={path:?}: {err}"))?;
3220 let d2t: Vec<u32> = if path.ends_with(".txt") {
3224 std::fs::read_to_string(&path)?
3225 .lines()
3226 .filter_map(|l| l.trim().parse::<u32>().ok())
3227 .collect()
3228 } else {
3229 let tg = GgufFile::open(&path)?;
3230 let d2t_t = tg
3231 .find("d2t")
3232 .expect("MEMRA_FRSPEC_TRIM file has no d2t tensor");
3233 let d2t_bytes = tg.tensor_data(d2t_t);
3234 match d2t_t.ggml_type {
3235 GgmlType::I32 => d2t_bytes
3236 .chunks_exact(4)
3237 .map(|c| i32::from_le_bytes(c.try_into().unwrap()) as u32)
3238 .collect(),
3239 GgmlType::I64 => d2t_bytes
3240 .chunks_exact(8)
3241 .map(|c| i64::from_le_bytes(c.try_into().unwrap()) as u32)
3242 .collect(),
3243 other => panic!("d2t must be I32/I64, got {other:?}"),
3244 }
3245 };
3246 let v = src
3247 .find("output.weight")
3248 .or_else(|| src.find("token_embd.weight"))
3249 .expect("model has no output.weight for FR-Spec trim");
3250 let out_f = v.ne[1] as usize;
3251 let row_bytes = v.bytes.len() / out_f;
3252 assert!(
3253 d2t.iter().all(|&t| (t as usize) < out_f),
3254 "d2t token id >= lm_head rows {out_f}"
3255 );
3256 let mut gathered = Vec::with_capacity(d2t.len() * row_bytes);
3257 for &t in &d2t {
3258 let off = t as usize * row_bytes;
3259 gathered.extend_from_slice(&v.bytes[off..off + row_bytes]);
3260 }
3261 let trimmed = GpuTensor::from_quant_bytes(
3262 e,
3263 &gathered,
3264 v.ggml_type,
3265 v.ne[0],
3266 d2t.len() as u64,
3267 match src.find("output.scale") {
3269 Some(sv) => f32::from_le_bytes(sv.bytes[..4].try_into().unwrap()),
3270 None => 1.0,
3271 },
3272 )?;
3273 eprintln!(
3274 "[frspec-trim] self-trimmed head: {} rows of main output.weight ({:?})",
3275 d2t.len(),
3276 v.ggml_type
3277 );
3278 head.shared_head_head = Some(trimmed);
3279 head.d2t = Some(d2t);
3280 head.d2t_from_target_head = true;
3281 Some(head)
3282 }
3283 (_, m) => m,
3284 };
3285 if mtp.as_ref().is_some_and(|head| head.d2t.is_some()) {
3286 mtp_extra.clear();
3287 }
3288 if !mtp_extra.is_empty() {
3289 if plan.draft_source != memra_gguf::model_plan::DraftSourcePlan::Embedded
3290 || plan.mtp_blocks.len() != 1 + mtp_extra.len()
3291 || plan
3292 .mtp_blocks
3293 .iter()
3294 .any(|block| !matches!(block.layer.mlp, MlpPlan::Dense(_)))
3295 || mtp
3296 .iter()
3297 .chain(mtp_extra.iter())
3298 .any(|head| !matches!(head.ffn, Ffn::Dense { .. }))
3299 {
3300 return Err(
3301 "multi-head MTP requires embedded dense canonical blocks and matching loaded heads"
3302 .into(),
3303 );
3304 }
3305 eprintln!(
3306 "[mtp-draft] embedded chain: heads={} blocks={}..={} scratch=per-head",
3307 1 + mtp_extra.len(),
3308 n_trunk,
3309 n_trunk + mtp_extra.len()
3310 );
3311 }
3312
3313 if let Some(ctx) = spill.as_ref() {
3314 eprintln!(
3315 "[spill] experts placed: {} pinned (Tier 1), {} mmap'd from disk (Tier 2, {} MiB)",
3316 ctx.n_pinned,
3317 ctx.n_mmap,
3318 ctx.mmap_bytes >> 20
3319 );
3320 }
3321
3322 if cfg.n_head_kv > 0 && cfg.n_head / cfg.n_head_kv > 8 {
3336 crate::FA_V4_MAX_DEFAULT.store(0, std::sync::atomic::Ordering::Relaxed);
3337 eprintln!(
3338 "[fa] v4 decode family disabled: gqa {} > fa_v4_smem capacity 8 (v3 lane serves)",
3339 cfg.n_head / cfg.n_head_kv
3340 );
3341 }
3342
3343 if gemma_program {
3344 crate::FA_VEC_MIN_DEFAULT.store(1, std::sync::atomic::Ordering::Relaxed);
3346 let real_moe = plan
3349 .trunk_operations()
3350 .contains(&memra_gguf::model_plan::OperationKind::MoeMlp);
3351 crate::FA_SPW_DEFAULT.store(
3352 if real_moe { 32 } else { 64 },
3353 std::sync::atomic::Ordering::Relaxed,
3354 );
3355 crate::FA_SP512_DEFAULT.store(
3357 if real_moe { 16 } else { 32 },
3358 std::sync::atomic::Ordering::Relaxed,
3359 );
3360 crate::FUSED_MR1_DEFAULT.store(!real_moe, std::sync::atomic::Ordering::Relaxed);
3370 crate::RMS_BLOCK_DEFAULT.store(1024, std::sync::atomic::Ordering::Relaxed);
3372 crate::FA_SP_GEMMA.store(true, std::sync::atomic::Ordering::Relaxed);
3374 }
3378 let force_embd_gpu = gemma_program;
3381 let gemma4_aux = if gemma_program {
3382 let rope_freqs = match src.find("rope_freqs.weight") {
3383 Some(t) => {
3384 let host = memra_gguf::dequant::dequantize(
3385 t.ggml_type,
3386 &t.bytes,
3387 t.ne.iter().product::<u64>() as usize,
3388 );
3389 let mut copies = Vec::new();
3390 if let Some(fence) = crate::pp::pp_cuts(n_trunk) {
3391 for s in 0..fence.len() - 1 {
3392 let owner = crate::pp::layer_engine(e, n_trunk, fence[s])?;
3393 let dev = owner.ctx().ordinal();
3394 if copies.iter().all(|(d, _)| *d != dev) {
3395 copies.push((dev, owner.htod(&host)?));
3396 }
3397 }
3398 } else {
3399 copies.push((e.ctx().ordinal(), e.htod(&host)?));
3400 }
3401 Some(copies)
3402 }
3403 None => {
3411 let g4 = cfg.gemma4.as_ref().unwrap();
3412 let n = (g4.rope_dims_global / 2) as usize;
3413 let keep =
3414 ((n as f32) * g4.partial_rotary_global.clamp(0.0, 1.0)).round() as usize;
3415 let host: Vec<f32> = (0..n)
3416 .map(|i| if i < keep { 1.0 } else { 1.0e30 })
3417 .collect();
3418 eprintln!(
3419 "[gemma4] rope_freqs.weight synthesized ({n} factors, first {keep} \
3420 rotate; source ships none — native checkpoint)"
3421 );
3422 let mut copies = Vec::new();
3423 if let Some(fence) = crate::pp::pp_cuts(n_trunk) {
3424 for s in 0..fence.len() - 1 {
3425 let owner = crate::pp::layer_engine(e, n_trunk, fence[s])?;
3426 let dev = owner.ctx().ordinal();
3427 if copies.iter().all(|(d, _)| *d != dev) {
3428 copies.push((dev, owner.htod(&host)?));
3429 }
3430 }
3431 } else {
3432 copies.push((e.ctx().ordinal(), e.htod(&host)?));
3433 }
3434 Some(copies)
3435 }
3436 };
3437 let e4b = match src.find("per_layer_token_embd.weight") {
3439 Some(t) => {
3440 let n_epl = cfg
3441 .gemma4
3442 .as_ref()
3443 .map(|g| g.n_embd_per_layer as usize)
3444 .unwrap_or(0);
3445 let row = t.ne[0] as usize; let row_bytes = t.bytes.len() / (t.ne[1] as usize);
3447 eprintln!(
3448 "[gemma4-e4b] per-layer-embed model detected (n_epl={n_epl}, row {row}) — \
3449 first-light forward (eager decode + prime); dc/graph/spec unwired \
3450 (HANDOVER-E4B.md)"
3451 );
3452 Some(crate::hybrid::Gemma4E4bModel {
3453 tok_tbl_gpu: std::sync::OnceLock::new(),
3454 tok_embd_bytes: t.bytes.to_vec(),
3455 tok_embd_qt: match t.ggml_type {
3456 memra_gguf::GgmlType::Q6_K => crate::QT_Q6_K,
3457 memra_gguf::GgmlType::Q8_0 => crate::QT_Q8_0,
3458 other => panic!("e4b per-layer tok embd: unhandled dtype {other:?}"),
3459 },
3460 tok_embd_row_bytes: row_bytes,
3461 model_proj: load_t(e, src, "per_layer_model_proj.weight")?,
3462 proj_norm: load_t(e, src, "per_layer_proj_norm.weight")?,
3463 n_epl,
3464 })
3465 }
3466 None => None,
3467 };
3468 let suppress_d = {
3469 let sup = &cfg.gemma4.as_ref().unwrap().suppress_tokens;
3470 if sup.is_empty() {
3471 None
3472 } else {
3473 let ids: Vec<i32> = sup.iter().map(|&x| x as i32).collect();
3474 eprintln!(
3475 "[gemma4] suppress_tokens: {} ids masked at sampling",
3476 ids.len()
3477 );
3478 Some((e.htod_i32(&ids)?, ids.len()))
3479 }
3480 };
3481 let ones_host = [1.0f32; 512];
3482 let mut ones = Vec::new();
3483 if let Some(fence) = crate::pp::pp_cuts(n_trunk) {
3484 for s in 0..fence.len() - 1 {
3485 let owner = crate::pp::layer_engine(e, n_trunk, fence[s])?;
3486 let dev = owner.ctx().ordinal();
3487 if ones.iter().all(|(d, _)| *d != dev) {
3488 ones.push((dev, owner.htod(&ones_host)?));
3489 }
3490 }
3491 } else {
3492 ones.push((e.ctx().ordinal(), e.htod(&ones_host)?));
3493 }
3494 Some(GemmaAux {
3495 rope_freqs,
3496 ones,
3497 suppress_d,
3498 e4b,
3499 })
3500 } else {
3501 None
3502 };
3503 let step35_aux = if sliding_gated_moe_program {
3507 let rope_freqs = match src.find("rope_freqs.weight") {
3508 Some(t) => {
3509 let host = memra_gguf::dequant::dequantize(
3510 t.ggml_type,
3511 &t.bytes,
3512 t.ne.iter().product::<u64>() as usize,
3513 );
3514 let mut copies = Vec::new();
3515 if let Some(fence) = crate::pp::pp_cuts(n_trunk) {
3516 for s in 0..fence.len() - 1 {
3517 let owner = crate::pp::layer_engine(e, n_trunk, fence[s])?;
3518 let dev = owner.ctx().ordinal();
3519 if copies.iter().all(|(d, _)| *d != dev) {
3520 copies.push((dev, owner.htod(&host)?));
3521 }
3522 }
3523 } else {
3524 copies.push((e.ctx().ordinal(), e.htod(&host)?));
3525 }
3526 Some(copies)
3527 }
3528 None => None,
3529 };
3530 Some(Step35Aux { rope_freqs })
3531 } else {
3532 None
3533 };
3534 let mut layers = layers;
3535 {
3542 let q8rp_on = match std::env::var("MEMRA_Q8RP").as_deref() {
3543 Ok("0") => false,
3544 Ok(_) => true,
3545 Err(_) => {
3553 cfg!(memra_hopper_mma) || {
3554 let q8b = |w: &crate::model::GpuTensor| -> usize {
3555 match w {
3556 crate::model::GpuTensor::Quant {
3557 bytes,
3558 qtype,
3559 row_bytes,
3560 ne,
3561 rp4: None,
3562 ..
3563 } if *qtype == crate::QT_Q8_0
3564 && ne.len() == 2
3565 && (ne[0] as usize) % 32 == 0
3566 && *row_bytes == (ne[0] as usize / 32) * 34 =>
3567 {
3568 bytes.len()
3569 }
3570 _ => 0,
3571 }
3572 };
3573 let mut need = q8b(&output);
3574 for layer in layers.iter() {
3575 match &layer.mixer {
3576 Mixer::Full(fa) => {
3577 for w in [&fa.wq, &fa.wk, &fa.wv, &fa.wo] {
3578 need += q8b(w);
3579 }
3580 }
3581 Mixer::Linear(la) => {
3582 for w in [
3583 &la.wqkv,
3584 &la.wqkv_gate,
3585 &la.ssm_beta,
3586 &la.ssm_alpha,
3587 &la.ssm_out,
3588 ] {
3589 need += q8b(w);
3590 }
3591 }
3592 Mixer::Mla(_) => {}
3593 }
3594 if let Ffn::Dense {
3595 ffn_gate,
3596 ffn_up,
3597 ffn_down,
3598 } = &layer.ffn
3599 {
3600 for w in [ffn_gate, ffn_up, ffn_down] {
3601 need += q8b(w);
3602 }
3603 }
3604 }
3605 need > 0
3606 && e.ctx()
3607 .mem_get_info()
3608 .map(|(free, _)| free >= need + (8usize << 30))
3609 .unwrap_or(false)
3610 }
3611 }
3612 };
3613 let kqrp_on = crate::Engine::kqrp_enabled() || {
3623 std::env::var("MEMRA_KQRP").is_err() && {
3624 let kqb = |w: &crate::model::GpuTensor| -> usize {
3625 match w {
3626 crate::model::GpuTensor::Quant {
3627 bytes,
3628 qtype,
3629 row_bytes,
3630 ne,
3631 rp4: None,
3632 ..
3633 } if ne.len() == 2 && (ne[0] as usize) % 256 == 0 => {
3634 let sb = if *qtype == crate::QT_Q4_K {
3635 144
3636 } else if *qtype == crate::QT_Q6_K {
3637 210
3638 } else {
3639 return 0;
3640 };
3641 if *row_bytes == (ne[0] as usize / 256) * sb {
3642 bytes.len()
3643 } else {
3644 0
3645 }
3646 }
3647 _ => 0,
3648 }
3649 };
3650 let mut need = kqb(&output);
3651 for layer in layers.iter() {
3652 if let Mixer::Full(fa) = &layer.mixer {
3653 for w in [&fa.wq, &fa.wk, &fa.wv, &fa.wo] {
3654 need += kqb(w);
3655 }
3656 }
3657 if let Ffn::Dense {
3658 ffn_gate,
3659 ffn_up,
3660 ffn_down,
3661 } = &layer.ffn
3662 {
3663 for w in [ffn_gate, ffn_up, ffn_down] {
3664 need += kqb(w);
3665 }
3666 }
3667 }
3668 need > 0
3669 && e.ctx()
3670 .mem_get_info()
3671 .map(|(free, _)| free >= need + (8usize << 30))
3672 .unwrap_or(false)
3673 }
3674 };
3675 if q8rp_on || kqrp_on {
3676 let f16_model_ok = gemma_program
3683 || plan
3684 .trunk_operations()
3685 .contains(&memra_gguf::model_plan::OperationKind::MoeMlp)
3686 || std::env::var("MEMRA_PP_F16").as_deref() == Ok("1");
3687 let mut nmir = 0usize;
3688 let mut mir = |e_ref: &crate::Engine,
3692 w: &mut crate::model::GpuTensor|
3693 -> Result<(), Box<dyn std::error::Error>> {
3694 let before = matches!(w, crate::model::GpuTensor::Quant { rp4: Some(_), .. });
3695 if q8rp_on {
3696 e_ref.build_q8_rp4(w)?;
3697 }
3698 if kqrp_on {
3699 e_ref.build_q4k_rp4(w)?;
3700 e_ref.build_q6k_rp4(w)?;
3701 }
3702 let q6k = matches!(w, crate::model::GpuTensor::Quant { qtype, .. }
3707 if *qtype == crate::QT_Q6_K);
3708 if q8rp_on && crate::f16_ffi::pp_f16_enabled() && (f16_model_ok || q6k) {
3709 e_ref.build_q8_f16(w)?;
3710 }
3711 if !before && matches!(w, crate::model::GpuTensor::Quant { rp4: Some(_), .. }) {
3712 nmir += 1;
3713 }
3714 Ok(())
3715 };
3716 for (il, layer) in layers.iter_mut().enumerate() {
3717 let el = crate::pp::layer_engine(e, n_trunk, il)?;
3718 match &mut layer.mixer {
3719 Mixer::Full(fa) => {
3720 for w in [&mut fa.wq, &mut fa.wk, &mut fa.wv, &mut fa.wo] {
3721 mir(el, w)?;
3722 }
3723 }
3724 Mixer::Linear(la) => {
3725 for w in [
3726 &mut la.wqkv,
3727 &mut la.wqkv_gate,
3728 &mut la.ssm_beta,
3729 &mut la.ssm_alpha,
3730 &mut la.ssm_out,
3731 ] {
3732 mir(el, w)?;
3733 }
3734 }
3735 Mixer::Mla(_) => {}
3738 }
3739 if let Ffn::Dense {
3740 ffn_gate,
3741 ffn_up,
3742 ffn_down,
3743 } = &mut layer.ffn
3744 {
3745 for w in [ffn_gate, ffn_up, ffn_down] {
3746 mir(el, w)?;
3747 }
3748 }
3749 }
3750 mir(e_head, &mut output)?;
3751 if nmir > 0 {
3752 eprintln!("[q8rp] split-plane decode mirrors built: {nmir} tensors");
3753 }
3754 if q8rp_on && crate::f16_ffi::pp_f16_enabled() {
3769 for (want, tag) in [(crate::QT_Q4_K, "q4kf16"), (crate::QT_Q5_K, "q5kf16")] {
3770 let (mut n4, mut b4) = (0usize, 0usize);
3771 let mut mirk =
3772 |e_ref: &crate::Engine,
3773 w: &mut crate::model::GpuTensor|
3774 -> Result<(), Box<dyn std::error::Error>> {
3775 if matches!(w, crate::model::GpuTensor::Quant { qtype, f16: None, .. }
3776 if *qtype == want)
3777 {
3778 e_ref.build_q8_f16(w)?;
3779 if let crate::model::GpuTensor::Quant { f16: Some(m), .. } = w {
3780 n4 += 1;
3781 b4 += m.len();
3782 }
3783 }
3784 Ok(())
3785 };
3786 for (il, layer) in layers.iter_mut().enumerate() {
3787 let el = crate::pp::layer_engine(e, n_trunk, il)?;
3788 match &mut layer.mixer {
3789 Mixer::Full(fa) => {
3790 for w in [&mut fa.wq, &mut fa.wk, &mut fa.wv, &mut fa.wo] {
3791 mirk(el, w)?;
3792 }
3793 }
3794 Mixer::Linear(la) => {
3795 for w in [
3796 &mut la.wqkv,
3797 &mut la.wqkv_gate,
3798 &mut la.ssm_beta,
3799 &mut la.ssm_alpha,
3800 &mut la.ssm_out,
3801 ] {
3802 mirk(el, w)?;
3803 }
3804 }
3805 Mixer::Mla(_) => {} }
3807 if let Ffn::Dense {
3808 ffn_gate,
3809 ffn_up,
3810 ffn_down,
3811 } = &mut layer.ffn
3812 {
3813 for w in [ffn_gate, ffn_up, ffn_down] {
3814 mirk(el, w)?;
3815 }
3816 }
3817 }
3818 mirk(e_head, &mut output)?;
3819 if n4 > 0 {
3820 eprintln!(
3821 "[{tag}] prefill fp16 mirrors built: {n4} tensors \
3822 ({} MB)",
3823 b4 >> 20
3824 );
3825 }
3826 }
3827 }
3828 }
3829 }
3830 if gemma_program && crate::Engine::q4rp_enabled() {
3837 let mut nmir = 0usize;
3838 for (il, layer) in layers.iter_mut().enumerate() {
3839 let e = crate::pp::layer_engine(e, n_trunk, il)?;
3841 let is_moe26 = layer.gemma4.as_ref().is_some_and(|g| g.moe_bits.is_some());
3850 let is_e4b = layer.gemma4.as_ref().is_some_and(|g| g.e4b.is_some());
3851 if !(is_moe26 || is_e4b) {
3852 continue;
3853 }
3854 if let Mixer::Full(fa) = &mut layer.mixer {
3855 for w in [&mut fa.wq, &mut fa.wk, &mut fa.wv, &mut fa.wo] {
3856 e.build_q4_rp4(w)?;
3857 nmir += 1;
3858 }
3859 }
3860 if is_e4b {
3861 let own_kv = layer
3863 .gemma4
3864 .as_ref()
3865 .unwrap()
3866 .e4b
3867 .as_ref()
3868 .is_some_and(|e4| e4.kv_share.is_none());
3869 if own_kv {
3870 if let Mixer::Full(fa) = &layer.mixer {
3871 if let Some(mut cat) = e.build_q4_out_concat3(&fa.wq, &fa.wk, &fa.wv)? {
3872 e.build_q4_rp4(&mut cat)?;
3873 nmir += 1;
3874 layer.gemma4.as_mut().unwrap().e4b.as_mut().unwrap().qkv_cat =
3875 Some(cat);
3876 }
3877 }
3878 }
3879 if let Ffn::Dense {
3880 ffn_gate,
3881 ffn_up,
3882 ffn_down,
3883 } = &mut layer.ffn
3884 {
3885 for w in [ffn_gate, ffn_up, ffn_down] {
3886 e.build_q4_rp4(w)?;
3887 nmir += 1;
3888 }
3889 }
3890 let e4 = layer.gemma4.as_mut().unwrap().e4b.as_mut().unwrap();
3891 for w in [&mut e4.inp_gate, &mut e4.proj] {
3892 e.build_q4_rp4(w)?;
3893 nmir += 1;
3894 }
3895 }
3896 if let Some(mb) = layer.gemma4.as_mut().unwrap().moe_bits.as_mut() {
3897 for w in [&mut mb.shared_gate, &mut mb.shared_up, &mut mb.shared_down] {
3898 e.build_q4_rp4(w)?;
3899 nmir += 1;
3900 }
3901 }
3902 }
3903 if nmir > 0 {
3904 eprintln!("[q4rp] split-plane decode mirrors built: {nmir} trunk tensors");
3905 }
3906 let fast_on = std::env::var("MEMRA_FAST").as_deref() != Ok("0");
3913 if fast_on {
3914 let mut nswap = 0usize;
3915 let mut nf16 = 0usize;
3916 let q4f16_model_ok = matches!(cfg.n_embd, 3840 | 5376); if let Ok(v) = std::env::var("MEMRA_Q4F16") {
3935 if v != "0" && v != "1" {
3936 return Err(format!(
3937 "MEMRA_Q4F16={v} is not 0 or 1 — this env selects the prefill \
3938 ARITHMETIC (fp16 mirrors vs int8 MMQ) and must never be guessed"
3939 )
3940 .into());
3941 }
3942 }
3943 let f16_need = {
3944 let f16b = |w: &crate::model::GpuTensor| -> usize {
3945 match w {
3946 crate::model::GpuTensor::Quant {
3947 qtype,
3948 ne,
3949 f16: None,
3950 ..
3951 } if ne.len() == 2
3952 && matches!(
3953 *qtype,
3954 crate::QT_Q8_0
3955 | crate::QT_Q4_0
3956 | crate::QT_Q6_K
3957 | crate::QT_Q4_K
3958 | crate::QT_Q5_K
3959 ) =>
3960 {
3961 (ne[0] as usize) * (ne[1] as usize) * 2
3962 }
3963 _ => 0,
3964 }
3965 };
3966 let mut need = 0usize;
3967 for layer in layers.iter() {
3968 if layer.gemma4.as_ref().is_none_or(|g| g.moe_bits.is_some()) {
3969 continue;
3970 }
3971 if let Mixer::Full(fa) = &layer.mixer {
3972 for w in [&fa.wq, &fa.wk, &fa.wv, &fa.wo] {
3973 need += f16b(w);
3974 }
3975 }
3976 if let Ffn::Dense {
3977 ffn_gate,
3978 ffn_up,
3979 ffn_down,
3980 } = &layer.ffn
3981 {
3982 for w in [ffn_gate, ffn_up, ffn_down] {
3983 need += f16b(w);
3984 }
3985 }
3986 }
3987 need
3988 };
3989 let f16_free = e.ctx().mem_get_info().map(|(free, _)| free).unwrap_or(0);
3990 let f16_auto = q4f16_model_ok
3991 && std::env::var("MEMRA_Q4F16").is_err()
3992 && crate::f16_ffi::pp_f16_capacity_ok(f16_free, f16_need);
3993 let (f16_on, f16_why) = match std::env::var("MEMRA_Q4F16").as_deref() {
3999 Ok("1") => (true, "env MEMRA_Q4F16=1"),
4000 Ok("0") => (false, "env MEMRA_Q4F16=0"),
4001 _ if crate::f16_ffi::pp_f16_enabled() && q4f16_model_ok => {
4002 (true, "env MEMRA_PP_F16")
4003 }
4004 _ if f16_auto => (true, "capacity-keyed auto (UNPINNED)"),
4005 _ if !q4f16_model_ok => (false, "model geometry not eligible"),
4006 _ => (false, "capacity-keyed auto REFUSED (UNPINNED)"),
4007 };
4008 eprintln!(
4015 "[q4f16] prefill program = {} (reason: {}); free {} MiB, mirror mass {} MiB, \
4016 capacity threshold {} MiB (mass + 8192 headroom) — SELECTS PREFILL ARITHMETIC",
4017 if f16_on {
4018 "FP16 MIRRORS"
4019 } else {
4020 "INT8 MMQ (no f16 mirrors)"
4021 },
4022 f16_why,
4023 f16_free >> 20,
4024 f16_need >> 20,
4025 (f16_need + (8usize << 30)) >> 20,
4026 );
4027 for (il, layer) in layers.iter_mut().enumerate() {
4028 let e = crate::pp::layer_engine(e, n_trunk, il)?;
4030 let dense_gemma = layer.gemma4.as_ref().is_some_and(|g| g.moe_bits.is_none());
4031 if !dense_gemma {
4032 continue;
4033 }
4034 if let Mixer::Full(fa) = &mut layer.mixer {
4035 for w in [&mut fa.wq, &mut fa.wk, &mut fa.wv, &mut fa.wo] {
4036 if f16_on {
4037 e.build_q8_f16(w)?;
4038 if matches!(w, crate::model::GpuTensor::Quant { f16: Some(_), .. })
4039 {
4040 nf16 += 1;
4041 }
4042 }
4043 if e.build_q4_rp_swap(w)? {
4044 nswap += 1;
4045 }
4046 }
4047 }
4048 if let Ffn::Dense {
4049 ffn_gate,
4050 ffn_up,
4051 ffn_down,
4052 } = &mut layer.ffn
4053 {
4054 for w in [ffn_gate, ffn_up, ffn_down] {
4055 if f16_on {
4056 e.build_q8_f16(w)?;
4057 if matches!(w, crate::model::GpuTensor::Quant { f16: Some(_), .. })
4058 {
4059 nf16 += 1;
4060 }
4061 }
4062 if e.build_q4_rp_swap(w)? {
4063 nswap += 1;
4064 }
4065 }
4066 }
4067 }
4068 if nswap > 0 {
4069 eprintln!("[q4rp] split-plane IN-PLACE swap: {nswap} dense trunk tensors");
4070 }
4071 if nf16 > 0 {
4072 eprintln!("[q4f16] prefill fp16 mirrors built: {nf16} dense trunk tensors");
4073 }
4074 }
4075 }
4076 let model = HybridModel {
4077 cfg,
4078 plan,
4079 rewrite_qualifications: None,
4080 embd,
4081 output_norm,
4082 output,
4083 layers,
4084 mtp,
4085 mtp_extra,
4086 embd_gpu: std::sync::OnceLock::new(),
4087 gemma4_aux,
4088 step35_aux,
4089 prime_slabs: std::sync::Mutex::new(std::collections::HashMap::new()),
4090 dspark_vgraphs: std::sync::Mutex::new(None),
4091 step_grouped_prefill: std::sync::Mutex::new(StepEpGroupedPrefill::default()),
4092 step35_token_graph: std::sync::Mutex::new(None),
4093 };
4094 e.configure_moe_cache_layout(model.moe_cache_block_sizes());
4095 if force_embd_gpu {
4096 let _ = model
4097 .embd_gpu
4098 .get_or_init(|| e.upload_u8(&model.embd.raw).expect("embed table upload"));
4099 }
4100 crate::pp::sync_stages_after_load(e, n_trunk)?;
4106 Ok(model)
4107 }
4108
4109 pub fn ensure_embed_resident(&self, e: &Engine) -> Result<(), Box<dyn std::error::Error>> {
4119 if std::env::var("MEMRA_EMBED_DEV").as_deref() == Ok("0") {
4120 return Ok(());
4121 }
4122 if self.embd_gpu.get().is_none() {
4123 let buf = e.upload_u8(&self.embd.raw)?;
4124 let _ = self.embd_gpu.set(buf); }
4126 Ok(())
4127 }
4128
4129 pub fn embed(
4130 &self,
4131 e: &Engine,
4132 tokens: &[u32],
4133 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4134 let n_embd = self.cfg.n_embd as usize;
4135 if std::env::var("MEMRA_EMBED_DEV").as_deref() != Ok("0") {
4141 let tbl = self
4142 .embd_gpu
4143 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
4144 let tok_d = e.htod_u32_v(tokens)?;
4145 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
4146 return e.embed_gather_device_td(tbl, &tok_d, tokens.len(), n_embd, qt, rb);
4147 }
4148 let x = self.embd.gather(n_embd, tokens);
4149 Ok(e.htod(&x)?)
4150 }
4151}
4152
4153#[cfg(test)]
4154mod step_expert_selection_tests {
4155 use super::{
4156 StepExpertArtifact, StepExpertLayout, StepParallelLoadConfig, StepParallelRuntimeRegistry,
4157 StepTpAttentionPlacement, select_step_expert_layout,
4158 };
4159 use crate::tp::StepEpLayerSpec;
4160
4161 fn spec(layer: usize, ranks: usize) -> StepEpLayerSpec {
4162 StepEpLayerSpec {
4163 layer,
4164 devices: (0..ranks).collect(),
4165 }
4166 }
4167
4168 #[test]
4169 fn tp2_keeps_projection_sharded_experts() {
4170 let selection = select_step_expert_layout(24, &[], &[spec(24, 2)])
4171 .unwrap()
4172 .unwrap();
4173 assert_eq!(selection.layout, StepExpertLayout::TensorParallel);
4174 assert!(selection.configured_by_tp);
4175 }
4176
4177 #[test]
4178 fn tp4_and_tp8_use_expert_ownership_without_a_second_flag() {
4179 for ranks in [4, 8] {
4180 let selection = select_step_expert_layout(24, &[], &[spec(24, ranks)])
4181 .unwrap()
4182 .unwrap();
4183 assert_eq!(selection.layout, StepExpertLayout::ExpertParallel);
4184 assert!(selection.configured_by_tp);
4185 assert_eq!(selection.spec.devices.len(), ranks);
4186 }
4187 }
4188
4189 #[test]
4190 fn explicit_ep_remains_expert_parallel() {
4191 let selection = select_step_expert_layout(24, &[spec(24, 2)], &[])
4192 .unwrap()
4193 .unwrap();
4194 assert_eq!(selection.layout, StepExpertLayout::ExpertParallel);
4195 assert!(!selection.configured_by_tp);
4196 }
4197
4198 #[test]
4199 fn conflicting_ep_and_tp_assignments_fail_closed() {
4200 let error = select_step_expert_layout(24, &[spec(24, 4)], &[spec(24, 4)]).unwrap_err();
4201 assert!(error.contains("cannot enable MEMRA_STEP_EP and MEMRA_STEP_TP together"));
4202 }
4203
4204 #[test]
4205 fn runtime_registry_owns_one_immutable_load_snapshot() {
4206 let mut source_specs = vec![spec(24, 8)];
4207 let registry = StepParallelRuntimeRegistry::with_config(StepParallelLoadConfig {
4208 ep_specs: Vec::new(),
4209 tp_specs: source_specs.clone(),
4210 native_p2p: true,
4211 ep_device_arithmetic: true,
4212 f32_mirror: true,
4213 bulk_p2p: true,
4214 expert_artifact: StepExpertArtifact::default(),
4215 });
4216 source_specs[0].devices.clear();
4217
4218 let stored = registry.tp_spec(24).unwrap();
4219 assert_eq!(stored.devices, (0..8).collect::<Vec<_>>());
4220 assert!(registry.config.native_p2p);
4221 assert!(registry.config.ep_device_arithmetic);
4222 assert!(registry.config.f32_mirror);
4223 assert!(registry.config.bulk_p2p);
4224 assert_eq!(
4225 registry.expert_selection(24).unwrap().unwrap().layout,
4226 StepExpertLayout::ExpertParallel
4227 );
4228
4229 let standalone = StepParallelRuntimeRegistry::default();
4230 assert!(standalone.tp_spec(24).is_none());
4231 assert!(!standalone.config.native_p2p);
4232 assert!(!standalone.config.ep_device_arithmetic);
4233 assert!(!standalone.config.f32_mirror);
4234 assert!(!standalone.config.bulk_p2p);
4235 }
4236
4237 #[test]
4238 fn rank_local_attention_uses_bounded_swa_rings_only_with_native_p2p() {
4239 assert_eq!(
4240 StepTpAttentionPlacement::resolve(true, None),
4241 StepTpAttentionPlacement::RankLocalGlobal
4242 );
4243 assert_eq!(
4244 StepTpAttentionPlacement::resolve(true, Some(512)),
4245 StepTpAttentionPlacement::RankLocalSwa
4246 );
4247 assert_eq!(
4248 StepTpAttentionPlacement::resolve(false, None),
4249 StepTpAttentionPlacement::OwnerTransportFallback
4250 );
4251 assert_eq!(
4252 StepTpAttentionPlacement::resolve(false, Some(512)),
4253 StepTpAttentionPlacement::OwnerSwa
4254 );
4255 }
4256}
4257
4258#[cfg(test)]
4259mod residency_tests {
4260 use super::{DevExpertFp8ProjectionScales, residency_bytes_by_device};
4261 use crate::model::HostExpertFp8BlockScales;
4262
4263 #[test]
4264 fn pp_residency_counts_only_each_devices_expert_slice() {
4265 let tensors = [
4266 ("blk.0.ffn_gate_exps.weight", 10usize),
4267 ("blk.0.ffn_up_exps.weight", 20),
4268 ("blk.1.ffn_down_exps.weight", 30),
4269 ("blk.2.ffn_gate_exps.weight", 40),
4270 ("blk.3.ffn_up_exps.weight", 50),
4271 ("blk.0.attn_q.weight", 7),
4272 ("output.weight", 11),
4273 ];
4274 let bytes = residency_bytes_by_device(tensors, &[0, 0, 1, 1], 0);
4275 assert_eq!(bytes.experts.get(&0), Some(&60));
4276 assert_eq!(bytes.experts.get(&1), Some(&90));
4277 assert_eq!(bytes.rest, 18);
4278 assert!(bytes.saw_experts);
4279 }
4280
4281 #[test]
4282 fn pp_residency_combines_stages_that_share_one_device() {
4283 let tensors = [
4284 ("blk.0.ffn_gate_exps.weight", 10usize),
4285 ("blk.1.ffn_gate_exps.weight", 20),
4286 ("blk.2.ffn_gate_exps.weight", 30),
4287 ("blk.3.ffn_gate_exps.weight", 40),
4288 ];
4289 let bytes = residency_bytes_by_device(tensors, &[0, 0, 0, 0], 0);
4290 assert_eq!(bytes.experts.get(&0), Some(&100));
4291 assert_eq!(bytes.experts.len(), 1);
4292 }
4293
4294 #[test]
4295 fn resident_fp8_scale_slab_must_match_every_expert() {
4296 let valid = HostExpertFp8BlockScales {
4297 scales: vec![1.0; 12],
4298 rows: 2,
4299 cols: 3,
4300 expert_stride: 6,
4301 };
4302 DevExpertFp8ProjectionScales::validate(&valid, 2).unwrap();
4303
4304 let short = HostExpertFp8BlockScales {
4305 scales: vec![1.0; 11],
4306 ..valid
4307 };
4308 assert_eq!(
4309 DevExpertFp8ProjectionScales::validate(&short, 2).unwrap_err(),
4310 "block-E4M3 scale slab length mismatch: got 11, want 2x6=12"
4311 );
4312 }
4313
4314 #[test]
4315 fn resident_fp8_scale_stride_must_match_its_grid() {
4316 let invalid = HostExpertFp8BlockScales {
4317 scales: vec![1.0; 8],
4318 rows: 2,
4319 cols: 2,
4320 expert_stride: 0,
4321 };
4322 assert_eq!(
4323 DevExpertFp8ProjectionScales::validate(&invalid, 2).unwrap_err(),
4324 "block-E4M3 expert scale stride must be nonzero"
4325 );
4326 }
4327}
4328
4329#[cfg(test)]
4330mod draft_head_tests {
4331 use super::draft_head_tensor;
4332
4333 const STEP37_DRAFTER: &[&str] = &[
4340 "output.weight",
4341 "output_norm.weight",
4342 "token_embd.weight",
4343 "blk.45.nextn.shared_head_norm.weight",
4344 "blk.45.nextn.shared_head_head.weight",
4345 "blk.46.nextn.shared_head_head.weight",
4346 "blk.47.nextn.shared_head_head.weight",
4347 ];
4348
4349 fn present(names: &'static [&'static str]) -> impl Fn(&str) -> bool {
4350 move |t: &str| names.contains(&t)
4351 }
4352
4353 #[test]
4361 fn step37_drafter_prefers_the_blocks_own_nextn_head_over_file_level_output() {
4362 assert_eq!(
4363 draft_head_tensor(present(STEP37_DRAFTER), 45),
4364 "blk.45.nextn.shared_head_head.weight"
4365 );
4366 }
4367
4368 #[test]
4372 fn each_nextn_block_selects_its_own_head() {
4373 for n in 45..=47u32 {
4374 assert_eq!(
4375 draft_head_tensor(present(STEP37_DRAFTER), n),
4376 format!("blk.{n}.nextn.shared_head_head.weight")
4377 );
4378 }
4379 }
4380
4381 #[test]
4385 fn draft_without_a_nextn_head_falls_back_to_file_level_output() {
4386 let fr_spec: &[&str] = &["output.weight", "output_norm.weight", "d2t.weight"];
4387 assert_eq!(draft_head_tensor(present(fr_spec), 45), "output.weight");
4388 }
4389
4390 #[test]
4395 fn legacy_shared_head_is_probed_but_loses_to_shared_head_head() {
4396 let legacy_only: &[&str] = &["output.weight", "blk.45.nextn.shared_head.weight"];
4397 assert_eq!(
4398 draft_head_tensor(present(legacy_only), 45),
4399 "blk.45.nextn.shared_head.weight"
4400 );
4401
4402 let both: &[&str] = &[
4403 "output.weight",
4404 "blk.45.nextn.shared_head.weight",
4405 "blk.45.nextn.shared_head_head.weight",
4406 ];
4407 assert_eq!(
4408 draft_head_tensor(present(both), 45),
4409 "blk.45.nextn.shared_head_head.weight"
4410 );
4411 }
4412
4413 #[test]
4417 fn a_different_blocks_nextn_head_is_never_borrowed() {
4418 let wrong_block: &[&str] = &[
4419 "output.weight",
4420 "blk.46.nextn.shared_head_head.weight",
4421 "blk.47.nextn.shared_head_head.weight",
4422 ];
4423 assert_eq!(draft_head_tensor(present(wrong_block), 45), "output.weight");
4424 }
4425}