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::{Arch, ModelConfig};
13
14/// The execution planner's supported rank envelope. Hardware qualification and tuned defaults
15/// remain model x rig evidence, but the placement/runtime contract must not stop at the three
16/// cards currently available on Pod B.
17pub const PRODUCT_MAX_CARDS: usize = 8;
18const STEP_FP8_BLOCK: usize = 128;
19
20#[derive(Debug, Clone, Copy, PartialEq, Eq)]
21pub enum HardwareTarget {
22    Rtx5090,
23    RtxPro6000Blackwell,
24}
25
26impl HardwareTarget {
27    fn max_cards(self) -> usize {
28        match self {
29            Self::Rtx5090 => 1,
30            Self::RtxPro6000Blackwell => PRODUCT_MAX_CARDS,
31        }
32    }
33
34    fn label(self) -> &'static str {
35        match self {
36            Self::Rtx5090 => "rtx-5090",
37            Self::RtxPro6000Blackwell => "rtx-pro-6000-blackwell",
38        }
39    }
40
41    fn from_device_name(name: &str) -> Result<Self, TopologyError> {
42        if name.contains("RTX PRO 6000") && name.contains("Blackwell") {
43            return Ok(Self::RtxPro6000Blackwell);
44        }
45        if name.contains("RTX 5090") {
46            return Ok(Self::Rtx5090);
47        }
48        Err(TopologyError::new(format!(
49            "unqualified CUDA device {name:?}; first-class targets are RTX 5090 and RTX PRO 6000 \
50             Blackwell"
51        )))
52    }
53}
54
55#[derive(Debug, Clone, Copy, PartialEq, Eq)]
56pub struct TopologyRequest {
57    pub pipeline: usize,
58    pub tensor: usize,
59    /// Routed experts are partitioned across the TP group. When false, each rank owns every
60    /// expert and tensor-shards the expert projections instead.
61    pub expert_parallel: bool,
62    pub available_devices: usize,
63    pub hardware: HardwareTarget,
64}
65
66impl TopologyRequest {
67    pub fn world_size(self) -> Result<usize, TopologyError> {
68        self.pipeline
69            .checked_mul(self.tensor)
70            .ok_or_else(|| TopologyError::new("PP x TP world size overflow"))
71    }
72}
73
74#[derive(Debug, Clone, PartialEq, Eq)]
75pub struct ModelParallelContract {
76    pub family: &'static str,
77    pub variant: String,
78    pub trunk_layers: usize,
79    pub mtp_layers: usize,
80    pub hidden_size: usize,
81    pub vocab_size: usize,
82    pub dense_ffn_size: usize,
83    pub dense_prefix_layers: usize,
84    pub head_dim: usize,
85    pub query_heads: Vec<usize>,
86    pub kv_heads: Vec<usize>,
87    pub expert_count: usize,
88    pub experts_per_token: usize,
89    pub expert_ffn_size: usize,
90    pub shared_expert_ffn_size: usize,
91    pub hardware_targets: Vec<HardwareTarget>,
92}
93
94impl ModelParallelContract {
95    /// Build the model-specific contract. Unregistered families refuse rather than inheriting a
96    /// generic transformer assumption.
97    pub fn from_model(cfg: &ModelConfig) -> Result<Self, TopologyError> {
98        match cfg.arch {
99            Arch::Step35 => Self::step35(cfg),
100            _ => Err(TopologyError::new(format!(
101                "no parallel contract registered for model family {:?}; loading/running does not \
102                 establish TP/EP support",
103                cfg.arch
104            ))),
105        }
106    }
107
108    fn step35(cfg: &ModelConfig) -> Result<Self, TopologyError> {
109        let step = cfg.step35.as_ref().ok_or_else(|| {
110            TopologyError::new(
111                "step35 parallel contract requires the model-specific per-layer Step geometry",
112            )
113        })?;
114        let total_layers = cfg.n_layer as usize;
115        let mtp_layers = cfg.nextn_predict_layers as usize;
116        let trunk_layers = total_layers.checked_sub(mtp_layers).ok_or_else(|| {
117            TopologyError::new(format!(
118                "step35 layer geometry invalid: total={total_layers} mtp={mtp_layers}"
119            ))
120        })?;
121        if trunk_layers == 0 {
122            return Err(TopologyError::new("step35 contract has no trunk layers"));
123        }
124        if step.head_count.len() < total_layers || step.head_count_kv.len() < total_layers {
125            return Err(TopologyError::new(format!(
126                "step35 per-layer head geometry incomplete: q={} kv={} need={total_layers}",
127                step.head_count.len(),
128                step.head_count_kv.len()
129            )));
130        }
131        let moe = cfg.moe.as_ref().ok_or_else(|| {
132            TopologyError::new("step35 parallel contract requires routed-expert geometry")
133        })?;
134        let query_heads: Vec<usize> = (0..total_layers)
135            .map(|il| cfg.n_head_at(il as u32) as usize)
136            .collect();
137        let kv_heads: Vec<usize> = (0..total_layers)
138            .map(|il| cfg.n_head_kv_at(il as u32) as usize)
139            .collect();
140        let is_step37 = trunk_layers == 45
141            && mtp_layers == 3
142            && cfg.n_embd == 4096
143            && cfg.n_ff == 11_264
144            && cfg.n_vocab == 128_896
145            && query_heads
146                .iter()
147                .enumerate()
148                .all(|(il, &heads)| heads == if il % 4 == 0 { 64 } else { 96 })
149            && kv_heads.iter().all(|&heads| heads == 8)
150            && moe.expert_count == 288
151            && moe.expert_used_count == 8
152            && moe.expert_ff_length == 1280
153            && moe.expert_shared_ff_length == 1280
154            && step.first_k_dense_replace == 3;
155        if !is_step37 {
156            return Err(TopologyError::new(format!(
157                "no qualified parallel contract for step35 variant {:?}: only the exact \
158                 Step-3.7-Flash geometry is registered",
159                cfg.name
160            )));
161        }
162
163        Ok(Self {
164            family: "step35",
165            variant: "Step-3.7-Flash-FP8".to_string(),
166            trunk_layers,
167            mtp_layers,
168            hidden_size: cfg.n_embd as usize,
169            vocab_size: cfg.n_vocab as usize,
170            dense_ffn_size: cfg.n_ff as usize,
171            dense_prefix_layers: step.first_k_dense_replace as usize,
172            head_dim: cfg.head_dim_k as usize,
173            query_heads,
174            kv_heads,
175            expert_count: moe.expert_count as usize,
176            experts_per_token: moe.expert_used_count as usize,
177            expert_ffn_size: moe.expert_ff_length as usize,
178            shared_expert_ffn_size: moe.expert_shared_ff_length as usize,
179            hardware_targets: vec![HardwareTarget::RtxPro6000Blackwell],
180        })
181    }
182
183    pub fn plan(&self, request: TopologyRequest) -> Result<ParallelPlan, TopologyError> {
184        let pp = request.pipeline;
185        let tp = request.tensor;
186        if !(1..=PRODUCT_MAX_CARDS).contains(&pp) {
187            return Err(TopologyError::new(format!(
188                "PP size {pp} outside product range 1..={PRODUCT_MAX_CARDS}"
189            )));
190        }
191        if !(1..=PRODUCT_MAX_CARDS).contains(&tp) {
192            return Err(TopologyError::new(format!(
193                "TP size {tp} outside product range 1..={PRODUCT_MAX_CARDS}"
194            )));
195        }
196        let world = request.world_size()?;
197        if world > PRODUCT_MAX_CARDS {
198            return Err(TopologyError::new(format!(
199                "PP={pp} x TP={tp} requires {world} cards; product envelope is \
200                 {PRODUCT_MAX_CARDS}"
201            )));
202        }
203        if !self.hardware_targets.contains(&request.hardware) {
204            return Err(TopologyError::new(format!(
205                "{} has no qualified {} contract",
206                self.variant,
207                request.hardware.label()
208            )));
209        }
210        if world > request.hardware.max_cards() {
211            return Err(TopologyError::new(format!(
212                "{} target permits at most {} card(s), requested {world}",
213                request.hardware.label(),
214                request.hardware.max_cards()
215            )));
216        }
217        if request.available_devices < world {
218            return Err(TopologyError::new(format!(
219                "PP={pp} x TP={tp} requires {world} cards, only {} available",
220                request.available_devices
221            )));
222        }
223        if pp > self.trunk_layers {
224            return Err(TopologyError::new(format!(
225                "PP={pp} exceeds {} trunk layers",
226                self.trunk_layers
227            )));
228        }
229        if request.expert_parallel && tp == 1 {
230            return Err(TopologyError::new(
231                "expert parallelism requires TP group size greater than one",
232            ));
233        }
234
235        // Check the family-specific, per-layer attention geometry before generic dimensions so a
236        // refused topology names the model program that actually makes it invalid.
237        for (il, (&q, &kv)) in self.query_heads.iter().zip(&self.kv_heads).enumerate() {
238            require_divisible(&format!("layer {il} query heads"), q, tp)?;
239            require_divisible(&format!("layer {il} KV heads"), kv, tp)?;
240        }
241        require_divisible("hidden size", self.hidden_size, tp)?;
242        require_divisible("vocabulary size", self.vocab_size, tp)?;
243        require_fp8_block_shard("dense FFN size", self.dense_ffn_size, tp)?;
244        if request.expert_parallel {
245            require_divisible("routed expert count", self.expert_count, tp)?;
246        } else {
247            require_fp8_block_shard("routed expert FFN size", self.expert_ffn_size, tp)?;
248        }
249
250        let stage_ranges = (0..pp)
251            .map(|stage| stage * self.trunk_layers / pp..(stage + 1) * self.trunk_layers / pp)
252            .collect();
253
254        Ok(ParallelPlan {
255            contract: self.clone(),
256            request,
257            world_size: world,
258            stage_ranges,
259            mtp_owner_stage: self.mtp_layers.gt(&0).then_some(pp - 1),
260            // Step's 1280-wide shared expert cannot be split four or eight ways without cutting
261            // through checkpoint 128-row E4M3 scale blocks. Replication is the exact program for
262            // those TP/EP layouts; only routed experts are distributed.
263            shared_expert_replicated: tp > 1 && self.shared_expert_ffn_size > 0,
264        })
265    }
266}
267
268#[derive(Debug, Clone, PartialEq, Eq)]
269pub struct ParallelPlan {
270    pub contract: ModelParallelContract,
271    pub request: TopologyRequest,
272    pub world_size: usize,
273    pub stage_ranges: Vec<Range<usize>>,
274    /// MTP layers are not pipeline stages of their own; the final PP stage owns them.
275    pub mtp_owner_stage: Option<usize>,
276    pub shared_expert_replicated: bool,
277}
278
279impl ParallelPlan {
280    pub fn global_rank(&self, pipeline_rank: usize, tensor_rank: usize) -> Option<usize> {
281        if pipeline_rank >= self.request.pipeline || tensor_rank >= self.request.tensor {
282            return None;
283        }
284        Some(pipeline_rank * self.request.tensor + tensor_rank)
285    }
286
287    pub fn query_head_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
288        split_range(
289            *self.contract.query_heads.get(layer)?,
290            self.request.tensor,
291            tensor_rank,
292        )
293    }
294
295    pub fn kv_head_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
296        split_range(
297            *self.contract.kv_heads.get(layer)?,
298            self.request.tensor,
299            tensor_rank,
300        )
301    }
302
303    /// Column-parallel Q output range and the matching row-parallel O input range.
304    pub fn query_feature_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
305        let heads = self.query_head_range(layer, tensor_rank)?;
306        Some(heads.start * self.contract.head_dim..heads.end * self.contract.head_dim)
307    }
308
309    /// Column-parallel K/V output range. Step-3.7 has eight KV heads, so TP2 and TP4 partition
310    /// them exactly; no KV-head replication is part of this registered contract.
311    pub fn kv_feature_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
312        let heads = self.kv_head_range(layer, tensor_rank)?;
313        Some(heads.start * self.contract.head_dim..heads.end * self.contract.head_dim)
314    }
315
316    /// Column-parallel dense gate/up output range and matching row-parallel down input range.
317    pub fn dense_ffn_range(&self, tensor_rank: usize) -> Option<Range<usize>> {
318        split_range(
319            self.contract.dense_ffn_size,
320            self.request.tensor,
321            tensor_rank,
322        )
323    }
324
325    pub fn routed_expert_range(&self, tensor_rank: usize) -> Option<Range<usize>> {
326        self.request
327            .expert_parallel
328            .then(|| split_range(self.contract.expert_count, self.request.tensor, tensor_rank))?
329    }
330
331    pub fn routed_expert_ffn_range(&self, tensor_rank: usize) -> Option<Range<usize>> {
332        (!self.request.expert_parallel).then(|| {
333            split_range(
334                self.contract.expert_ffn_size,
335                self.request.tensor,
336                tensor_rank,
337            )
338        })?
339    }
340}
341
342/// Validate the live Step PP request before the loader allocates CUDA state. Checkpoint tensor
343/// census is deliberately a separate loader gate: topology legality must remain testable without
344/// opening model files, while serving requires both gates.
345pub fn validate_step_pp_request(cfg: &ModelConfig) -> Result<Option<ParallelPlan>, TopologyError> {
346    let pp = match std::env::var("MEMRA_PP_STAGES") {
347        Err(_) => return Ok(None),
348        Ok(value) if value.is_empty() || value == "0" || value == "1" => return Ok(None),
349        Ok(value) => value.parse::<usize>().map_err(|_| {
350            TopologyError::new(format!("MEMRA_PP_STAGES={value} is not a positive integer"))
351        })?,
352    };
353    let devices = selected_pp_devices(pp)?;
354    let hardware = detect_uniform_hardware(&devices)?;
355    let contract = ModelParallelContract::from_model(cfg)?;
356    let trunk_layers = contract.trunk_layers;
357    let plan = contract.plan(TopologyRequest {
358        pipeline: pp,
359        tensor: 1,
360        expert_parallel: false,
361        available_devices: devices.len(),
362        hardware,
363    })?;
364    let fence = crate::pp::pp_cuts(trunk_layers).ok_or_else(|| {
365        TopologyError::new(format!(
366            "Step PP={pp} has no valid runtime stage fence over {trunk_layers} trunk layers"
367        ))
368    })?;
369    let plan = apply_stage_fence(plan, &fence)?;
370    Ok(Some(plan))
371}
372
373fn apply_stage_fence(
374    mut plan: ParallelPlan,
375    fence: &[usize],
376) -> Result<ParallelPlan, TopologyError> {
377    let expected = plan.request.pipeline + 1;
378    if fence.len() != expected
379        || fence.first() != Some(&0)
380        || fence.last() != Some(&plan.contract.trunk_layers)
381        || fence.windows(2).any(|window| window[0] >= window[1])
382    {
383        return Err(TopologyError::new(format!(
384            "invalid PP fence {fence:?} for {} stages over {} trunk layers",
385            plan.request.pipeline, plan.contract.trunk_layers
386        )));
387    }
388    plan.stage_ranges = fence
389        .windows(2)
390        .map(|window| window[0]..window[1])
391        .collect();
392    Ok(plan)
393}
394
395fn selected_pp_devices(pp: usize) -> Result<Vec<usize>, TopologyError> {
396    let raw = std::env::var("MEMRA_PP_DEVICES").map_err(|_| {
397        TopologyError::new(format!(
398            "Step PP={pp} requires explicit MEMRA_PP_DEVICES with one distinct CUDA ordinal per \
399             stage; same-device diagnostics do not qualify the multi-card product"
400        ))
401    })?;
402    let devices: Result<Vec<usize>, _> = raw
403        .split(',')
404        .map(|part| part.trim().parse::<usize>())
405        .collect();
406    let devices = devices.map_err(|_| {
407        TopologyError::new(format!(
408            "MEMRA_PP_DEVICES={raw:?} is not a comma-separated CUDA ordinal list"
409        ))
410    })?;
411    if devices.len() != pp {
412        return Err(TopologyError::new(format!(
413            "MEMRA_PP_DEVICES lists {} devices but MEMRA_PP_STAGES={pp}",
414            devices.len()
415        )));
416    }
417    let mut unique = devices.clone();
418    unique.sort_unstable();
419    unique.dedup();
420    if unique.len() != devices.len() {
421        return Err(TopologyError::new(format!(
422            "Step PP={pp} requires {pp} distinct devices; MEMRA_PP_DEVICES={raw:?} repeats an \
423             ordinal"
424        )));
425    }
426    Ok(devices)
427}
428
429fn detect_uniform_hardware(devices: &[usize]) -> Result<HardwareTarget, TopologyError> {
430    cudarc::driver::result::init().map_err(|error| {
431        TopologyError::new(format!("CUDA driver initialization failed: {error}"))
432    })?;
433    let mut target = None;
434    for &ordinal in devices {
435        let device = cudarc::driver::result::device::get(ordinal as i32).map_err(|error| {
436            TopologyError::new(format!("CUDA device {ordinal} lookup failed: {error}"))
437        })?;
438        let name = cudarc::driver::result::device::get_name(device).map_err(|error| {
439            TopologyError::new(format!("CUDA device {ordinal} name lookup failed: {error}"))
440        })?;
441        let current = HardwareTarget::from_device_name(&name)?;
442        if let Some(expected) = target {
443            if current != expected {
444                return Err(TopologyError::new(format!(
445                    "mixed hardware targets in MEMRA_PP_DEVICES: expected {}, device {ordinal} is \
446                     {}",
447                    expected.label(),
448                    current.label()
449                )));
450            }
451        } else {
452            target = Some(current);
453        }
454    }
455    target.ok_or_else(|| TopologyError::new("MEMRA_PP_DEVICES is empty"))
456}
457
458fn require_divisible(label: &str, value: usize, parts: usize) -> Result<(), TopologyError> {
459    if value == 0 {
460        return Err(TopologyError::new(format!("{label} is zero")));
461    }
462    if value % parts != 0 {
463        return Err(TopologyError::new(format!(
464            "{label} {value} is not divisible by TP={parts}"
465        )));
466    }
467    Ok(())
468}
469
470fn require_fp8_block_shard(label: &str, value: usize, parts: usize) -> Result<(), TopologyError> {
471    require_divisible(label, value, parts)?;
472    let local = value / parts;
473    if local % STEP_FP8_BLOCK != 0 {
474        return Err(TopologyError::new(format!(
475            "{label} shard {local} for TP={parts} cuts through the Step E4M3 block size \
476             {STEP_FP8_BLOCK}"
477        )));
478    }
479    Ok(())
480}
481
482fn split_range(total: usize, parts: usize, rank: usize) -> Option<Range<usize>> {
483    if parts == 0 || rank >= parts || total % parts != 0 {
484        return None;
485    }
486    let width = total / parts;
487    Some(rank * width..(rank + 1) * width)
488}
489
490#[derive(Debug, Clone, PartialEq, Eq)]
491pub struct TopologyError {
492    message: String,
493}
494
495impl TopologyError {
496    fn new(message: impl Into<String>) -> Self {
497        Self {
498            message: message.into(),
499        }
500    }
501}
502
503impl fmt::Display for TopologyError {
504    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
505        self.message.fmt(f)
506    }
507}
508
509impl std::error::Error for TopologyError {}
510
511#[cfg(test)]
512mod tests {
513    use super::*;
514    use memra_gguf::config::{MoeConfig, Step35Config};
515
516    fn step37_contract() -> ModelParallelContract {
517        let total_layers = 48;
518        ModelParallelContract {
519            family: "step35",
520            variant: "Step-3.7-Flash-FP8".to_string(),
521            trunk_layers: 45,
522            mtp_layers: 3,
523            hidden_size: 4096,
524            vocab_size: 128_896,
525            dense_ffn_size: 11_264,
526            dense_prefix_layers: 3,
527            head_dim: 128,
528            query_heads: (0..total_layers)
529                .map(|il| if il % 4 == 0 { 64 } else { 96 })
530                .collect(),
531            kv_heads: vec![8; total_layers],
532            expert_count: 288,
533            experts_per_token: 8,
534            expert_ffn_size: 1280,
535            shared_expert_ffn_size: 1280,
536            hardware_targets: vec![HardwareTarget::RtxPro6000Blackwell],
537        }
538    }
539
540    fn step37_model_config() -> ModelConfig {
541        let total_layers = 48;
542        let head_count: Vec<u32> = (0..total_layers)
543            .map(|il| if il % 4 == 0 { 64 } else { 96 })
544            .collect();
545        ModelConfig {
546            arch: Arch::Step35,
547            name: "Step-3.7-Flash-FP8".to_string(),
548            n_layer: total_layers,
549            n_embd: 4096,
550            n_head: 96,
551            n_head_kv: 8,
552            head_dim_k: 128,
553            head_dim_v: 128,
554            n_ff: 11_264,
555            n_vocab: 128_896,
556            context_length: 262_144,
557            rms_eps: 1e-6,
558            rope_freq_base: 5_000_000.0,
559            rope_dim_count: 128,
560            rope_sections: Vec::new(),
561            full_attention_interval: 0,
562            ssm: None,
563            moe: Some(MoeConfig {
564                expert_count: 288,
565                expert_used_count: 8,
566                expert_ff_length: 1280,
567                expert_shared_ff_length: 1280,
568            }),
569            m3: None,
570            hy3: None,
571            gemma4: None,
572            mla: None,
573            step35: Some(Step35Config {
574                head_count,
575                head_count_kv: vec![8; total_layers as usize],
576                swa_pattern: (0..total_layers).map(|il| il % 4 != 0).collect(),
577                sliding_window: 512,
578                rope_base_global: 5_000_000.0,
579                rope_base_swa: 10_000.0,
580                rope_dims_full: 64,
581                rope_dims_swa: 128,
582                swiglu_clamp_exp: vec![0.0; total_layers as usize],
583                swiglu_clamp_shexp: vec![0.0; total_layers as usize],
584                sigmoid_routing: true,
585                routed_scaling_factor: 3.0,
586                route_norm: true,
587                first_k_dense_replace: 3,
588            }),
589            geometry: None,
590            nextn_predict_layers: 3,
591            n_layer_total: total_layers,
592        }
593    }
594
595    fn request(pp: usize, tp: usize, expert_parallel: bool) -> TopologyRequest {
596        TopologyRequest {
597            pipeline: pp,
598            tensor: tp,
599            expert_parallel,
600            available_devices: pp * tp,
601            hardware: HardwareTarget::RtxPro6000Blackwell,
602        }
603    }
604
605    #[test]
606    fn step_pp3_maps_fifteen_trunk_layers_per_card() {
607        let plan = step37_contract().plan(request(3, 1, false)).unwrap();
608        assert_eq!(plan.world_size, 3);
609        assert_eq!(plan.stage_ranges, vec![0..15, 15..30, 30..45]);
610        assert_eq!(plan.mtp_owner_stage, Some(2));
611    }
612
613    #[test]
614    fn step_pp_marker_uses_the_runtime_stage_fence() {
615        let plan = step37_contract().plan(request(3, 1, false)).unwrap();
616        let plan = apply_stage_fence(plan, &[0, 10, 28, 45]).unwrap();
617        assert_eq!(plan.stage_ranges, vec![0..10, 10..28, 28..45]);
618    }
619
620    #[test]
621    fn step_contract_is_extracted_from_model_specific_geometry() {
622        let contract = ModelParallelContract::from_model(&step37_model_config()).unwrap();
623        assert_eq!(contract.family, "step35");
624        assert_eq!(contract.trunk_layers, 45);
625        assert_eq!(contract.mtp_layers, 3);
626        assert_eq!(contract.query_heads[0], 64);
627        assert_eq!(contract.query_heads[1], 96);
628        assert_eq!(contract.kv_heads[47], 8);
629        assert_eq!(contract.expert_count, 288);
630        assert_eq!(contract.experts_per_token, 8);
631    }
632
633    #[test]
634    fn step_sibling_does_not_inherit_the_step37_contract() {
635        let mut sibling = step37_model_config();
636        sibling.name = "Step-3.5-Flash".to_string();
637        sibling.n_vocab = 128_000;
638        let error = ModelParallelContract::from_model(&sibling).unwrap_err();
639        assert!(
640            error
641                .to_string()
642                .contains("only the exact Step-3.7-Flash geometry is registered")
643        );
644    }
645
646    #[test]
647    fn step_without_the_official_mtp_geometry_does_not_inherit_the_contract() {
648        let mut stripped = step37_model_config();
649        stripped.nextn_predict_layers = 0;
650        let error = ModelParallelContract::from_model(&stripped).unwrap_err();
651        assert!(
652            error
653                .to_string()
654                .contains("only the exact Step-3.7-Flash geometry is registered")
655        );
656    }
657
658    #[test]
659    fn hardware_target_classification_is_exact() {
660        assert_eq!(
661            HardwareTarget::from_device_name("NVIDIA RTX PRO 6000 Blackwell Server Edition")
662                .unwrap(),
663            HardwareTarget::RtxPro6000Blackwell
664        );
665        assert_eq!(
666            HardwareTarget::from_device_name("NVIDIA GeForce RTX 5090 Laptop GPU").unwrap(),
667            HardwareTarget::Rtx5090
668        );
669        assert!(HardwareTarget::from_device_name("NVIDIA H100 80GB HBM3").is_err());
670    }
671
672    #[test]
673    fn step_tp2_tp4_tp8_and_hybrid_plans_are_geometry_valid() {
674        let tp2 = step37_contract().plan(request(1, 2, true)).unwrap();
675        assert_eq!(tp2.query_head_range(0, 1), Some(32..64));
676        assert_eq!(tp2.query_head_range(1, 1), Some(48..96));
677        assert_eq!(tp2.kv_head_range(0, 1), Some(4..8));
678        assert_eq!(tp2.routed_expert_range(1), Some(144..288));
679
680        let tp4 = step37_contract().plan(request(1, 4, true)).unwrap();
681        assert_eq!(tp4.query_head_range(0, 3), Some(48..64));
682        assert_eq!(tp4.query_head_range(1, 3), Some(72..96));
683        assert_eq!(tp4.kv_head_range(0, 3), Some(6..8));
684        assert_eq!(tp4.query_feature_range(0, 3), Some(6144..8192));
685        assert_eq!(tp4.query_feature_range(1, 3), Some(9216..12_288));
686        assert_eq!(tp4.kv_feature_range(0, 3), Some(768..1024));
687        assert_eq!(tp4.dense_ffn_range(3), Some(8448..11_264));
688        assert_eq!(tp4.routed_expert_range(3), Some(216..288));
689        assert!(tp4.shared_expert_replicated);
690
691        let tp8 = step37_contract().plan(request(1, 8, true)).unwrap();
692        assert_eq!(tp8.query_head_range(0, 7), Some(56..64));
693        assert_eq!(tp8.query_head_range(1, 7), Some(84..96));
694        assert_eq!(tp8.kv_head_range(0, 7), Some(7..8));
695        assert_eq!(tp8.dense_ffn_range(7), Some(9856..11_264));
696        assert_eq!(tp8.routed_expert_range(7), Some(252..288));
697        assert!(tp8.shared_expert_replicated);
698
699        let hybrid = step37_contract().plan(request(2, 4, true)).unwrap();
700        assert_eq!(hybrid.world_size, 8);
701        assert_eq!(hybrid.stage_ranges, vec![0..22, 22..45]);
702        assert_eq!(hybrid.global_rank(1, 3), Some(7));
703        assert_eq!(hybrid.global_rank(2, 0), None);
704    }
705
706    #[test]
707    fn step_tp4_requires_whole_expert_parallelism() {
708        let error = step37_contract().plan(request(1, 4, false)).unwrap_err();
709        assert!(
710            error
711                .to_string()
712                .contains("routed expert FFN size shard 320")
713        );
714        let tp2 = step37_contract().plan(request(1, 2, false)).unwrap();
715        assert_eq!(tp2.routed_expert_ffn_range(1), Some(640..1280));
716        assert!(tp2.shared_expert_replicated);
717    }
718
719    #[test]
720    fn step_tp3_refuses_the_real_per_layer_head_geometry() {
721        let error = step37_contract()
722            .plan(TopologyRequest {
723                pipeline: 1,
724                tensor: 3,
725                expert_parallel: true,
726                available_devices: 3,
727                hardware: HardwareTarget::RtxPro6000Blackwell,
728            })
729            .unwrap_err();
730        assert!(error.to_string().contains("layer 0 query heads 64"));
731    }
732
733    #[test]
734    fn product_envelope_accepts_eight_and_refuses_more() {
735        let pp8 = step37_contract().plan(request(8, 1, false)).unwrap();
736        assert_eq!(pp8.world_size, 8);
737        assert_eq!(pp8.stage_ranges.len(), 8);
738        assert!(pp8.stage_ranges.iter().all(|range| !range.is_empty()));
739
740        let error = step37_contract()
741            .plan(TopologyRequest {
742                pipeline: 3,
743                tensor: 4,
744                expert_parallel: true,
745                available_devices: 12,
746                hardware: HardwareTarget::RtxPro6000Blackwell,
747            })
748            .unwrap_err();
749        assert!(error.to_string().contains("product envelope is 8"));
750    }
751
752    #[test]
753    fn expert_parallel_requires_a_multi_rank_tp_group() {
754        let error = step37_contract().plan(request(3, 1, true)).unwrap_err();
755        assert!(
756            error
757                .to_string()
758                .contains("expert parallelism requires TP group size greater than one")
759        );
760    }
761
762    #[test]
763    fn step_does_not_inherit_the_5090_hardware_contract() {
764        let error = step37_contract()
765            .plan(TopologyRequest {
766                pipeline: 1,
767                tensor: 1,
768                expert_parallel: false,
769                available_devices: 1,
770                hardware: HardwareTarget::Rtx5090,
771            })
772            .unwrap_err();
773        assert!(
774            error
775                .to_string()
776                .contains("has no qualified rtx-5090 contract")
777        );
778    }
779}