Skip to main content

memra_engine/
parallel.rs

1//! ModelPlan-driven parallel topology and artifact placement contracts.
2//!
3//! The automatic loader does not select a family-specific TP/EP recipe. It compiles the canonical
4//! operation plan, binds the source tensor census, estimates the checkpoint residency of every
5//! legal program, and then selects a registered numeric backend. Family packs remain responsible
6//! for semantic tensor/config validation; they do not carry per-layer placement lists.
7
8use std::fmt;
9use std::ops::Range;
10
11use memra_gguf::config::ModelConfig;
12use memra_gguf::model_plan::{MlpPlan, ModelPlan};
13use memra_gguf::placement::{LayerPlacementCost, PlacementRequest, plan_contiguous_stages};
14use memra_gguf::source::{ExpertActivationPrecision, TensorSource};
15use memra_gguf::tensor_contract::{
16    ContractOptions, LayerTensor, OutputHead, TensorContract, TensorId, TensorOwner,
17};
18
19/// The execution planner's supported rank envelope. Hardware qualification and tuned defaults
20/// remain model x rig evidence, but the placement/runtime contract must not stop at earlier
21/// three-card qualification cells.
22pub const PRODUCT_MAX_CARDS: usize = 8;
23pub const AUTO_PARALLEL_MAX_CARDS: usize = 4;
24pub const STEP37_TRUNK_LAYERS: usize = 45;
25const STEP_FP8_BLOCK: usize = 128;
26const AUTO_PARALLEL_RESERVE_MB_DEFAULT: u64 = 6_144;
27
28#[derive(Debug, Clone, Copy, PartialEq, Eq)]
29pub enum HardwareTarget {
30    Rtx5090,
31    RtxPro6000Blackwell,
32}
33
34impl HardwareTarget {
35    fn max_cards(self) -> usize {
36        match self {
37            Self::Rtx5090 => 1,
38            Self::RtxPro6000Blackwell => PRODUCT_MAX_CARDS,
39        }
40    }
41
42    fn label(self) -> &'static str {
43        match self {
44            Self::Rtx5090 => "rtx-5090",
45            Self::RtxPro6000Blackwell => "rtx-pro-6000-blackwell",
46        }
47    }
48
49    fn from_device_name(name: &str) -> Result<Self, TopologyError> {
50        if name.contains("RTX PRO 6000") && name.contains("Blackwell") {
51            return Ok(Self::RtxPro6000Blackwell);
52        }
53        if name.contains("RTX 5090") {
54            return Ok(Self::Rtx5090);
55        }
56        Err(TopologyError::new(format!(
57            "unqualified CUDA device {name:?}; first-class targets are RTX 5090 and RTX PRO 6000 \
58             Blackwell"
59        )))
60    }
61}
62
63#[derive(Debug, Clone, Copy, PartialEq, Eq)]
64pub struct TopologyRequest {
65    pub pipeline: usize,
66    pub tensor: usize,
67    /// Routed experts are partitioned across the TP group. When false, each rank owns every
68    /// expert and tensor-shards the expert projections instead.
69    pub expert_parallel: bool,
70    pub available_devices: usize,
71    pub hardware: HardwareTarget,
72}
73
74impl TopologyRequest {
75    pub fn world_size(self) -> Result<usize, TopologyError> {
76        self.pipeline
77            .checked_mul(self.tensor)
78            .ok_or_else(|| TopologyError::new("PP x TP world size overflow"))
79    }
80}
81
82/// One pipeline stage's model-specific tensor/expert group.
83///
84/// Stage groups make odd physical card counts useful without pretending PP is TP. For example,
85/// three cards can run a TP1 dense-prefix stage followed by a TP2/EP2 MoE stage. The layer range is
86/// explicit because memory-balanced Step placement is not necessarily an equal layer split.
87#[derive(Debug, Clone, PartialEq, Eq)]
88pub struct StageGroupRequest {
89    pub layers: Range<usize>,
90    pub tensor: usize,
91    pub expert_parallel: bool,
92}
93
94#[derive(Debug, Clone, PartialEq, Eq)]
95pub struct GroupedTopologyRequest {
96    pub stages: Vec<StageGroupRequest>,
97    pub available_devices: usize,
98    pub hardware: HardwareTarget,
99}
100
101#[derive(Debug, Clone, PartialEq, Eq)]
102pub struct ModelParallelContract {
103    pub family: &'static str,
104    pub variant: String,
105    pub trunk_layers: usize,
106    pub mtp_layers: usize,
107    pub hidden_size: usize,
108    pub vocab_size: usize,
109    pub dense_ffn_size: usize,
110    pub dense_prefix_layers: usize,
111    pub head_dim: usize,
112    pub query_heads: Vec<usize>,
113    pub kv_heads: Vec<usize>,
114    /// True when every layer exposes a Full/Sliding attention geometry that the generic TP
115    /// planner can shard. Expert-only EP does not require this.
116    pub tensor_attention_supported: bool,
117    pub expert_count: usize,
118    pub experts_per_token: usize,
119    pub expert_ffn_size: usize,
120    pub shared_expert_ffn_size: usize,
121    /// Trunk layer indices whose ModelPlan MLP is routed MoE. This is the automatic EP scope.
122    pub routed_layers: Vec<usize>,
123    pub partition_boundaries: Vec<usize>,
124    pub hardware_targets: Vec<HardwareTarget>,
125}
126
127#[derive(Debug, Clone, Copy, PartialEq, Eq)]
128pub(crate) enum AutoParallelBackend {
129    Pipeline,
130    ExpertParallel,
131}
132
133#[derive(Debug, Clone, PartialEq, Eq)]
134pub(crate) struct AutoParallelPlacement {
135    pub backend: AutoParallelBackend,
136    pub devices: Vec<usize>,
137    pub routed_layers: Vec<usize>,
138    pub pipeline_splits: Vec<usize>,
139    pub checkpoint_peak_bytes: u64,
140    pub expert_root_bytes: u64,
141    pub expert_peer_bytes: u64,
142    pub reserve_bytes: u64,
143    pub device_capacity_bytes: Vec<u64>,
144}
145
146#[derive(Debug, Clone, PartialEq, Eq)]
147struct AutoArtifactCosts {
148    layers: Vec<LayerPlacementCost>,
149    first_fixed_bytes: u64,
150    last_fixed_bytes: u64,
151    trunk_expert_bytes: u64,
152    non_distributed_bytes: u64,
153}
154
155fn placement_first_stage_tensor(id: &TensorId) -> bool {
156    match id {
157        TensorId::TokenEmbedding | TensorId::RopeFactors | TensorId::Vision { .. } => true,
158        TensorId::QuantAux { tensor, .. } => placement_first_stage_tensor(tensor),
159        _ => false,
160    }
161}
162
163fn routed_expert_tensor(id: &TensorId) -> bool {
164    match id {
165        TensorId::Expert { .. } => true,
166        TensorId::Layer {
167            tensor:
168                LayerTensor::MoeExpertGateUpBank
169                | LayerTensor::MoeExpertGateBank
170                | LayerTensor::MoeExpertUpBank
171                | LayerTensor::MoeExpertDownBank
172                | LayerTensor::MoeExpertOutputScale,
173            ..
174        } => true,
175        TensorId::QuantAux { tensor, .. } => routed_expert_tensor(tensor),
176        _ => false,
177    }
178}
179
180fn checked_add_bytes(total: &mut u64, bytes: u64, label: &str) -> Result<(), TopologyError> {
181    *total = total
182        .checked_add(bytes)
183        .ok_or_else(|| TopologyError::new(format!("{label} byte total overflows u64")))?;
184    Ok(())
185}
186
187fn artifact_costs(
188    src: &dyn TensorSource,
189    cfg: &ModelConfig,
190    plan: &ModelPlan,
191) -> Result<AutoArtifactCosts, TopologyError> {
192    let census = src.tensor_census().map_err(|error| {
193        TopologyError::new(format!(
194            "automatic parallel placement requires a source tensor census: {error}"
195        ))
196    })?;
197    let output_head = if census
198        .tensors
199        .iter()
200        .any(|row| row.entry.name == "lm_head.weight" || row.entry.name == "output.weight")
201    {
202        OutputHead::Separate
203    } else {
204        OutputHead::TiedToEmbedding
205    };
206    let contract = match memra_gguf::model_packs::for_config(cfg) {
207        Some(pack) => {
208            pack.compile_tensor_contract(cfg, plan, census.dialect, ContractOptions { output_head })
209        }
210        None => TensorContract::for_plan(plan, census.dialect, ContractOptions { output_head }),
211    }
212    .map_err(|error| {
213        TopologyError::new(format!(
214            "cannot compile automatic parallel tensor contract: {error}"
215        ))
216    })?;
217    let entries = census
218        .tensors
219        .iter()
220        .map(|row| row.entry.clone())
221        .collect::<Vec<_>>();
222    let binding = contract.bind(&entries).map_err(|error| {
223        TopologyError::new(format!(
224            "cannot bind automatic parallel tensor census: {error}"
225        ))
226    })?;
227
228    let mut layers = vec![LayerPlacementCost::default(); plan.layers.len()];
229    let mut first_fixed_bytes = 0u64;
230    let mut last_fixed_bytes = 0u64;
231    let mut trunk_expert_bytes = 0u64;
232    let mut total_bytes = 0u64;
233    for (id, tensor) in &binding.tensors {
234        checked_add_bytes(
235            &mut total_bytes,
236            tensor.physical_bytes,
237            "automatic placement checkpoint",
238        )?;
239        match tensor.owner {
240            TensorOwner::Layer(layer) if (layer as usize) < layers.len() => {
241                checked_add_bytes(
242                    &mut layers[layer as usize].weight_bytes,
243                    tensor.physical_bytes,
244                    "automatic placement layer",
245                )?;
246            }
247            // Some legacy contracts retain the physical MTP index rather than rewriting the
248            // owner to TensorOwner::Mtp. It executes with the tail/head stage either way.
249            TensorOwner::Layer(_) => checked_add_bytes(
250                &mut last_fixed_bytes,
251                tensor.physical_bytes,
252                "automatic placement head stage",
253            )?,
254            TensorOwner::Vision(_) => checked_add_bytes(
255                &mut first_fixed_bytes,
256                tensor.physical_bytes,
257                "automatic placement first stage",
258            )?,
259            TensorOwner::Global if placement_first_stage_tensor(id) => checked_add_bytes(
260                &mut first_fixed_bytes,
261                tensor.physical_bytes,
262                "automatic placement first stage",
263            )?,
264            TensorOwner::Global | TensorOwner::Mtp(_) => checked_add_bytes(
265                &mut last_fixed_bytes,
266                tensor.physical_bytes,
267                "automatic placement head stage",
268            )?,
269        }
270
271        let trunk_expert = routed_expert_tensor(id)
272            && matches!(
273                tensor.owner,
274                TensorOwner::Layer(layer)
275                    if (layer as usize) < plan.layers.len()
276                        && matches!(plan.layers[layer as usize].mlp, MlpPlan::Moe(_))
277            );
278        if trunk_expert {
279            checked_add_bytes(
280                &mut trunk_expert_bytes,
281                tensor.physical_bytes,
282                "automatic placement trunk experts",
283            )?;
284        }
285    }
286    let non_distributed_bytes = total_bytes
287        .checked_sub(trunk_expert_bytes)
288        .ok_or_else(|| TopologyError::new("automatic placement expert bytes exceed total bytes"))?;
289    Ok(AutoArtifactCosts {
290        layers,
291        first_fixed_bytes,
292        last_fixed_bytes,
293        trunk_expert_bytes,
294        non_distributed_bytes,
295    })
296}
297
298fn auto_parallel_reserve_bytes() -> Result<u64, TopologyError> {
299    let reserve_mb = match std::env::var("MEMRA_PARALLEL_RESERVE_MB") {
300        Ok(raw) => raw.parse::<u64>().map_err(|_| {
301            TopologyError::new(format!(
302                "MEMRA_PARALLEL_RESERVE_MB={raw:?} is not an unsigned integer"
303            ))
304        })?,
305        Err(std::env::VarError::NotPresent) => AUTO_PARALLEL_RESERVE_MB_DEFAULT,
306        Err(error) => {
307            return Err(TopologyError::new(format!(
308                "cannot read MEMRA_PARALLEL_RESERVE_MB: {error}"
309            )));
310        }
311    };
312    reserve_mb
313        .checked_mul(1024 * 1024)
314        .ok_or_else(|| TopologyError::new("MEMRA_PARALLEL_RESERVE_MB overflows bytes"))
315}
316
317fn device_capacity_bytes(devices: &[usize]) -> Result<Vec<u64>, TopologyError> {
318    cudarc::driver::result::init().map_err(|error| {
319        TopologyError::new(format!("CUDA driver initialization failed: {error}"))
320    })?;
321    devices
322        .iter()
323        .map(|&ordinal| {
324            let device = cudarc::driver::result::device::get(ordinal as i32).map_err(|error| {
325                TopologyError::new(format!("CUDA device {ordinal} lookup failed: {error}"))
326            })?;
327            // SAFETY: `device` was returned by CUDA for this exact process-local ordinal.
328            let bytes =
329                unsafe { cudarc::driver::result::device::total_mem(device) }.map_err(|error| {
330                    TopologyError::new(format!(
331                        "CUDA device {ordinal} memory query failed: {error}"
332                    ))
333                })?;
334            u64::try_from(bytes).map_err(|_| {
335                TopologyError::new(format!("CUDA device {ordinal} memory exceeds u64"))
336            })
337        })
338        .collect()
339}
340
341fn fits_capacity(bytes: u64, reserve: u64, capacity: u64) -> bool {
342    bytes
343        .checked_add(reserve)
344        .is_some_and(|required| required <= capacity)
345}
346
347fn choose_auto_parallel_placement(
348    costs: &AutoArtifactCosts,
349    contract: &ModelParallelContract,
350    activation: ExpertActivationPrecision,
351    devices: &[usize],
352    capacity_bytes: &[u64],
353    reserve_bytes: u64,
354) -> Result<AutoParallelPlacement, TopologyError> {
355    if devices.len() != capacity_bytes.len() {
356        return Err(TopologyError::new(format!(
357            "automatic placement has {} devices but {} capacity rows",
358            devices.len(),
359            capacity_bytes.len()
360        )));
361    }
362    let mut fixed = vec![0u64; devices.len()];
363    fixed[0] = costs.first_fixed_bytes;
364    fixed[devices.len() - 1] = fixed[devices.len() - 1]
365        .checked_add(costs.last_fixed_bytes)
366        .ok_or_else(|| TopologyError::new("automatic placement fixed bytes overflow"))?;
367    let pipeline = plan_contiguous_stages(PlacementRequest {
368        layers: &costs.layers,
369        fixed_stage_bytes: &fixed,
370        context_tokens: 0,
371        devices,
372        legal_boundaries: &contract.partition_boundaries,
373    })
374    .map_err(|error| TopologyError::new(format!("automatic PP placement failed: {error}")))?;
375    let pipeline_fits = pipeline
376        .stages
377        .iter()
378        .enumerate()
379        .all(|(stage, placement)| {
380            fits_capacity(
381                placement.cost.total_bytes,
382                reserve_bytes,
383                capacity_bytes[stage],
384            )
385        });
386
387    let world = devices.len() as u64;
388    let expert_peer_bytes = if contract.expert_count == 0 {
389        0
390    } else {
391        let expert_count = contract.expert_count as u64;
392        let bytes_per_expert = costs.trunk_expert_bytes.div_ceil(expert_count);
393        bytes_per_expert
394            .checked_mul(expert_count.div_ceil(world))
395            .ok_or_else(|| TopologyError::new("automatic EP peer byte total overflows"))?
396    };
397    let expert_root_bytes = costs
398        .non_distributed_bytes
399        .checked_add(expert_peer_bytes)
400        .ok_or_else(|| TopologyError::new("automatic EP root byte total overflows"))?;
401    let expert_fits = !contract.routed_layers.is_empty()
402        && activation == ExpertActivationPrecision::Bf16
403        && capacity_bytes.iter().enumerate().all(|(rank, &capacity)| {
404            let bytes = if rank == 0 {
405                expert_root_bytes
406            } else {
407                expert_peer_bytes
408            };
409            fits_capacity(bytes, reserve_bytes, capacity)
410        });
411
412    if expert_fits {
413        return Ok(AutoParallelPlacement {
414            backend: AutoParallelBackend::ExpertParallel,
415            devices: devices.to_vec(),
416            routed_layers: contract.routed_layers.clone(),
417            pipeline_splits: Vec::new(),
418            checkpoint_peak_bytes: expert_root_bytes,
419            expert_root_bytes,
420            expert_peer_bytes,
421            reserve_bytes,
422            device_capacity_bytes: capacity_bytes.to_vec(),
423        });
424    }
425    if pipeline_fits {
426        let pipeline_splits = pipeline
427            .stages
428            .iter()
429            .take(pipeline.stages.len() - 1)
430            .map(|stage| stage.layers.end)
431            .collect();
432        return Ok(AutoParallelPlacement {
433            backend: AutoParallelBackend::Pipeline,
434            devices: devices.to_vec(),
435            routed_layers: contract.routed_layers.clone(),
436            pipeline_splits,
437            checkpoint_peak_bytes: pipeline.max_stage_bytes,
438            expert_root_bytes,
439            expert_peer_bytes,
440            reserve_bytes,
441            device_capacity_bytes: capacity_bytes.to_vec(),
442        });
443    }
444
445    Err(TopologyError::new(format!(
446        "automatic placement found no capacity-safe program: PP peak={} bytes, EP root={} bytes, \
447         reserve={} bytes, device capacities={capacity_bytes:?}",
448        pipeline.max_stage_bytes, expert_root_bytes, reserve_bytes,
449    )))
450}
451
452pub(crate) fn plan_auto_parallel(
453    src: &dyn TensorSource,
454    cfg: &ModelConfig,
455    plan: &ModelPlan,
456    devices: &[usize],
457) -> Result<AutoParallelPlacement, TopologyError> {
458    let hardware = detect_uniform_hardware(devices)?;
459    let contract = ModelParallelContract::from_plan(cfg, plan)?;
460    if !contract.hardware_targets.contains(&hardware) {
461        return Err(TopologyError::new(format!(
462            "{} has no qualified {} automatic placement contract",
463            contract.variant,
464            hardware.label()
465        )));
466    }
467    let costs = artifact_costs(src, cfg, plan)?;
468    let capacity_bytes = device_capacity_bytes(devices)?;
469    let reserve_bytes = auto_parallel_reserve_bytes()?;
470    choose_auto_parallel_placement(
471        &costs,
472        &contract,
473        src.expert_activation_precision(),
474        devices,
475        &capacity_bytes,
476        reserve_bytes,
477    )
478}
479
480#[derive(Debug, Clone, Copy, PartialEq, Eq)]
481pub(crate) enum StepTpExpertLayout {
482    AttentionOnly,
483    TensorParallel,
484    ExpertParallel,
485}
486
487#[derive(Debug, Clone, PartialEq, Eq)]
488pub(crate) struct StepTpLayerPlan {
489    pub layer: usize,
490    pub devices: Vec<usize>,
491    pub owner_device: usize,
492    pub expert_layout: StepTpExpertLayout,
493}
494
495#[derive(Debug, Clone, PartialEq, Eq)]
496pub(crate) struct StepTpPreflightPlan {
497    pub layers: Vec<StepTpLayerPlan>,
498    pub runtime_groups: Vec<Vec<usize>>,
499    pub full_trunk: bool,
500}
501
502impl StepTpPreflightPlan {
503    pub fn dense_attention_layers(&self) -> usize {
504        self.layers
505            .iter()
506            .filter(|layer| layer.expert_layout == StepTpExpertLayout::AttentionOnly)
507            .count()
508    }
509
510    pub fn tensor_parallel_expert_layers(&self) -> usize {
511        self.layers
512            .iter()
513            .filter(|layer| layer.expert_layout == StepTpExpertLayout::TensorParallel)
514            .count()
515    }
516
517    pub fn expert_parallel_layers(&self) -> usize {
518        self.layers
519            .iter()
520            .filter(|layer| layer.expert_layout == StepTpExpertLayout::ExpertParallel)
521            .count()
522    }
523}
524
525impl ModelParallelContract {
526    /// Build the structural contract from the canonical operation plan. Unsupported operations
527    /// refuse during ModelPlan compilation; family names do not select placement.
528    pub fn from_model(cfg: &ModelConfig) -> Result<Self, TopologyError> {
529        let plan = memra_gguf::model_plan::ModelPlan::compile(cfg).map_err(|error| {
530            TopologyError::new(format!("cannot compile parallel ModelPlan: {error}"))
531        })?;
532        Self::from_plan(cfg, &plan)
533    }
534
535    fn from_plan(
536        cfg: &ModelConfig,
537        plan: &memra_gguf::model_plan::ModelPlan,
538    ) -> Result<Self, TopologyError> {
539        use memra_gguf::model_plan::{AttentionPlan, MlpPlan};
540
541        let trunk_layers = plan.layers.len();
542        let mtp_layers = plan.mtp_blocks.len();
543        if trunk_layers == 0 {
544            return Err(TopologyError::new("parallel contract has no trunk layers"));
545        }
546        let layers: Vec<_> = plan
547            .layers
548            .iter()
549            .chain(plan.mtp_blocks.iter().map(|block| &block.layer))
550            .collect();
551        let attention_geometry = layers
552            .iter()
553            .map(|layer| match &layer.attention {
554                AttentionPlan::Full(attention) | AttentionPlan::SlidingWindow { attention, .. } => {
555                    Some((
556                        attention.query_heads as usize,
557                        attention.kv_heads as usize,
558                        attention.key_head_dim as usize,
559                    ))
560                }
561                _ => None,
562            })
563            .collect::<Vec<_>>();
564        let query_heads: Vec<_> = attention_geometry
565            .iter()
566            .map(|geometry| geometry.map_or(0, |geometry| geometry.0))
567            .collect();
568        let kv_heads: Vec<_> = attention_geometry
569            .iter()
570            .map(|geometry| geometry.map_or(0, |geometry| geometry.1))
571            .collect();
572        let head_dim = attention_geometry
573            .iter()
574            .flatten()
575            .map(|geometry| geometry.2)
576            .next()
577            .unwrap_or(cfg.head_dim_k as usize);
578        let tensor_attention_supported = attention_geometry
579            .iter()
580            .all(|geometry| geometry.is_some_and(|geometry| geometry.2 == head_dim));
581        let dense_prefix_layers = plan
582            .layers
583            .iter()
584            .take_while(|layer| matches!(layer.mlp, MlpPlan::Dense(_)))
585            .count();
586        if plan.layers[dense_prefix_layers..]
587            .iter()
588            .any(|layer| matches!(layer.mlp, MlpPlan::Dense(_)))
589        {
590            return Err(TopologyError::new(
591                "generic parallel loader requires dense layers to form one prefix before routed \
592                 MoE layers",
593            ));
594        }
595        let dense_sizes = layers
596            .iter()
597            .filter_map(|layer| match &layer.mlp {
598                MlpPlan::Dense(dense) => Some(dense.intermediate_size as usize),
599                _ => None,
600            })
601            .collect::<std::collections::BTreeSet<_>>();
602        if dense_sizes.len() > 1 {
603            return Err(TopologyError::new(format!(
604                "generic parallel loader requires one dense FFN width, got {dense_sizes:?}"
605            )));
606        }
607        let dense_ffn_size = dense_sizes.iter().next().copied().unwrap_or(0);
608        let routed_layers = plan
609            .layers
610            .iter()
611            .enumerate()
612            .filter_map(|(layer, plan)| match plan.mlp {
613                MlpPlan::Moe(_) => Some(layer),
614                MlpPlan::Dense(_) => None,
615            })
616            .collect::<Vec<_>>();
617        let moe_layers = layers
618            .iter()
619            .filter_map(|layer| match &layer.mlp {
620                MlpPlan::Moe(moe) => Some(moe),
621                MlpPlan::Dense(_) => None,
622            })
623            .collect::<Vec<_>>();
624        let (expert_count, experts_per_token, expert_ffn_size, shared_expert_ffn_size) =
625            if let Some(first) = moe_layers.first() {
626                let shared = first
627                    .shared
628                    .as_ref()
629                    .map_or(0, |shared| shared.intermediate_size as usize);
630                if moe_layers.iter().any(|moe| {
631                    moe.expert_count != first.expert_count
632                        || moe.experts_per_token != first.experts_per_token
633                        || moe.expert_intermediate_size != first.expert_intermediate_size
634                        || moe
635                            .shared
636                            .as_ref()
637                            .map_or(0, |shared| shared.intermediate_size as usize)
638                            != shared
639                }) {
640                    return Err(TopologyError::new(
641                        "generic parallel loader requires one routed-expert geometry across the \
642                         selected model plan",
643                    ));
644                }
645                (
646                    first.expert_count as usize,
647                    first.experts_per_token as usize,
648                    first.expert_intermediate_size as usize,
649                    shared,
650                )
651            } else {
652                (0, 0, 0, 0)
653            };
654        if routed_layers.is_empty() && dense_ffn_size == 0 {
655            return Err(TopologyError::new(
656                "generic parallel loader found neither dense nor routed MLP layers",
657            ));
658        };
659
660        Ok(Self {
661            family: if routed_layers.is_empty() {
662                "dense-transformer"
663            } else {
664                "routed-moe"
665            },
666            variant: cfg.name.clone(),
667            trunk_layers,
668            mtp_layers,
669            hidden_size: cfg.n_embd as usize,
670            vocab_size: cfg.n_vocab as usize,
671            dense_ffn_size,
672            dense_prefix_layers,
673            head_dim,
674            query_heads,
675            kv_heads,
676            tensor_attention_supported,
677            expert_count,
678            experts_per_token,
679            expert_ffn_size,
680            shared_expert_ffn_size,
681            routed_layers,
682            partition_boundaries: plan.partition_boundaries.clone(),
683            hardware_targets: vec![HardwareTarget::RtxPro6000Blackwell],
684        })
685    }
686
687    pub fn plan(&self, request: TopologyRequest) -> Result<ParallelPlan, TopologyError> {
688        let pp = request.pipeline;
689        let tp = request.tensor;
690        if !(1..=PRODUCT_MAX_CARDS).contains(&pp) {
691            return Err(TopologyError::new(format!(
692                "PP size {pp} outside product range 1..={PRODUCT_MAX_CARDS}"
693            )));
694        }
695        if !(1..=PRODUCT_MAX_CARDS).contains(&tp) {
696            return Err(TopologyError::new(format!(
697                "TP size {tp} outside product range 1..={PRODUCT_MAX_CARDS}"
698            )));
699        }
700        let world = request.world_size()?;
701        if world > PRODUCT_MAX_CARDS {
702            return Err(TopologyError::new(format!(
703                "PP={pp} x TP={tp} requires {world} cards; product envelope is \
704                 {PRODUCT_MAX_CARDS}"
705            )));
706        }
707        if !self.hardware_targets.contains(&request.hardware) {
708            return Err(TopologyError::new(format!(
709                "{} has no qualified {} contract",
710                self.variant,
711                request.hardware.label()
712            )));
713        }
714        if world > request.hardware.max_cards() {
715            return Err(TopologyError::new(format!(
716                "{} target permits at most {} card(s), requested {world}",
717                request.hardware.label(),
718                request.hardware.max_cards()
719            )));
720        }
721        if request.available_devices < world {
722            return Err(TopologyError::new(format!(
723                "PP={pp} x TP={tp} requires {world} cards, only {} available",
724                request.available_devices
725            )));
726        }
727        if pp > self.trunk_layers {
728            return Err(TopologyError::new(format!(
729                "PP={pp} exceeds {} trunk layers",
730                self.trunk_layers
731            )));
732        }
733        if request.expert_parallel && tp == 1 {
734            return Err(TopologyError::new(
735                "expert parallelism requires TP group size greater than one",
736            ));
737        }
738        if request.expert_parallel && self.expert_count == 0 {
739            return Err(TopologyError::new(
740                "expert parallelism requested for a dense-only ModelPlan",
741            ));
742        }
743        if tp > 1 && !self.tensor_attention_supported {
744            return Err(TopologyError::new(format!(
745                "{} has attention operations without a generic TP shard contract; expert-only \
746                 EP may still be selected independently",
747                self.variant
748            )));
749        }
750
751        // Check the plan-derived per-layer attention geometry before generic dimensions so a
752        // refused topology names the operation that actually makes it invalid.
753        for (il, (&q, &kv)) in self.query_heads.iter().zip(&self.kv_heads).enumerate() {
754            require_divisible(&format!("layer {il} query heads"), q, tp)?;
755            require_divisible(&format!("layer {il} KV heads"), kv, tp)?;
756        }
757        require_divisible("hidden size", self.hidden_size, tp)?;
758        require_divisible("vocabulary size", self.vocab_size, tp)?;
759        if self.dense_ffn_size > 0 {
760            require_divisible("dense FFN size", self.dense_ffn_size, tp)?;
761        }
762        if self.expert_count > 0 {
763            if request.expert_parallel {
764                require_divisible("routed expert count", self.expert_count, tp)?;
765            } else {
766                require_divisible("routed expert FFN size", self.expert_ffn_size, tp)?;
767            }
768        }
769
770        let stage_ranges = (0..pp)
771            .map(|stage| stage * self.trunk_layers / pp..(stage + 1) * self.trunk_layers / pp)
772            .collect();
773
774        Ok(ParallelPlan {
775            contract: self.clone(),
776            request,
777            world_size: world,
778            stage_ranges,
779            mtp_owner_stage: self.mtp_layers.gt(&0).then_some(pp - 1),
780            // Shared experts remain replicated until a backend advertises a sharded shared branch.
781            shared_expert_replicated: tp > 1 && self.shared_expert_ffn_size > 0,
782        })
783    }
784
785    /// Validate every selected Step TP layer before the loader opens its first weight tensor.
786    ///
787    /// Physical device availability is checked separately because this pure plan is also the
788    /// topology oracle for tests and offline launch preparation.
789    pub(crate) fn preflight_step_tp_specs<'a>(
790        &self,
791        specs: impl IntoIterator<Item = (usize, &'a [usize])>,
792        layer_owners: &[usize],
793    ) -> Result<StepTpPreflightPlan, TopologyError> {
794        if layer_owners.len() != self.trunk_layers {
795            return Err(TopologyError::new(format!(
796                "Step TP owner map has {} layers, expected {}",
797                layer_owners.len(),
798                self.trunk_layers
799            )));
800        }
801
802        let mut seen = vec![false; self.trunk_layers];
803        let mut layers = Vec::new();
804        let mut runtime_groups: Vec<Vec<usize>> = Vec::new();
805        for (layer, devices) in specs {
806            if layer >= self.trunk_layers {
807                return Err(TopologyError::new(format!(
808                    "Step TP layer {layer} is outside trunk layers 0..{}",
809                    self.trunk_layers
810                )));
811            }
812            if seen[layer] {
813                return Err(TopologyError::new(format!(
814                    "Step TP preflight assigns layer {layer} more than once"
815                )));
816            }
817            if !(2..=PRODUCT_MAX_CARDS).contains(&devices.len()) {
818                return Err(TopologyError::new(format!(
819                    "Step TP layer {layer} requires 2..={PRODUCT_MAX_CARDS} devices, got {}",
820                    devices.len()
821                )));
822            }
823            let mut unique = devices.to_vec();
824            unique.sort_unstable();
825            unique.dedup();
826            if unique.len() != devices.len() {
827                return Err(TopologyError::new(format!(
828                    "Step TP layer {layer} devices must be distinct, got {devices:?}"
829                )));
830            }
831            let owner_device = layer_owners[layer];
832            if devices.first().copied() != Some(owner_device) {
833                return Err(TopologyError::new(format!(
834                    "Step TP layer {layer} owning PP device {owner_device} must be the first rank, \
835                     got {devices:?}"
836                )));
837            }
838
839            let expert_layout = if layer < self.dense_prefix_layers {
840                StepTpExpertLayout::AttentionOnly
841            } else if devices.len() > 2 {
842                StepTpExpertLayout::ExpertParallel
843            } else {
844                StepTpExpertLayout::TensorParallel
845            };
846            let plan = self.plan(TopologyRequest {
847                pipeline: 1,
848                tensor: devices.len(),
849                // Dense-prefix layers have no routed expert bank, but the full-model TP4/TP8
850                // contract still uses the EP geometry arm so the irrelevant 1280-wide expert
851                // projection is not falsely tensor-sharded during topology validation.
852                expert_parallel: devices.len() > 2,
853                available_devices: devices.len(),
854                hardware: HardwareTarget::RtxPro6000Blackwell,
855            })?;
856            for rank in 0..devices.len() {
857                let query = plan.query_head_range(layer, rank).ok_or_else(|| {
858                    TopologyError::new(format!(
859                        "Step TP layer {layer} has no query-head range for rank {rank}"
860                    ))
861                })?;
862                let kv = plan.kv_head_range(layer, rank).ok_or_else(|| {
863                    TopologyError::new(format!(
864                        "Step TP layer {layer} has no KV-head range for rank {rank}"
865                    ))
866                })?;
867                if query.is_empty() || kv.is_empty() {
868                    return Err(TopologyError::new(format!(
869                        "Step TP layer {layer} rank {rank} has an empty attention shard"
870                    )));
871                }
872            }
873
874            if !runtime_groups.iter().any(|group| group == devices) {
875                runtime_groups.push(devices.to_vec());
876            }
877            seen[layer] = true;
878            layers.push(StepTpLayerPlan {
879                layer,
880                devices: devices.to_vec(),
881                owner_device,
882                expert_layout,
883            });
884        }
885        layers.sort_unstable_by_key(|layer| layer.layer);
886
887        Ok(StepTpPreflightPlan {
888            layers,
889            runtime_groups,
890            full_trunk: seen.into_iter().all(|selected| selected),
891        })
892    }
893
894    /// Plan an explicit sequence of PP stages whose TP/EP widths may differ.
895    ///
896    /// This is the placement contract for arbitrary 1-8 card counts. It validates only the layers
897    /// assigned to each group, so a TP1 dense-prefix stage can coexist with a TP2/TP4/TP8 MoE
898    /// stage. Logical ranks are contiguous per stage; physical-device binding is a separate
899    /// runtime concern and must preserve each group's rank order.
900    pub fn plan_grouped(
901        &self,
902        request: GroupedTopologyRequest,
903    ) -> Result<GroupedParallelPlan, TopologyError> {
904        if request.stages.is_empty() {
905            return Err(TopologyError::new(
906                "grouped Step topology requires at least one stage",
907            ));
908        }
909        if request.stages.len() > self.trunk_layers {
910            return Err(TopologyError::new(format!(
911                "{} grouped stages exceed {} trunk layers",
912                request.stages.len(),
913                self.trunk_layers
914            )));
915        }
916        if !self.hardware_targets.contains(&request.hardware) {
917            return Err(TopologyError::new(format!(
918                "{} has no qualified {} contract",
919                self.variant,
920                request.hardware.label()
921            )));
922        }
923
924        let mut world_size = 0usize;
925        let mut expected_layer = 0usize;
926        let mut rank_groups = Vec::with_capacity(request.stages.len());
927        for (stage, group) in request.stages.iter().enumerate() {
928            if group.layers.start != expected_layer
929                || group.layers.start >= group.layers.end
930                || group.layers.end > self.trunk_layers
931            {
932                return Err(TopologyError::new(format!(
933                    "grouped stage {stage} layers {:?} do not continue the exact 0..{} trunk \
934                     partition at layer {expected_layer}",
935                    group.layers, self.trunk_layers
936                )));
937            }
938            if !(1..=PRODUCT_MAX_CARDS).contains(&group.tensor) {
939                return Err(TopologyError::new(format!(
940                    "grouped stage {stage} TP={} outside product range 1..={PRODUCT_MAX_CARDS}",
941                    group.tensor
942                )));
943            }
944            if group.expert_parallel && group.tensor == 1 {
945                return Err(TopologyError::new(format!(
946                    "grouped stage {stage} expert parallelism requires more than one rank"
947                )));
948            }
949
950            validate_group_geometry(self, stage, group)?;
951            let rank_start = world_size;
952            world_size = world_size
953                .checked_add(group.tensor)
954                .ok_or_else(|| TopologyError::new("grouped topology world size overflow"))?;
955            rank_groups.push(StageRankGroup {
956                stage,
957                layers: group.layers.clone(),
958                global_ranks: rank_start..world_size,
959                tensor: group.tensor,
960                expert_parallel: group.expert_parallel,
961                shared_expert_replicated: group.tensor > 1 && self.shared_expert_ffn_size > 0,
962            });
963            expected_layer = group.layers.end;
964        }
965        if expected_layer != self.trunk_layers {
966            return Err(TopologyError::new(format!(
967                "grouped Step topology ends at layer {expected_layer}, expected {}",
968                self.trunk_layers
969            )));
970        }
971        if world_size > PRODUCT_MAX_CARDS {
972            return Err(TopologyError::new(format!(
973                "grouped Step topology requires {world_size} cards; product envelope is \
974                 {PRODUCT_MAX_CARDS}"
975            )));
976        }
977        if world_size > request.hardware.max_cards() {
978            return Err(TopologyError::new(format!(
979                "{} target permits at most {} card(s), requested {world_size}",
980                request.hardware.label(),
981                request.hardware.max_cards()
982            )));
983        }
984        if request.available_devices < world_size {
985            return Err(TopologyError::new(format!(
986                "grouped Step topology requires {world_size} cards, only {} available",
987                request.available_devices
988            )));
989        }
990
991        let mtp_owner_stage = self.mtp_layers.gt(&0).then_some(expected_layer_stage(
992            self.trunk_layers - 1,
993            &request.stages,
994        )?);
995        Ok(GroupedParallelPlan {
996            contract: self.clone(),
997            request,
998            world_size,
999            rank_groups,
1000            mtp_owner_stage,
1001        })
1002    }
1003}
1004
1005#[derive(Debug, Clone, PartialEq, Eq)]
1006pub struct ParallelPlan {
1007    pub contract: ModelParallelContract,
1008    pub request: TopologyRequest,
1009    pub world_size: usize,
1010    pub stage_ranges: Vec<Range<usize>>,
1011    /// MTP layers are not pipeline stages of their own; the final PP stage owns them.
1012    pub mtp_owner_stage: Option<usize>,
1013    pub shared_expert_replicated: bool,
1014}
1015
1016impl ParallelPlan {
1017    pub fn global_rank(&self, pipeline_rank: usize, tensor_rank: usize) -> Option<usize> {
1018        if pipeline_rank >= self.request.pipeline || tensor_rank >= self.request.tensor {
1019            return None;
1020        }
1021        Some(pipeline_rank * self.request.tensor + tensor_rank)
1022    }
1023
1024    pub fn query_head_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
1025        split_range(
1026            *self.contract.query_heads.get(layer)?,
1027            self.request.tensor,
1028            tensor_rank,
1029        )
1030    }
1031
1032    pub fn kv_head_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
1033        split_range(
1034            *self.contract.kv_heads.get(layer)?,
1035            self.request.tensor,
1036            tensor_rank,
1037        )
1038    }
1039
1040    /// Column-parallel Q output range and the matching row-parallel O input range.
1041    pub fn query_feature_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
1042        let heads = self.query_head_range(layer, tensor_rank)?;
1043        Some(heads.start * self.contract.head_dim..heads.end * self.contract.head_dim)
1044    }
1045
1046    /// Column-parallel K/V output range. Step-3.7 has eight KV heads, so TP2 and TP4 partition
1047    /// them exactly; no KV-head replication is part of this registered contract.
1048    pub fn kv_feature_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
1049        let heads = self.kv_head_range(layer, tensor_rank)?;
1050        Some(heads.start * self.contract.head_dim..heads.end * self.contract.head_dim)
1051    }
1052
1053    /// Column-parallel dense gate/up output range and matching row-parallel down input range.
1054    pub fn dense_ffn_range(&self, tensor_rank: usize) -> Option<Range<usize>> {
1055        split_range(
1056            self.contract.dense_ffn_size,
1057            self.request.tensor,
1058            tensor_rank,
1059        )
1060    }
1061
1062    pub fn routed_expert_range(&self, tensor_rank: usize) -> Option<Range<usize>> {
1063        self.request
1064            .expert_parallel
1065            .then(|| split_range(self.contract.expert_count, self.request.tensor, tensor_rank))?
1066    }
1067
1068    pub fn routed_expert_ffn_range(&self, tensor_rank: usize) -> Option<Range<usize>> {
1069        (!self.request.expert_parallel).then(|| {
1070            split_range(
1071                self.contract.expert_ffn_size,
1072                self.request.tensor,
1073                tensor_rank,
1074            )
1075        })?
1076    }
1077}
1078
1079#[derive(Debug, Clone, PartialEq, Eq)]
1080pub struct StageRankGroup {
1081    pub stage: usize,
1082    pub layers: Range<usize>,
1083    pub global_ranks: Range<usize>,
1084    pub tensor: usize,
1085    pub expert_parallel: bool,
1086    pub shared_expert_replicated: bool,
1087}
1088
1089#[derive(Debug, Clone, PartialEq, Eq)]
1090pub struct GroupedParallelPlan {
1091    pub contract: ModelParallelContract,
1092    pub request: GroupedTopologyRequest,
1093    pub world_size: usize,
1094    pub rank_groups: Vec<StageRankGroup>,
1095    pub mtp_owner_stage: Option<usize>,
1096}
1097
1098impl GroupedParallelPlan {
1099    pub fn group_for_layer(&self, layer: usize) -> Option<&StageRankGroup> {
1100        self.rank_groups
1101            .iter()
1102            .find(|group| group.layers.contains(&layer))
1103    }
1104
1105    pub fn group_for_global_rank(&self, rank: usize) -> Option<&StageRankGroup> {
1106        self.rank_groups
1107            .iter()
1108            .find(|group| group.global_ranks.contains(&rank))
1109    }
1110
1111    pub fn global_rank(&self, stage: usize, tensor_rank: usize) -> Option<usize> {
1112        let group = self.rank_groups.get(stage)?;
1113        (tensor_rank < group.tensor).then_some(group.global_ranks.start + tensor_rank)
1114    }
1115
1116    pub fn query_head_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
1117        let group = self.group_for_layer(layer)?;
1118        split_range(
1119            *self.contract.query_heads.get(layer)?,
1120            group.tensor,
1121            tensor_rank,
1122        )
1123    }
1124
1125    pub fn kv_head_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
1126        let group = self.group_for_layer(layer)?;
1127        split_range(
1128            *self.contract.kv_heads.get(layer)?,
1129            group.tensor,
1130            tensor_rank,
1131        )
1132    }
1133
1134    pub fn routed_expert_range(&self, layer: usize, tensor_rank: usize) -> Option<Range<usize>> {
1135        let group = self.group_for_layer(layer)?;
1136        group
1137            .expert_parallel
1138            .then(|| split_range(self.contract.expert_count, group.tensor, tensor_rank))?
1139    }
1140}
1141
1142/// Validate the live Step PP request before the loader allocates CUDA state. Checkpoint tensor
1143/// census is deliberately a separate loader gate: topology legality must remain testable without
1144/// opening model files, while serving requires both gates.
1145pub fn validate_step_pp_request(cfg: &ModelConfig) -> Result<Option<ParallelPlan>, TopologyError> {
1146    let pp = match std::env::var("MEMRA_PP_STAGES") {
1147        Err(_) => return Ok(None),
1148        Ok(value) if value.is_empty() || value == "0" || value == "1" => return Ok(None),
1149        Ok(value) => value.parse::<usize>().map_err(|_| {
1150            TopologyError::new(format!("MEMRA_PP_STAGES={value} is not a positive integer"))
1151        })?,
1152    };
1153    let devices = selected_pp_devices(pp)?;
1154    let hardware = detect_uniform_hardware(&devices)?;
1155    let contract = ModelParallelContract::from_model(cfg)?;
1156    let trunk_layers = contract.trunk_layers;
1157    let plan = contract.plan(TopologyRequest {
1158        pipeline: pp,
1159        tensor: 1,
1160        expert_parallel: false,
1161        available_devices: devices.len(),
1162        hardware,
1163    })?;
1164    let fence = crate::pp::pp_cuts(trunk_layers).ok_or_else(|| {
1165        TopologyError::new(format!(
1166            "Step PP={pp} has no valid runtime stage fence over {trunk_layers} trunk layers"
1167        ))
1168    })?;
1169    let plan = apply_stage_fence(plan, &fence)?;
1170    Ok(Some(plan))
1171}
1172
1173/// Prove that every routed layer in the structural contract exposes native stacked block-128
1174/// E4M3 expert banks. Converted and per-tensor artifacts do not inherit this backend.
1175pub fn validate_fp8_expert_checkpoint(
1176    src: &dyn TensorSource,
1177    contract: &ModelParallelContract,
1178) -> Result<usize, TopologyError> {
1179    if src.st_dir().is_none() {
1180        return Err(TopologyError::new(
1181            "native E4M3 expert parallelism requires a safetensors checkpoint source; a \
1182             converted artifact cannot inherit this backend",
1183        ));
1184    }
1185
1186    let projections = [
1187        (
1188            "ffn_gate_exps",
1189            contract.hidden_size,
1190            contract.expert_ffn_size,
1191        ),
1192        (
1193            "ffn_up_exps",
1194            contract.hidden_size,
1195            contract.expert_ffn_size,
1196        ),
1197        (
1198            "ffn_down_exps",
1199            contract.expert_ffn_size,
1200            contract.hidden_size,
1201        ),
1202    ];
1203    let mut qualified = 0usize;
1204    for layer in contract.dense_prefix_layers..contract.trunk_layers {
1205        for &(projection, expected_in, expected_out) in &projections {
1206            let name = format!("blk.{layer}.{projection}.weight");
1207            let fp8 = src.find_fp8_stacked_native(&name).ok_or_else(|| {
1208                TopologyError::new(format!(
1209                    "{name} is not a checkpoint-faithful stacked block-128 E4M3 bank"
1210                ))
1211            })?;
1212            if fp8.n_expert != contract.expert_count {
1213                return Err(TopologyError::new(format!(
1214                    "{name} carries {} experts, expected {}",
1215                    fp8.n_expert, contract.expert_count
1216                )));
1217            }
1218            if fp8.in_f != expected_in || fp8.out_f != expected_out {
1219                return Err(TopologyError::new(format!(
1220                    "{name} expert shape {}x{} != expected {expected_out}x{expected_in}",
1221                    fp8.out_f, fp8.in_f
1222                )));
1223            }
1224            let expected_rows = expected_out.div_ceil(STEP_FP8_BLOCK);
1225            let expected_cols = expected_in.div_ceil(STEP_FP8_BLOCK);
1226            let expected_scales = contract.expert_count * expected_rows * expected_cols;
1227            if fp8.scale_rows != expected_rows
1228                || fp8.scale_cols != expected_cols
1229                || fp8.scales.len() != expected_scales
1230            {
1231                return Err(TopologyError::new(format!(
1232                    "{name} block-128 E4M3 grid {}x{} ({} scales) != expected {} experts x \
1233                     {expected_rows}x{expected_cols} ({expected_scales} scales)",
1234                    fp8.scale_rows,
1235                    fp8.scale_cols,
1236                    fp8.scales.len(),
1237                    contract.expert_count
1238                )));
1239            }
1240            qualified += fp8.n_expert;
1241        }
1242    }
1243
1244    let expected = (contract.trunk_layers - contract.dense_prefix_layers)
1245        * contract.expert_count
1246        * projections.len();
1247    if qualified != expected {
1248        return Err(TopologyError::new(format!(
1249            "E4M3 expert tensor census qualified {qualified}, expected {expected}"
1250        )));
1251    }
1252    Ok(qualified)
1253}
1254
1255/// Prove that every routed layer in the structural contract exposes native ModelOpt NVFP4
1256/// experts: either one stacked bank or the Hugging Face per-expert layout. Both carry packed
1257/// e2m1 codes, per-16 UE4M3 scales, and finite-positive per-expert macros.
1258pub fn validate_nvfp4_expert_checkpoint(
1259    src: &dyn TensorSource,
1260    contract: &ModelParallelContract,
1261) -> Result<usize, TopologyError> {
1262    if src.st_dir().is_none() {
1263        return Err(TopologyError::new(
1264            "native NVFP4 expert parallelism requires a safetensors checkpoint source; a \
1265             converted artifact cannot inherit this backend",
1266        ));
1267    }
1268
1269    let projections = [
1270        (
1271            "ffn_gate_exps",
1272            contract.hidden_size,
1273            contract.expert_ffn_size,
1274        ),
1275        (
1276            "ffn_up_exps",
1277            contract.hidden_size,
1278            contract.expert_ffn_size,
1279        ),
1280        (
1281            "ffn_down_exps",
1282            contract.expert_ffn_size,
1283            contract.hidden_size,
1284        ),
1285    ];
1286    let mut qualified = 0usize;
1287    for layer in contract.dense_prefix_layers..contract.trunk_layers {
1288        for &(projection, expected_in, expected_out) in &projections {
1289            let name = format!("blk.{layer}.{projection}.weight");
1290            if let Some(bank) = src.find_nvfp4_stacked_native(&name) {
1291                if bank.n_expert != contract.expert_count {
1292                    return Err(TopologyError::new(format!(
1293                        "{name} carries {} experts, expected {}",
1294                        bank.n_expert, contract.expert_count
1295                    )));
1296                }
1297                if bank.in_f != expected_in || bank.out_f != expected_out {
1298                    return Err(TopologyError::new(format!(
1299                        "{name} expert shape {}x{} != expected {expected_out}x{expected_in}",
1300                        bank.out_f, bank.in_f
1301                    )));
1302                }
1303                if bank.in_f % 64 != 0 {
1304                    return Err(TopologyError::new(format!(
1305                        "{name} in_features {} is not 64-aligned; memra block_nvfp4 kernels \
1306                         require whole 64-element superblocks",
1307                        bank.in_f
1308                    )));
1309                }
1310                if bank.macros.len() != contract.expert_count {
1311                    return Err(TopologyError::new(format!(
1312                        "{name} carries {} weight_scale_2 macros, expected {}",
1313                        bank.macros.len(),
1314                        contract.expert_count
1315                    )));
1316                }
1317                qualified += bank.n_expert;
1318                continue;
1319            }
1320
1321            for expert in 0..contract.expert_count {
1322                let expert_name = format!("blk.{layer}.{projection}.{expert}.weight");
1323                let tensor = src.find_nvfp4_native(&expert_name).ok_or_else(|| {
1324                    TopologyError::new(format!(
1325                        "{name} is neither a checkpoint-faithful stacked modelopt NVFP4 bank nor \
1326                         a complete per-expert NVFP4 set; missing {expert_name}"
1327                    ))
1328                })?;
1329                if tensor.in_f != expected_in || tensor.out_f != expected_out {
1330                    return Err(TopologyError::new(format!(
1331                        "{expert_name} shape {}x{} != expected {expected_out}x{expected_in}",
1332                        tensor.out_f, tensor.in_f
1333                    )));
1334                }
1335                if tensor.in_f % 64 != 0 {
1336                    return Err(TopologyError::new(format!(
1337                        "{expert_name} in_features {} is not 64-aligned; memra block_nvfp4 \
1338                         kernels require whole 64-element superblocks",
1339                        tensor.in_f
1340                    )));
1341                }
1342                qualified += 1;
1343            }
1344        }
1345    }
1346
1347    let expected = (contract.trunk_layers - contract.dense_prefix_layers)
1348        * contract.expert_count
1349        * projections.len();
1350    if qualified != expected {
1351        return Err(TopologyError::new(format!(
1352            "NVFP4 expert tensor census qualified {qualified}, expected {expected}"
1353        )));
1354    }
1355    Ok(qualified)
1356}
1357
1358/// Legacy API retained for existing focused gates.
1359pub fn validate_step_fp8_checkpoint(
1360    src: &dyn TensorSource,
1361    contract: &ModelParallelContract,
1362) -> Result<usize, TopologyError> {
1363    validate_fp8_expert_checkpoint(src, contract)
1364}
1365
1366/// Legacy API retained for existing focused gates.
1367pub fn validate_step_nvfp4_checkpoint(
1368    src: &dyn TensorSource,
1369    contract: &ModelParallelContract,
1370) -> Result<usize, TopologyError> {
1371    validate_nvfp4_expert_checkpoint(src, contract)
1372}
1373
1374fn apply_stage_fence(
1375    mut plan: ParallelPlan,
1376    fence: &[usize],
1377) -> Result<ParallelPlan, TopologyError> {
1378    let expected = plan.request.pipeline + 1;
1379    if fence.len() != expected
1380        || fence.first() != Some(&0)
1381        || fence.last() != Some(&plan.contract.trunk_layers)
1382        || fence.windows(2).any(|window| window[0] >= window[1])
1383        || fence[1..fence.len() - 1]
1384            .iter()
1385            .any(|boundary| !plan.contract.partition_boundaries.contains(boundary))
1386    {
1387        return Err(TopologyError::new(format!(
1388            "invalid PP fence {fence:?} for {} stages over {} trunk layers",
1389            plan.request.pipeline, plan.contract.trunk_layers
1390        )));
1391    }
1392    plan.stage_ranges = fence
1393        .windows(2)
1394        .map(|window| window[0]..window[1])
1395        .collect();
1396    Ok(plan)
1397}
1398
1399fn selected_pp_devices(pp: usize) -> Result<Vec<usize>, TopologyError> {
1400    let raw = std::env::var("MEMRA_PP_DEVICES").map_err(|_| {
1401        TopologyError::new(format!(
1402            "Step PP={pp} requires explicit MEMRA_PP_DEVICES with one distinct CUDA ordinal per \
1403             stage; same-device diagnostics do not qualify the multi-card product"
1404        ))
1405    })?;
1406    let devices: Result<Vec<usize>, _> = raw
1407        .split(',')
1408        .map(|part| part.trim().parse::<usize>())
1409        .collect();
1410    let devices = devices.map_err(|_| {
1411        TopologyError::new(format!(
1412            "MEMRA_PP_DEVICES={raw:?} is not a comma-separated CUDA ordinal list"
1413        ))
1414    })?;
1415    if devices.len() != pp {
1416        return Err(TopologyError::new(format!(
1417            "MEMRA_PP_DEVICES lists {} devices but MEMRA_PP_STAGES={pp}",
1418            devices.len()
1419        )));
1420    }
1421    let mut unique = devices.clone();
1422    unique.sort_unstable();
1423    unique.dedup();
1424    if unique.len() != devices.len() {
1425        return Err(TopologyError::new(format!(
1426            "Step PP={pp} requires {pp} distinct devices; MEMRA_PP_DEVICES={raw:?} repeats an \
1427             ordinal"
1428        )));
1429    }
1430    Ok(devices)
1431}
1432
1433pub(crate) fn detect_uniform_hardware(devices: &[usize]) -> Result<HardwareTarget, TopologyError> {
1434    cudarc::driver::result::init().map_err(|error| {
1435        TopologyError::new(format!("CUDA driver initialization failed: {error}"))
1436    })?;
1437    let mut target = None;
1438    for &ordinal in devices {
1439        let device = cudarc::driver::result::device::get(ordinal as i32).map_err(|error| {
1440            TopologyError::new(format!("CUDA device {ordinal} lookup failed: {error}"))
1441        })?;
1442        let name = cudarc::driver::result::device::get_name(device).map_err(|error| {
1443            TopologyError::new(format!("CUDA device {ordinal} name lookup failed: {error}"))
1444        })?;
1445        let current = HardwareTarget::from_device_name(&name)?;
1446        if let Some(expected) = target {
1447            if current != expected {
1448                return Err(TopologyError::new(format!(
1449                    "mixed hardware targets in MEMRA_PP_DEVICES: expected {}, device {ordinal} is \
1450                     {}",
1451                    expected.label(),
1452                    current.label()
1453                )));
1454            }
1455        } else {
1456            target = Some(current);
1457        }
1458    }
1459    target.ok_or_else(|| TopologyError::new("MEMRA_PP_DEVICES is empty"))
1460}
1461
1462#[allow(clippy::manual_is_multiple_of)] // allow: divisor is runtime-derived; the modulo form keeps a zero divisor loud (a panic), where is_multiple_of would return false silently
1463fn require_divisible(label: &str, value: usize, parts: usize) -> Result<(), TopologyError> {
1464    if value == 0 {
1465        return Err(TopologyError::new(format!("{label} is zero")));
1466    }
1467    if value % parts != 0 {
1468        return Err(TopologyError::new(format!(
1469            "{label} {value} is not divisible by TP={parts}"
1470        )));
1471    }
1472    Ok(())
1473}
1474
1475fn require_fp8_block_shard(label: &str, value: usize, parts: usize) -> Result<(), TopologyError> {
1476    require_divisible(label, value, parts)?;
1477    let local = value / parts;
1478    if !local.is_multiple_of(STEP_FP8_BLOCK) {
1479        return Err(TopologyError::new(format!(
1480            "{label} shard {local} for TP={parts} cuts through the Step E4M3 block size \
1481             {STEP_FP8_BLOCK}"
1482        )));
1483    }
1484    Ok(())
1485}
1486
1487fn validate_group_geometry(
1488    contract: &ModelParallelContract,
1489    stage: usize,
1490    group: &StageGroupRequest,
1491) -> Result<(), TopologyError> {
1492    let tp = group.tensor;
1493    for layer in group.layers.clone() {
1494        require_divisible(
1495            &format!("stage {stage} layer {layer} query heads"),
1496            contract.query_heads[layer],
1497            tp,
1498        )?;
1499        require_divisible(
1500            &format!("stage {stage} layer {layer} KV heads"),
1501            contract.kv_heads[layer],
1502            tp,
1503        )?;
1504    }
1505    require_divisible(
1506        &format!("stage {stage} hidden size"),
1507        contract.hidden_size,
1508        tp,
1509    )?;
1510    if group.layers.start < contract.dense_prefix_layers {
1511        require_fp8_block_shard(
1512            &format!("stage {stage} dense FFN size"),
1513            contract.dense_ffn_size,
1514            tp,
1515        )?;
1516    }
1517    if group.layers.end > contract.dense_prefix_layers {
1518        if group.expert_parallel {
1519            require_divisible(
1520                &format!("stage {stage} routed expert count"),
1521                contract.expert_count,
1522                tp,
1523            )?;
1524        } else {
1525            require_fp8_block_shard(
1526                &format!("stage {stage} routed expert FFN size"),
1527                contract.expert_ffn_size,
1528                tp,
1529            )?;
1530        }
1531    }
1532    if group.layers.end == contract.trunk_layers {
1533        require_divisible(
1534            &format!("stage {stage} vocabulary size"),
1535            contract.vocab_size,
1536            tp,
1537        )?;
1538    }
1539    Ok(())
1540}
1541
1542fn expected_layer_stage(
1543    layer: usize,
1544    stages: &[StageGroupRequest],
1545) -> Result<usize, TopologyError> {
1546    stages
1547        .iter()
1548        .position(|stage| stage.layers.contains(&layer))
1549        .ok_or_else(|| TopologyError::new(format!("no grouped stage owns layer {layer}")))
1550}
1551
1552#[allow(clippy::manual_is_multiple_of)] // allow: divisor is runtime-derived; the modulo form keeps a zero divisor loud (a panic), where is_multiple_of would return false silently
1553fn split_range(total: usize, parts: usize, rank: usize) -> Option<Range<usize>> {
1554    if parts == 0 || rank >= parts || total % parts != 0 {
1555        return None;
1556    }
1557    let width = total / parts;
1558    Some(rank * width..(rank + 1) * width)
1559}
1560
1561#[derive(Debug, Clone, PartialEq, Eq)]
1562pub struct TopologyError {
1563    message: String,
1564}
1565
1566impl TopologyError {
1567    fn new(message: impl Into<String>) -> Self {
1568        Self {
1569            message: message.into(),
1570        }
1571    }
1572}
1573
1574impl fmt::Display for TopologyError {
1575    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1576        self.message.fmt(f)
1577    }
1578}
1579
1580impl std::error::Error for TopologyError {}
1581
1582#[cfg(test)]
1583mod tests {
1584    use super::*;
1585    use memra_gguf::config::{Arch, HfConfig, MoeConfig, Step35Config};
1586    use memra_gguf::source::{Fp8StackedNative, TensorView};
1587    use std::path::Path;
1588
1589    fn step37_contract() -> ModelParallelContract {
1590        let total_layers = 48;
1591        ModelParallelContract {
1592            family: "sliding-gated-moe",
1593            variant: "Step-3.7-Flash-FP8".to_string(),
1594            trunk_layers: 45,
1595            mtp_layers: 3,
1596            hidden_size: 4096,
1597            vocab_size: 128_896,
1598            dense_ffn_size: 11_264,
1599            dense_prefix_layers: 3,
1600            head_dim: 128,
1601            query_heads: (0..total_layers)
1602                .map(|il| if il % 4 == 0 { 64 } else { 96 })
1603                .collect(),
1604            kv_heads: vec![8; total_layers],
1605            tensor_attention_supported: true,
1606            expert_count: 288,
1607            experts_per_token: 8,
1608            expert_ffn_size: 1280,
1609            shared_expert_ffn_size: 1280,
1610            routed_layers: (3..45).collect(),
1611            partition_boundaries: (1..45).collect(),
1612            hardware_targets: vec![HardwareTarget::RtxPro6000Blackwell],
1613        }
1614    }
1615
1616    fn step37_model_config() -> ModelConfig {
1617        let total_layers = 48;
1618        let head_count: Vec<u32> = (0..total_layers)
1619            .map(|il| if il % 4 == 0 { 64 } else { 96 })
1620            .collect();
1621        ModelConfig {
1622            arch: Arch::Step35,
1623            name: "Step-3.7-Flash-FP8".to_string(),
1624            n_layer: total_layers,
1625            n_embd: 4096,
1626            n_head: 96,
1627            n_head_kv: 8,
1628            head_dim_k: 128,
1629            head_dim_v: 128,
1630            n_ff: 11_264,
1631            n_vocab: 128_896,
1632            context_length: 262_144,
1633            rms_eps: 1e-6,
1634            rope_freq_base: 5_000_000.0,
1635            rope_dim_count: 128,
1636            rope_sections: Vec::new(),
1637            full_attention_interval: 0,
1638            ssm: None,
1639            moe: Some(MoeConfig {
1640                expert_count: 288,
1641                expert_used_count: 8,
1642                expert_ff_length: 1280,
1643                expert_shared_ff_length: 1280,
1644            }),
1645            m3: None,
1646            hy3: None,
1647            gemma4: None,
1648            vision: None,
1649            vision_glm5: None,
1650            multimodal: None,
1651            mla: None,
1652            dsv4: None,
1653            qwen4exp: None,
1654            rope_yarn: None,
1655            glm5: None,
1656            step35: Some(Step35Config {
1657                head_count,
1658                head_count_kv: vec![8; total_layers as usize],
1659                swa_pattern: (0..total_layers).map(|il| il % 4 != 0).collect(),
1660                sliding_window: 512,
1661                rope_base_global: 5_000_000.0,
1662                rope_base_swa: 10_000.0,
1663                rope_dims_full: 64,
1664                rope_dims_swa: 128,
1665                rope_freq_factors: None,
1666                swiglu_clamp_exp: vec![0.0; total_layers as usize],
1667                swiglu_clamp_shexp: vec![0.0; total_layers as usize],
1668                sigmoid_routing: true,
1669                routed_scaling_factor: 3.0,
1670                route_norm: true,
1671                first_k_dense_replace: 3,
1672            }),
1673            geometry: None,
1674            nextn_predict_layers: 3,
1675            n_layer_total: total_layers,
1676        }
1677    }
1678
1679    fn hy3_model_config() -> ModelConfig {
1680        ModelConfig::from_hf(&HfConfig::parse(
1681            r#"{
1682                "model_type":"hy_v3",
1683                "num_hidden_layers":80,
1684                "num_nextn_predict_layers":1,
1685                "hidden_size":4096,
1686                "num_attention_heads":64,
1687                "num_key_value_heads":8,
1688                "head_dim":128,
1689                "intermediate_size":13312,
1690                "vocab_size":120832,
1691                "max_position_embeddings":262144,
1692                "first_k_dense_replace":1,
1693                "num_experts":192,
1694                "num_experts_per_tok":8,
1695                "moe_intermediate_size":1536,
1696                "num_shared_experts":1,
1697                "moe_router_use_sigmoid":true,
1698                "moe_router_enable_expert_bias":true,
1699                "route_norm":true,
1700                "router_scaling_factor":2.826,
1701                "qk_norm":true
1702            }"#,
1703        ))
1704    }
1705
1706    fn dense_model_config() -> ModelConfig {
1707        ModelConfig::from_hf(&HfConfig::parse(
1708            r#"{
1709                "model_type":"qwen3",
1710                "num_hidden_layers":4,
1711                "hidden_size":4096,
1712                "num_attention_heads":32,
1713                "num_key_value_heads":8,
1714                "head_dim":128,
1715                "intermediate_size":12288,
1716                "vocab_size":131072,
1717                "max_position_embeddings":32768
1718            }"#,
1719        ))
1720    }
1721
1722    fn synthetic_auto_contract(routed: bool) -> ModelParallelContract {
1723        ModelParallelContract {
1724            family: if routed {
1725                "routed-moe"
1726            } else {
1727                "dense-transformer"
1728            },
1729            variant: "synthetic-auto".to_string(),
1730            trunk_layers: 4,
1731            mtp_layers: 0,
1732            hidden_size: 64,
1733            vocab_size: 128,
1734            dense_ffn_size: 128,
1735            dense_prefix_layers: if routed { 1 } else { 4 },
1736            head_dim: 32,
1737            query_heads: vec![2; 4],
1738            kv_heads: vec![1; 4],
1739            tensor_attention_supported: true,
1740            expert_count: if routed { 8 } else { 0 },
1741            experts_per_token: if routed { 2 } else { 0 },
1742            expert_ffn_size: if routed { 32 } else { 0 },
1743            shared_expert_ffn_size: 0,
1744            routed_layers: if routed { vec![1, 2, 3] } else { Vec::new() },
1745            partition_boundaries: vec![1, 2, 3],
1746            hardware_targets: vec![HardwareTarget::RtxPro6000Blackwell],
1747        }
1748    }
1749
1750    fn synthetic_auto_costs(trunk_expert_bytes: u64) -> AutoArtifactCosts {
1751        AutoArtifactCosts {
1752            layers: vec![
1753                LayerPlacementCost {
1754                    weight_bytes: 25,
1755                    kv_bytes_per_token: 0,
1756                };
1757                4
1758            ],
1759            first_fixed_bytes: 0,
1760            last_fixed_bytes: 0,
1761            trunk_expert_bytes,
1762            non_distributed_bytes: 20,
1763        }
1764    }
1765
1766    fn request(pp: usize, tp: usize, expert_parallel: bool) -> TopologyRequest {
1767        TopologyRequest {
1768            pipeline: pp,
1769            tensor: tp,
1770            expert_parallel,
1771            available_devices: pp * tp,
1772            hardware: HardwareTarget::RtxPro6000Blackwell,
1773        }
1774    }
1775
1776    struct MockStepFp8Source {
1777        safetensors: bool,
1778        block_scales: bool,
1779    }
1780
1781    impl TensorSource for MockStepFp8Source {
1782        fn config(&self) -> ModelConfig {
1783            step37_model_config()
1784        }
1785
1786        fn find(&self, _ggml_name: &str) -> Option<TensorView<'_>> {
1787            None
1788        }
1789
1790        fn st_dir(&self) -> Option<&Path> {
1791            self.safetensors.then(|| Path::new("/mock-step-fp8"))
1792        }
1793
1794        fn find_fp8_stacked_native(&self, name: &str) -> Option<Fp8StackedNative<'_>> {
1795            let (in_f, out_f): (usize, usize) = if name.contains("ffn_down_exps") {
1796                (1280, 4096)
1797            } else if name.contains("ffn_gate_exps") || name.contains("ffn_up_exps") {
1798                (4096, 1280)
1799            } else {
1800                return None;
1801            };
1802            let (scale_rows, scale_cols) = if self.block_scales {
1803                (
1804                    out_f.div_ceil(STEP_FP8_BLOCK),
1805                    in_f.div_ceil(STEP_FP8_BLOCK),
1806                )
1807            } else {
1808                (1, 1)
1809            };
1810            Some(Fp8StackedNative {
1811                bytes: &[],
1812                scales: vec![1.0; 288 * scale_rows * scale_cols],
1813                n_expert: 288,
1814                out_f,
1815                in_f,
1816                scale_rows,
1817                scale_cols,
1818            })
1819        }
1820    }
1821
1822    #[test]
1823    fn step_fp8_checkpoint_census_covers_every_routed_projection() {
1824        let source = MockStepFp8Source {
1825            safetensors: true,
1826            block_scales: true,
1827        };
1828        let qualified =
1829            validate_step_fp8_checkpoint(&source, &step37_contract()).expect("valid FP8 source");
1830        assert_eq!(qualified, 42 * 288 * 3);
1831    }
1832
1833    #[test]
1834    fn step_fp8_checkpoint_census_refuses_conversion_and_wrong_scale_class() {
1835        let converted = MockStepFp8Source {
1836            safetensors: false,
1837            block_scales: true,
1838        };
1839        assert!(
1840            validate_step_fp8_checkpoint(&converted, &step37_contract())
1841                .unwrap_err()
1842                .to_string()
1843                .contains("safetensors checkpoint source")
1844        );
1845
1846        let per_tensor = MockStepFp8Source {
1847            safetensors: true,
1848            block_scales: false,
1849        };
1850        assert!(
1851            validate_step_fp8_checkpoint(&per_tensor, &step37_contract())
1852                .unwrap_err()
1853                .to_string()
1854                .contains("block-128 E4M3")
1855        );
1856    }
1857
1858    #[test]
1859    fn step_pp3_maps_fifteen_trunk_layers_per_card() {
1860        let plan = step37_contract().plan(request(3, 1, false)).unwrap();
1861        assert_eq!(plan.world_size, 3);
1862        assert_eq!(plan.stage_ranges, vec![0..15, 15..30, 30..45]);
1863        assert_eq!(plan.mtp_owner_stage, Some(2));
1864    }
1865
1866    #[test]
1867    fn step_pp_marker_uses_the_runtime_stage_fence() {
1868        let plan = step37_contract().plan(request(3, 1, false)).unwrap();
1869        let plan = apply_stage_fence(plan, &[0, 10, 28, 45]).unwrap();
1870        assert_eq!(plan.stage_ranges, vec![0..10, 10..28, 28..45]);
1871    }
1872
1873    #[test]
1874    fn stage_fence_must_use_model_plan_partition_boundaries() {
1875        let mut contract = step37_contract();
1876        contract
1877            .partition_boundaries
1878            .retain(|&boundary| boundary != 10);
1879        let plan = contract.plan(request(3, 1, false)).unwrap();
1880        let error = apply_stage_fence(plan, &[0, 10, 28, 45]).unwrap_err();
1881        assert!(error.to_string().contains("invalid PP fence"));
1882    }
1883
1884    #[test]
1885    fn step_contract_is_extracted_from_model_specific_geometry() {
1886        let contract = ModelParallelContract::from_model(&step37_model_config()).unwrap();
1887        assert_eq!(contract.family, "routed-moe");
1888        assert_eq!(contract.trunk_layers, 45);
1889        assert_eq!(contract.mtp_layers, 3);
1890        assert_eq!(contract.query_heads[0], 64);
1891        assert_eq!(contract.query_heads[1], 96);
1892        assert_eq!(contract.kv_heads[47], 8);
1893        assert_eq!(contract.expert_count, 288);
1894        assert_eq!(contract.experts_per_token, 8);
1895    }
1896
1897    #[test]
1898    fn hy3_contract_is_extracted_from_exact_full_sigmoid_moe_geometry() {
1899        let contract = ModelParallelContract::from_model(&hy3_model_config()).unwrap();
1900        assert_eq!(contract.family, "routed-moe");
1901        assert_eq!(contract.trunk_layers, 80);
1902        assert_eq!(contract.mtp_layers, 1);
1903        assert_eq!(contract.query_heads, vec![64; 81]);
1904        assert_eq!(contract.kv_heads, vec![8; 81]);
1905        assert_eq!(contract.dense_prefix_layers, 1);
1906        assert_eq!(contract.expert_count, 192);
1907        assert_eq!(contract.experts_per_token, 8);
1908        assert_eq!(contract.expert_ffn_size, 1536);
1909        assert_eq!(contract.routed_layers, (1..80).collect::<Vec<_>>());
1910    }
1911
1912    #[test]
1913    fn hy3_sibling_geometry_is_derived_without_a_family_loader() {
1914        let mut sibling = hy3_model_config();
1915        sibling.n_vocab += 1;
1916        let contract = ModelParallelContract::from_model(&sibling).unwrap();
1917        assert_eq!(contract.vocab_size, 120_833);
1918        assert_eq!(contract.routed_layers.len(), 79);
1919    }
1920
1921    #[test]
1922    fn dense_transformer_contract_and_tp_geometry_are_plan_derived() {
1923        let contract = ModelParallelContract::from_model(&dense_model_config()).unwrap();
1924        assert_eq!(contract.family, "dense-transformer");
1925        assert!(contract.routed_layers.is_empty());
1926        assert_eq!(contract.dense_prefix_layers, 4);
1927        assert_eq!(contract.dense_ffn_size, 12_288);
1928        let tp4 = contract.plan(request(1, 4, false)).unwrap();
1929        assert_eq!(tp4.query_feature_range(0, 3), Some(3072..4096));
1930        assert_eq!(tp4.dense_ffn_range(3), Some(9216..12_288));
1931    }
1932
1933    #[test]
1934    fn automatic_placement_uses_capacity_not_family_recipes() {
1935        let routed = synthetic_auto_contract(true);
1936        let costs = synthetic_auto_costs(80);
1937
1938        let pp2 = choose_auto_parallel_placement(
1939            &costs,
1940            &routed,
1941            ExpertActivationPrecision::Bf16,
1942            &[0, 1],
1943            &[60, 60],
1944            6,
1945        )
1946        .unwrap();
1947        assert_eq!(pp2.backend, AutoParallelBackend::Pipeline);
1948        assert_eq!(pp2.pipeline_splits, vec![2]);
1949        assert_eq!(pp2.checkpoint_peak_bytes, 50);
1950
1951        let ep3 = choose_auto_parallel_placement(
1952            &costs,
1953            &routed,
1954            ExpertActivationPrecision::Bf16,
1955            &[0, 1, 2],
1956            &[60, 60, 60],
1957            6,
1958        )
1959        .unwrap();
1960        assert_eq!(ep3.backend, AutoParallelBackend::ExpertParallel);
1961        assert_eq!(ep3.expert_root_bytes, 50);
1962        assert_eq!(ep3.expert_peer_bytes, 30);
1963
1964        let ep4 = choose_auto_parallel_placement(
1965            &costs,
1966            &routed,
1967            ExpertActivationPrecision::Bf16,
1968            &[0, 1, 2, 3],
1969            &[60; 4],
1970            6,
1971        )
1972        .unwrap();
1973        assert_eq!(ep4.backend, AutoParallelBackend::ExpertParallel);
1974        assert_eq!(ep4.expert_root_bytes, 40);
1975        assert_eq!(ep4.expert_peer_bytes, 20);
1976    }
1977
1978    #[test]
1979    fn automatic_placement_routes_dense_and_non_w4a16_plans_to_pipeline() {
1980        let costs = synthetic_auto_costs(80);
1981        let dense = choose_auto_parallel_placement(
1982            &costs,
1983            &synthetic_auto_contract(false),
1984            ExpertActivationPrecision::Bf16,
1985            &[0, 1, 2, 3],
1986            &[60; 4],
1987            6,
1988        )
1989        .unwrap();
1990        assert_eq!(dense.backend, AutoParallelBackend::Pipeline);
1991
1992        let activation_quantized = choose_auto_parallel_placement(
1993            &costs,
1994            &synthetic_auto_contract(true),
1995            ExpertActivationPrecision::Quantized,
1996            &[0, 1, 2, 3],
1997            &[60; 4],
1998            6,
1999        )
2000        .unwrap();
2001        assert_eq!(activation_quantized.backend, AutoParallelBackend::Pipeline);
2002    }
2003
2004    #[test]
2005    fn automatic_placement_refuses_when_no_program_preserves_reserve() {
2006        let error = choose_auto_parallel_placement(
2007            &synthetic_auto_costs(80),
2008            &synthetic_auto_contract(true),
2009            ExpertActivationPrecision::Bf16,
2010            &[0, 1],
2011            &[55, 55],
2012            6,
2013        )
2014        .unwrap_err();
2015        assert!(error.to_string().contains("no capacity-safe program"));
2016    }
2017
2018    #[test]
2019    fn step_sibling_geometry_is_derived_without_a_family_loader() {
2020        let mut sibling = step37_model_config();
2021        sibling.name = "Step-3.5-Flash".to_string();
2022        sibling.n_vocab = 128_000;
2023        let contract = ModelParallelContract::from_model(&sibling).unwrap();
2024        assert_eq!(contract.variant, "Step-3.5-Flash");
2025        assert_eq!(contract.vocab_size, 128_000);
2026    }
2027
2028    #[test]
2029    fn step_without_mtp_keeps_the_same_structural_parallel_contract() {
2030        let mut stripped = step37_model_config();
2031        stripped.nextn_predict_layers = 0;
2032        let contract = ModelParallelContract::from_model(&stripped).unwrap();
2033        assert_eq!(contract.mtp_layers, 0);
2034        assert_eq!(contract.routed_layers, (3..48).collect::<Vec<_>>());
2035    }
2036
2037    #[test]
2038    fn hardware_target_classification_is_exact() {
2039        assert_eq!(
2040            HardwareTarget::from_device_name("NVIDIA RTX PRO 6000 Blackwell Server Edition")
2041                .unwrap(),
2042            HardwareTarget::RtxPro6000Blackwell
2043        );
2044        assert_eq!(
2045            HardwareTarget::from_device_name("NVIDIA GeForce RTX 5090 Laptop GPU").unwrap(),
2046            HardwareTarget::Rtx5090
2047        );
2048        assert!(HardwareTarget::from_device_name("NVIDIA H100 80GB HBM3").is_err());
2049    }
2050
2051    #[test]
2052    fn step_tp2_tp4_tp8_and_hybrid_plans_are_geometry_valid() {
2053        let tp2 = step37_contract().plan(request(1, 2, true)).unwrap();
2054        assert_eq!(tp2.query_head_range(0, 1), Some(32..64));
2055        assert_eq!(tp2.query_head_range(1, 1), Some(48..96));
2056        assert_eq!(tp2.kv_head_range(0, 1), Some(4..8));
2057        assert_eq!(tp2.routed_expert_range(1), Some(144..288));
2058
2059        let tp4 = step37_contract().plan(request(1, 4, true)).unwrap();
2060        assert_eq!(tp4.query_head_range(0, 3), Some(48..64));
2061        assert_eq!(tp4.query_head_range(1, 3), Some(72..96));
2062        assert_eq!(tp4.kv_head_range(0, 3), Some(6..8));
2063        assert_eq!(tp4.query_feature_range(0, 3), Some(6144..8192));
2064        assert_eq!(tp4.query_feature_range(1, 3), Some(9216..12_288));
2065        assert_eq!(tp4.kv_feature_range(0, 3), Some(768..1024));
2066        assert_eq!(tp4.dense_ffn_range(3), Some(8448..11_264));
2067        assert_eq!(tp4.routed_expert_range(3), Some(216..288));
2068        assert!(tp4.shared_expert_replicated);
2069
2070        let tp8 = step37_contract().plan(request(1, 8, true)).unwrap();
2071        assert_eq!(tp8.query_head_range(0, 7), Some(56..64));
2072        assert_eq!(tp8.query_head_range(1, 7), Some(84..96));
2073        assert_eq!(tp8.kv_head_range(0, 7), Some(7..8));
2074        assert_eq!(tp8.dense_ffn_range(7), Some(9856..11_264));
2075        assert_eq!(tp8.routed_expert_range(7), Some(252..288));
2076        assert!(tp8.shared_expert_replicated);
2077
2078        let hybrid = step37_contract().plan(request(2, 4, true)).unwrap();
2079        assert_eq!(hybrid.world_size, 8);
2080        assert_eq!(hybrid.stage_ranges, vec![0..22, 22..45]);
2081        assert_eq!(hybrid.global_rank(1, 3), Some(7));
2082        assert_eq!(hybrid.global_rank(2, 0), None);
2083    }
2084
2085    #[test]
2086    fn grouped_three_card_plan_is_pp1_then_tp2_ep2() {
2087        let plan = step37_contract()
2088            .plan_grouped(GroupedTopologyRequest {
2089                stages: vec![
2090                    StageGroupRequest {
2091                        layers: 0..15,
2092                        tensor: 1,
2093                        expert_parallel: false,
2094                    },
2095                    StageGroupRequest {
2096                        layers: 15..45,
2097                        tensor: 2,
2098                        expert_parallel: true,
2099                    },
2100                ],
2101                available_devices: 3,
2102                hardware: HardwareTarget::RtxPro6000Blackwell,
2103            })
2104            .unwrap();
2105
2106        assert_eq!(plan.world_size, 3);
2107        assert_eq!(plan.rank_groups[0].global_ranks, 0..1);
2108        assert_eq!(plan.rank_groups[1].global_ranks, 1..3);
2109        assert_eq!(plan.global_rank(0, 0), Some(0));
2110        assert_eq!(plan.global_rank(1, 0), Some(1));
2111        assert_eq!(plan.global_rank(1, 1), Some(2));
2112        assert_eq!(plan.query_head_range(16, 1), Some(32..64));
2113        assert_eq!(plan.query_head_range(17, 1), Some(48..96));
2114        assert_eq!(plan.kv_head_range(16, 1), Some(4..8));
2115        assert_eq!(plan.routed_expert_range(16, 1), Some(144..288));
2116        assert_eq!(plan.mtp_owner_stage, Some(1));
2117    }
2118
2119    #[test]
2120    fn grouped_step_layouts_cover_every_card_count_through_eight() {
2121        let layouts: Vec<Vec<usize>> = vec![
2122            vec![1],
2123            vec![2],
2124            vec![1, 2],
2125            vec![4],
2126            vec![1, 4],
2127            vec![2, 4],
2128            vec![1, 2, 4],
2129            vec![8],
2130        ];
2131        for (index, widths) in layouts.into_iter().enumerate() {
2132            let cards = index + 1;
2133            let cuts: Vec<usize> = match widths.len() {
2134                1 => vec![0, 45],
2135                2 => vec![0, 3, 45],
2136                3 => vec![0, 3, 15, 45],
2137                _ => unreachable!(),
2138            };
2139            let stages = widths
2140                .iter()
2141                .enumerate()
2142                .map(|(stage, &tensor)| StageGroupRequest {
2143                    layers: cuts[stage]..cuts[stage + 1],
2144                    tensor,
2145                    expert_parallel: tensor > 1 && cuts[stage + 1] > 3,
2146                })
2147                .collect();
2148            let plan = step37_contract()
2149                .plan_grouped(GroupedTopologyRequest {
2150                    stages,
2151                    available_devices: cards,
2152                    hardware: HardwareTarget::RtxPro6000Blackwell,
2153                })
2154                .unwrap_or_else(|error| panic!("{cards}-card grouped plan failed: {error}"));
2155            assert_eq!(plan.world_size, cards);
2156            assert_eq!(plan.rank_groups.last().unwrap().layers.end, 45);
2157        }
2158    }
2159
2160    #[test]
2161    fn grouped_step_layout_refuses_gaps_overlap_and_invalid_stage_tp() {
2162        for stages in [
2163            vec![
2164                StageGroupRequest {
2165                    layers: 0..3,
2166                    tensor: 1,
2167                    expert_parallel: false,
2168                },
2169                StageGroupRequest {
2170                    layers: 4..45,
2171                    tensor: 2,
2172                    expert_parallel: true,
2173                },
2174            ],
2175            vec![
2176                StageGroupRequest {
2177                    layers: 0..16,
2178                    tensor: 1,
2179                    expert_parallel: false,
2180                },
2181                StageGroupRequest {
2182                    layers: 15..45,
2183                    tensor: 2,
2184                    expert_parallel: true,
2185                },
2186            ],
2187        ] {
2188            let error = step37_contract()
2189                .plan_grouped(GroupedTopologyRequest {
2190                    stages,
2191                    available_devices: 3,
2192                    hardware: HardwareTarget::RtxPro6000Blackwell,
2193                })
2194                .unwrap_err();
2195            assert!(error.to_string().contains("do not continue"));
2196        }
2197
2198        let tp3 = step37_contract()
2199            .plan_grouped(GroupedTopologyRequest {
2200                stages: vec![StageGroupRequest {
2201                    layers: 0..45,
2202                    tensor: 3,
2203                    expert_parallel: true,
2204                }],
2205                available_devices: 3,
2206                hardware: HardwareTarget::RtxPro6000Blackwell,
2207            })
2208            .unwrap_err();
2209        assert!(tp3.to_string().contains("layer 0 query heads 64"));
2210    }
2211
2212    #[test]
2213    fn full_model_step_tp8_preflight_binds_one_runtime_group() {
2214        let contract = step37_contract();
2215        let devices = (0..8).collect::<Vec<_>>();
2216        let owners = vec![0; contract.trunk_layers];
2217        let plan = contract
2218            .preflight_step_tp_specs(
2219                (0..contract.trunk_layers).map(|layer| (layer, devices.as_slice())),
2220                &owners,
2221            )
2222            .unwrap();
2223
2224        assert!(plan.full_trunk);
2225        assert_eq!(plan.layers.len(), STEP37_TRUNK_LAYERS);
2226        assert_eq!(plan.runtime_groups, vec![devices.clone()]);
2227        assert_eq!(plan.dense_attention_layers(), 3);
2228        assert_eq!(plan.tensor_parallel_expert_layers(), 0);
2229        assert_eq!(plan.expert_parallel_layers(), 42);
2230        assert_eq!(plan.layers.first().unwrap().layer, 0);
2231        assert_eq!(plan.layers.last().unwrap().layer, 44);
2232        assert!(
2233            plan.layers
2234                .iter()
2235                .all(|layer| layer.owner_device == 0 && layer.devices == devices)
2236        );
2237    }
2238
2239    #[test]
2240    fn step_tp_preflight_is_partial_for_tp2_and_fails_closed_on_invalid_specs() {
2241        let contract = step37_contract();
2242        let owners = vec![0; contract.trunk_layers];
2243        let tp2 = vec![0, 1];
2244        let partial = contract
2245            .preflight_step_tp_specs([(3, tp2.as_slice()), (44, tp2.as_slice())], &owners)
2246            .unwrap();
2247        assert!(!partial.full_trunk);
2248        assert_eq!(partial.runtime_groups, vec![tp2]);
2249        assert_eq!(partial.tensor_parallel_expert_layers(), 2);
2250        assert_eq!(partial.expert_parallel_layers(), 0);
2251
2252        let wrong_owner = vec![1, 2];
2253        assert!(
2254            contract
2255                .preflight_step_tp_specs([(24, wrong_owner.as_slice())], &owners)
2256                .unwrap_err()
2257                .to_string()
2258                .contains("owning PP device 0 must be the first rank")
2259        );
2260
2261        let tp3 = vec![0, 1, 2];
2262        assert!(
2263            contract
2264                .preflight_step_tp_specs([(24, tp3.as_slice())], &owners)
2265                .unwrap_err()
2266                .to_string()
2267                .contains("layer 0 query heads 64")
2268        );
2269
2270        assert!(
2271            contract
2272                .preflight_step_tp_specs(
2273                    [(24, [0, 1].as_slice()), (24, [0, 1].as_slice())],
2274                    &owners,
2275                )
2276                .unwrap_err()
2277                .to_string()
2278                .contains("assigns layer 24 more than once")
2279        );
2280    }
2281
2282    #[test]
2283    fn structural_tp4_defers_quant_block_legality_to_the_artifact_backend() {
2284        let tp4 = step37_contract().plan(request(1, 4, false)).unwrap();
2285        assert_eq!(tp4.routed_expert_ffn_range(1), Some(320..640));
2286        let tp2 = step37_contract().plan(request(1, 2, false)).unwrap();
2287        assert_eq!(tp2.routed_expert_ffn_range(1), Some(640..1280));
2288        assert!(tp2.shared_expert_replicated);
2289    }
2290
2291    #[test]
2292    fn step_tp3_refuses_the_real_per_layer_head_geometry() {
2293        let error = step37_contract()
2294            .plan(TopologyRequest {
2295                pipeline: 1,
2296                tensor: 3,
2297                expert_parallel: true,
2298                available_devices: 3,
2299                hardware: HardwareTarget::RtxPro6000Blackwell,
2300            })
2301            .unwrap_err();
2302        assert!(error.to_string().contains("layer 0 query heads 64"));
2303    }
2304
2305    #[test]
2306    fn product_envelope_accepts_eight_and_refuses_more() {
2307        let pp8 = step37_contract().plan(request(8, 1, false)).unwrap();
2308        assert_eq!(pp8.world_size, 8);
2309        assert_eq!(pp8.stage_ranges.len(), 8);
2310        assert!(pp8.stage_ranges.iter().all(|range| !range.is_empty()));
2311
2312        let error = step37_contract()
2313            .plan(TopologyRequest {
2314                pipeline: 3,
2315                tensor: 4,
2316                expert_parallel: true,
2317                available_devices: 12,
2318                hardware: HardwareTarget::RtxPro6000Blackwell,
2319            })
2320            .unwrap_err();
2321        assert!(error.to_string().contains("product envelope is 8"));
2322    }
2323
2324    #[test]
2325    fn expert_parallel_requires_a_multi_rank_tp_group() {
2326        let error = step37_contract().plan(request(3, 1, true)).unwrap_err();
2327        assert!(
2328            error
2329                .to_string()
2330                .contains("expert parallelism requires TP group size greater than one")
2331        );
2332    }
2333
2334    #[test]
2335    fn step_does_not_inherit_the_5090_hardware_contract() {
2336        let error = step37_contract()
2337            .plan(TopologyRequest {
2338                pipeline: 1,
2339                tensor: 1,
2340                expert_parallel: false,
2341                available_devices: 1,
2342                hardware: HardwareTarget::Rtx5090,
2343            })
2344            .unwrap_err();
2345        assert!(
2346            error
2347                .to_string()
2348                .contains("has no qualified rtx-5090 contract")
2349        );
2350    }
2351}