Skip to main content

memra_engine/
parallel.rs

1//! Model-specific parallel topology contracts.
2//!
3//! The rank planner is reusable, but model support is never inferred from a loader or a few
4//! scalar dimensions. Each family must register the complete geometry that its TP/EP program
5//! shards. Step-3.7-Flash is the first registered contract because its query-head count varies by
6//! layer (64 full-attention / 96 sliding-attention), while KV heads stay at 8. Step-3.5 and other
7//! siblings do not inherit this contract merely because they share the `step35` architecture tag.
8
9use std::fmt;
10use std::ops::Range;
11
12use memra_gguf::config::ModelConfig;
13use memra_gguf::source::TensorSource;
14
15/// The execution planner's supported rank envelope. Hardware qualification and tuned defaults
16/// remain model x rig evidence, but the placement/runtime contract must not stop at earlier
17/// three-card qualification cells.
18pub const PRODUCT_MAX_CARDS: usize = 8;
19pub const STEP37_TRUNK_LAYERS: usize = 45;
20const STEP_FP8_BLOCK: usize = 128;
21
22#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum HardwareTarget {
24    Rtx5090,
25    RtxPro6000Blackwell,
26}
27
28impl HardwareTarget {
29    fn max_cards(self) -> usize {
30        match self {
31            Self::Rtx5090 => 1,
32            Self::RtxPro6000Blackwell => PRODUCT_MAX_CARDS,
33        }
34    }
35
36    fn label(self) -> &'static str {
37        match self {
38            Self::Rtx5090 => "rtx-5090",
39            Self::RtxPro6000Blackwell => "rtx-pro-6000-blackwell",
40        }
41    }
42
43    fn from_device_name(name: &str) -> Result<Self, TopologyError> {
44        if name.contains("RTX PRO 6000") && name.contains("Blackwell") {
45            return Ok(Self::RtxPro6000Blackwell);
46        }
47        if name.contains("RTX 5090") {
48            return Ok(Self::Rtx5090);
49        }
50        Err(TopologyError::new(format!(
51            "unqualified CUDA device {name:?}; first-class targets are RTX 5090 and RTX PRO 6000 \
52             Blackwell"
53        )))
54    }
55}
56
57#[derive(Debug, Clone, Copy, PartialEq, Eq)]
58pub struct TopologyRequest {
59    pub pipeline: usize,
60    pub tensor: usize,
61    /// Routed experts are partitioned across the TP group. When false, each rank owns every
62    /// expert and tensor-shards the expert projections instead.
63    pub expert_parallel: bool,
64    pub available_devices: usize,
65    pub hardware: HardwareTarget,
66}
67
68impl TopologyRequest {
69    pub fn world_size(self) -> Result<usize, TopologyError> {
70        self.pipeline
71            .checked_mul(self.tensor)
72            .ok_or_else(|| TopologyError::new("PP x TP world size overflow"))
73    }
74}
75
76/// One pipeline stage's model-specific tensor/expert group.
77///
78/// Stage groups make odd physical card counts useful without pretending PP is TP. For example,
79/// three cards can run a TP1 dense-prefix stage followed by a TP2/EP2 MoE stage. The layer range is
80/// explicit because memory-balanced Step placement is not necessarily an equal layer split.
81#[derive(Debug, Clone, PartialEq, Eq)]
82pub struct StageGroupRequest {
83    pub layers: Range<usize>,
84    pub tensor: usize,
85    pub expert_parallel: bool,
86}
87
88#[derive(Debug, Clone, PartialEq, Eq)]
89pub struct GroupedTopologyRequest {
90    pub stages: Vec<StageGroupRequest>,
91    pub available_devices: usize,
92    pub hardware: HardwareTarget,
93}
94
95#[derive(Debug, Clone, PartialEq, Eq)]
96pub struct ModelParallelContract {
97    pub family: &'static str,
98    pub variant: String,
99    pub trunk_layers: usize,
100    pub mtp_layers: usize,
101    pub hidden_size: usize,
102    pub vocab_size: usize,
103    pub dense_ffn_size: usize,
104    pub dense_prefix_layers: usize,
105    pub head_dim: usize,
106    pub query_heads: Vec<usize>,
107    pub kv_heads: Vec<usize>,
108    pub expert_count: usize,
109    pub experts_per_token: usize,
110    pub expert_ffn_size: usize,
111    pub shared_expert_ffn_size: usize,
112    pub partition_boundaries: Vec<usize>,
113    pub hardware_targets: Vec<HardwareTarget>,
114}
115
116#[derive(Debug, Clone, Copy, PartialEq, Eq)]
117pub(crate) enum StepTpExpertLayout {
118    AttentionOnly,
119    TensorParallel,
120    ExpertParallel,
121}
122
123#[derive(Debug, Clone, PartialEq, Eq)]
124pub(crate) struct StepTpLayerPlan {
125    pub layer: usize,
126    pub devices: Vec<usize>,
127    pub owner_device: usize,
128    pub expert_layout: StepTpExpertLayout,
129}
130
131#[derive(Debug, Clone, PartialEq, Eq)]
132pub(crate) struct StepTpPreflightPlan {
133    pub layers: Vec<StepTpLayerPlan>,
134    pub runtime_groups: Vec<Vec<usize>>,
135    pub full_trunk: bool,
136}
137
138impl StepTpPreflightPlan {
139    pub fn dense_attention_layers(&self) -> usize {
140        self.layers
141            .iter()
142            .filter(|layer| layer.expert_layout == StepTpExpertLayout::AttentionOnly)
143            .count()
144    }
145
146    pub fn tensor_parallel_expert_layers(&self) -> usize {
147        self.layers
148            .iter()
149            .filter(|layer| layer.expert_layout == StepTpExpertLayout::TensorParallel)
150            .count()
151    }
152
153    pub fn expert_parallel_layers(&self) -> usize {
154        self.layers
155            .iter()
156            .filter(|layer| layer.expert_layout == StepTpExpertLayout::ExpertParallel)
157            .count()
158    }
159}
160
161impl ModelParallelContract {
162    /// Build the model-specific contract. Unregistered families refuse rather than inheriting a
163    /// generic transformer assumption.
164    pub fn from_model(cfg: &ModelConfig) -> Result<Self, TopologyError> {
165        let plan = memra_gguf::model_plan::ModelPlan::compile(cfg).map_err(|error| {
166            TopologyError::new(format!("cannot compile parallel ModelPlan: {error}"))
167        })?;
168        Self::from_plan(cfg, &plan)
169    }
170
171    fn from_plan(
172        cfg: &ModelConfig,
173        plan: &memra_gguf::model_plan::ModelPlan,
174    ) -> Result<Self, TopologyError> {
175        use memra_gguf::model_plan::{AttentionPlan, MlpPlan};
176
177        if crate::plan_backend::decode_batch_program(plan)
178            != crate::plan_backend::DecodeBatchProgram::SlidingGatedMoe
179        {
180            return Err(TopologyError::new(format!(
181                "no parallel contract registered for plan operations {:?}; loading/running does not establish TP/EP support",
182                plan.trunk_operations()
183            )));
184        }
185        // TRUNK scope, like every other text-serving surface (worker.rs precedent): the
186        // TP/pipeline contract governs the text trunk; a checkpoint that also carries a
187        // vision encoder (step37-flash) must not have its TEXT decode blocked by vision
188        // operations the pipeline program never runs. Vision serving gates on its own
189        // surface (multimodal_prefill_capabilities), not here.
190        let pipeline = crate::plan_backend::PIPELINE
191            .trunk_capabilities(plan)
192            .pipeline;
193        if !pipeline.supported {
194            return Err(TopologyError::new(format!(
195                "pipeline program {} does not implement plan operations {:?}",
196                crate::plan_backend::PIPELINE.name,
197                pipeline.blockers
198            )));
199        }
200        let trunk_layers = plan.layers.len();
201        let mtp_layers = plan.mtp_blocks.len();
202        if trunk_layers == 0 {
203            return Err(TopologyError::new("parallel contract has no trunk layers"));
204        }
205        let layers: Vec<_> = plan
206            .layers
207            .iter()
208            .chain(plan.mtp_blocks.iter().map(|block| &block.layer))
209            .collect();
210        let attention_geometry = layers
211            .iter()
212            .map(|layer| match &layer.attention {
213                AttentionPlan::Full(attention) | AttentionPlan::SlidingWindow { attention, .. } => {
214                    Ok((
215                        attention.query_heads as usize,
216                        attention.kv_heads as usize,
217                        attention.key_head_dim as usize,
218                    ))
219                }
220                _ => Err(TopologyError::new(format!(
221                    "parallel contract has unsupported attention at layer {}",
222                    layer.index
223                ))),
224            })
225            .collect::<Result<Vec<_>, _>>()?;
226        let query_heads: Vec<_> = attention_geometry
227            .iter()
228            .map(|geometry| geometry.0)
229            .collect();
230        let kv_heads: Vec<_> = attention_geometry
231            .iter()
232            .map(|geometry| geometry.1)
233            .collect();
234        let head_dim = attention_geometry[0].2;
235        if attention_geometry
236            .iter()
237            .any(|geometry| geometry.2 != head_dim)
238        {
239            return Err(TopologyError::new(
240                "parallel contract requires one sharding head dimension",
241            ));
242        }
243        let dense_prefix_layers = plan
244            .layers
245            .iter()
246            .take_while(|layer| matches!(layer.mlp, MlpPlan::Dense(_)))
247            .count();
248        let dense_ffn_size = plan
249            .layers
250            .iter()
251            .find_map(|layer| match &layer.mlp {
252                MlpPlan::Dense(dense) => Some(dense.intermediate_size as usize),
253                _ => None,
254            })
255            .ok_or_else(|| TopologyError::new("parallel contract requires a dense prefix"))?;
256        let moe = layers
257            .iter()
258            .find_map(|layer| match &layer.mlp {
259                MlpPlan::Moe(moe) => Some(moe),
260                _ => None,
261            })
262            .ok_or_else(|| TopologyError::new("parallel contract requires routed experts"))?;
263        let shared_expert_ffn_size = moe
264            .shared
265            .as_ref()
266            .map_or(0, |shared| shared.intermediate_size as usize);
267        let is_step37 = trunk_layers == STEP37_TRUNK_LAYERS
268            && mtp_layers == 3
269            && plan.hidden_size == 4096
270            && dense_ffn_size == 11_264
271            && plan.vocab_size == 128_896
272            && query_heads
273                .iter()
274                .enumerate()
275                .all(|(il, &heads)| heads == if il % 4 == 0 { 64 } else { 96 })
276            && kv_heads.iter().all(|&heads| heads == 8)
277            && moe.expert_count == 288
278            && moe.experts_per_token == 8
279            && moe.expert_intermediate_size == 1280
280            && shared_expert_ffn_size == 1280
281            && dense_prefix_layers == 3;
282        if !is_step37 {
283            return Err(TopologyError::new(format!(
284                "no qualified parallel contract for variant {:?}: only the exact Step-3.7-Flash geometry is registered; derived trunk={trunk_layers} mtp={mtp_layers} hidden={} vocab={} dense_ff={dense_ffn_size} dense_prefix={dense_prefix_layers} head_dim={head_dim} q_heads={query_heads:?} kv_heads={kv_heads:?} experts={}/{}/{} shared={shared_expert_ffn_size}",
285                cfg.name,
286                plan.hidden_size,
287                plan.vocab_size,
288                moe.expert_count,
289                moe.experts_per_token,
290                moe.expert_intermediate_size,
291            )));
292        }
293
294        Ok(Self {
295            family: "sliding-gated-moe",
296            variant: "Step-3.7-Flash-FP8".to_string(),
297            trunk_layers,
298            mtp_layers,
299            hidden_size: cfg.n_embd as usize,
300            vocab_size: cfg.n_vocab as usize,
301            dense_ffn_size,
302            dense_prefix_layers,
303            head_dim,
304            query_heads,
305            kv_heads,
306            expert_count: moe.expert_count as usize,
307            experts_per_token: moe.experts_per_token as usize,
308            expert_ffn_size: moe.expert_intermediate_size as usize,
309            shared_expert_ffn_size,
310            partition_boundaries: plan.partition_boundaries.clone(),
311            hardware_targets: vec![HardwareTarget::RtxPro6000Blackwell],
312        })
313    }
314
315    pub fn plan(&self, request: TopologyRequest) -> Result<ParallelPlan, TopologyError> {
316        let pp = request.pipeline;
317        let tp = request.tensor;
318        if !(1..=PRODUCT_MAX_CARDS).contains(&pp) {
319            return Err(TopologyError::new(format!(
320                "PP size {pp} outside product range 1..={PRODUCT_MAX_CARDS}"
321            )));
322        }
323        if !(1..=PRODUCT_MAX_CARDS).contains(&tp) {
324            return Err(TopologyError::new(format!(
325                "TP size {tp} outside product range 1..={PRODUCT_MAX_CARDS}"
326            )));
327        }
328        let world = request.world_size()?;
329        if world > PRODUCT_MAX_CARDS {
330            return Err(TopologyError::new(format!(
331                "PP={pp} x TP={tp} requires {world} cards; product envelope is \
332                 {PRODUCT_MAX_CARDS}"
333            )));
334        }
335        if !self.hardware_targets.contains(&request.hardware) {
336            return Err(TopologyError::new(format!(
337                "{} has no qualified {} contract",
338                self.variant,
339                request.hardware.label()
340            )));
341        }
342        if world > request.hardware.max_cards() {
343            return Err(TopologyError::new(format!(
344                "{} target permits at most {} card(s), requested {world}",
345                request.hardware.label(),
346                request.hardware.max_cards()
347            )));
348        }
349        if request.available_devices < world {
350            return Err(TopologyError::new(format!(
351                "PP={pp} x TP={tp} requires {world} cards, only {} available",
352                request.available_devices
353            )));
354        }
355        if pp > self.trunk_layers {
356            return Err(TopologyError::new(format!(
357                "PP={pp} exceeds {} trunk layers",
358                self.trunk_layers
359            )));
360        }
361        if request.expert_parallel && tp == 1 {
362            return Err(TopologyError::new(
363                "expert parallelism requires TP group size greater than one",
364            ));
365        }
366
367        // Check the family-specific, per-layer attention geometry before generic dimensions so a
368        // refused topology names the model program that actually makes it invalid.
369        for (il, (&q, &kv)) in self.query_heads.iter().zip(&self.kv_heads).enumerate() {
370            require_divisible(&format!("layer {il} query heads"), q, tp)?;
371            require_divisible(&format!("layer {il} KV heads"), kv, tp)?;
372        }
373        require_divisible("hidden size", self.hidden_size, tp)?;
374        require_divisible("vocabulary size", self.vocab_size, tp)?;
375        require_fp8_block_shard("dense FFN size", self.dense_ffn_size, tp)?;
376        if request.expert_parallel {
377            require_divisible("routed expert count", self.expert_count, tp)?;
378        } else {
379            require_fp8_block_shard("routed expert FFN size", self.expert_ffn_size, tp)?;
380        }
381
382        let stage_ranges = (0..pp)
383            .map(|stage| stage * self.trunk_layers / pp..(stage + 1) * self.trunk_layers / pp)
384            .collect();
385
386        Ok(ParallelPlan {
387            contract: self.clone(),
388            request,
389            world_size: world,
390            stage_ranges,
391            mtp_owner_stage: self.mtp_layers.gt(&0).then_some(pp - 1),
392            // Step's 1280-wide shared expert cannot be split four or eight ways without cutting
393            // through checkpoint 128-row E4M3 scale blocks. Replication is the exact program for
394            // those TP/EP layouts; only routed experts are distributed.
395            shared_expert_replicated: tp > 1 && self.shared_expert_ffn_size > 0,
396        })
397    }
398
399    /// Validate every selected Step TP layer before the loader opens its first weight tensor.
400    ///
401    /// Physical device availability is checked separately because this pure plan is also the
402    /// topology oracle for tests and offline launch preparation.
403    pub(crate) fn preflight_step_tp_specs<'a>(
404        &self,
405        specs: impl IntoIterator<Item = (usize, &'a [usize])>,
406        layer_owners: &[usize],
407    ) -> Result<StepTpPreflightPlan, TopologyError> {
408        if layer_owners.len() != self.trunk_layers {
409            return Err(TopologyError::new(format!(
410                "Step TP owner map has {} layers, expected {}",
411                layer_owners.len(),
412                self.trunk_layers
413            )));
414        }
415
416        let mut seen = vec![false; self.trunk_layers];
417        let mut layers = Vec::new();
418        let mut runtime_groups: Vec<Vec<usize>> = Vec::new();
419        for (layer, devices) in specs {
420            if layer >= self.trunk_layers {
421                return Err(TopologyError::new(format!(
422                    "Step TP layer {layer} is outside trunk layers 0..{}",
423                    self.trunk_layers
424                )));
425            }
426            if seen[layer] {
427                return Err(TopologyError::new(format!(
428                    "Step TP preflight assigns layer {layer} more than once"
429                )));
430            }
431            if !(2..=PRODUCT_MAX_CARDS).contains(&devices.len()) {
432                return Err(TopologyError::new(format!(
433                    "Step TP layer {layer} requires 2..={PRODUCT_MAX_CARDS} devices, got {}",
434                    devices.len()
435                )));
436            }
437            let mut unique = devices.to_vec();
438            unique.sort_unstable();
439            unique.dedup();
440            if unique.len() != devices.len() {
441                return Err(TopologyError::new(format!(
442                    "Step TP layer {layer} devices must be distinct, got {devices:?}"
443                )));
444            }
445            let owner_device = layer_owners[layer];
446            if devices.first().copied() != Some(owner_device) {
447                return Err(TopologyError::new(format!(
448                    "Step TP layer {layer} owning PP device {owner_device} must be the first rank, \
449                     got {devices:?}"
450                )));
451            }
452
453            let expert_layout = if layer < self.dense_prefix_layers {
454                StepTpExpertLayout::AttentionOnly
455            } else if devices.len() > 2 {
456                StepTpExpertLayout::ExpertParallel
457            } else {
458                StepTpExpertLayout::TensorParallel
459            };
460            let plan = self.plan(TopologyRequest {
461                pipeline: 1,
462                tensor: devices.len(),
463                // Dense-prefix layers have no routed expert bank, but the full-model TP4/TP8
464                // contract still uses the EP geometry arm so the irrelevant 1280-wide expert
465                // projection is not falsely tensor-sharded during topology validation.
466                expert_parallel: devices.len() > 2,
467                available_devices: devices.len(),
468                hardware: HardwareTarget::RtxPro6000Blackwell,
469            })?;
470            for rank in 0..devices.len() {
471                let query = plan.query_head_range(layer, rank).ok_or_else(|| {
472                    TopologyError::new(format!(
473                        "Step TP layer {layer} has no query-head range for rank {rank}"
474                    ))
475                })?;
476                let kv = plan.kv_head_range(layer, rank).ok_or_else(|| {
477                    TopologyError::new(format!(
478                        "Step TP layer {layer} has no KV-head range for rank {rank}"
479                    ))
480                })?;
481                if query.is_empty() || kv.is_empty() {
482                    return Err(TopologyError::new(format!(
483                        "Step TP layer {layer} rank {rank} has an empty attention shard"
484                    )));
485                }
486            }
487
488            if !runtime_groups.iter().any(|group| group == devices) {
489                runtime_groups.push(devices.to_vec());
490            }
491            seen[layer] = true;
492            layers.push(StepTpLayerPlan {
493                layer,
494                devices: devices.to_vec(),
495                owner_device,
496                expert_layout,
497            });
498        }
499        layers.sort_unstable_by_key(|layer| layer.layer);
500
501        Ok(StepTpPreflightPlan {
502            layers,
503            runtime_groups,
504            full_trunk: seen.into_iter().all(|selected| selected),
505        })
506    }
507
508    /// Plan an explicit sequence of PP stages whose TP/EP widths may differ.
509    ///
510    /// This is the placement contract for arbitrary 1-8 card counts. It validates only the layers
511    /// assigned to each group, so a TP1 dense-prefix stage can coexist with a TP2/TP4/TP8 MoE
512    /// stage. Logical ranks are contiguous per stage; physical-device binding is a separate
513    /// runtime concern and must preserve each group's rank order.
514    pub fn plan_grouped(
515        &self,
516        request: GroupedTopologyRequest,
517    ) -> Result<GroupedParallelPlan, TopologyError> {
518        if request.stages.is_empty() {
519            return Err(TopologyError::new(
520                "grouped Step topology requires at least one stage",
521            ));
522        }
523        if request.stages.len() > self.trunk_layers {
524            return Err(TopologyError::new(format!(
525                "{} grouped stages exceed {} trunk layers",
526                request.stages.len(),
527                self.trunk_layers
528            )));
529        }
530        if !self.hardware_targets.contains(&request.hardware) {
531            return Err(TopologyError::new(format!(
532                "{} has no qualified {} contract",
533                self.variant,
534                request.hardware.label()
535            )));
536        }
537
538        let mut world_size = 0usize;
539        let mut expected_layer = 0usize;
540        let mut rank_groups = Vec::with_capacity(request.stages.len());
541        for (stage, group) in request.stages.iter().enumerate() {
542            if group.layers.start != expected_layer
543                || group.layers.start >= group.layers.end
544                || group.layers.end > self.trunk_layers
545            {
546                return Err(TopologyError::new(format!(
547                    "grouped stage {stage} layers {:?} do not continue the exact 0..{} trunk \
548                     partition at layer {expected_layer}",
549                    group.layers, self.trunk_layers
550                )));
551            }
552            if !(1..=PRODUCT_MAX_CARDS).contains(&group.tensor) {
553                return Err(TopologyError::new(format!(
554                    "grouped stage {stage} TP={} outside product range 1..={PRODUCT_MAX_CARDS}",
555                    group.tensor
556                )));
557            }
558            if group.expert_parallel && group.tensor == 1 {
559                return Err(TopologyError::new(format!(
560                    "grouped stage {stage} expert parallelism requires more than one rank"
561                )));
562            }
563
564            validate_group_geometry(self, stage, group)?;
565            let rank_start = world_size;
566            world_size = world_size
567                .checked_add(group.tensor)
568                .ok_or_else(|| TopologyError::new("grouped topology world size overflow"))?;
569            rank_groups.push(StageRankGroup {
570                stage,
571                layers: group.layers.clone(),
572                global_ranks: rank_start..world_size,
573                tensor: group.tensor,
574                expert_parallel: group.expert_parallel,
575                shared_expert_replicated: group.tensor > 1 && self.shared_expert_ffn_size > 0,
576            });
577            expected_layer = group.layers.end;
578        }
579        if expected_layer != self.trunk_layers {
580            return Err(TopologyError::new(format!(
581                "grouped Step topology ends at layer {expected_layer}, expected {}",
582                self.trunk_layers
583            )));
584        }
585        if world_size > PRODUCT_MAX_CARDS {
586            return Err(TopologyError::new(format!(
587                "grouped Step topology requires {world_size} cards; product envelope is \
588                 {PRODUCT_MAX_CARDS}"
589            )));
590        }
591        if world_size > request.hardware.max_cards() {
592            return Err(TopologyError::new(format!(
593                "{} target permits at most {} card(s), requested {world_size}",
594                request.hardware.label(),
595                request.hardware.max_cards()
596            )));
597        }
598        if request.available_devices < world_size {
599            return Err(TopologyError::new(format!(
600                "grouped Step topology requires {world_size} cards, only {} available",
601                request.available_devices
602            )));
603        }
604
605        let mtp_owner_stage = self.mtp_layers.gt(&0).then_some(expected_layer_stage(
606            self.trunk_layers - 1,
607            &request.stages,
608        )?);
609        Ok(GroupedParallelPlan {
610            contract: self.clone(),
611            request,
612            world_size,
613            rank_groups,
614            mtp_owner_stage,
615        })
616    }
617}
618
619#[derive(Debug, Clone, PartialEq, Eq)]
620pub struct ParallelPlan {
621    pub contract: ModelParallelContract,
622    pub request: TopologyRequest,
623    pub world_size: usize,
624    pub stage_ranges: Vec<Range<usize>>,
625    /// MTP layers are not pipeline stages of their own; the final PP stage owns them.
626    pub mtp_owner_stage: Option<usize>,
627    pub shared_expert_replicated: bool,
628}
629
630impl ParallelPlan {
631    pub fn global_rank(&self, pipeline_rank: usize, tensor_rank: usize) -> Option<usize> {
632        if pipeline_rank >= self.request.pipeline || tensor_rank >= self.request.tensor {
633            return None;
634        }
635        Some(pipeline_rank * self.request.tensor + tensor_rank)
636    }
637
638    pub fn query_head_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
639        split_range(
640            *self.contract.query_heads.get(layer)?,
641            self.request.tensor,
642            tensor_rank,
643        )
644    }
645
646    pub fn kv_head_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
647        split_range(
648            *self.contract.kv_heads.get(layer)?,
649            self.request.tensor,
650            tensor_rank,
651        )
652    }
653
654    /// Column-parallel Q output range and the matching row-parallel O input range.
655    pub fn query_feature_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
656        let heads = self.query_head_range(layer, tensor_rank)?;
657        Some(heads.start * self.contract.head_dim..heads.end * self.contract.head_dim)
658    }
659
660    /// Column-parallel K/V output range. Step-3.7 has eight KV heads, so TP2 and TP4 partition
661    /// them exactly; no KV-head replication is part of this registered contract.
662    pub fn kv_feature_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
663        let heads = self.kv_head_range(layer, tensor_rank)?;
664        Some(heads.start * self.contract.head_dim..heads.end * self.contract.head_dim)
665    }
666
667    /// Column-parallel dense gate/up output range and matching row-parallel down input range.
668    pub fn dense_ffn_range(&self, tensor_rank: usize) -> Option<Range<usize>> {
669        split_range(
670            self.contract.dense_ffn_size,
671            self.request.tensor,
672            tensor_rank,
673        )
674    }
675
676    pub fn routed_expert_range(&self, tensor_rank: usize) -> Option<Range<usize>> {
677        self.request
678            .expert_parallel
679            .then(|| split_range(self.contract.expert_count, self.request.tensor, tensor_rank))?
680    }
681
682    pub fn routed_expert_ffn_range(&self, tensor_rank: usize) -> Option<Range<usize>> {
683        (!self.request.expert_parallel).then(|| {
684            split_range(
685                self.contract.expert_ffn_size,
686                self.request.tensor,
687                tensor_rank,
688            )
689        })?
690    }
691}
692
693#[derive(Debug, Clone, PartialEq, Eq)]
694pub struct StageRankGroup {
695    pub stage: usize,
696    pub layers: Range<usize>,
697    pub global_ranks: Range<usize>,
698    pub tensor: usize,
699    pub expert_parallel: bool,
700    pub shared_expert_replicated: bool,
701}
702
703#[derive(Debug, Clone, PartialEq, Eq)]
704pub struct GroupedParallelPlan {
705    pub contract: ModelParallelContract,
706    pub request: GroupedTopologyRequest,
707    pub world_size: usize,
708    pub rank_groups: Vec<StageRankGroup>,
709    pub mtp_owner_stage: Option<usize>,
710}
711
712impl GroupedParallelPlan {
713    pub fn group_for_layer(&self, layer: usize) -> Option<&StageRankGroup> {
714        self.rank_groups
715            .iter()
716            .find(|group| group.layers.contains(&layer))
717    }
718
719    pub fn group_for_global_rank(&self, rank: usize) -> Option<&StageRankGroup> {
720        self.rank_groups
721            .iter()
722            .find(|group| group.global_ranks.contains(&rank))
723    }
724
725    pub fn global_rank(&self, stage: usize, tensor_rank: usize) -> Option<usize> {
726        let group = self.rank_groups.get(stage)?;
727        (tensor_rank < group.tensor).then_some(group.global_ranks.start + tensor_rank)
728    }
729
730    pub fn query_head_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
731        let group = self.group_for_layer(layer)?;
732        split_range(
733            *self.contract.query_heads.get(layer)?,
734            group.tensor,
735            tensor_rank,
736        )
737    }
738
739    pub fn kv_head_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
740        let group = self.group_for_layer(layer)?;
741        split_range(
742            *self.contract.kv_heads.get(layer)?,
743            group.tensor,
744            tensor_rank,
745        )
746    }
747
748    pub fn routed_expert_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
749        let group = self.group_for_layer(layer)?;
750        group
751            .expert_parallel
752            .then(|| split_range(self.contract.expert_count, group.tensor, tensor_rank))?
753    }
754}
755
756/// Validate the live Step PP request before the loader allocates CUDA state. Checkpoint tensor
757/// census is deliberately a separate loader gate: topology legality must remain testable without
758/// opening model files, while serving requires both gates.
759pub fn validate_step_pp_request(cfg: &ModelConfig) -> Result<Option<ParallelPlan>, TopologyError> {
760    let pp = match std::env::var("MEMRA_PP_STAGES") {
761        Err(_) => return Ok(None),
762        Ok(value) if value.is_empty() || value == "0" || value == "1" => return Ok(None),
763        Ok(value) => value.parse::<usize>().map_err(|_| {
764            TopologyError::new(format!("MEMRA_PP_STAGES={value} is not a positive integer"))
765        })?,
766    };
767    let devices = selected_pp_devices(pp)?;
768    let hardware = detect_uniform_hardware(&devices)?;
769    let contract = ModelParallelContract::from_model(cfg)?;
770    let trunk_layers = contract.trunk_layers;
771    let plan = contract.plan(TopologyRequest {
772        pipeline: pp,
773        tensor: 1,
774        expert_parallel: false,
775        available_devices: devices.len(),
776        hardware,
777    })?;
778    let fence = crate::pp::pp_cuts(trunk_layers).ok_or_else(|| {
779        TopologyError::new(format!(
780            "Step PP={pp} has no valid runtime stage fence over {trunk_layers} trunk layers"
781        ))
782    })?;
783    let plan = apply_stage_fence(plan, &fence)?;
784    Ok(Some(plan))
785}
786
787/// Prove that the official Step checkpoint exposes every routed expert projection as a native
788/// stacked block-128 E4M3 bank. Converted and per-tensor artifacts do not inherit this contract.
789pub fn validate_step_fp8_checkpoint(
790    src: &dyn TensorSource,
791    contract: &ModelParallelContract,
792) -> Result<usize, TopologyError> {
793    if src.st_dir().is_none() {
794        return Err(TopologyError::new(
795            "Step-3.7-Flash-FP8 topology qualification requires the official safetensors \
796             checkpoint source; a converted artifact cannot inherit this contract",
797        ));
798    }
799
800    let projections = [
801        (
802            "ffn_gate_exps",
803            contract.hidden_size,
804            contract.expert_ffn_size,
805        ),
806        (
807            "ffn_up_exps",
808            contract.hidden_size,
809            contract.expert_ffn_size,
810        ),
811        (
812            "ffn_down_exps",
813            contract.expert_ffn_size,
814            contract.hidden_size,
815        ),
816    ];
817    let mut qualified = 0usize;
818    for layer in contract.dense_prefix_layers..contract.trunk_layers {
819        for &(projection, expected_in, expected_out) in &projections {
820            let name = format!("blk.{layer}.{projection}.weight");
821            let fp8 = src.find_fp8_stacked_native(&name).ok_or_else(|| {
822                TopologyError::new(format!(
823                    "{name} is not a checkpoint-faithful stacked block-128 E4M3 bank"
824                ))
825            })?;
826            if fp8.n_expert != contract.expert_count {
827                return Err(TopologyError::new(format!(
828                    "{name} carries {} experts, expected {}",
829                    fp8.n_expert, contract.expert_count
830                )));
831            }
832            if fp8.in_f != expected_in || fp8.out_f != expected_out {
833                return Err(TopologyError::new(format!(
834                    "{name} expert shape {}x{} != expected {expected_out}x{expected_in}",
835                    fp8.out_f, fp8.in_f
836                )));
837            }
838            let expected_rows = expected_out.div_ceil(STEP_FP8_BLOCK);
839            let expected_cols = expected_in.div_ceil(STEP_FP8_BLOCK);
840            let expected_scales = contract.expert_count * expected_rows * expected_cols;
841            if fp8.scale_rows != expected_rows
842                || fp8.scale_cols != expected_cols
843                || fp8.scales.len() != expected_scales
844            {
845                return Err(TopologyError::new(format!(
846                    "{name} block-128 E4M3 grid {}x{} ({} scales) != expected {} experts x \
847                     {expected_rows}x{expected_cols} ({expected_scales} scales)",
848                    fp8.scale_rows,
849                    fp8.scale_cols,
850                    fp8.scales.len(),
851                    contract.expert_count
852                )));
853            }
854            qualified += fp8.n_expert;
855        }
856    }
857
858    let expected = (contract.trunk_layers - contract.dense_prefix_layers)
859        * contract.expert_count
860        * projections.len();
861    if qualified != expected {
862        return Err(TopologyError::new(format!(
863            "Step E4M3 tensor census qualified {qualified}, expected {expected}"
864        )));
865    }
866    Ok(qualified)
867}
868
869/// Prove that the official Step NVFP4 checkpoint exposes every routed expert projection as a
870/// native stacked modelopt NVFP4 bank: packed e2m1 codes `[E, out, in/2]`, per-16 UE4M3 scales
871/// `[E, out, in/16]`, and a finite positive per-expert `weight_scale_2` macro. Converted and
872/// per-tensor artifacts do not inherit this contract. The macro census matters: those values run
873/// ~1e-5..1e-4 in the official artifact and dropping them silently produces garbage, so a bank
874/// whose macros fail the finite-positive check refuses here rather than at first decode.
875pub fn validate_step_nvfp4_checkpoint(
876    src: &dyn TensorSource,
877    contract: &ModelParallelContract,
878) -> Result<usize, TopologyError> {
879    if src.st_dir().is_none() {
880        return Err(TopologyError::new(
881            "Step-3.7-Flash-NVFP4 topology qualification requires the official safetensors \
882             checkpoint source; a converted artifact cannot inherit this contract",
883        ));
884    }
885
886    let projections = [
887        (
888            "ffn_gate_exps",
889            contract.hidden_size,
890            contract.expert_ffn_size,
891        ),
892        (
893            "ffn_up_exps",
894            contract.hidden_size,
895            contract.expert_ffn_size,
896        ),
897        (
898            "ffn_down_exps",
899            contract.expert_ffn_size,
900            contract.hidden_size,
901        ),
902    ];
903    let mut qualified = 0usize;
904    for layer in contract.dense_prefix_layers..contract.trunk_layers {
905        for &(projection, expected_in, expected_out) in &projections {
906            let name = format!("blk.{layer}.{projection}.weight");
907            let bank = src.find_nvfp4_stacked_native(&name).ok_or_else(|| {
908                TopologyError::new(format!(
909                    "{name} is not a checkpoint-faithful stacked modelopt NVFP4 bank \
910                     (packed e2m1 codes + per-16 UE4M3 scales + per-expert macro)"
911                ))
912            })?;
913            if bank.n_expert != contract.expert_count {
914                return Err(TopologyError::new(format!(
915                    "{name} carries {} experts, expected {}",
916                    bank.n_expert, contract.expert_count
917                )));
918            }
919            if bank.in_f != expected_in || bank.out_f != expected_out {
920                return Err(TopologyError::new(format!(
921                    "{name} expert shape {}x{} != expected {expected_out}x{expected_in}",
922                    bank.out_f, bank.in_f
923                )));
924            }
925            if bank.in_f % 64 != 0 {
926                return Err(TopologyError::new(format!(
927                    "{name} in_features {} is not 64-aligned; memra block_nvfp4 kernels \
928                     require whole 64-element superblocks",
929                    bank.in_f
930                )));
931            }
932            if bank.macros.len() != contract.expert_count {
933                return Err(TopologyError::new(format!(
934                    "{name} carries {} weight_scale_2 macros, expected {}",
935                    bank.macros.len(),
936                    contract.expert_count
937                )));
938            }
939            qualified += bank.n_expert;
940        }
941    }
942
943    let expected = (contract.trunk_layers - contract.dense_prefix_layers)
944        * contract.expert_count
945        * projections.len();
946    if qualified != expected {
947        return Err(TopologyError::new(format!(
948            "Step NVFP4 tensor census qualified {qualified}, expected {expected}"
949        )));
950    }
951    Ok(qualified)
952}
953
954fn apply_stage_fence(
955    mut plan: ParallelPlan,
956    fence: &[usize],
957) -> Result<ParallelPlan, TopologyError> {
958    let expected = plan.request.pipeline + 1;
959    if fence.len() != expected
960        || fence.first() != Some(&0)
961        || fence.last() != Some(&plan.contract.trunk_layers)
962        || fence.windows(2).any(|window| window[0] >= window[1])
963        || fence[1..fence.len() - 1]
964            .iter()
965            .any(|boundary| !plan.contract.partition_boundaries.contains(boundary))
966    {
967        return Err(TopologyError::new(format!(
968            "invalid PP fence {fence:?} for {} stages over {} trunk layers",
969            plan.request.pipeline, plan.contract.trunk_layers
970        )));
971    }
972    plan.stage_ranges = fence
973        .windows(2)
974        .map(|window| window[0]..window[1])
975        .collect();
976    Ok(plan)
977}
978
979fn selected_pp_devices(pp: usize) -> Result<Vec<usize>, TopologyError> {
980    let raw = std::env::var("MEMRA_PP_DEVICES").map_err(|_| {
981        TopologyError::new(format!(
982            "Step PP={pp} requires explicit MEMRA_PP_DEVICES with one distinct CUDA ordinal per \
983             stage; same-device diagnostics do not qualify the multi-card product"
984        ))
985    })?;
986    let devices: Result<Vec<usize>, _> = raw
987        .split(',')
988        .map(|part| part.trim().parse::<usize>())
989        .collect();
990    let devices = devices.map_err(|_| {
991        TopologyError::new(format!(
992            "MEMRA_PP_DEVICES={raw:?} is not a comma-separated CUDA ordinal list"
993        ))
994    })?;
995    if devices.len() != pp {
996        return Err(TopologyError::new(format!(
997            "MEMRA_PP_DEVICES lists {} devices but MEMRA_PP_STAGES={pp}",
998            devices.len()
999        )));
1000    }
1001    let mut unique = devices.clone();
1002    unique.sort_unstable();
1003    unique.dedup();
1004    if unique.len() != devices.len() {
1005        return Err(TopologyError::new(format!(
1006            "Step PP={pp} requires {pp} distinct devices; MEMRA_PP_DEVICES={raw:?} repeats an \
1007             ordinal"
1008        )));
1009    }
1010    Ok(devices)
1011}
1012
1013pub(crate) fn detect_uniform_hardware(devices: &[usize]) -> Result<HardwareTarget, TopologyError> {
1014    cudarc::driver::result::init().map_err(|error| {
1015        TopologyError::new(format!("CUDA driver initialization failed: {error}"))
1016    })?;
1017    let mut target = None;
1018    for &ordinal in devices {
1019        let device = cudarc::driver::result::device::get(ordinal as i32).map_err(|error| {
1020            TopologyError::new(format!("CUDA device {ordinal} lookup failed: {error}"))
1021        })?;
1022        let name = cudarc::driver::result::device::get_name(device).map_err(|error| {
1023            TopologyError::new(format!("CUDA device {ordinal} name lookup failed: {error}"))
1024        })?;
1025        let current = HardwareTarget::from_device_name(&name)?;
1026        if let Some(expected) = target {
1027            if current != expected {
1028                return Err(TopologyError::new(format!(
1029                    "mixed hardware targets in MEMRA_PP_DEVICES: expected {}, device {ordinal} is \
1030                     {}",
1031                    expected.label(),
1032                    current.label()
1033                )));
1034            }
1035        } else {
1036            target = Some(current);
1037        }
1038    }
1039    target.ok_or_else(|| TopologyError::new("MEMRA_PP_DEVICES is empty"))
1040}
1041
1042fn require_divisible(label: &str, value: usize, parts: usize) -> Result<(), TopologyError> {
1043    if value == 0 {
1044        return Err(TopologyError::new(format!("{label} is zero")));
1045    }
1046    if value % parts != 0 {
1047        return Err(TopologyError::new(format!(
1048            "{label} {value} is not divisible by TP={parts}"
1049        )));
1050    }
1051    Ok(())
1052}
1053
1054fn require_fp8_block_shard(label: &str, value: usize, parts: usize) -> Result<(), TopologyError> {
1055    require_divisible(label, value, parts)?;
1056    let local = value / parts;
1057    if local % STEP_FP8_BLOCK != 0 {
1058        return Err(TopologyError::new(format!(
1059            "{label} shard {local} for TP={parts} cuts through the Step E4M3 block size \
1060             {STEP_FP8_BLOCK}"
1061        )));
1062    }
1063    Ok(())
1064}
1065
1066fn validate_group_geometry(
1067    contract: &ModelParallelContract,
1068    stage: usize,
1069    group: &StageGroupRequest,
1070) -> Result<(), TopologyError> {
1071    let tp = group.tensor;
1072    for layer in group.layers.clone() {
1073        require_divisible(
1074            &format!("stage {stage} layer {layer} query heads"),
1075            contract.query_heads[layer],
1076            tp,
1077        )?;
1078        require_divisible(
1079            &format!("stage {stage} layer {layer} KV heads"),
1080            contract.kv_heads[layer],
1081            tp,
1082        )?;
1083    }
1084    require_divisible(
1085        &format!("stage {stage} hidden size"),
1086        contract.hidden_size,
1087        tp,
1088    )?;
1089    if group.layers.start < contract.dense_prefix_layers {
1090        require_fp8_block_shard(
1091            &format!("stage {stage} dense FFN size"),
1092            contract.dense_ffn_size,
1093            tp,
1094        )?;
1095    }
1096    if group.layers.end > contract.dense_prefix_layers {
1097        if group.expert_parallel {
1098            require_divisible(
1099                &format!("stage {stage} routed expert count"),
1100                contract.expert_count,
1101                tp,
1102            )?;
1103        } else {
1104            require_fp8_block_shard(
1105                &format!("stage {stage} routed expert FFN size"),
1106                contract.expert_ffn_size,
1107                tp,
1108            )?;
1109        }
1110    }
1111    if group.layers.end == contract.trunk_layers {
1112        require_divisible(
1113            &format!("stage {stage} vocabulary size"),
1114            contract.vocab_size,
1115            tp,
1116        )?;
1117    }
1118    Ok(())
1119}
1120
1121fn expected_layer_stage(
1122    layer: usize,
1123    stages: &[StageGroupRequest],
1124) -> Result<usize, TopologyError> {
1125    stages
1126        .iter()
1127        .position(|stage| stage.layers.contains(&layer))
1128        .ok_or_else(|| TopologyError::new(format!("no grouped stage owns layer {layer}")))
1129}
1130
1131fn split_range(total: usize, parts: usize, rank: usize) -> Option<Range<usize>> {
1132    if parts == 0 || rank >= parts || total % parts != 0 {
1133        return None;
1134    }
1135    let width = total / parts;
1136    Some(rank * width..(rank + 1) * width)
1137}
1138
1139#[derive(Debug, Clone, PartialEq, Eq)]
1140pub struct TopologyError {
1141    message: String,
1142}
1143
1144impl TopologyError {
1145    fn new(message: impl Into<String>) -> Self {
1146        Self {
1147            message: message.into(),
1148        }
1149    }
1150}
1151
1152impl fmt::Display for TopologyError {
1153    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1154        self.message.fmt(f)
1155    }
1156}
1157
1158impl std::error::Error for TopologyError {}
1159
1160#[cfg(test)]
1161mod tests {
1162    use super::*;
1163    use memra_gguf::config::{Arch, MoeConfig, Step35Config};
1164    use memra_gguf::source::{Fp8StackedNative, TensorView};
1165    use std::path::Path;
1166
1167    fn step37_contract() -> ModelParallelContract {
1168        let total_layers = 48;
1169        ModelParallelContract {
1170            family: "sliding-gated-moe",
1171            variant: "Step-3.7-Flash-FP8".to_string(),
1172            trunk_layers: 45,
1173            mtp_layers: 3,
1174            hidden_size: 4096,
1175            vocab_size: 128_896,
1176            dense_ffn_size: 11_264,
1177            dense_prefix_layers: 3,
1178            head_dim: 128,
1179            query_heads: (0..total_layers)
1180                .map(|il| if il % 4 == 0 { 64 } else { 96 })
1181                .collect(),
1182            kv_heads: vec![8; total_layers],
1183            expert_count: 288,
1184            experts_per_token: 8,
1185            expert_ffn_size: 1280,
1186            shared_expert_ffn_size: 1280,
1187            partition_boundaries: (1..45).collect(),
1188            hardware_targets: vec![HardwareTarget::RtxPro6000Blackwell],
1189        }
1190    }
1191
1192    fn step37_model_config() -> ModelConfig {
1193        let total_layers = 48;
1194        let head_count: Vec<u32> = (0..total_layers)
1195            .map(|il| if il % 4 == 0 { 64 } else { 96 })
1196            .collect();
1197        ModelConfig {
1198            arch: Arch::Step35,
1199            name: "Step-3.7-Flash-FP8".to_string(),
1200            n_layer: total_layers,
1201            n_embd: 4096,
1202            n_head: 96,
1203            n_head_kv: 8,
1204            head_dim_k: 128,
1205            head_dim_v: 128,
1206            n_ff: 11_264,
1207            n_vocab: 128_896,
1208            context_length: 262_144,
1209            rms_eps: 1e-6,
1210            rope_freq_base: 5_000_000.0,
1211            rope_dim_count: 128,
1212            rope_sections: Vec::new(),
1213            full_attention_interval: 0,
1214            ssm: None,
1215            moe: Some(MoeConfig {
1216                expert_count: 288,
1217                expert_used_count: 8,
1218                expert_ff_length: 1280,
1219                expert_shared_ff_length: 1280,
1220            }),
1221            m3: None,
1222            hy3: None,
1223            gemma4: None,
1224            vision: None,
1225            multimodal: None,
1226            mla: None,
1227            dsv4: None,
1228            step35: Some(Step35Config {
1229                head_count,
1230                head_count_kv: vec![8; total_layers as usize],
1231                swa_pattern: (0..total_layers).map(|il| il % 4 != 0).collect(),
1232                sliding_window: 512,
1233                rope_base_global: 5_000_000.0,
1234                rope_base_swa: 10_000.0,
1235                rope_dims_full: 64,
1236                rope_dims_swa: 128,
1237                rope_freq_factors: None,
1238                swiglu_clamp_exp: vec![0.0; total_layers as usize],
1239                swiglu_clamp_shexp: vec![0.0; total_layers as usize],
1240                sigmoid_routing: true,
1241                routed_scaling_factor: 3.0,
1242                route_norm: true,
1243                first_k_dense_replace: 3,
1244            }),
1245            geometry: None,
1246            nextn_predict_layers: 3,
1247            n_layer_total: total_layers,
1248        }
1249    }
1250
1251    fn request(pp: usize, tp: usize, expert_parallel: bool) -> TopologyRequest {
1252        TopologyRequest {
1253            pipeline: pp,
1254            tensor: tp,
1255            expert_parallel,
1256            available_devices: pp * tp,
1257            hardware: HardwareTarget::RtxPro6000Blackwell,
1258        }
1259    }
1260
1261    struct MockStepFp8Source {
1262        safetensors: bool,
1263        block_scales: bool,
1264    }
1265
1266    impl TensorSource for MockStepFp8Source {
1267        fn config(&self) -> ModelConfig {
1268            step37_model_config()
1269        }
1270
1271        fn find(&self, _ggml_name: &str) -> Option<TensorView<'_>> {
1272            None
1273        }
1274
1275        fn st_dir(&self) -> Option<&Path> {
1276            self.safetensors.then(|| Path::new("/mock-step-fp8"))
1277        }
1278
1279        fn find_fp8_stacked_native(&self, name: &str) -> Option<Fp8StackedNative<'_>> {
1280            let (in_f, out_f): (usize, usize) = if name.contains("ffn_down_exps") {
1281                (1280, 4096)
1282            } else if name.contains("ffn_gate_exps") || name.contains("ffn_up_exps") {
1283                (4096, 1280)
1284            } else {
1285                return None;
1286            };
1287            let (scale_rows, scale_cols) = if self.block_scales {
1288                (
1289                    out_f.div_ceil(STEP_FP8_BLOCK),
1290                    in_f.div_ceil(STEP_FP8_BLOCK),
1291                )
1292            } else {
1293                (1, 1)
1294            };
1295            Some(Fp8StackedNative {
1296                bytes: &[],
1297                scales: vec![1.0; 288 * scale_rows * scale_cols],
1298                n_expert: 288,
1299                out_f,
1300                in_f,
1301                scale_rows,
1302                scale_cols,
1303            })
1304        }
1305    }
1306
1307    #[test]
1308    fn step_fp8_checkpoint_census_covers_every_routed_projection() {
1309        let source = MockStepFp8Source {
1310            safetensors: true,
1311            block_scales: true,
1312        };
1313        let qualified =
1314            validate_step_fp8_checkpoint(&source, &step37_contract()).expect("valid FP8 source");
1315        assert_eq!(qualified, 42 * 288 * 3);
1316    }
1317
1318    #[test]
1319    fn step_fp8_checkpoint_census_refuses_conversion_and_wrong_scale_class() {
1320        let converted = MockStepFp8Source {
1321            safetensors: false,
1322            block_scales: true,
1323        };
1324        assert!(
1325            validate_step_fp8_checkpoint(&converted, &step37_contract())
1326                .unwrap_err()
1327                .to_string()
1328                .contains("official safetensors")
1329        );
1330
1331        let per_tensor = MockStepFp8Source {
1332            safetensors: true,
1333            block_scales: false,
1334        };
1335        assert!(
1336            validate_step_fp8_checkpoint(&per_tensor, &step37_contract())
1337                .unwrap_err()
1338                .to_string()
1339                .contains("block-128 E4M3")
1340        );
1341    }
1342
1343    #[test]
1344    fn step_pp3_maps_fifteen_trunk_layers_per_card() {
1345        let plan = step37_contract().plan(request(3, 1, false)).unwrap();
1346        assert_eq!(plan.world_size, 3);
1347        assert_eq!(plan.stage_ranges, vec![0..15, 15..30, 30..45]);
1348        assert_eq!(plan.mtp_owner_stage, Some(2));
1349    }
1350
1351    #[test]
1352    fn step_pp_marker_uses_the_runtime_stage_fence() {
1353        let plan = step37_contract().plan(request(3, 1, false)).unwrap();
1354        let plan = apply_stage_fence(plan, &[0, 10, 28, 45]).unwrap();
1355        assert_eq!(plan.stage_ranges, vec![0..10, 10..28, 28..45]);
1356    }
1357
1358    #[test]
1359    fn stage_fence_must_use_model_plan_partition_boundaries() {
1360        let mut contract = step37_contract();
1361        contract
1362            .partition_boundaries
1363            .retain(|&boundary| boundary != 10);
1364        let plan = contract.plan(request(3, 1, false)).unwrap();
1365        let error = apply_stage_fence(plan, &[0, 10, 28, 45]).unwrap_err();
1366        assert!(error.to_string().contains("invalid PP fence"));
1367    }
1368
1369    #[test]
1370    fn step_contract_is_extracted_from_model_specific_geometry() {
1371        let contract = ModelParallelContract::from_model(&step37_model_config()).unwrap();
1372        assert_eq!(contract.family, "sliding-gated-moe");
1373        assert_eq!(contract.trunk_layers, 45);
1374        assert_eq!(contract.mtp_layers, 3);
1375        assert_eq!(contract.query_heads[0], 64);
1376        assert_eq!(contract.query_heads[1], 96);
1377        assert_eq!(contract.kv_heads[47], 8);
1378        assert_eq!(contract.expert_count, 288);
1379        assert_eq!(contract.experts_per_token, 8);
1380    }
1381
1382    #[test]
1383    fn step_sibling_does_not_inherit_the_step37_contract() {
1384        let mut sibling = step37_model_config();
1385        sibling.name = "Step-3.5-Flash".to_string();
1386        sibling.n_vocab = 128_000;
1387        let error = ModelParallelContract::from_model(&sibling).unwrap_err();
1388        assert!(
1389            error
1390                .to_string()
1391                .contains("only the exact Step-3.7-Flash geometry is registered")
1392        );
1393    }
1394
1395    #[test]
1396    fn step_without_the_official_mtp_geometry_does_not_inherit_the_contract() {
1397        let mut stripped = step37_model_config();
1398        stripped.nextn_predict_layers = 0;
1399        let error = ModelParallelContract::from_model(&stripped).unwrap_err();
1400        assert!(
1401            error
1402                .to_string()
1403                .contains("only the exact Step-3.7-Flash geometry is registered")
1404        );
1405    }
1406
1407    #[test]
1408    fn hardware_target_classification_is_exact() {
1409        assert_eq!(
1410            HardwareTarget::from_device_name("NVIDIA RTX PRO 6000 Blackwell Server Edition")
1411                .unwrap(),
1412            HardwareTarget::RtxPro6000Blackwell
1413        );
1414        assert_eq!(
1415            HardwareTarget::from_device_name("NVIDIA GeForce RTX 5090 Laptop GPU").unwrap(),
1416            HardwareTarget::Rtx5090
1417        );
1418        assert!(HardwareTarget::from_device_name("NVIDIA H100 80GB HBM3").is_err());
1419    }
1420
1421    #[test]
1422    fn step_tp2_tp4_tp8_and_hybrid_plans_are_geometry_valid() {
1423        let tp2 = step37_contract().plan(request(1, 2, true)).unwrap();
1424        assert_eq!(tp2.query_head_range(0, 1), Some(32..64));
1425        assert_eq!(tp2.query_head_range(1, 1), Some(48..96));
1426        assert_eq!(tp2.kv_head_range(0, 1), Some(4..8));
1427        assert_eq!(tp2.routed_expert_range(1), Some(144..288));
1428
1429        let tp4 = step37_contract().plan(request(1, 4, true)).unwrap();
1430        assert_eq!(tp4.query_head_range(0, 3), Some(48..64));
1431        assert_eq!(tp4.query_head_range(1, 3), Some(72..96));
1432        assert_eq!(tp4.kv_head_range(0, 3), Some(6..8));
1433        assert_eq!(tp4.query_feature_range(0, 3), Some(6144..8192));
1434        assert_eq!(tp4.query_feature_range(1, 3), Some(9216..12_288));
1435        assert_eq!(tp4.kv_feature_range(0, 3), Some(768..1024));
1436        assert_eq!(tp4.dense_ffn_range(3), Some(8448..11_264));
1437        assert_eq!(tp4.routed_expert_range(3), Some(216..288));
1438        assert!(tp4.shared_expert_replicated);
1439
1440        let tp8 = step37_contract().plan(request(1, 8, true)).unwrap();
1441        assert_eq!(tp8.query_head_range(0, 7), Some(56..64));
1442        assert_eq!(tp8.query_head_range(1, 7), Some(84..96));
1443        assert_eq!(tp8.kv_head_range(0, 7), Some(7..8));
1444        assert_eq!(tp8.dense_ffn_range(7), Some(9856..11_264));
1445        assert_eq!(tp8.routed_expert_range(7), Some(252..288));
1446        assert!(tp8.shared_expert_replicated);
1447
1448        let hybrid = step37_contract().plan(request(2, 4, true)).unwrap();
1449        assert_eq!(hybrid.world_size, 8);
1450        assert_eq!(hybrid.stage_ranges, vec![0..22, 22..45]);
1451        assert_eq!(hybrid.global_rank(1, 3), Some(7));
1452        assert_eq!(hybrid.global_rank(2, 0), None);
1453    }
1454
1455    #[test]
1456    fn grouped_three_card_plan_is_pp1_then_tp2_ep2() {
1457        let plan = step37_contract()
1458            .plan_grouped(GroupedTopologyRequest {
1459                stages: vec![
1460                    StageGroupRequest {
1461                        layers: 0..15,
1462                        tensor: 1,
1463                        expert_parallel: false,
1464                    },
1465                    StageGroupRequest {
1466                        layers: 15..45,
1467                        tensor: 2,
1468                        expert_parallel: true,
1469                    },
1470                ],
1471                available_devices: 3,
1472                hardware: HardwareTarget::RtxPro6000Blackwell,
1473            })
1474            .unwrap();
1475
1476        assert_eq!(plan.world_size, 3);
1477        assert_eq!(plan.rank_groups[0].global_ranks, 0..1);
1478        assert_eq!(plan.rank_groups[1].global_ranks, 1..3);
1479        assert_eq!(plan.global_rank(0, 0), Some(0));
1480        assert_eq!(plan.global_rank(1, 0), Some(1));
1481        assert_eq!(plan.global_rank(1, 1), Some(2));
1482        assert_eq!(plan.query_head_range(16, 1), Some(32..64));
1483        assert_eq!(plan.query_head_range(17, 1), Some(48..96));
1484        assert_eq!(plan.kv_head_range(16, 1), Some(4..8));
1485        assert_eq!(plan.routed_expert_range(16, 1), Some(144..288));
1486        assert_eq!(plan.mtp_owner_stage, Some(1));
1487    }
1488
1489    #[test]
1490    fn grouped_step_layouts_cover_every_card_count_through_eight() {
1491        let layouts: Vec<Vec<usize>> = vec![
1492            vec![1],
1493            vec![2],
1494            vec![1, 2],
1495            vec![4],
1496            vec![1, 4],
1497            vec![2, 4],
1498            vec![1, 2, 4],
1499            vec![8],
1500        ];
1501        for (index, widths) in layouts.into_iter().enumerate() {
1502            let cards = index + 1;
1503            let cuts: Vec<usize> = match widths.len() {
1504                1 => vec![0, 45],
1505                2 => vec![0, 3, 45],
1506                3 => vec![0, 3, 15, 45],
1507                _ => unreachable!(),
1508            };
1509            let stages = widths
1510                .iter()
1511                .enumerate()
1512                .map(|(stage, &tensor)| StageGroupRequest {
1513                    layers: cuts[stage]..cuts[stage + 1],
1514                    tensor,
1515                    expert_parallel: tensor > 1 && cuts[stage + 1] > 3,
1516                })
1517                .collect();
1518            let plan = step37_contract()
1519                .plan_grouped(GroupedTopologyRequest {
1520                    stages,
1521                    available_devices: cards,
1522                    hardware: HardwareTarget::RtxPro6000Blackwell,
1523                })
1524                .unwrap_or_else(|error| panic!("{cards}-card grouped plan failed: {error}"));
1525            assert_eq!(plan.world_size, cards);
1526            assert_eq!(plan.rank_groups.last().unwrap().layers.end, 45);
1527        }
1528    }
1529
1530    #[test]
1531    fn grouped_step_layout_refuses_gaps_overlap_and_invalid_stage_tp() {
1532        for stages in [
1533            vec![
1534                StageGroupRequest {
1535                    layers: 0..3,
1536                    tensor: 1,
1537                    expert_parallel: false,
1538                },
1539                StageGroupRequest {
1540                    layers: 4..45,
1541                    tensor: 2,
1542                    expert_parallel: true,
1543                },
1544            ],
1545            vec![
1546                StageGroupRequest {
1547                    layers: 0..16,
1548                    tensor: 1,
1549                    expert_parallel: false,
1550                },
1551                StageGroupRequest {
1552                    layers: 15..45,
1553                    tensor: 2,
1554                    expert_parallel: true,
1555                },
1556            ],
1557        ] {
1558            let error = step37_contract()
1559                .plan_grouped(GroupedTopologyRequest {
1560                    stages,
1561                    available_devices: 3,
1562                    hardware: HardwareTarget::RtxPro6000Blackwell,
1563                })
1564                .unwrap_err();
1565            assert!(error.to_string().contains("do not continue"));
1566        }
1567
1568        let tp3 = step37_contract()
1569            .plan_grouped(GroupedTopologyRequest {
1570                stages: vec![StageGroupRequest {
1571                    layers: 0..45,
1572                    tensor: 3,
1573                    expert_parallel: true,
1574                }],
1575                available_devices: 3,
1576                hardware: HardwareTarget::RtxPro6000Blackwell,
1577            })
1578            .unwrap_err();
1579        assert!(tp3.to_string().contains("layer 0 query heads 64"));
1580    }
1581
1582    #[test]
1583    fn full_model_step_tp8_preflight_binds_one_runtime_group() {
1584        let contract = step37_contract();
1585        let devices = (0..8).collect::<Vec<_>>();
1586        let owners = vec![0; contract.trunk_layers];
1587        let plan = contract
1588            .preflight_step_tp_specs(
1589                (0..contract.trunk_layers).map(|layer| (layer, devices.as_slice())),
1590                &owners,
1591            )
1592            .unwrap();
1593
1594        assert!(plan.full_trunk);
1595        assert_eq!(plan.layers.len(), STEP37_TRUNK_LAYERS);
1596        assert_eq!(plan.runtime_groups, vec![devices.clone()]);
1597        assert_eq!(plan.dense_attention_layers(), 3);
1598        assert_eq!(plan.tensor_parallel_expert_layers(), 0);
1599        assert_eq!(plan.expert_parallel_layers(), 42);
1600        assert_eq!(plan.layers.first().unwrap().layer, 0);
1601        assert_eq!(plan.layers.last().unwrap().layer, 44);
1602        assert!(
1603            plan.layers
1604                .iter()
1605                .all(|layer| layer.owner_device == 0 && layer.devices == devices)
1606        );
1607    }
1608
1609    #[test]
1610    fn step_tp_preflight_is_partial_for_tp2_and_fails_closed_on_invalid_specs() {
1611        let contract = step37_contract();
1612        let owners = vec![0; contract.trunk_layers];
1613        let tp2 = vec![0, 1];
1614        let partial = contract
1615            .preflight_step_tp_specs([(3, tp2.as_slice()), (44, tp2.as_slice())], &owners)
1616            .unwrap();
1617        assert!(!partial.full_trunk);
1618        assert_eq!(partial.runtime_groups, vec![tp2]);
1619        assert_eq!(partial.tensor_parallel_expert_layers(), 2);
1620        assert_eq!(partial.expert_parallel_layers(), 0);
1621
1622        let wrong_owner = vec![1, 2];
1623        assert!(
1624            contract
1625                .preflight_step_tp_specs([(24, wrong_owner.as_slice())], &owners)
1626                .unwrap_err()
1627                .to_string()
1628                .contains("owning PP device 0 must be the first rank")
1629        );
1630
1631        let tp3 = vec![0, 1, 2];
1632        assert!(
1633            contract
1634                .preflight_step_tp_specs([(24, tp3.as_slice())], &owners)
1635                .unwrap_err()
1636                .to_string()
1637                .contains("layer 0 query heads 64")
1638        );
1639
1640        assert!(
1641            contract
1642                .preflight_step_tp_specs(
1643                    [(24, [0, 1].as_slice()), (24, [0, 1].as_slice())],
1644                    &owners,
1645                )
1646                .unwrap_err()
1647                .to_string()
1648                .contains("assigns layer 24 more than once")
1649        );
1650    }
1651
1652    #[test]
1653    fn step_tp4_requires_whole_expert_parallelism() {
1654        let error = step37_contract().plan(request(1, 4, false)).unwrap_err();
1655        assert!(
1656            error
1657                .to_string()
1658                .contains("routed expert FFN size shard 320")
1659        );
1660        let tp2 = step37_contract().plan(request(1, 2, false)).unwrap();
1661        assert_eq!(tp2.routed_expert_ffn_range(1), Some(640..1280));
1662        assert!(tp2.shared_expert_replicated);
1663    }
1664
1665    #[test]
1666    fn step_tp3_refuses_the_real_per_layer_head_geometry() {
1667        let error = step37_contract()
1668            .plan(TopologyRequest {
1669                pipeline: 1,
1670                tensor: 3,
1671                expert_parallel: true,
1672                available_devices: 3,
1673                hardware: HardwareTarget::RtxPro6000Blackwell,
1674            })
1675            .unwrap_err();
1676        assert!(error.to_string().contains("layer 0 query heads 64"));
1677    }
1678
1679    #[test]
1680    fn product_envelope_accepts_eight_and_refuses_more() {
1681        let pp8 = step37_contract().plan(request(8, 1, false)).unwrap();
1682        assert_eq!(pp8.world_size, 8);
1683        assert_eq!(pp8.stage_ranges.len(), 8);
1684        assert!(pp8.stage_ranges.iter().all(|range| !range.is_empty()));
1685
1686        let error = step37_contract()
1687            .plan(TopologyRequest {
1688                pipeline: 3,
1689                tensor: 4,
1690                expert_parallel: true,
1691                available_devices: 12,
1692                hardware: HardwareTarget::RtxPro6000Blackwell,
1693            })
1694            .unwrap_err();
1695        assert!(error.to_string().contains("product envelope is 8"));
1696    }
1697
1698    #[test]
1699    fn expert_parallel_requires_a_multi_rank_tp_group() {
1700        let error = step37_contract().plan(request(3, 1, true)).unwrap_err();
1701        assert!(
1702            error
1703                .to_string()
1704                .contains("expert parallelism requires TP group size greater than one")
1705        );
1706    }
1707
1708    #[test]
1709    fn step_does_not_inherit_the_5090_hardware_contract() {
1710        let error = step37_contract()
1711            .plan(TopologyRequest {
1712                pipeline: 1,
1713                tensor: 1,
1714                expert_parallel: false,
1715                available_devices: 1,
1716                hardware: HardwareTarget::Rtx5090,
1717            })
1718            .unwrap_err();
1719        assert!(
1720            error
1721                .to_string()
1722                .contains("has no qualified rtx-5090 contract")
1723        );
1724    }
1725}