use serde::{Deserialize, Serialize};
use crate::common::enums::{BackendKind, CommQuantMode, DatabaseMode, MoeQuantMode};
use crate::common::error::AicError;
use crate::common::system_spec::SystemSpec;
use crate::operators::base::{PerformanceResult, SolComponents, Source};
use crate::operators::communication::{CustomAllReduceOp, NcclOp};
use crate::operators::util_empirical::{self, UtilGrid};
use crate::perf_database::PerfDatabase;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub enum DispatchFlavor {
CustomAllReduce,
TrtllmAlltoall,
RetiredDeepEp,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct MoEDispatchOp {
pub name: String,
pub scale_factor: f64,
pub hidden_size: u32,
pub topk: u32,
pub num_experts: u32,
pub moe_tp_size: u32,
pub moe_ep_size: u32,
pub attention_dp_size: u32,
pub pre_dispatch: bool,
#[serde(default)]
pub attn_ar_modeled: bool,
pub backend: BackendKind,
pub flavor: DispatchFlavor,
pub comm_quant: CommQuantMode,
pub moe_quant: MoeQuantMode,
#[serde(default = "crate::operators::gemm::default_seq_split")]
pub attn_cp_size: u32,
#[serde(default)]
pub is_context: bool,
#[serde(default = "default_sms")]
pub sms: u32,
#[serde(default = "crate::operators::gemm::default_seq_split")]
pub scale_num_tokens: u32,
}
fn default_sms() -> u32 {
12
}
impl MoEDispatchOp {
pub fn new(
name: impl Into<String>,
hidden_size: u32,
topk: u32,
num_experts: u32,
moe_tp_size: u32,
moe_ep_size: u32,
attention_dp_size: u32,
pre_dispatch: bool,
backend: BackendKind,
flavor: DispatchFlavor,
) -> Self {
Self {
name: name.into(),
scale_factor: 1.0,
hidden_size,
topk,
num_experts,
moe_tp_size,
moe_ep_size,
attention_dp_size,
pre_dispatch,
backend,
flavor,
comm_quant: CommQuantMode::Half,
moe_quant: MoeQuantMode::Bfloat16,
attn_cp_size: 1,
is_context: false,
sms: default_sms(),
scale_num_tokens: 1,
attn_ar_modeled: false,
}
}
fn attention_tp_size(&self) -> u32 {
let total = self.moe_tp_size * self.moe_ep_size;
(total / self.attention_dp_size.max(1)).max(1)
}
pub fn query(&self, db: &PerfDatabase, num_tokens: u32) -> Result<PerformanceResult, AicError> {
let spec: &SystemSpec = &db.system_spec;
match self.flavor {
DispatchFlavor::RetiredDeepEp => {
return Err(AicError::InvalidEngineConfig(format!(
"MoEDispatch '{}' (moe_backend='deepep_moe') has no native evaluation \
(retired with AIC-1601; large-EP comm is modeled by MoeAllToAll)",
self.name
)));
}
DispatchFlavor::CustomAllReduce => {
let num_gpus = (self.moe_tp_size * self.moe_ep_size).max(1);
let attn_tp = self.attention_tp_size();
let attn_dp = self.attention_dp_size.max(1);
let pre = self.pre_dispatch;
let comm_latency_ms = match self.backend {
BackendKind::Vllm => {
let mut total = 0.0;
if attn_tp > 1 && !(pre && self.attn_ar_modeled) {
let ar =
CustomAllReduceOp::new(&self.name, 1.0, self.hidden_size, num_gpus);
total += ar.query(db, num_tokens)?.latency_ms;
}
if attn_dp > 1 {
let op_name = if pre { "all_gather" } else { "reduce_scatter" };
let nccl = NcclOp::new(
&self.name,
1.0,
self.hidden_size as f64,
num_gpus,
op_name,
);
total += nccl.query(db, num_tokens * attn_dp)?.latency_ms;
}
total
}
BackendKind::Sglang => {
let combined_tp_dp = attn_tp > 1 && attn_dp > 1;
if combined_tp_dp {
let (op1, gpus1, tokens1, op2, gpus2, tokens2) = if pre {
(
"reduce_scatter",
attn_tp,
num_tokens,
"all_gather",
num_gpus,
num_tokens * attn_dp,
)
} else {
(
"reduce_scatter",
num_gpus,
num_tokens * attn_dp,
"all_gather",
attn_tp,
num_tokens,
)
};
let n1 =
NcclOp::new(&self.name, 1.0, self.hidden_size as f64, gpus1, op1);
let n2 =
NcclOp::new(&self.name, 1.0, self.hidden_size as f64, gpus2, op2);
n1.query(db, tokens1)?.latency_ms + n2.query(db, tokens2)?.latency_ms
} else if self.attn_cp_size > 1 {
if self.is_context {
let op_name = if pre { "all_gather" } else { "reduce_scatter" };
let nccl = NcclOp::new(
&self.name,
1.0,
self.hidden_size as f64,
num_gpus,
op_name,
);
nccl.query(db, num_tokens)?.latency_ms
} else if pre {
0.0
} else {
let ar = CustomAllReduceOp::new(
&self.name,
1.0,
self.hidden_size,
num_gpus,
);
ar.query(db, num_tokens)?.latency_ms
}
} else if attn_tp > 1 {
let ar =
CustomAllReduceOp::new(&self.name, 1.0, self.hidden_size, num_gpus);
ar.query(db, num_tokens)?.latency_ms
} else if attn_dp > 1 {
let op_name = if pre { "all_gather" } else { "reduce_scatter" };
let nccl = NcclOp::new(
&self.name,
1.0,
self.hidden_size as f64,
num_gpus,
op_name,
);
nccl.query(db, num_tokens * attn_dp)?.latency_ms
} else {
0.0
}
}
BackendKind::Trtllm => {
let ar = CustomAllReduceOp::new(&self.name, 1.0, self.hidden_size, attn_tp);
ar.query(db, num_tokens)?.latency_ms
}
};
Ok(PerformanceResult::new(comm_latency_ms, Source::Silicon)
.clamp_non_negative()
.scaled(self.scale_factor))
}
DispatchFlavor::TrtllmAlltoall => {
let num_gpus = (self.moe_tp_size * self.moe_ep_size).max(1);
let attention_tp = self.attention_tp_size();
let pre = self.pre_dispatch;
let sm_version = spec.gpu.sm_version.map(i64::from).unwrap_or(-1);
let comm_latency_ms = if sm_version == 100 {
let is_nvl72 = spec.node.num_gpus_per_node >= 72;
let enable_alltoall =
self.attention_dp_size > 1 && self.moe_tp_size == 1 && is_nvl72;
let (x_factor, sf_factor) = match self.moe_quant {
MoeQuantMode::Nvfp4 => (0.25, 0.25 / 8.0),
MoeQuantMode::Fp8 | MoeQuantMode::Fp8Block => (0.5, 0.0),
_ => (1.0, 0.0),
};
if enable_alltoall {
let op_name = if pre {
"alltoall_dispatch"
} else {
"alltoall_combine"
};
query_alltoall_table(
db,
op_name,
num_tokens,
self.hidden_size,
self.topk,
self.num_experts,
self.moe_ep_size,
self.moe_quant,
None,
None,
)?
.latency_ms
} else if self.attention_dp_size > 1 {
if pre {
let nccl = NcclOp::new(
&self.name,
1.0,
self.hidden_size as f64 * (x_factor + sf_factor),
num_gpus,
"all_gather",
);
nccl.query(db, num_tokens * self.attention_dp_size)?
.latency_ms
} else {
let nccl = NcclOp::new(
&self.name,
1.0,
self.hidden_size as f64,
num_gpus,
"reduce_scatter",
);
nccl.query(db, num_tokens * self.attention_dp_size)?
.latency_ms
}
} else if attention_tp > 1 {
let ar =
CustomAllReduceOp::new(&self.name, 1.0, self.hidden_size, num_gpus);
ar.query(db, num_tokens)?.latency_ms
} else {
0.0
}
} else if attention_tp > 1 {
let ar = CustomAllReduceOp::new(&self.name, 1.0, self.hidden_size, num_gpus);
ar.query(db, num_tokens)?.latency_ms
} else if self.attention_dp_size > 1 {
let op_name = if pre { "all_gather" } else { "reduce_scatter" };
let nccl =
NcclOp::new(&self.name, 1.0, self.hidden_size as f64, num_gpus, op_name);
nccl.query(db, num_tokens * self.attention_dp_size)?
.latency_ms
} else {
0.0
};
Ok(PerformanceResult::new(comm_latency_ms, Source::Silicon)
.clamp_non_negative()
.scaled(self.scale_factor))
}
}
}
}
fn normalize_alltoall_quant_for_table(quant: MoeQuantMode) -> MoeQuantMode {
if quant == MoeQuantMode::Fp8Block {
MoeQuantMode::Fp8
} else {
quant
}
}
fn select_alltoall_kernel(
spec: &SystemSpec,
moe_ep_size: u32,
topk: u32,
moe_backend: Option<&str>,
) -> &'static str {
if let Some(backend) = moe_backend {
let upper = backend.to_uppercase();
if upper == "DEEPGEMM" || upper == "CUTE_DSL" {
return "NotEnabled";
}
}
let sm_version = spec.gpu.sm_version.unwrap_or(0);
let num_gpus_per_node = spec.node.num_gpus_per_node;
let is_inter_node = moe_ep_size > num_gpus_per_node;
let is_wideep = moe_backend.is_some_and(|b| b.to_uppercase() == "WIDEEP");
let supports_mnnvl = sm_version >= 100;
if is_wideep {
if supports_mnnvl {
"NVLinkTwoSided"
} else {
let deepep_feasible = moe_ep_size > 1 && topk <= 8;
if deepep_feasible && is_inter_node {
"DeepEP"
} else if deepep_feasible {
"DeepEPLowLatency"
} else {
"NotEnabled"
}
}
} else if supports_mnnvl {
"NVLinkOneSided"
} else {
"NotEnabled"
}
}
#[allow(clippy::too_many_arguments)]
fn alltoall_sol_ms(
spec: &SystemSpec,
op_name: &str,
quant: MoeQuantMode,
node_num: u32,
num_tokens: f64,
hidden_size: u32,
topk: u32,
num_experts: u32,
moe_ep_size: u32,
) -> f64 {
let is_inter_node = node_num > 1;
let bw = if is_inter_node {
spec.node.inter_node_bw
} else {
spec.node.intra_node_bw
};
let remote_ranks = topk.min(num_experts).min(moe_ep_size.saturating_sub(1)) as f64;
let data_bytes = if op_name == "alltoall_prepare" {
num_tokens * topk as f64 * 4.0 } else if op_name.contains("combine") {
let bytes_per_element = if op_name.contains("low_precision") {
0.5
} else {
2.0
};
num_tokens * remote_ranks * hidden_size as f64 * bytes_per_element
} else {
num_tokens * remote_ranks * hidden_size as f64 * quant.mapping().memory
};
data_bytes / bw * 1000.0
}
#[allow(clippy::too_many_arguments)]
fn query_alltoall_table(
db: &PerfDatabase,
op_name: &str,
num_tokens: u32,
hidden_size: u32,
topk: u32,
num_experts: u32,
moe_ep_size: u32,
quant: MoeQuantMode,
node_num: Option<u32>,
moe_backend: Option<&str>,
) -> Result<PerformanceResult, AicError> {
let table_quant = normalize_alltoall_quant_for_table(quant);
let node_num = node_num.unwrap_or(if moe_ep_size < 4 { 1 } else { moe_ep_size / 4 });
const VALID_OP_NAMES: [&str; 4] = [
"alltoall_prepare",
"alltoall_dispatch",
"alltoall_combine",
"alltoall_combine_low_precision",
];
if !VALID_OP_NAMES.contains(&op_name) {
return Err(AicError::InvalidEngineConfig(format!(
"Invalid op_name '{op_name}'. Must be one of {VALID_OP_NAMES:?}"
)));
}
let kernel_source = select_alltoall_kernel(&db.system_spec, moe_ep_size, topk, moe_backend);
if kernel_source == "NotEnabled" {
if matches!(db.database_mode, DatabaseMode::Sol | DatabaseMode::SolFull) {
return Ok(PerformanceResult::sol(SolComponents::new(0.0, 0.0)));
}
return Ok(PerformanceResult::new(0.0, Source::Empirical));
}
let silicon = || {
db.trtllm_alltoall.query_trtllm_alltoall(
&db.system_spec,
op_name,
num_tokens,
hidden_size,
topk,
num_experts,
moe_ep_size,
quant,
moe_backend,
)
};
let empirical = || {
alltoall_empirical(
db,
kernel_source,
op_name,
quant,
table_quant,
node_num,
num_tokens,
hidden_size,
topk,
num_experts,
moe_ep_size,
)
};
match db.database_mode {
DatabaseMode::Sol | DatabaseMode::SolFull => {
let sol_comm = alltoall_sol_ms(
&db.system_spec,
op_name,
quant,
node_num,
num_tokens as f64,
hidden_size,
topk,
num_experts,
moe_ep_size,
);
Ok(PerformanceResult::sol(SolComponents::new(sol_comm, 0.0)))
}
DatabaseMode::Empirical => Ok(PerformanceResult::new(empirical()?, Source::Empirical)),
DatabaseMode::Hybrid => match silicon() {
Ok(latency) => Ok(PerformanceResult::new(latency, Source::Silicon)),
Err(err) if err.is_missing_perf_data() => {
Ok(PerformanceResult::new(empirical()?, Source::Empirical))
}
Err(err) => Err(err),
},
_ => Ok(PerformanceResult::new(silicon()?, Source::Silicon)),
}
}
#[allow(clippy::too_many_arguments)]
fn alltoall_empirical(
db: &PerfDatabase,
kernel_source: &str,
op_name: &str,
quant: MoeQuantMode,
table_quant: MoeQuantMode,
node_num: u32,
num_tokens: u32,
hidden_size: u32,
topk: u32,
num_experts: u32,
moe_ep_size: u32,
) -> Result<f64, AicError> {
let spec = &db.system_spec;
let sol = |c: &[f64]| {
alltoall_sol_ms(
spec,
op_name,
quant,
node_num,
c[0],
hidden_size,
topk,
num_experts,
moe_ep_size,
)
};
let sol_time = sol(&[num_tokens as f64]);
let key = format!(
"alltoall:{kernel_source}:{op_name}:{}:{node_num}:{hidden_size}:{topk}:{num_experts}:{moe_ep_size}",
table_quant.name(),
);
let grid = db.util_grids.get_or_try_build(&key, || {
match db.trtllm_alltoall.alltoall_slice_points(
kernel_source,
op_name,
table_quant,
node_num,
hidden_size,
topk,
num_experts,
moe_ep_size,
) {
Ok(points) => Ok(Some(UtilGrid::new(util_empirical::build_samples(
points.into_iter().map(|(t, lat)| (vec![t as f64], lat)),
sol,
)))),
Err(err) if err.is_missing_perf_data() => Ok(None),
Err(err) => Err(err),
}
})?;
let query = [num_tokens as f64];
let (latency, _) = util_empirical::estimate(sol_time, &query, grid.as_deref(), 1.0)?;
db.note_provenance(util_empirical::ProvenanceTier::Empirical);
Ok(latency)
}
#[cfg(test)]
mod tests {
use super::*;
use std::path::PathBuf;
fn b200_sglang_db() -> PerfDatabase {
let root = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("../..")
.join("python/aisimulate/src/aiconfigurator_core/systems");
PerfDatabase::load(&root, "b200_sxm", "sglang", "0.5.10").expect("db loads")
}
fn cp_dispatch(pre_dispatch: bool, is_context: bool) -> MoEDispatchOp {
let mut op = MoEDispatchOp::new(
"moe_dispatch",
7168,
8,
256,
1, 8, 1, pre_dispatch,
BackendKind::Sglang,
DispatchFlavor::CustomAllReduce,
);
op.attn_cp_size = 8;
op.is_context = is_context;
op
}
#[test]
fn cp_decode_combine_is_custom_allreduce_not_zero() {
let db = b200_sglang_db();
let num_tokens = 64;
let combine = cp_dispatch(false, false)
.query(&db, num_tokens)
.expect("decode combine query");
let reference = CustomAllReduceOp::new("moe_dispatch", 1.0, 7168, 8)
.query(&db, num_tokens)
.expect("allreduce reference query");
assert!(
combine.latency_ms > 0.0,
"decode combine under CP must not be zeroed, got {}",
combine.latency_ms
);
assert!(
(combine.latency_ms - reference.latency_ms).abs() < 1e-12,
"decode combine ({}) must equal custom_allreduce(num_gpus=8, volume=64*7168) ({})",
combine.latency_ms,
reference.latency_ms
);
let pre = cp_dispatch(true, false)
.query(&db, num_tokens)
.expect("decode pre query");
assert_eq!(pre.latency_ms, 0.0, "decode pre-dispatch under CP is local");
}
use crate::common::enums::TransferPolicy;
fn gb200_trtllm_db(mode: DatabaseMode) -> PerfDatabase {
let root = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("../..")
.join("python/aisimulate/src/aiconfigurator_core/systems");
PerfDatabase::load(&root, "gb200", "trtllm", "1.3.0rc10")
.expect("db loads")
.with_mode(mode, TransferPolicy::ALL)
}
fn h100_sglang_db(mode: DatabaseMode) -> PerfDatabase {
let root = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("../..")
.join("python/aisimulate/src/aiconfigurator_core/systems");
PerfDatabase::load(&root, "h100_sxm", "sglang", "0.5.6.post2")
.expect("db loads")
.with_mode(mode, TransferPolicy::ALL)
}
fn a2a(
db: &PerfDatabase,
op_name: &str,
num_tokens: u32,
quant: MoeQuantMode,
moe_backend: Option<&str>,
) -> Result<PerformanceResult, AicError> {
query_alltoall_table(
db,
op_name,
num_tokens,
7168,
8,
256,
8,
quant,
None,
moe_backend,
)
}
fn assert_oracle(result: &PerformanceResult, expected: f64, source: Source, label: &str) {
assert!(
(result.latency_ms - expected).abs() < 1e-9,
"{label}: expected {expected}, got {}",
result.latency_ms
);
assert_eq!(result.source, source, "{label}: wrong source");
}
#[test]
fn alltoall_empirical_matches_python_oracles() {
let db = gb200_trtllm_db(DatabaseMode::Empirical);
let hit = a2a(&db, "alltoall_dispatch", 64, MoeQuantMode::Nvfp4, None).expect("exact hit");
assert_oracle(
&hit,
0.018886399269104005,
Source::Empirical,
"emp_dispatch_t64",
);
let off = a2a(&db, "alltoall_dispatch", 333, MoeQuantMode::Nvfp4, None).expect("off-grid");
assert_oracle(
&off,
0.033548976838374114,
Source::Empirical,
"emp_dispatch_t333",
);
let combine =
a2a(&db, "alltoall_combine", 333, MoeQuantMode::Nvfp4, None).expect("combine");
assert_oracle(
&combine,
0.07118654040018774,
Source::Empirical,
"emp_combine_t333",
);
}
#[test]
fn alltoall_hybrid_prefers_silicon_when_covered() {
let db = gb200_trtllm_db(DatabaseMode::Hybrid);
let hit = a2a(&db, "alltoall_dispatch", 64, MoeQuantMode::Nvfp4, None).expect("exact hit");
assert_oracle(
&hit,
0.018886399269104005,
Source::Silicon,
"hyb_dispatch_t64",
);
let off = a2a(&db, "alltoall_dispatch", 333, MoeQuantMode::Nvfp4, None).expect("off-grid");
assert_oracle(
&off,
0.033547499962151055,
Source::Silicon,
"hyb_dispatch_t333",
);
let combine =
a2a(&db, "alltoall_combine", 333, MoeQuantMode::Nvfp4, None).expect("combine");
assert_oracle(
&combine,
0.07116495203226805,
Source::Silicon,
"hyb_combine_t333",
);
}
#[test]
fn alltoall_missing_slice_is_typed_empirical_miss() {
let emp = gb200_trtllm_db(DatabaseMode::Empirical);
for quant in [MoeQuantMode::Fp8, MoeQuantMode::Fp8Block] {
let result = a2a(&emp, "alltoall_dispatch", 333, quant, None);
assert!(
matches!(result, Err(AicError::EmpiricalNotImplemented(_))),
"EMPIRICAL {quant:?} must be a typed empirical miss, got {result:?}"
);
}
let hyb = gb200_trtllm_db(DatabaseMode::Hybrid);
let result = a2a(&hyb, "alltoall_dispatch", 333, MoeQuantMode::Fp8, None);
assert!(
matches!(result, Err(AicError::EmpiricalNotImplemented(_))),
"HYBRID fallback on a data-less slice must be a typed empirical miss, got {result:?}"
);
}
#[test]
fn alltoall_kernel_selection_matches_python() {
let db = gb200_trtllm_db(DatabaseMode::Hybrid);
let spec = &db.system_spec; assert_eq!(select_alltoall_kernel(spec, 8, 8, None), "NVLinkOneSided");
assert_eq!(
select_alltoall_kernel(spec, 8, 8, Some("wideep")),
"NVLinkTwoSided"
);
assert_eq!(
select_alltoall_kernel(spec, 8, 8, Some("DEEPGEMM")),
"NotEnabled"
);
assert_eq!(
select_alltoall_kernel(spec, 8, 8, Some("cute_dsl")),
"NotEnabled"
);
let mut hopper = db.system_spec.clone();
hopper.gpu.sm_version = Some(90);
assert_eq!(
select_alltoall_kernel(&hopper, 8, 8, Some("wideep")),
"DeepEP"
);
assert_eq!(
select_alltoall_kernel(&hopper, 4, 8, Some("wideep")),
"DeepEPLowLatency"
);
assert_eq!(
select_alltoall_kernel(&hopper, 8, 16, Some("wideep")),
"NotEnabled"
);
assert_eq!(select_alltoall_kernel(&hopper, 8, 8, None), "NotEnabled");
let zero = a2a(
&db,
"alltoall_dispatch",
333,
MoeQuantMode::Nvfp4,
Some("deepgemm"),
)
.expect("NotEnabled early return");
assert_oracle(&zero, 0.0, Source::Empirical, "not_enabled_zero");
}
}