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 geom: Option<DraftGeom>,
2172 pub step35: Option<Step35MtpGeom>,
2177}
2178
2179#[derive(Debug, Clone)]
2193pub struct Step35MtpGeom {
2194 pub il: u32,
2196 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>,
2207}
2208
2209impl Step35MtpGeom {
2210 pub fn from_plan(layer: &memra_gguf::model_plan::LayerPlan) -> Result<Self, String> {
2212 use memra_gguf::model_plan::{ActivationPlan, AttentionPlan};
2213
2214 let (attention, window) = match &layer.attention {
2215 AttentionPlan::Full(attention) => (attention, None),
2216 AttentionPlan::SlidingWindow { attention, window } => (attention, Some(*window)),
2217 other => {
2218 return Err(format!(
2219 "MTP block {} has unsupported tuned attention {other:?}",
2220 layer.index
2221 ));
2222 }
2223 };
2224 if attention.output_gate != memra_gguf::config::AttentionGateKind::SeparateHead {
2225 return Err(format!(
2226 "MTP block {} does not declare a separate attention gate",
2227 layer.index
2228 ));
2229 }
2230 let activation = match &layer.mlp {
2231 MlpPlan::Dense(dense) => &dense.activation,
2232 MlpPlan::Moe(moe) => &moe.activation,
2233 };
2234 let clamp_shexp = match activation {
2235 ActivationPlan::SwiGluClamped { limit } if *limit > 0.0 => Some(*limit),
2236 _ => None,
2237 };
2238 Ok(Step35MtpGeom {
2239 il: layer.index,
2240 n_head: attention.query_heads as usize,
2241 n_head_kv: attention.kv_heads as usize,
2242 n_rot: attention.rope.dimensions as usize,
2243 rope_base: attention.rope.base,
2244 swa: window.is_some(),
2245 window: window.unwrap_or(0) as usize,
2246 clamp_shexp,
2247 })
2248 }
2249}
2250
2251pub struct DraftGeom {
2253 pub d_inner: usize, pub n_head: usize, pub n_head_kv: usize,
2256 pub out_up: GpuTensor, }
2258
2259pub fn draft_head_tensor(has: impl Fn(&str) -> bool, n: u32) -> String {
2268 let own = format!("blk.{n}.nextn.shared_head_head.weight");
2269 if has(&own) {
2270 return own;
2271 }
2272 let legacy = format!("blk.{n}.nextn.shared_head.weight");
2275 if has(&legacy) {
2276 return legacy;
2277 }
2278 "output.weight".to_string()
2280}
2281
2282impl MtpHead {
2283 pub fn load_draft(
2290 e: &Engine,
2291 g: &GgufFile,
2292 main_cfg: &ModelConfig,
2293 ) -> Result<Self, Box<dyn std::error::Error>> {
2294 let src = GgufSource(g);
2295 let dcfg = src.config();
2296 let draft_plan = match memra_gguf::model_packs::for_config(&dcfg) {
2297 Some(pack) => pack.compile_plan(&dcfg)?,
2298 None => memra_gguf::model_plan::ModelPlan::compile(&dcfg)?,
2299 };
2300 let main_plan = match memra_gguf::model_packs::for_config(main_cfg) {
2301 Some(pack) => pack.compile_plan(main_cfg)?,
2302 None => memra_gguf::model_plan::ModelPlan::compile(main_cfg)?,
2303 };
2304 if dcfg.nextn_predict_layers == 0 {
2309 return Err(format!(
2310 "draft GGUF has no nextn_predict_layers (arch {:?}) — not a NextN/MTP regime \
2311 draft; gemma assistant drafters attach via MEMRA_DRAFT, not '+draft'",
2312 g.arch()
2313 )
2314 .into());
2315 }
2316 let n = dcfg.n_layer - dcfg.nextn_predict_layers;
2317 let draft_block = draft_plan
2318 .mtp_blocks
2319 .iter()
2320 .find(|block| block.layer.index == n)
2321 .ok_or_else(|| format!("draft ModelPlan has no MTP block {n}"))?;
2322 let p = |s: &str| format!("blk.{n}.{s}");
2323
2324 let student = src.has(&p("nextn.out_up.weight"));
2328 assert_eq!(dcfg.n_embd, main_cfg.n_embd, "draft n_embd != model n_embd");
2329 assert_eq!(
2330 dcfg.head_dim_k, main_cfg.head_dim_k,
2331 "draft head_dim != model head_dim"
2332 );
2333 let main_sliding_gated = crate::plan_backend::decode_batch_program(&main_plan)
2340 == crate::plan_backend::DecodeBatchProgram::SlidingGatedMoe;
2341 let draft_sliding_gated = crate::plan_backend::decode_batch_program(&draft_plan)
2342 == crate::plan_backend::DecodeBatchProgram::SlidingGatedMoe;
2343 let step35 = match (main_sliding_gated, draft_sliding_gated) {
2344 (true, true) => {
2345 let g = Step35MtpGeom::from_plan(&draft_block.layer)?;
2346 let out_f = |t: &str| -> Option<usize> {
2348 src.find(&p(t))
2349 .and_then(|v| v.ne.get(1).copied())
2350 .map(|x| x as usize)
2351 };
2352 let hd = dcfg.head_dim_k as usize;
2353 let wq_out =
2354 out_f("attn_q.weight").ok_or("step35 draft block has no attn_q.weight")?;
2355 assert_eq!(
2356 wq_out,
2357 g.n_head * hd,
2358 "step35 draft blk.{n}: attn_q out {wq_out} != n_head({}) * head_dim({hd}) — \
2359 the draft file's head_count array disagrees with its own tensors",
2360 g.n_head
2361 );
2362 let wg_out = out_f("attn_gate.weight")
2365 .ok_or("step35 draft block has no attn_gate.weight (head-wise gate)")?;
2366 assert_eq!(
2367 wg_out, g.n_head,
2368 "step35 draft blk.{n}: attn_gate out {wg_out} != n_head({})",
2369 g.n_head
2370 );
2371 assert_eq!(
2375 g.n_head_kv, main_cfg.n_head_kv as usize,
2376 "step35 draft blk.{n} KV heads {} != trunk n_head_kv {} — the MTP scratch \
2377 rows are sized from the trunk cfg, so a differing draft KV width would \
2378 write past the row",
2379 g.n_head_kv, main_cfg.n_head_kv
2380 );
2381 eprintln!(
2382 "[mtp-draft] step35 MTP geometry blk.{n}: n_head={} n_head_kv={} n_rot={} \
2383 rope_base={:.0} swa={} window={}",
2384 g.n_head, g.n_head_kv, g.n_rot, g.rope_base, g.swa, g.window
2385 );
2386 Some(g)
2387 }
2388 (true, false) => {
2389 return Err(format!(
2390 "MEMRA_MTP_DRAFT operations are incompatible with the model's \
2391 sliding-gated-MoE program (draft arch {:?})",
2392 g.arch()
2393 )
2394 .into());
2395 }
2396 (false, true) => {
2397 return Err(
2398 "MEMRA_MTP_DRAFT requires sliding-gated-MoE operations but the model does not"
2399 .into(),
2400 );
2401 }
2402 (false, false) => None,
2403 };
2404 if step35.is_none() && !student {
2405 assert_eq!(dcfg.n_head, main_cfg.n_head, "draft n_head != model n_head");
2408 assert_eq!(
2409 dcfg.n_head_kv, main_cfg.n_head_kv,
2410 "draft n_head_kv != model n_head_kv"
2411 );
2412 }
2413
2414 let head_name = draft_head_tensor(|t| src.has(t), n);
2441 let head = load_t(e, &src, &head_name)?;
2442 let head_norm = match load_opt(e, &src, &p("nextn.shared_head_norm.weight"))? {
2443 Some(t) => Some(t),
2444 None => load_opt(e, &src, "output_norm.weight")?,
2445 };
2446
2447 let d2t: Option<Vec<u32>> = g.find("d2t").map(|t| {
2449 let bytes = g.tensor_data(t);
2450 match t.ggml_type {
2451 GgmlType::I32 => bytes
2452 .chunks_exact(4)
2453 .map(|c| i32::from_le_bytes(c.try_into().unwrap()) as u32)
2454 .collect(),
2455 GgmlType::I64 => bytes
2456 .chunks_exact(8)
2457 .map(|c| i64::from_le_bytes(c.try_into().unwrap()) as u32)
2458 .collect(),
2459 other => panic!("d2t must be I32/I64, got {other:?}"),
2460 }
2461 });
2462 if let Some(map) = &d2t {
2463 assert_eq!(
2464 map.len(),
2465 head.out_features(),
2466 "d2t len {} != draft head rows {}",
2467 map.len(),
2468 head.out_features()
2469 );
2470 let n_vocab = main_cfg.n_vocab as u64;
2471 assert!(
2472 map.iter().all(|&t| (t as u64) < n_vocab),
2473 "d2t contains token id >= model n_vocab {n_vocab}"
2474 );
2475 }
2476 let eh_proj = load_t(e, &src, &p("nextn.eh_proj.weight"))?;
2477 assert_eq!(
2480 eh_proj.in_features(),
2481 2 * main_cfg.n_embd as usize,
2482 "eh_proj in dim != 2*n_embd"
2483 );
2484 let geom = if student {
2485 let out_up = load_t(e, &src, &p("nextn.out_up.weight"))?;
2486 let d_inner = eh_proj.out_features();
2487 assert_eq!(
2488 out_up.out_features(),
2489 main_cfg.n_embd as usize,
2490 "out_up out dim != n_embd"
2491 );
2492 assert_eq!(
2493 out_up.in_features(),
2494 d_inner,
2495 "out_up in dim != eh_proj out dim (d_inner)"
2496 );
2497 assert!(
2498 dcfg.n_head >= 1 && dcfg.n_head_kv >= 1 && dcfg.n_head % dcfg.n_head_kv == 0,
2499 "student head counts malformed ({}/{})",
2500 dcfg.n_head,
2501 dcfg.n_head_kv
2502 );
2503 Some(DraftGeom {
2504 d_inner,
2505 n_head: dcfg.n_head as usize,
2506 n_head_kv: dcfg.n_head_kv as usize,
2507 out_up,
2508 })
2509 } else {
2510 None
2511 };
2512 let blk_prefix = format!("blk.{n}.");
2516 let head_src = head_name.strip_prefix(&blk_prefix).unwrap_or(&head_name);
2517 eprintln!(
2518 "[mtp-draft] external draft head: blk.{n}, source={}, head_vocab={}{}{}",
2519 head_src,
2520 head.out_features(),
2521 if d2t.is_some() {
2522 " (trimmed, d2t map)"
2523 } else {
2524 " (full)"
2525 },
2526 match &geom {
2527 Some(g) => format!(
2528 " (student d_inner={} heads={}/{})",
2529 g.d_inner, g.n_head, g.n_head_kv
2530 ),
2531 None => String::new(),
2532 }
2533 );
2534
2535 let mut resident = ResidentPlan::unsharded(e, &src, &dcfg);
2536 let mut step_runtimes = StepParallelRuntimeRegistry::default();
2537 Ok(MtpHead {
2538 enorm: load_t(e, &src, &p("nextn.enorm.weight"))?,
2539 hnorm: load_t(e, &src, &p("nextn.hnorm.weight"))?,
2540 eh_proj,
2541 attn_norm: load_t(e, &src, &p("attn_norm.weight"))?,
2542 post_attn_norm: load_opt(e, &src, &p("post_attention_norm.weight"))?
2543 .or(load_opt(e, &src, &p("ffn_norm.weight"))?)
2544 .expect("draft NextN block needs post_attention_norm or ffn_norm"),
2545 mixer: load_mixer_kind(
2546 e,
2547 &src,
2548 &dcfg,
2549 n,
2550 &draft_block.layer.attention,
2551 &mut step_runtimes,
2552 )?,
2553 ffn: load_ffn(
2554 e,
2555 &src,
2556 &dcfg,
2557 &draft_block.layer.mlp,
2558 n,
2559 None,
2560 &mut resident,
2561 &mut step_runtimes,
2562 )?,
2563 shared_head_norm: head_norm,
2564 shared_head_head: Some(head),
2565 d2t,
2566 geom,
2567 step35,
2568 })
2569 }
2570}
2571
2572pub struct GemmaAux {
2574 pub rope_freqs: Option<Vec<(usize, CudaSlice<f32>)>>,
2577 pub ones: Vec<(usize, CudaSlice<f32>)>,
2580 pub suppress_d: Option<(CudaSlice<i32>, usize)>,
2583 pub e4b: Option<Gemma4E4bModel>,
2585}
2586
2587impl GemmaAux {
2588 pub fn rope_freqs(&self, e: &Engine) -> Option<&CudaSlice<f32>> {
2589 self.rope_freqs.as_ref().map(|copies| {
2590 let dev = e.ctx().ordinal();
2591 &copies
2592 .iter()
2593 .find(|(d, _)| *d == dev)
2594 .unwrap_or_else(|| panic!("gemma4 rope_freqs has no local copy for device {dev}"))
2595 .1
2596 })
2597 }
2598
2599 pub fn ones(&self, e: &Engine) -> &CudaSlice<f32> {
2600 let dev = e.ctx().ordinal();
2601 &self
2602 .ones
2603 .iter()
2604 .find(|(d, _)| *d == dev)
2605 .unwrap_or_else(|| panic!("gemma4 ones has no local copy for device {dev}"))
2606 .1
2607 }
2608}
2609
2610pub struct Step35Aux {
2613 pub rope_freqs: Option<Vec<(usize, CudaSlice<f32>)>>,
2619}
2620
2621impl Step35Aux {
2622 pub fn rope_freqs(&self, e: &Engine) -> Option<&CudaSlice<f32>> {
2623 self.rope_freqs.as_ref().map(|copies| {
2624 let dev = e.ctx().ordinal();
2625 &copies
2626 .iter()
2627 .find(|(d, _)| *d == dev)
2628 .unwrap_or_else(|| panic!("step35 rope_freqs has no local copy for device {dev}"))
2629 .1
2630 })
2631 }
2632}
2633
2634pub struct HybridModel {
2635 pub cfg: ModelConfig,
2636 pub plan: memra_gguf::model_plan::ModelPlan,
2637 pub rewrite_qualifications: Option<memra_gguf::execution_manifest::RewriteQualifications>,
2638 pub embd: EmbedHost,
2639 pub output_norm: GpuTensor,
2640 pub output: GpuTensor,
2641 pub layers: Vec<HybridLayer>,
2642 pub mtp: Option<MtpHead>, pub mtp_extra: Vec<MtpHead>,
2646 pub embd_gpu: std::sync::OnceLock<cudarc::driver::CudaSlice<u8>>,
2649 pub gemma4_aux: Option<GemmaAux>,
2650 pub step35_aux: Option<Step35Aux>,
2652 pub prime_slabs: std::sync::Mutex<
2660 std::collections::HashMap<
2661 usize,
2662 std::sync::Arc<std::sync::Mutex<crate::hybrid_forward::PrimeSlabs>>,
2663 >,
2664 >,
2665 pub(crate) dspark_vgraphs: std::sync::Mutex<Option<crate::spec::DsparkVerifyGraphs>>,
2678 pub(crate) step_grouped_prefill: std::sync::Mutex<StepEpGroupedPrefill>,
2684 pub(crate) step35_token_graph:
2687 std::sync::Mutex<Option<crate::hybrid_forward::Step35TokenGraphState>>,
2688}
2689
2690impl HybridModel {
2691 pub fn install_rewrite_bundle(
2692 &mut self,
2693 bundle: &std::path::Path,
2694 ) -> Result<(), Box<dyn std::error::Error>> {
2695 self.rewrite_qualifications = Some(
2696 memra_gguf::execution_manifest::RewriteQualifications::load(bundle, &self.plan)
2697 .map_err(|error| format!("rewrite qualification: {error}"))?,
2698 );
2699 Ok(())
2700 }
2701
2702 pub fn rewrite_allowed(&self, surface: memra_gguf::execution_manifest::RewriteSurface) -> bool {
2703 self.rewrite_qualifications
2704 .as_ref()
2705 .is_none_or(|qualifications| qualifications.allows(surface))
2706 }
2707
2708 pub fn step_tp_unmaterialized_kv_bytes(
2714 &self,
2715 cache: Option<&crate::cache::Cache>,
2716 capacity: usize,
2717 ) -> Result<Vec<StepTpKvDeviceAdmission>, String> {
2718 if let Some(cache) = cache
2719 && cache.tp_kv.len() < self.layers.len()
2720 {
2721 return Err(format!(
2722 "Step TP admission cache has {} layers, model trunk has {}",
2723 cache.tp_kv.len(),
2724 self.layers.len()
2725 ));
2726 }
2727
2728 let mut by_device: HashMap<usize, usize> = HashMap::new();
2729 for (layer, weights) in self.layers.iter().enumerate() {
2730 let Mixer::Full(attention) = &weights.mixer else {
2731 continue;
2732 };
2733 let Some(tp) = attention
2734 .step_tp_qkv
2735 .as_ref()
2736 .filter(|tp| tp.attention.is_some())
2737 else {
2738 continue;
2739 };
2740 if cache.is_some_and(|cache| cache.tp_kv[layer].is_some()) {
2741 continue;
2742 }
2743 let geometry = self.cfg.full_attention_geometry_at(layer as u32);
2744 let shape = crate::cache::tp_kv_rank_allocation_shape(
2745 geometry.n_head_kv as usize * geometry.head_dim_k as usize,
2746 geometry.n_head_kv as usize * geometry.head_dim_v as usize,
2747 tp.devices.len(),
2748 )?;
2749 let physical_rows = geometry
2750 .window
2751 .map(|window| crate::cache::swa_ring_rows(window as usize, capacity))
2752 .unwrap_or(capacity);
2753 let bytes = shape.allocation_bytes(physical_rows);
2754 for &device in &tp.devices {
2755 let total = by_device.entry(device).or_default();
2756 *total = total.saturating_add(bytes);
2757 }
2758 }
2759
2760 let mut out: Vec<_> = by_device
2761 .into_iter()
2762 .map(|(device, bytes)| StepTpKvDeviceAdmission { device, bytes })
2763 .collect();
2764 out.sort_unstable_by_key(|charge| charge.device);
2765 Ok(out)
2766 }
2767
2768 pub fn step_tp_rank_engine(&self, device: usize) -> Option<&Engine> {
2770 self.layers.iter().find_map(|weights| {
2771 let Mixer::Full(attention) = &weights.mixer else {
2772 return None;
2773 };
2774 let tp = attention.step_tp_qkv.as_ref()?;
2775 let rank = tp
2776 .runtime
2777 .devices()
2778 .iter()
2779 .position(|&rank| rank == device)?;
2780 tp.runtime.rank_engine(rank)
2781 })
2782 }
2783
2784 pub(crate) fn step_tp_runtime_for_layer(
2785 &self,
2786 layer: usize,
2787 ) -> Option<&crate::tp::TpE4m3HostBounce> {
2788 let Mixer::Full(attention) = &self.layers.get(layer)?.mixer else {
2789 return None;
2790 };
2791 let tp = attention.step_tp_qkv.as_ref()?;
2792 tp.attention.as_ref()?;
2793 Some(tp.runtime.as_ref())
2794 }
2795
2796 pub fn decode_batch_program(&self) -> crate::plan_backend::DecodeBatchProgram {
2797 crate::plan_backend::decode_batch_program(&self.plan)
2798 }
2799
2800 pub fn uses_gemma_program(&self) -> bool {
2801 self.decode_batch_program() == crate::plan_backend::DecodeBatchProgram::Gemma
2802 }
2803
2804 pub fn uses_sliding_gated_moe_program(&self) -> bool {
2805 self.decode_batch_program() == crate::plan_backend::DecodeBatchProgram::SlidingGatedMoe
2806 }
2807
2808 pub fn has_plan_operation(&self, operation: memra_gguf::model_plan::OperationKind) -> bool {
2809 self.plan.trunk_operations().contains(&operation)
2810 }
2811
2812 pub fn load(e: &Engine, g: &GgufFile) -> Result<Self, Box<dyn std::error::Error>> {
2814 Self::load_from_source(e, &GgufSource(g))
2815 }
2816
2817 pub fn load_without_mtp(e: &Engine, g: &GgufFile) -> Result<Self, Box<dyn std::error::Error>> {
2820 Self::load_from_source_impl(e, &GgufSource(g), false)
2821 }
2822
2823 pub fn load_from_source(
2827 e: &Engine,
2828 src: &dyn TensorSource,
2829 ) -> Result<Self, Box<dyn std::error::Error>> {
2830 Self::load_from_source_impl(e, src, true)
2831 }
2832
2833 pub fn load_from_source_without_mtp(
2835 e: &Engine,
2836 src: &dyn TensorSource,
2837 ) -> Result<Self, Box<dyn std::error::Error>> {
2838 Self::load_from_source_impl(e, src, false)
2839 }
2840
2841 fn load_from_source_impl(
2842 e: &Engine,
2843 src: &dyn TensorSource,
2844 load_mtp: bool,
2845 ) -> Result<Self, Box<dyn std::error::Error>> {
2846 let cfg = src.config();
2847 let plan = match memra_gguf::model_packs::for_config(&cfg) {
2848 Some(pack) => pack.compile_plan(&cfg)?,
2849 None => memra_gguf::model_plan::ModelPlan::compile(&cfg)?,
2850 };
2851 let batch_program = crate::plan_backend::decode_batch_program(&plan);
2852 let gemma_program = batch_program == crate::plan_backend::DecodeBatchProgram::Gemma;
2853 let sliding_gated_moe_program =
2854 batch_program == crate::plan_backend::DecodeBatchProgram::SlidingGatedMoe;
2855 cfg.validate_attention_gate_layout()?;
2860 if cfg.sigmoid_router().is_some() {
2867 let host_oracle = std::env::var("MEMRA_SIG_ROUTER").as_deref() == Ok("0");
2868 match crate::sigrouter_contract::verify_host_expf() {
2869 Ok(()) => {}
2870 Err(e) if host_oracle => return Err(e.into()),
2871 Err(e) => eprintln!(
2872 "[sigrouter] WARN: host expf probe mismatch ({e}); device routing is \
2873 unaffected, but host-oracle replay/comparison cells are invalid on this host"
2874 ),
2875 }
2876 }
2877 if std::env::var("MEMRA_DRAFT").is_ok() && std::env::var("MEMRA_MMQ_SK").is_err() {
2886 let force = if cfg.n_embd >= 3500 { 0i8 } else { -1i8 };
2887 crate::MMQ_SK_FORCE.store(force, std::sync::atomic::Ordering::Relaxed);
2888 }
2889 crate::KV_FP8_FORCE.store(0, std::sync::atomic::Ordering::Relaxed);
2893
2894 let n_trunk = (cfg.n_layer - cfg.nextn_predict_layers) as usize;
2899 crate::pp::init_model_transport(e, &cfg, n_trunk)?;
2900 let step_parallel = prepare_step_parallel_load(e, src, &cfg, n_trunk)?;
2901 let embd = EmbedHost::from_source(src, "token_embd.weight");
2902 let e_head = crate::pp::layer_engine(e, n_trunk, n_trunk - 1)?;
2906 let output_norm = load_t(e_head, src, "output_norm.weight")?;
2907 let mut output = if src.has("output.weight") {
2909 load_t(e_head, src, "output.weight")?
2910 } else {
2911 load_t(e_head, src, "token_embd.weight")?
2912 };
2913 let mut resident = ResidentPlan::pp(e, src, &cfg, n_trunk)?;
2914 let mut step_runtimes = StepParallelRuntimeRegistry::with_config(step_parallel);
2915
2916 let gguf: Option<&GgufFile> = src.gguf();
2923 let mut spill: Option<crate::spill::SpillCtx> = if cfg
2926 .moe
2927 .as_ref()
2928 .is_some_and(|m| m.expert_count > 0)
2929 && crate::spill::disk_tier_enabled()
2930 && gguf.is_some()
2931 {
2932 let budget = crate::spill::MemBudget::probe(e)?;
2933 let ctx = crate::spill::SpillCtx::open(gguf.unwrap(), &budget)?;
2934 eprintln!(
2935 "[spill] disk tier ON: free_vram={} MiB free_pinnable_ram={} MiB (MemAvailable*resolved_frac)",
2936 budget.free_vram >> 20,
2937 budget.free_pinnable_ram >> 20
2938 );
2939 Some(ctx)
2940 } else {
2941 None
2942 };
2943
2944 let mut layers = Vec::with_capacity(n_trunk);
2947 for il in 0..n_trunk as u32 {
2948 let p = |s: &str| format!("blk.{il}.{s}");
2949 let layer_plan = plan
2950 .layers
2951 .get(il as usize)
2952 .ok_or_else(|| format!("ModelPlan has no trunk layer {il}"))?;
2953 let e = crate::pp::layer_engine(e, n_trunk, il as usize)?;
2957 layers.push(HybridLayer {
2959 attn_norm: load_t(e, src, &p("attn_norm.weight"))?,
2960 post_attn_norm: load_opt(e, src, &p("post_attention_norm.weight"))?
2961 .or(load_opt(e, src, &p("ffn_norm.weight"))?)
2962 .expect("need post_attention_norm or ffn_norm"),
2963 mixer: {
2964 let g4_shared = cfg.gemma4.as_ref().map(|g| g.shared_kv_layers).unwrap_or(0);
2968 let kv_from = n_trunk as u32 - g4_shared;
2969 if g4_shared > 0
2970 && il >= kv_from
2971 && !src.has(&format!("blk.{il}.attn_k.weight"))
2972 {
2973 let g4 = cfg.gemma4.as_ref().unwrap();
2974 let swa = g4.swa_pattern.get(il as usize).copied().unwrap_or(true);
2975 let tgt = kv_from - if swa { 2 } else { 1 };
2976 let tp = |s: &str| format!("blk.{tgt}.{s}");
2977 Mixer::Full(FullAttnLayer {
2978 wq: load_t(e, src, &p("attn_q.weight"))?,
2979 wk: load_t(e, src, &tp("attn_k.weight"))?,
2980 wv: load_t(e, src, &tp("attn_v.weight"))?,
2981 wo: load_t(e, src, &p("attn_output.weight"))?,
2982 q_norm: load_t(e, src, &p("attn_q_norm.weight"))?,
2983 k_norm: load_t(e, src, &tp("attn_k_norm.weight"))?,
2984 attn_gate: None, step_tp_qkv: None,
2986 })
2987 } else {
2988 load_mixer_kind(
2989 e,
2990 src,
2991 &cfg,
2992 il,
2993 &layer_plan.attention,
2994 &mut step_runtimes,
2995 )?
2996 }
2997 },
2998 ffn: load_ffn(
2999 e,
3000 src,
3001 &cfg,
3002 &layer_plan.mlp,
3003 il,
3004 spill.as_mut().map(|c| (gguf.unwrap(), c)),
3005 &mut resident,
3006 &mut step_runtimes,
3007 )?,
3008 gemma4: if gemma_program {
3009 let scalar = |n: &str| -> f32 {
3010 let t = src.find(&p(n)).unwrap_or_else(|| panic!("missing {n}"));
3011 memra_gguf::dequant::dequantize(t.ggml_type, &t.bytes, 1)[0]
3012 };
3013 let vecf = |n: &str| -> Vec<f32> {
3014 let t = src.find(&p(n)).unwrap_or_else(|| panic!("missing {n}"));
3015 memra_gguf::dequant::dequantize(
3016 t.ggml_type,
3017 &t.bytes,
3018 t.ne.iter().product::<u64>() as usize,
3019 )
3020 };
3021 let moe_bits = if src.find(&p("ffn_gate_inp.scale")).is_some() {
3022 Some(crate::hybrid::Gemma4MoeBits {
3023 post_ffw_norm_1: load_t(e, src, &p("post_ffw_norm_1.weight"))?,
3024 pre_ffw_norm_2: load_t(e, src, &p("pre_ffw_norm_2.weight"))?,
3025 post_ffw_norm_2: load_t(e, src, &p("post_ffw_norm_2.weight"))?,
3026 shared_gate: load_t(e, src, &p("ffn_gate.weight"))?,
3027 shared_up: load_t(e, src, &p("ffn_up.weight"))?,
3028 shared_down: load_t(e, src, &p("ffn_down.weight"))?,
3029 router_scale_pre: {
3030 let inv = 1.0 / (cfg.n_embd as f32).sqrt();
3031 let v: Vec<f32> =
3032 vecf("ffn_gate_inp.scale").iter().map(|x| x * inv).collect();
3033 e.htod(&v)?
3034 },
3035 per_expert_scale: vecf("ffn_down_exps.scale"),
3036 per_expert_scale_d: e.htod(&vecf("ffn_down_exps.scale"))?,
3037 })
3038 } else {
3039 None
3040 };
3041 let e4b = if src.has(&p("inp_gate.weight")) {
3043 let g4 = cfg.gemma4.as_ref().unwrap();
3044 let kv_from = n_trunk as u32 - g4.shared_kv_layers;
3045 let kv_share = if g4.shared_kv_layers > 0 && il >= kv_from {
3046 let swa = g4.swa_pattern.get(il as usize).copied().unwrap_or(true);
3047 Some(kv_from - if swa { 2 } else { 1 })
3048 } else {
3049 None
3050 };
3051 Some(crate::hybrid::Gemma4E4bLayer {
3052 inp_gate: load_t(e, src, &p("inp_gate.weight"))?,
3053 proj: load_t(e, src, &p("proj.weight"))?,
3054 post_norm: load_t(e, src, &p("post_norm.weight"))?,
3055 kv_share,
3056 qkv_cat: None, })
3058 } else {
3059 None
3060 };
3061 Some(Gemma4LayerBits {
3062 ffn_norm: load_t(e, src, &p("ffn_norm.weight"))?,
3063 post_ffw_norm: load_t(e, src, &p("post_ffw_norm.weight"))?,
3064 moe_bits,
3065 layer_scale: scalar("layer_output_scale.weight"),
3066 e4b,
3067 })
3068 } else {
3069 None
3070 },
3071 });
3072 }
3073
3074 let external_mtp_requested =
3078 load_mtp && std::env::var("MEMRA_MTP_DRAFT").is_ok_and(|path| !path.is_empty());
3079 let trim_mtp_requested = load_mtp
3080 && !crate::model::full_prec_enabled()
3081 && std::env::var("MEMRA_FRSPEC_TRIM").is_ok_and(|path| !path.is_empty());
3082 let embedded_head_count = if external_mtp_requested {
3083 0
3084 } else if trim_mtp_requested {
3085 1
3086 } else {
3087 cfg.nextn_predict_layers
3088 };
3089 let mut embedded_mtp = Vec::new();
3090 if load_mtp && embedded_head_count > 0 {
3091 for offset in 0..embedded_head_count {
3092 let n = n_trunk as u32 + offset;
3093 let p = |s: &str| format!("blk.{n}.{s}");
3094 let mtp_plan = plan
3095 .mtp_blocks
3096 .iter()
3097 .find(|block| block.layer.index == n)
3098 .ok_or_else(|| format!("ModelPlan has no embedded MTP block {n}"))?;
3099 if !src.has(&p("nextn.eh_proj.weight")) {
3100 if offset == 0 {
3101 break;
3102 }
3103 return Err(format!(
3104 "embedded MTP chain declares {} heads but blk.{n} has no \
3105 nextn.eh_proj.weight",
3106 cfg.nextn_predict_layers
3107 )
3108 .into());
3109 }
3110 embedded_mtp.push(MtpHead {
3111 enorm: load_t(e, src, &p("nextn.enorm.weight"))?,
3112 hnorm: load_t(e, src, &p("nextn.hnorm.weight"))?,
3113 eh_proj: load_t(e, src, &p("nextn.eh_proj.weight"))?,
3114 attn_norm: load_t(e, src, &p("attn_norm.weight"))?,
3115 post_attn_norm: load_opt(e, src, &p("post_attention_norm.weight"))?
3116 .or(load_opt(e, src, &p("ffn_norm.weight"))?)
3117 .expect("MTP block needs post_attention_norm or ffn_norm"),
3118 mixer: load_mixer_kind(
3119 e,
3120 src,
3121 &cfg,
3122 n,
3123 &mtp_plan.layer.attention,
3124 &mut step_runtimes,
3125 )?,
3126 ffn: load_ffn(
3127 e,
3128 src,
3129 &cfg,
3130 &mtp_plan.layer.mlp,
3131 n,
3132 spill.as_mut().map(|c| (gguf.unwrap(), c)),
3133 &mut resident,
3134 &mut step_runtimes,
3135 )?,
3136 shared_head_norm: load_opt(e, src, &p("nextn.shared_head_norm.weight"))?,
3137 shared_head_head: load_opt(e, src, &p("nextn.shared_head_head.weight"))?
3146 .or(load_opt(e, src, &p("nextn.shared_head.weight"))?),
3147 d2t: None,
3148 geom: None,
3149 step35: if sliding_gated_moe_program {
3150 Some(Step35MtpGeom::from_plan(&mtp_plan.layer)?)
3151 } else {
3152 None
3153 },
3154 });
3155 }
3156 }
3157 let mut embedded_mtp = embedded_mtp.into_iter();
3158 let mut mtp = embedded_mtp.next();
3159 let mut mtp_extra: Vec<MtpHead> = embedded_mtp.collect();
3160
3161 mtp = if load_mtp {
3165 match std::env::var("MEMRA_MTP_DRAFT") {
3166 Ok(path) if !path.is_empty() => {
3167 eprintln!("[mtp-draft] loading external MTP draft: {path}");
3168 let dg = GgufFile::open(&path)?;
3169 mtp_extra.clear();
3170 Some(MtpHead::load_draft(e, &dg, &cfg)?)
3171 }
3172 _ => mtp,
3173 }
3174 } else {
3175 None
3176 };
3177
3178 let trim_env = if load_mtp {
3189 std::env::var("MEMRA_FRSPEC_TRIM")
3190 } else {
3191 Err(std::env::VarError::NotPresent)
3192 };
3193 if crate::model::full_prec_enabled()
3194 && trim_env.as_deref().map(|p| !p.is_empty()).unwrap_or(false)
3195 {
3196 eprintln!(
3197 "[frspec-trim] DISABLED under MEMRA_FULL_PREC — using the natural full MTP head"
3198 );
3199 }
3200 mtp = match (
3201 if crate::model::full_prec_enabled() {
3202 Err(std::env::VarError::NotPresent)
3203 } else {
3204 trim_env
3205 },
3206 mtp,
3207 ) {
3208 (Ok(path), Some(mut head)) if !path.is_empty() => {
3209 let d2t: Vec<u32> = if path.ends_with(".txt") {
3213 std::fs::read_to_string(&path)?
3214 .lines()
3215 .filter_map(|l| l.trim().parse::<u32>().ok())
3216 .collect()
3217 } else {
3218 let tg = GgufFile::open(&path)?;
3219 let d2t_t = tg
3220 .find("d2t")
3221 .expect("MEMRA_FRSPEC_TRIM file has no d2t tensor");
3222 let d2t_bytes = tg.tensor_data(d2t_t);
3223 match d2t_t.ggml_type {
3224 GgmlType::I32 => d2t_bytes
3225 .chunks_exact(4)
3226 .map(|c| i32::from_le_bytes(c.try_into().unwrap()) as u32)
3227 .collect(),
3228 GgmlType::I64 => d2t_bytes
3229 .chunks_exact(8)
3230 .map(|c| i64::from_le_bytes(c.try_into().unwrap()) as u32)
3231 .collect(),
3232 other => panic!("d2t must be I32/I64, got {other:?}"),
3233 }
3234 };
3235 let v = src
3236 .find("output.weight")
3237 .or_else(|| src.find("token_embd.weight"))
3238 .expect("model has no output.weight for FR-Spec trim");
3239 let out_f = v.ne[1] as usize;
3240 let row_bytes = v.bytes.len() / out_f;
3241 assert!(
3242 d2t.iter().all(|&t| (t as usize) < out_f),
3243 "d2t token id >= lm_head rows {out_f}"
3244 );
3245 let mut gathered = Vec::with_capacity(d2t.len() * row_bytes);
3246 for &t in &d2t {
3247 let off = t as usize * row_bytes;
3248 gathered.extend_from_slice(&v.bytes[off..off + row_bytes]);
3249 }
3250 let trimmed = GpuTensor::from_quant_bytes(
3251 e,
3252 &gathered,
3253 v.ggml_type,
3254 v.ne[0],
3255 d2t.len() as u64,
3256 match src.find("output.scale") {
3258 Some(sv) => f32::from_le_bytes(sv.bytes[..4].try_into().unwrap()),
3259 None => 1.0,
3260 },
3261 )?;
3262 eprintln!(
3263 "[frspec-trim] self-trimmed head: {} rows of main output.weight ({:?})",
3264 d2t.len(),
3265 v.ggml_type
3266 );
3267 head.shared_head_head = Some(trimmed);
3268 head.d2t = Some(d2t);
3269 Some(head)
3270 }
3271 (_, m) => m,
3272 };
3273 if mtp.as_ref().is_some_and(|head| head.d2t.is_some()) {
3274 mtp_extra.clear();
3275 }
3276 if !mtp_extra.is_empty() {
3277 if plan.draft_source != memra_gguf::model_plan::DraftSourcePlan::Embedded
3278 || plan.mtp_blocks.len() != 1 + mtp_extra.len()
3279 || plan
3280 .mtp_blocks
3281 .iter()
3282 .any(|block| !matches!(block.layer.mlp, MlpPlan::Dense(_)))
3283 || mtp
3284 .iter()
3285 .chain(mtp_extra.iter())
3286 .any(|head| !matches!(head.ffn, Ffn::Dense { .. }))
3287 {
3288 return Err(
3289 "multi-head MTP requires embedded dense canonical blocks and matching loaded heads"
3290 .into(),
3291 );
3292 }
3293 eprintln!(
3294 "[mtp-draft] embedded chain: heads={} blocks={}..={} scratch=per-head",
3295 1 + mtp_extra.len(),
3296 n_trunk,
3297 n_trunk + mtp_extra.len()
3298 );
3299 }
3300
3301 if let Some(ctx) = spill.as_ref() {
3302 eprintln!(
3303 "[spill] experts placed: {} pinned (Tier 1), {} mmap'd from disk (Tier 2, {} MiB)",
3304 ctx.n_pinned,
3305 ctx.n_mmap,
3306 ctx.mmap_bytes >> 20
3307 );
3308 }
3309
3310 if cfg.n_head_kv > 0 && cfg.n_head / cfg.n_head_kv > 8 {
3324 crate::FA_V4_MAX_DEFAULT.store(0, std::sync::atomic::Ordering::Relaxed);
3325 eprintln!(
3326 "[fa] v4 decode family disabled: gqa {} > fa_v4_smem capacity 8 (v3 lane serves)",
3327 cfg.n_head / cfg.n_head_kv
3328 );
3329 }
3330
3331 if gemma_program {
3332 crate::FA_VEC_MIN_DEFAULT.store(1, std::sync::atomic::Ordering::Relaxed);
3334 let real_moe = plan
3337 .trunk_operations()
3338 .contains(&memra_gguf::model_plan::OperationKind::MoeMlp);
3339 crate::FA_SPW_DEFAULT.store(
3340 if real_moe { 32 } else { 64 },
3341 std::sync::atomic::Ordering::Relaxed,
3342 );
3343 crate::FA_SP512_DEFAULT.store(
3345 if real_moe { 16 } else { 32 },
3346 std::sync::atomic::Ordering::Relaxed,
3347 );
3348 crate::FUSED_MR1_DEFAULT.store(!real_moe, std::sync::atomic::Ordering::Relaxed);
3358 crate::RMS_BLOCK_DEFAULT.store(1024, std::sync::atomic::Ordering::Relaxed);
3360 crate::FA_SP_GEMMA.store(true, std::sync::atomic::Ordering::Relaxed);
3362 }
3366 let force_embd_gpu = gemma_program;
3369 let gemma4_aux = if gemma_program {
3370 let rope_freqs = match src.find("rope_freqs.weight") {
3371 Some(t) => {
3372 let host = memra_gguf::dequant::dequantize(
3373 t.ggml_type,
3374 &t.bytes,
3375 t.ne.iter().product::<u64>() as usize,
3376 );
3377 let mut copies = Vec::new();
3378 if let Some(fence) = crate::pp::pp_cuts(n_trunk) {
3379 for s in 0..fence.len() - 1 {
3380 let owner = crate::pp::layer_engine(e, n_trunk, fence[s])?;
3381 let dev = owner.ctx().ordinal();
3382 if copies.iter().all(|(d, _)| *d != dev) {
3383 copies.push((dev, owner.htod(&host)?));
3384 }
3385 }
3386 } else {
3387 copies.push((e.ctx().ordinal(), e.htod(&host)?));
3388 }
3389 Some(copies)
3390 }
3391 None => {
3399 let g4 = cfg.gemma4.as_ref().unwrap();
3400 let n = (g4.rope_dims_global / 2) as usize;
3401 let keep =
3402 ((n as f32) * g4.partial_rotary_global.clamp(0.0, 1.0)).round() as usize;
3403 let host: Vec<f32> = (0..n)
3404 .map(|i| if i < keep { 1.0 } else { 1.0e30 })
3405 .collect();
3406 eprintln!(
3407 "[gemma4] rope_freqs.weight synthesized ({n} factors, first {keep} \
3408 rotate; source ships none — native checkpoint)"
3409 );
3410 let mut copies = Vec::new();
3411 if let Some(fence) = crate::pp::pp_cuts(n_trunk) {
3412 for s in 0..fence.len() - 1 {
3413 let owner = crate::pp::layer_engine(e, n_trunk, fence[s])?;
3414 let dev = owner.ctx().ordinal();
3415 if copies.iter().all(|(d, _)| *d != dev) {
3416 copies.push((dev, owner.htod(&host)?));
3417 }
3418 }
3419 } else {
3420 copies.push((e.ctx().ordinal(), e.htod(&host)?));
3421 }
3422 Some(copies)
3423 }
3424 };
3425 let e4b = match src.find("per_layer_token_embd.weight") {
3427 Some(t) => {
3428 let n_epl = cfg
3429 .gemma4
3430 .as_ref()
3431 .map(|g| g.n_embd_per_layer as usize)
3432 .unwrap_or(0);
3433 let row = t.ne[0] as usize; let row_bytes = t.bytes.len() / (t.ne[1] as usize);
3435 eprintln!(
3436 "[gemma4-e4b] per-layer-embed model detected (n_epl={n_epl}, row {row}) — \
3437 first-light forward (eager decode + prime); dc/graph/spec unwired \
3438 (HANDOVER-E4B.md)"
3439 );
3440 Some(crate::hybrid::Gemma4E4bModel {
3441 tok_tbl_gpu: std::sync::OnceLock::new(),
3442 tok_embd_bytes: t.bytes.to_vec(),
3443 tok_embd_qt: match t.ggml_type {
3444 memra_gguf::GgmlType::Q6_K => crate::QT_Q6_K,
3445 memra_gguf::GgmlType::Q8_0 => crate::QT_Q8_0,
3446 other => panic!("e4b per-layer tok embd: unhandled dtype {other:?}"),
3447 },
3448 tok_embd_row_bytes: row_bytes,
3449 model_proj: load_t(e, src, "per_layer_model_proj.weight")?,
3450 proj_norm: load_t(e, src, "per_layer_proj_norm.weight")?,
3451 n_epl,
3452 })
3453 }
3454 None => None,
3455 };
3456 let suppress_d = {
3457 let sup = &cfg.gemma4.as_ref().unwrap().suppress_tokens;
3458 if sup.is_empty() {
3459 None
3460 } else {
3461 let ids: Vec<i32> = sup.iter().map(|&x| x as i32).collect();
3462 eprintln!(
3463 "[gemma4] suppress_tokens: {} ids masked at sampling",
3464 ids.len()
3465 );
3466 Some((e.htod_i32(&ids)?, ids.len()))
3467 }
3468 };
3469 let ones_host = [1.0f32; 512];
3470 let mut ones = Vec::new();
3471 if let Some(fence) = crate::pp::pp_cuts(n_trunk) {
3472 for s in 0..fence.len() - 1 {
3473 let owner = crate::pp::layer_engine(e, n_trunk, fence[s])?;
3474 let dev = owner.ctx().ordinal();
3475 if ones.iter().all(|(d, _)| *d != dev) {
3476 ones.push((dev, owner.htod(&ones_host)?));
3477 }
3478 }
3479 } else {
3480 ones.push((e.ctx().ordinal(), e.htod(&ones_host)?));
3481 }
3482 Some(GemmaAux {
3483 rope_freqs,
3484 ones,
3485 suppress_d,
3486 e4b,
3487 })
3488 } else {
3489 None
3490 };
3491 let step35_aux = if sliding_gated_moe_program {
3495 let rope_freqs = match src.find("rope_freqs.weight") {
3496 Some(t) => {
3497 let host = memra_gguf::dequant::dequantize(
3498 t.ggml_type,
3499 &t.bytes,
3500 t.ne.iter().product::<u64>() as usize,
3501 );
3502 let mut copies = Vec::new();
3503 if let Some(fence) = crate::pp::pp_cuts(n_trunk) {
3504 for s in 0..fence.len() - 1 {
3505 let owner = crate::pp::layer_engine(e, n_trunk, fence[s])?;
3506 let dev = owner.ctx().ordinal();
3507 if copies.iter().all(|(d, _)| *d != dev) {
3508 copies.push((dev, owner.htod(&host)?));
3509 }
3510 }
3511 } else {
3512 copies.push((e.ctx().ordinal(), e.htod(&host)?));
3513 }
3514 Some(copies)
3515 }
3516 None => None,
3517 };
3518 Some(Step35Aux { rope_freqs })
3519 } else {
3520 None
3521 };
3522 let mut layers = layers;
3523 {
3530 let q8rp_on = match std::env::var("MEMRA_Q8RP").as_deref() {
3531 Ok("0") => false,
3532 Ok(_) => true,
3533 Err(_) => {
3541 cfg!(memra_hopper_mma) || {
3542 let q8b = |w: &crate::model::GpuTensor| -> usize {
3543 match w {
3544 crate::model::GpuTensor::Quant {
3545 bytes,
3546 qtype,
3547 row_bytes,
3548 ne,
3549 rp4: None,
3550 ..
3551 } if *qtype == crate::QT_Q8_0
3552 && ne.len() == 2
3553 && (ne[0] as usize) % 32 == 0
3554 && *row_bytes == (ne[0] as usize / 32) * 34 =>
3555 {
3556 bytes.len()
3557 }
3558 _ => 0,
3559 }
3560 };
3561 let mut need = q8b(&output);
3562 for layer in layers.iter() {
3563 match &layer.mixer {
3564 Mixer::Full(fa) => {
3565 for w in [&fa.wq, &fa.wk, &fa.wv, &fa.wo] {
3566 need += q8b(w);
3567 }
3568 }
3569 Mixer::Linear(la) => {
3570 for w in [
3571 &la.wqkv,
3572 &la.wqkv_gate,
3573 &la.ssm_beta,
3574 &la.ssm_alpha,
3575 &la.ssm_out,
3576 ] {
3577 need += q8b(w);
3578 }
3579 }
3580 Mixer::Mla(_) => {}
3581 }
3582 if let Ffn::Dense {
3583 ffn_gate,
3584 ffn_up,
3585 ffn_down,
3586 } = &layer.ffn
3587 {
3588 for w in [ffn_gate, ffn_up, ffn_down] {
3589 need += q8b(w);
3590 }
3591 }
3592 }
3593 need > 0
3594 && e.ctx()
3595 .mem_get_info()
3596 .map(|(free, _)| free >= need + (8usize << 30))
3597 .unwrap_or(false)
3598 }
3599 }
3600 };
3601 let kqrp_on = crate::Engine::kqrp_enabled() || {
3611 std::env::var("MEMRA_KQRP").is_err() && {
3612 let kqb = |w: &crate::model::GpuTensor| -> usize {
3613 match w {
3614 crate::model::GpuTensor::Quant {
3615 bytes,
3616 qtype,
3617 row_bytes,
3618 ne,
3619 rp4: None,
3620 ..
3621 } if ne.len() == 2 && (ne[0] as usize) % 256 == 0 => {
3622 let sb = if *qtype == crate::QT_Q4_K {
3623 144
3624 } else if *qtype == crate::QT_Q6_K {
3625 210
3626 } else {
3627 return 0;
3628 };
3629 if *row_bytes == (ne[0] as usize / 256) * sb {
3630 bytes.len()
3631 } else {
3632 0
3633 }
3634 }
3635 _ => 0,
3636 }
3637 };
3638 let mut need = kqb(&output);
3639 for layer in layers.iter() {
3640 if let Mixer::Full(fa) = &layer.mixer {
3641 for w in [&fa.wq, &fa.wk, &fa.wv, &fa.wo] {
3642 need += kqb(w);
3643 }
3644 }
3645 if let Ffn::Dense {
3646 ffn_gate,
3647 ffn_up,
3648 ffn_down,
3649 } = &layer.ffn
3650 {
3651 for w in [ffn_gate, ffn_up, ffn_down] {
3652 need += kqb(w);
3653 }
3654 }
3655 }
3656 need > 0
3657 && e.ctx()
3658 .mem_get_info()
3659 .map(|(free, _)| free >= need + (8usize << 30))
3660 .unwrap_or(false)
3661 }
3662 };
3663 if q8rp_on || kqrp_on {
3664 let f16_model_ok = gemma_program
3671 || plan
3672 .trunk_operations()
3673 .contains(&memra_gguf::model_plan::OperationKind::MoeMlp)
3674 || std::env::var("MEMRA_PP_F16").as_deref() == Ok("1");
3675 let mut nmir = 0usize;
3676 let mut mir = |e_ref: &crate::Engine,
3680 w: &mut crate::model::GpuTensor|
3681 -> Result<(), Box<dyn std::error::Error>> {
3682 let before = matches!(w, crate::model::GpuTensor::Quant { rp4: Some(_), .. });
3683 if q8rp_on {
3684 e_ref.build_q8_rp4(w)?;
3685 }
3686 if kqrp_on {
3687 e_ref.build_q4k_rp4(w)?;
3688 e_ref.build_q6k_rp4(w)?;
3689 }
3690 let q6k = matches!(w, crate::model::GpuTensor::Quant { qtype, .. }
3695 if *qtype == crate::QT_Q6_K);
3696 if q8rp_on && crate::f16_ffi::pp_f16_enabled() && (f16_model_ok || q6k) {
3697 e_ref.build_q8_f16(w)?;
3698 }
3699 if !before && matches!(w, crate::model::GpuTensor::Quant { rp4: Some(_), .. }) {
3700 nmir += 1;
3701 }
3702 Ok(())
3703 };
3704 for (il, layer) in layers.iter_mut().enumerate() {
3705 let el = crate::pp::layer_engine(e, n_trunk, il)?;
3706 match &mut layer.mixer {
3707 Mixer::Full(fa) => {
3708 for w in [&mut fa.wq, &mut fa.wk, &mut fa.wv, &mut fa.wo] {
3709 mir(el, w)?;
3710 }
3711 }
3712 Mixer::Linear(la) => {
3713 for w in [
3714 &mut la.wqkv,
3715 &mut la.wqkv_gate,
3716 &mut la.ssm_beta,
3717 &mut la.ssm_alpha,
3718 &mut la.ssm_out,
3719 ] {
3720 mir(el, w)?;
3721 }
3722 }
3723 Mixer::Mla(_) => {}
3726 }
3727 if let Ffn::Dense {
3728 ffn_gate,
3729 ffn_up,
3730 ffn_down,
3731 } = &mut layer.ffn
3732 {
3733 for w in [ffn_gate, ffn_up, ffn_down] {
3734 mir(el, w)?;
3735 }
3736 }
3737 }
3738 mir(e_head, &mut output)?;
3739 if nmir > 0 {
3740 eprintln!("[q8rp] split-plane decode mirrors built: {nmir} tensors");
3741 }
3742 if q8rp_on && crate::f16_ffi::pp_f16_enabled() {
3757 for (want, tag) in [(crate::QT_Q4_K, "q4kf16"), (crate::QT_Q5_K, "q5kf16")] {
3758 let (mut n4, mut b4) = (0usize, 0usize);
3759 let mut mirk =
3760 |e_ref: &crate::Engine,
3761 w: &mut crate::model::GpuTensor|
3762 -> Result<(), Box<dyn std::error::Error>> {
3763 if matches!(w, crate::model::GpuTensor::Quant { qtype, f16: None, .. }
3764 if *qtype == want)
3765 {
3766 e_ref.build_q8_f16(w)?;
3767 if let crate::model::GpuTensor::Quant { f16: Some(m), .. } = w {
3768 n4 += 1;
3769 b4 += m.len();
3770 }
3771 }
3772 Ok(())
3773 };
3774 for (il, layer) in layers.iter_mut().enumerate() {
3775 let el = crate::pp::layer_engine(e, n_trunk, il)?;
3776 match &mut layer.mixer {
3777 Mixer::Full(fa) => {
3778 for w in [&mut fa.wq, &mut fa.wk, &mut fa.wv, &mut fa.wo] {
3779 mirk(el, w)?;
3780 }
3781 }
3782 Mixer::Linear(la) => {
3783 for w in [
3784 &mut la.wqkv,
3785 &mut la.wqkv_gate,
3786 &mut la.ssm_beta,
3787 &mut la.ssm_alpha,
3788 &mut la.ssm_out,
3789 ] {
3790 mirk(el, w)?;
3791 }
3792 }
3793 Mixer::Mla(_) => {} }
3795 if let Ffn::Dense {
3796 ffn_gate,
3797 ffn_up,
3798 ffn_down,
3799 } = &mut layer.ffn
3800 {
3801 for w in [ffn_gate, ffn_up, ffn_down] {
3802 mirk(el, w)?;
3803 }
3804 }
3805 }
3806 mirk(e_head, &mut output)?;
3807 if n4 > 0 {
3808 eprintln!(
3809 "[{tag}] prefill fp16 mirrors built: {n4} tensors \
3810 ({} MB)",
3811 b4 >> 20
3812 );
3813 }
3814 }
3815 }
3816 }
3817 }
3818 if gemma_program && crate::Engine::q4rp_enabled() {
3825 let mut nmir = 0usize;
3826 for (il, layer) in layers.iter_mut().enumerate() {
3827 let e = crate::pp::layer_engine(e, n_trunk, il)?;
3829 let is_moe26 = layer.gemma4.as_ref().is_some_and(|g| g.moe_bits.is_some());
3838 let is_e4b = layer.gemma4.as_ref().is_some_and(|g| g.e4b.is_some());
3839 if !(is_moe26 || is_e4b) {
3840 continue;
3841 }
3842 if let Mixer::Full(fa) = &mut layer.mixer {
3843 for w in [&mut fa.wq, &mut fa.wk, &mut fa.wv, &mut fa.wo] {
3844 e.build_q4_rp4(w)?;
3845 nmir += 1;
3846 }
3847 }
3848 if is_e4b {
3849 let own_kv = layer
3851 .gemma4
3852 .as_ref()
3853 .unwrap()
3854 .e4b
3855 .as_ref()
3856 .is_some_and(|e4| e4.kv_share.is_none());
3857 if own_kv {
3858 if let Mixer::Full(fa) = &layer.mixer {
3859 if let Some(mut cat) = e.build_q4_out_concat3(&fa.wq, &fa.wk, &fa.wv)? {
3860 e.build_q4_rp4(&mut cat)?;
3861 nmir += 1;
3862 layer.gemma4.as_mut().unwrap().e4b.as_mut().unwrap().qkv_cat =
3863 Some(cat);
3864 }
3865 }
3866 }
3867 if let Ffn::Dense {
3868 ffn_gate,
3869 ffn_up,
3870 ffn_down,
3871 } = &mut layer.ffn
3872 {
3873 for w in [ffn_gate, ffn_up, ffn_down] {
3874 e.build_q4_rp4(w)?;
3875 nmir += 1;
3876 }
3877 }
3878 let e4 = layer.gemma4.as_mut().unwrap().e4b.as_mut().unwrap();
3879 for w in [&mut e4.inp_gate, &mut e4.proj] {
3880 e.build_q4_rp4(w)?;
3881 nmir += 1;
3882 }
3883 }
3884 if let Some(mb) = layer.gemma4.as_mut().unwrap().moe_bits.as_mut() {
3885 for w in [&mut mb.shared_gate, &mut mb.shared_up, &mut mb.shared_down] {
3886 e.build_q4_rp4(w)?;
3887 nmir += 1;
3888 }
3889 }
3890 }
3891 if nmir > 0 {
3892 eprintln!("[q4rp] split-plane decode mirrors built: {nmir} trunk tensors");
3893 }
3894 let fast_on = std::env::var("MEMRA_FAST").as_deref() != Ok("0");
3901 if fast_on {
3902 let mut nswap = 0usize;
3903 let mut nf16 = 0usize;
3904 let q4f16_model_ok = matches!(cfg.n_embd, 3840 | 5376); if let Ok(v) = std::env::var("MEMRA_Q4F16") {
3923 if v != "0" && v != "1" {
3924 return Err(format!(
3925 "MEMRA_Q4F16={v} is not 0 or 1 — this env selects the prefill \
3926 ARITHMETIC (fp16 mirrors vs int8 MMQ) and must never be guessed"
3927 )
3928 .into());
3929 }
3930 }
3931 let f16_need = {
3932 let f16b = |w: &crate::model::GpuTensor| -> usize {
3933 match w {
3934 crate::model::GpuTensor::Quant {
3935 qtype,
3936 ne,
3937 f16: None,
3938 ..
3939 } if ne.len() == 2
3940 && matches!(
3941 *qtype,
3942 crate::QT_Q8_0
3943 | crate::QT_Q4_0
3944 | crate::QT_Q6_K
3945 | crate::QT_Q4_K
3946 | crate::QT_Q5_K
3947 ) =>
3948 {
3949 (ne[0] as usize) * (ne[1] as usize) * 2
3950 }
3951 _ => 0,
3952 }
3953 };
3954 let mut need = 0usize;
3955 for layer in layers.iter() {
3956 if layer.gemma4.as_ref().is_none_or(|g| g.moe_bits.is_some()) {
3957 continue;
3958 }
3959 if let Mixer::Full(fa) = &layer.mixer {
3960 for w in [&fa.wq, &fa.wk, &fa.wv, &fa.wo] {
3961 need += f16b(w);
3962 }
3963 }
3964 if let Ffn::Dense {
3965 ffn_gate,
3966 ffn_up,
3967 ffn_down,
3968 } = &layer.ffn
3969 {
3970 for w in [ffn_gate, ffn_up, ffn_down] {
3971 need += f16b(w);
3972 }
3973 }
3974 }
3975 need
3976 };
3977 let f16_free = e.ctx().mem_get_info().map(|(free, _)| free).unwrap_or(0);
3978 let f16_auto = q4f16_model_ok
3979 && std::env::var("MEMRA_Q4F16").is_err()
3980 && crate::f16_ffi::pp_f16_capacity_ok(f16_free, f16_need);
3981 let (f16_on, f16_why) = match std::env::var("MEMRA_Q4F16").as_deref() {
3987 Ok("1") => (true, "env MEMRA_Q4F16=1"),
3988 Ok("0") => (false, "env MEMRA_Q4F16=0"),
3989 _ if crate::f16_ffi::pp_f16_enabled() && q4f16_model_ok => {
3990 (true, "env MEMRA_PP_F16")
3991 }
3992 _ if f16_auto => (true, "capacity-keyed auto (UNPINNED)"),
3993 _ if !q4f16_model_ok => (false, "model geometry not eligible"),
3994 _ => (false, "capacity-keyed auto REFUSED (UNPINNED)"),
3995 };
3996 eprintln!(
4003 "[q4f16] prefill program = {} (reason: {}); free {} MiB, mirror mass {} MiB, \
4004 capacity threshold {} MiB (mass + 8192 headroom) — SELECTS PREFILL ARITHMETIC",
4005 if f16_on {
4006 "FP16 MIRRORS"
4007 } else {
4008 "INT8 MMQ (no f16 mirrors)"
4009 },
4010 f16_why,
4011 f16_free >> 20,
4012 f16_need >> 20,
4013 (f16_need + (8usize << 30)) >> 20,
4014 );
4015 for (il, layer) in layers.iter_mut().enumerate() {
4016 let e = crate::pp::layer_engine(e, n_trunk, il)?;
4018 let dense_gemma = layer.gemma4.as_ref().is_some_and(|g| g.moe_bits.is_none());
4019 if !dense_gemma {
4020 continue;
4021 }
4022 if let Mixer::Full(fa) = &mut layer.mixer {
4023 for w in [&mut fa.wq, &mut fa.wk, &mut fa.wv, &mut fa.wo] {
4024 if f16_on {
4025 e.build_q8_f16(w)?;
4026 if matches!(w, crate::model::GpuTensor::Quant { f16: Some(_), .. })
4027 {
4028 nf16 += 1;
4029 }
4030 }
4031 if e.build_q4_rp_swap(w)? {
4032 nswap += 1;
4033 }
4034 }
4035 }
4036 if let Ffn::Dense {
4037 ffn_gate,
4038 ffn_up,
4039 ffn_down,
4040 } = &mut layer.ffn
4041 {
4042 for w in [ffn_gate, ffn_up, ffn_down] {
4043 if f16_on {
4044 e.build_q8_f16(w)?;
4045 if matches!(w, crate::model::GpuTensor::Quant { f16: Some(_), .. })
4046 {
4047 nf16 += 1;
4048 }
4049 }
4050 if e.build_q4_rp_swap(w)? {
4051 nswap += 1;
4052 }
4053 }
4054 }
4055 }
4056 if nswap > 0 {
4057 eprintln!("[q4rp] split-plane IN-PLACE swap: {nswap} dense trunk tensors");
4058 }
4059 if nf16 > 0 {
4060 eprintln!("[q4f16] prefill fp16 mirrors built: {nf16} dense trunk tensors");
4061 }
4062 }
4063 }
4064 let model = HybridModel {
4065 cfg,
4066 plan,
4067 rewrite_qualifications: None,
4068 embd,
4069 output_norm,
4070 output,
4071 layers,
4072 mtp,
4073 mtp_extra,
4074 embd_gpu: std::sync::OnceLock::new(),
4075 gemma4_aux,
4076 step35_aux,
4077 prime_slabs: std::sync::Mutex::new(std::collections::HashMap::new()),
4078 dspark_vgraphs: std::sync::Mutex::new(None),
4079 step_grouped_prefill: std::sync::Mutex::new(StepEpGroupedPrefill::default()),
4080 step35_token_graph: std::sync::Mutex::new(None),
4081 };
4082 e.configure_moe_cache_layout(model.moe_cache_block_sizes());
4083 if force_embd_gpu {
4084 let _ = model
4085 .embd_gpu
4086 .get_or_init(|| e.upload_u8(&model.embd.raw).expect("embed table upload"));
4087 }
4088 crate::pp::sync_stages_after_load(e, n_trunk)?;
4094 Ok(model)
4095 }
4096
4097 pub fn ensure_embed_resident(&self, e: &Engine) -> Result<(), Box<dyn std::error::Error>> {
4107 if std::env::var("MEMRA_EMBED_DEV").as_deref() == Ok("0") {
4108 return Ok(());
4109 }
4110 if self.embd_gpu.get().is_none() {
4111 let buf = e.upload_u8(&self.embd.raw)?;
4112 let _ = self.embd_gpu.set(buf); }
4114 Ok(())
4115 }
4116
4117 pub fn embed(
4118 &self,
4119 e: &Engine,
4120 tokens: &[u32],
4121 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4122 let n_embd = self.cfg.n_embd as usize;
4123 if std::env::var("MEMRA_EMBED_DEV").as_deref() != Ok("0") {
4129 let tbl = self
4130 .embd_gpu
4131 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
4132 let tok_d = e.htod_u32_v(tokens)?;
4133 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
4134 return e.embed_gather_device_td(tbl, &tok_d, tokens.len(), n_embd, qt, rb);
4135 }
4136 let x = self.embd.gather(n_embd, tokens);
4137 Ok(e.htod(&x)?)
4138 }
4139}
4140
4141#[cfg(test)]
4142mod step_expert_selection_tests {
4143 use super::{
4144 StepExpertArtifact, StepExpertLayout, StepParallelLoadConfig, StepParallelRuntimeRegistry,
4145 StepTpAttentionPlacement, select_step_expert_layout,
4146 };
4147 use crate::tp::StepEpLayerSpec;
4148
4149 fn spec(layer: usize, ranks: usize) -> StepEpLayerSpec {
4150 StepEpLayerSpec {
4151 layer,
4152 devices: (0..ranks).collect(),
4153 }
4154 }
4155
4156 #[test]
4157 fn tp2_keeps_projection_sharded_experts() {
4158 let selection = select_step_expert_layout(24, &[], &[spec(24, 2)])
4159 .unwrap()
4160 .unwrap();
4161 assert_eq!(selection.layout, StepExpertLayout::TensorParallel);
4162 assert!(selection.configured_by_tp);
4163 }
4164
4165 #[test]
4166 fn tp4_and_tp8_use_expert_ownership_without_a_second_flag() {
4167 for ranks in [4, 8] {
4168 let selection = select_step_expert_layout(24, &[], &[spec(24, ranks)])
4169 .unwrap()
4170 .unwrap();
4171 assert_eq!(selection.layout, StepExpertLayout::ExpertParallel);
4172 assert!(selection.configured_by_tp);
4173 assert_eq!(selection.spec.devices.len(), ranks);
4174 }
4175 }
4176
4177 #[test]
4178 fn explicit_ep_remains_expert_parallel() {
4179 let selection = select_step_expert_layout(24, &[spec(24, 2)], &[])
4180 .unwrap()
4181 .unwrap();
4182 assert_eq!(selection.layout, StepExpertLayout::ExpertParallel);
4183 assert!(!selection.configured_by_tp);
4184 }
4185
4186 #[test]
4187 fn conflicting_ep_and_tp_assignments_fail_closed() {
4188 let error = select_step_expert_layout(24, &[spec(24, 4)], &[spec(24, 4)]).unwrap_err();
4189 assert!(error.contains("cannot enable MEMRA_STEP_EP and MEMRA_STEP_TP together"));
4190 }
4191
4192 #[test]
4193 fn runtime_registry_owns_one_immutable_load_snapshot() {
4194 let mut source_specs = vec![spec(24, 8)];
4195 let registry = StepParallelRuntimeRegistry::with_config(StepParallelLoadConfig {
4196 ep_specs: Vec::new(),
4197 tp_specs: source_specs.clone(),
4198 native_p2p: true,
4199 ep_device_arithmetic: true,
4200 f32_mirror: true,
4201 bulk_p2p: true,
4202 expert_artifact: StepExpertArtifact::default(),
4203 });
4204 source_specs[0].devices.clear();
4205
4206 let stored = registry.tp_spec(24).unwrap();
4207 assert_eq!(stored.devices, (0..8).collect::<Vec<_>>());
4208 assert!(registry.config.native_p2p);
4209 assert!(registry.config.ep_device_arithmetic);
4210 assert!(registry.config.f32_mirror);
4211 assert!(registry.config.bulk_p2p);
4212 assert_eq!(
4213 registry.expert_selection(24).unwrap().unwrap().layout,
4214 StepExpertLayout::ExpertParallel
4215 );
4216
4217 let standalone = StepParallelRuntimeRegistry::default();
4218 assert!(standalone.tp_spec(24).is_none());
4219 assert!(!standalone.config.native_p2p);
4220 assert!(!standalone.config.ep_device_arithmetic);
4221 assert!(!standalone.config.f32_mirror);
4222 assert!(!standalone.config.bulk_p2p);
4223 }
4224
4225 #[test]
4226 fn rank_local_attention_uses_bounded_swa_rings_only_with_native_p2p() {
4227 assert_eq!(
4228 StepTpAttentionPlacement::resolve(true, None),
4229 StepTpAttentionPlacement::RankLocalGlobal
4230 );
4231 assert_eq!(
4232 StepTpAttentionPlacement::resolve(true, Some(512)),
4233 StepTpAttentionPlacement::RankLocalSwa
4234 );
4235 assert_eq!(
4236 StepTpAttentionPlacement::resolve(false, None),
4237 StepTpAttentionPlacement::OwnerTransportFallback
4238 );
4239 assert_eq!(
4240 StepTpAttentionPlacement::resolve(false, Some(512)),
4241 StepTpAttentionPlacement::OwnerSwa
4242 );
4243 }
4244}
4245
4246#[cfg(test)]
4247mod residency_tests {
4248 use super::{DevExpertFp8ProjectionScales, residency_bytes_by_device};
4249 use crate::model::HostExpertFp8BlockScales;
4250
4251 #[test]
4252 fn pp_residency_counts_only_each_devices_expert_slice() {
4253 let tensors = [
4254 ("blk.0.ffn_gate_exps.weight", 10usize),
4255 ("blk.0.ffn_up_exps.weight", 20),
4256 ("blk.1.ffn_down_exps.weight", 30),
4257 ("blk.2.ffn_gate_exps.weight", 40),
4258 ("blk.3.ffn_up_exps.weight", 50),
4259 ("blk.0.attn_q.weight", 7),
4260 ("output.weight", 11),
4261 ];
4262 let bytes = residency_bytes_by_device(tensors, &[0, 0, 1, 1], 0);
4263 assert_eq!(bytes.experts.get(&0), Some(&60));
4264 assert_eq!(bytes.experts.get(&1), Some(&90));
4265 assert_eq!(bytes.rest, 18);
4266 assert!(bytes.saw_experts);
4267 }
4268
4269 #[test]
4270 fn pp_residency_combines_stages_that_share_one_device() {
4271 let tensors = [
4272 ("blk.0.ffn_gate_exps.weight", 10usize),
4273 ("blk.1.ffn_gate_exps.weight", 20),
4274 ("blk.2.ffn_gate_exps.weight", 30),
4275 ("blk.3.ffn_gate_exps.weight", 40),
4276 ];
4277 let bytes = residency_bytes_by_device(tensors, &[0, 0, 0, 0], 0);
4278 assert_eq!(bytes.experts.get(&0), Some(&100));
4279 assert_eq!(bytes.experts.len(), 1);
4280 }
4281
4282 #[test]
4283 fn resident_fp8_scale_slab_must_match_every_expert() {
4284 let valid = HostExpertFp8BlockScales {
4285 scales: vec![1.0; 12],
4286 rows: 2,
4287 cols: 3,
4288 expert_stride: 6,
4289 };
4290 DevExpertFp8ProjectionScales::validate(&valid, 2).unwrap();
4291
4292 let short = HostExpertFp8BlockScales {
4293 scales: vec![1.0; 11],
4294 ..valid
4295 };
4296 assert_eq!(
4297 DevExpertFp8ProjectionScales::validate(&short, 2).unwrap_err(),
4298 "block-E4M3 scale slab length mismatch: got 11, want 2x6=12"
4299 );
4300 }
4301
4302 #[test]
4303 fn resident_fp8_scale_stride_must_match_its_grid() {
4304 let invalid = HostExpertFp8BlockScales {
4305 scales: vec![1.0; 8],
4306 rows: 2,
4307 cols: 2,
4308 expert_stride: 0,
4309 };
4310 assert_eq!(
4311 DevExpertFp8ProjectionScales::validate(&invalid, 2).unwrap_err(),
4312 "block-E4M3 expert scale stride must be nonzero"
4313 );
4314 }
4315}
4316
4317#[cfg(test)]
4318mod draft_head_tests {
4319 use super::draft_head_tensor;
4320
4321 const STEP37_DRAFTER: &[&str] = &[
4328 "output.weight",
4329 "output_norm.weight",
4330 "token_embd.weight",
4331 "blk.45.nextn.shared_head_norm.weight",
4332 "blk.45.nextn.shared_head_head.weight",
4333 "blk.46.nextn.shared_head_head.weight",
4334 "blk.47.nextn.shared_head_head.weight",
4335 ];
4336
4337 fn present(names: &'static [&'static str]) -> impl Fn(&str) -> bool {
4338 move |t: &str| names.contains(&t)
4339 }
4340
4341 #[test]
4349 fn step37_drafter_prefers_the_blocks_own_nextn_head_over_file_level_output() {
4350 assert_eq!(
4351 draft_head_tensor(present(STEP37_DRAFTER), 45),
4352 "blk.45.nextn.shared_head_head.weight"
4353 );
4354 }
4355
4356 #[test]
4360 fn each_nextn_block_selects_its_own_head() {
4361 for n in 45..=47u32 {
4362 assert_eq!(
4363 draft_head_tensor(present(STEP37_DRAFTER), n),
4364 format!("blk.{n}.nextn.shared_head_head.weight")
4365 );
4366 }
4367 }
4368
4369 #[test]
4373 fn draft_without_a_nextn_head_falls_back_to_file_level_output() {
4374 let fr_spec: &[&str] = &["output.weight", "output_norm.weight", "d2t.weight"];
4375 assert_eq!(draft_head_tensor(present(fr_spec), 45), "output.weight");
4376 }
4377
4378 #[test]
4383 fn legacy_shared_head_is_probed_but_loses_to_shared_head_head() {
4384 let legacy_only: &[&str] = &["output.weight", "blk.45.nextn.shared_head.weight"];
4385 assert_eq!(
4386 draft_head_tensor(present(legacy_only), 45),
4387 "blk.45.nextn.shared_head.weight"
4388 );
4389
4390 let both: &[&str] = &[
4391 "output.weight",
4392 "blk.45.nextn.shared_head.weight",
4393 "blk.45.nextn.shared_head_head.weight",
4394 ];
4395 assert_eq!(
4396 draft_head_tensor(present(both), 45),
4397 "blk.45.nextn.shared_head_head.weight"
4398 );
4399 }
4400
4401 #[test]
4405 fn a_different_blocks_nextn_head_is_never_borrowed() {
4406 let wrong_block: &[&str] = &[
4407 "output.weight",
4408 "blk.46.nextn.shared_head_head.weight",
4409 "blk.47.nextn.shared_head_head.weight",
4410 ];
4411 assert_eq!(draft_head_tensor(present(wrong_block), 45), "output.weight");
4412 }
4413}