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