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