use serde::{Deserialize, Serialize};
use crate::common::enums::DatabaseMode;
use crate::common::error::AicError;
use crate::operators::base::{MoeCommFallback, PerformanceResult, Source};
use crate::perf_database::PerfDatabase;
const MOE_A2A_BACKENDS: [(&str, &[&str]); 8] = [
("deepep_ht", &["dispatch", "combine"]),
("deepep_ll", &["dispatch", "combine"]),
("deepep_v2_context", &["dispatch", "combine"]),
("deepep_v2_generation", &["dispatch", "combine"]),
("trtllm_deepep_ht", &["dispatch", "combine"]),
("trtllm_deepep_ll", &["dispatch", "combine"]),
("nvlink_two_sided", &["prepare", "dispatch", "combine"]),
("nvlink_one_sided", &["dispatch", "combine"]),
];
const A2A_PHASES: [&str; 3] = ["prepare", "dispatch", "combine"];
fn validate_a2a_request(comm_backend: &str, phase: &str) -> Result<(), AicError> {
if !MOE_A2A_BACKENDS
.iter()
.any(|(name, _)| *name == comm_backend)
{
let known: Vec<&str> = MOE_A2A_BACKENDS.iter().map(|(name, _)| *name).collect();
return Err(AicError::InvalidEngineConfig(format!(
"Invalid comm_backend '{comm_backend}'. Must be one of {known:?}"
)));
}
if !A2A_PHASES.contains(&phase) {
return Err(AicError::InvalidEngineConfig(format!(
"Invalid phase '{phase}'. Must be one of {A2A_PHASES:?}"
)));
}
let supported = MOE_A2A_BACKENDS
.iter()
.find(|(name, _)| *name == comm_backend)
.map(|(_, phases)| *phases)
.expect("backend membership checked above");
if !supported.contains(&phase) {
return Err(AicError::InvalidEngineConfig(format!(
"comm_backend '{comm_backend}' does not implement phase '{phase}'; supported: {supported:?}"
)));
}
Ok(())
}
fn default_comm_dtype() -> String {
"default".to_string()
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct MoeAllToAllOp {
pub name: String,
pub scale_factor: f64,
pub phase: String,
pub comm_backend: String,
#[serde(default = "default_comm_dtype")]
pub comm_dtype: String,
pub hidden_size: u32,
pub topk: u32,
pub num_experts: u32,
pub moe_ep_size: u32,
pub node_num: u32,
#[serde(default)]
pub sms: u32,
#[serde(default = "crate::operators::gemm::default_seq_split")]
pub attention_tp_size: u32,
}
impl MoeAllToAllOp {
pub fn query(&self, db: &PerfDatabase, num_tokens: u32) -> Result<PerformanceResult, AicError> {
let tokens = num_tokens / self.attention_tp_size.max(1);
validate_a2a_request(&self.comm_backend, &self.phase)?;
match db.database_mode {
DatabaseMode::Silicon | DatabaseMode::Hybrid => {}
mode => {
return Err(AicError::EmpiricalNotImplemented(format!(
"{mode:?} mode is not available for moe_a2a {}/{}: silicon data required \
(estimation tier is a planned follow-up).",
self.comm_backend, self.phase
)));
}
}
let exact_shape = db.moe_a2a.has_shape(
&self.comm_backend,
&self.phase,
&self.comm_dtype,
self.moe_ep_size,
self.node_num,
self.hidden_size,
self.topk,
self.num_experts,
)?;
let use_node1_fallback = matches!(
self.comm_backend.as_str(),
"deepep_ht" | "deepep_ll" | "trtllm_deepep_ht" | "trtllm_deepep_ll"
) && self.node_num > 1
&& !exact_shape;
let (lookup_ep_size, lookup_node_num, source) = if use_node1_fallback {
let donor_ep = if db.backend == "sglang" {
crate::perf_database::moe_a2a::legacy_deepep_ep_size(1)
} else {
db.system_spec.node.num_gpus_per_node
};
(donor_ep, 1, Source::Estimated)
} else {
(self.moe_ep_size, self.node_num, Source::Silicon)
};
let latency = self
.silicon_latency_at(db, tokens, lookup_ep_size, lookup_node_num)
.map_err(|err| {
if db.database_mode == DatabaseMode::Hybrid && err.is_missing_perf_data() {
AicError::EmpiricalNotImplemented(format!(
"HYBRID empirical fallback is not available for moe_a2a {}/{}: silicon data \
required (estimation tier is a planned follow-up). Silicon miss: {err}",
self.comm_backend, self.phase
))
} else {
err
}
})?;
let result = PerformanceResult::new(latency, source).scaled(self.scale_factor);
if use_node1_fallback {
let comm_backend = match self.comm_backend.as_str() {
"deepep_ht" => "deepep_ht",
"deepep_ll" => "deepep_ll",
"trtllm_deepep_ht" => "trtllm_deepep_ht",
"trtllm_deepep_ll" => "trtllm_deepep_ll",
_ => unreachable!("node-1 fallback is restricted to DeepEP HT/LL backends"),
};
Ok(result.with_moe_comm_fallback(MoeCommFallback {
comm_backend,
requested_ep_size: self.moe_ep_size,
requested_node_num: self.node_num,
measurement_ep_size: lookup_ep_size,
measurement_node_num: lookup_node_num,
}))
} else {
Ok(result)
}
}
fn silicon_latency_at(
&self,
db: &PerfDatabase,
tokens: u32,
ep_size: u32,
node_num: u32,
) -> Result<f64, AicError> {
db.moe_a2a.query(
&self.comm_backend,
&self.phase,
&self.comm_dtype,
ep_size,
node_num,
self.hidden_size,
self.topk,
self.num_experts,
tokens,
self.sms,
)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::common::enums::TransferPolicy;
use crate::operators::{FallbackOp, Op, OverlapOp, RuntimeContext};
use crate::perf_database::MoeA2aTable;
use parquet::data_type::{ByteArray, ByteArrayType, DoubleType, Int64Type};
use parquet::file::properties::WriterProperties;
use parquet::file::writer::{SerializedFileWriter, SerializedRowGroupWriter};
use parquet::schema::parser::parse_message_type;
use std::collections::BTreeMap;
use std::fs::File;
use std::path::{Path, PathBuf};
use std::sync::Arc;
fn systems_root() -> PathBuf {
PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("../../python/aisimulate/src/aiconfigurator_core/systems")
}
fn write_column<T: parquet::data_type::DataType>(
rg: &mut SerializedRowGroupWriter<'_, File>,
values: &[T::T],
) {
let mut col = rg.next_column().unwrap().unwrap();
col.typed::<T>().write_batch(values, None, None).unwrap();
col.close().unwrap();
}
struct A2aRow {
comm_backend: &'static str,
phase: &'static str,
ep_size: i64,
node_num: i64,
sms: i64,
num_tokens: i64,
latency_us: f64,
}
fn a2a_row(
comm_backend: &'static str,
phase: &'static str,
num_tokens: i64,
latency_us: f64,
) -> A2aRow {
A2aRow {
comm_backend,
phase,
ep_size: 16,
node_num: 2,
sms: 0,
num_tokens,
latency_us,
}
}
fn a2a_row_at(
comm_backend: &'static str,
phase: &'static str,
ep_size: i64,
node_num: i64,
sms: i64,
num_tokens: i64,
latency_us: f64,
) -> A2aRow {
A2aRow {
comm_backend,
phase,
ep_size,
node_num,
sms,
num_tokens,
latency_us,
}
}
fn write_moe_a2a_parquet(path: &Path, rows: &[A2aRow]) {
let schema = Arc::new(
parse_message_type(
"message a2a {
REQUIRED BYTE_ARRAY comm_backend (UTF8);
REQUIRED BYTE_ARRAY phase (UTF8);
REQUIRED BYTE_ARRAY comm_dtype (UTF8);
REQUIRED INT64 ep_size;
REQUIRED INT64 node_num;
REQUIRED INT64 hidden_size;
REQUIRED INT64 topk;
REQUIRED INT64 num_experts;
REQUIRED INT64 sms;
REQUIRED INT64 num_tokens;
REQUIRED DOUBLE latency;
}",
)
.unwrap(),
);
let file = File::create(path).unwrap();
let mut writer =
SerializedFileWriter::new(file, schema, Arc::new(WriterProperties::builder().build()))
.unwrap();
let mut rg = writer.next_row_group().unwrap();
let n = rows.len();
write_column::<ByteArrayType>(
&mut rg,
&rows
.iter()
.map(|r| ByteArray::from(r.comm_backend))
.collect::<Vec<_>>(),
);
write_column::<ByteArrayType>(
&mut rg,
&rows
.iter()
.map(|r| ByteArray::from(r.phase))
.collect::<Vec<_>>(),
);
write_column::<ByteArrayType>(&mut rg, &vec![ByteArray::from("default"); n]);
write_column::<Int64Type>(&mut rg, &rows.iter().map(|r| r.ep_size).collect::<Vec<_>>());
write_column::<Int64Type>(
&mut rg,
&rows.iter().map(|r| r.node_num).collect::<Vec<_>>(),
);
write_column::<Int64Type>(&mut rg, &vec![7168_i64; n]);
write_column::<Int64Type>(&mut rg, &vec![8_i64; n]);
write_column::<Int64Type>(&mut rg, &vec![256_i64; n]);
write_column::<Int64Type>(&mut rg, &rows.iter().map(|r| r.sms).collect::<Vec<_>>());
write_column::<Int64Type>(
&mut rg,
&rows.iter().map(|r| r.num_tokens).collect::<Vec<_>>(),
);
write_column::<DoubleType>(
&mut rg,
&rows.iter().map(|r| r.latency_us).collect::<Vec<_>>(),
);
rg.close().unwrap();
writer.close().unwrap();
}
fn synthetic_db(mode: DatabaseMode) -> (tempfile::TempDir, PerfDatabase) {
let tmp = tempfile::tempdir().expect("tmpdir");
write_moe_a2a_parquet(
&tmp.path().join("moe_a2a_perf.parquet"),
&[
a2a_row("deepep_ht", "dispatch", 31, 310.0),
a2a_row("deepep_ht", "dispatch", 32, 320.0),
a2a_row("deepep_ht", "dispatch", 63, 630.0),
a2a_row("deepep_ht", "dispatch", 64, 640.0),
a2a_row("deepep_ht", "combine", 31, 3100.0),
a2a_row("deepep_ht", "combine", 63, 6300.0),
a2a_row("deepep_ht", "combine", 64, 6400.0),
a2a_row("deepep_ll", "dispatch", 63, 63000.0),
a2a_row("deepep_ll", "dispatch", 64, 64000.0),
a2a_row_at("deepep_ht", "dispatch", 8, 1, 20, 64, 1280.0),
a2a_row_at("deepep_ht", "combine", 8, 1, 20, 64, 2560.0),
a2a_row_at("deepep_ll", "dispatch", 8, 1, 0, 64, 12800.0),
a2a_row_at("trtllm_deepep_ht", "dispatch", 8, 1, 20, 64, 5120.0),
a2a_row_at("trtllm_deepep_ll", "dispatch", 8, 1, 0, 64, 25600.0),
],
);
let mut db = PerfDatabase::load(&systems_root(), "h200_sxm", "sglang", "0.5.6.post2")
.expect("h200_sxm/sglang/0.5.6.post2 must load")
.with_mode(mode, TransferPolicy::ALL);
db.tables_mut().moe_a2a = MoeA2aTable::new(tmp.path().to_path_buf());
(tmp, db)
}
fn op(phase: &str, comm_backend: &str, attention_tp_size: u32) -> MoeAllToAllOp {
MoeAllToAllOp {
name: format!("moe_{phase}"),
scale_factor: 1.0,
phase: phase.to_string(),
comm_backend: comm_backend.to_string(),
comm_dtype: "default".into(),
hidden_size: 7168,
topk: 8,
num_experts: 256,
moe_ep_size: 16,
node_num: 2,
sms: 0,
attention_tp_size,
}
}
#[test]
fn tp2_context_divides_and_tp1_generation_does_not() {
let (_tmp, db) = synthetic_db(DatabaseMode::Silicon);
let ctx = op("dispatch", "deepep_ht", 2)
.query(&db, 128)
.expect("context dispatch");
assert!((ctx.latency_ms - 0.640).abs() < 1e-12, "got {ctx:?}");
let generation = op("dispatch", "deepep_ll", 1)
.query(&db, 64)
.expect("generation dispatch");
assert!(
(generation.latency_ms - 64.0).abs() < 1e-12,
"got {generation:?}"
);
let halved = op("dispatch", "deepep_ll", 2)
.query(&db, 64)
.expect("tp2 generation");
assert!(
(halved.latency_ms - 64.0).abs() > 1e-6,
"tp=2 must move the token key, got {halved:?}"
);
}
#[test]
fn floor_division_is_exact_at_odd_token_counts() {
let (_tmp, db) = synthetic_db(DatabaseMode::Silicon);
let got = op("dispatch", "deepep_ht", 2)
.query(&db, 63)
.expect("odd x");
assert!(
(got.latency_ms - 0.310).abs() < 1e-12,
"x=63, tp=2 must key 31 (0.310 ms), got {got:?}"
);
let up = op("dispatch", "deepep_ht", 2)
.query(&db, 64)
.expect("even x");
assert!((up.latency_ms - 0.320).abs() < 1e-12, "got {up:?}");
}
#[test]
fn zero_token_key_is_reachable_and_not_clamped_to_one() {
let (_tmp, db) = synthetic_db(DatabaseMode::Silicon);
let missed = op("dispatch", "deepep_ht", 4).query(&db, 3);
assert!(
matches!(missed, Err(AicError::PerfDatabase(_))),
"0 tokens must stay a typed miss, got {missed:?}"
);
let guarded = op("dispatch", "deepep_ht", 4)
.query(&db, 4)
.expect("token 1 resolves through the below-range hold");
assert!(
guarded.latency_ms > 0.0,
"a max(1, ...) guard would have returned {} for x=3",
guarded.latency_ms
);
}
#[test]
fn scale_factor_multiplies_the_resolved_latency() {
let (_tmp, db) = synthetic_db(DatabaseMode::Silicon);
let mut scaled = op("combine", "deepep_ht", 1);
scaled.scale_factor = 61.0;
let got = scaled.query(&db, 64).expect("combine");
assert!((got.latency_ms - 6.400 * 61.0).abs() < 1e-12, "got {got:?}");
assert_eq!(got.source, Source::Silicon);
assert!(got.moe_comm_fallbacks.is_empty());
}
#[test]
fn sglang_deepep_node1_substitutes_for_missing_multi_node_scale_as_estimated() {
let (_tmp, db) = synthetic_db(DatabaseMode::Silicon);
let mut fallback = op("dispatch", "deepep_ht", 1);
fallback.moe_ep_size = 128;
fallback.node_num = 32;
fallback.sms = 20;
let got = fallback.query(&db, 64).expect("node-1 fallback");
assert!((got.latency_ms - 1.280).abs() < 1e-12, "got {got:?}");
assert_eq!(got.source, Source::Estimated);
assert_eq!(
got.moe_comm_fallbacks.iter().copied().collect::<Vec<_>>(),
vec![MoeCommFallback {
comm_backend: "deepep_ht",
requested_ep_size: 128,
requested_node_num: 32,
measurement_ep_size: 8,
measurement_node_num: 1,
}]
);
}
#[test]
fn trtllm_deepep_node1_substitutes_for_missing_multi_node_scale_as_estimated() {
let (_tmp, mut db) = synthetic_db(DatabaseMode::Silicon);
db.tables_mut().backend = "trtllm".to_string();
let mut fallback = op("dispatch", "trtllm_deepep_ht", 1);
fallback.moe_ep_size = 64;
fallback.node_num = 8;
fallback.sms = 20;
let got = fallback.query(&db, 64).expect("node-1 fallback");
assert!((got.latency_ms - 5.120).abs() < 1e-12, "got {got:?}");
assert_eq!(got.source, Source::Estimated);
assert_eq!(
got.moe_comm_fallbacks.iter().copied().collect::<Vec<_>>(),
vec![MoeCommFallback {
comm_backend: "trtllm_deepep_ht",
requested_ep_size: 64,
requested_node_num: 8,
measurement_ep_size: 8,
measurement_node_num: 1,
}]
);
}
#[test]
fn vllm_deepep_node1_substitutes_for_missing_multi_node_scale_as_estimated() {
let (_tmp, mut db) = synthetic_db(DatabaseMode::Silicon);
db.tables_mut().backend = "vllm".to_string();
let mut fallback = op("dispatch", "deepep_ll", 1);
fallback.moe_ep_size = 64;
fallback.node_num = 8;
fallback.sms = 0;
let got = fallback.query(&db, 64).expect("node-1 fallback");
assert!((got.latency_ms - 12.800).abs() < 1e-12, "got {got:?}");
assert_eq!(got.source, Source::Estimated);
assert_eq!(
got.moe_comm_fallbacks.iter().copied().collect::<Vec<_>>(),
vec![MoeCommFallback {
comm_backend: "deepep_ll",
requested_ep_size: 64,
requested_node_num: 8,
measurement_ep_size: 8,
measurement_node_num: 1,
}]
);
}
#[test]
fn exact_requested_topology_wins_over_available_node1_donor() {
let (_tmp, db) = synthetic_db(DatabaseMode::Silicon);
let got = op("dispatch", "deepep_ht", 1)
.query(&db, 64)
.expect("exact requested topology");
assert!((got.latency_ms - 0.640).abs() < 1e-12, "got {got:?}");
assert_eq!(got.source, Source::Silicon);
assert!(got.moe_comm_fallbacks.is_empty());
}
#[test]
fn missing_node1_donor_remains_a_data_error() {
let (_tmp, db) = synthetic_db(DatabaseMode::Silicon);
let mut missing = op("dispatch", "deepep_ht", 1);
missing.moe_ep_size = 128;
missing.node_num = 32;
missing.hidden_size = 9999;
missing.sms = 20;
let result = missing.query(&db, 64);
assert!(
matches!(result, Err(AicError::PerfDatabase(_))),
"an unavailable donor must remain a typed data miss, got {result:?}"
);
}
#[test]
fn overlap_and_fallback_composites_preserve_every_executed_a2a_substitution() {
let (_tmp, db) = synthetic_db(DatabaseMode::Silicon);
let mut context = op("dispatch", "deepep_ht", 1);
context.moe_ep_size = 128;
context.node_num = 32;
context.sms = 20;
let mut generation = op("dispatch", "deepep_ll", 1);
generation.moe_ep_size = 128;
generation.node_num = 32;
let overlap = Op::Overlap(OverlapOp::new(
"a2a_overlap",
vec![Op::MoeAllToAll(context.clone())],
vec![Op::MoeAllToAll(generation.clone())],
));
let overlap_result = overlap
.query(
&db,
&RuntimeContext {
num_tokens: 64,
..RuntimeContext::default()
},
)
.expect("overlap query");
assert_eq!(
overlap_result
.moe_comm_fallbacks
.iter()
.map(|fallback| fallback.comm_backend)
.collect::<Vec<_>>(),
vec!["deepep_ht", "deepep_ll"]
);
let mut missing_primary = context.clone();
missing_primary.hidden_size = 9999;
let fallback = Op::Fallback(FallbackOp::new(
"a2a_fallback",
Op::MoeAllToAll(missing_primary),
vec![Op::MoeAllToAll(context), Op::MoeAllToAll(generation)],
));
let fallback_result = fallback
.query(
&db,
&RuntimeContext {
num_tokens: 64,
..RuntimeContext::default()
},
)
.expect("fallback query");
assert_eq!(
fallback_result
.moe_comm_fallbacks
.iter()
.map(|record| record.comm_backend)
.collect::<Vec<_>>(),
vec!["deepep_ht", "deepep_ll"]
);
}
#[test]
fn unknown_backend_and_phase_are_config_errors_not_data_misses() {
let (_tmp, db) = synthetic_db(DatabaseMode::Silicon);
let bad_backend = op("dispatch", "deepep_xl", 1).query(&db, 64);
assert!(
matches!(&bad_backend, Err(AicError::InvalidEngineConfig(msg)) if msg.contains("comm_backend")),
"got {bad_backend:?}"
);
assert!(
!bad_backend.unwrap_err().is_missing_perf_data(),
"a registry ValueError must not be a missing-data signal"
);
let bad_phase = op("scatter", "deepep_ht", 1).query(&db, 64);
assert!(
matches!(&bad_phase, Err(AicError::InvalidEngineConfig(msg)) if msg.contains("phase")),
"got {bad_phase:?}"
);
}
#[test]
fn uncollected_shape_on_a_known_backend_is_a_data_miss() {
let (_tmp, db) = synthetic_db(DatabaseMode::Silicon);
let mut missing = op("dispatch", "deepep_ht", 1);
missing.hidden_size = 9999;
let result = missing.query(&db, 64);
assert!(
matches!(result, Err(AicError::PerfDatabase(_))),
"got {result:?}"
);
}
#[test]
fn undeclared_phase_is_a_config_error_not_a_data_miss() {
let (_tmp, db) = synthetic_db(DatabaseMode::Silicon);
let result = op("prepare", "deepep_ht", 1).query(&db, 64);
assert!(
matches!(result, Err(AicError::InvalidEngineConfig(_))),
"an undeclared phase must be a config error, got {result:?}"
);
let declared: &[&str] = MOE_A2A_BACKENDS
.iter()
.find(|(name, _)| *name == "nvlink_two_sided")
.map(|(_, phases)| *phases)
.expect("registry entry");
assert!(declared.contains(&"prepare"));
}
#[test]
fn declared_comm_phases_are_a_subset_of_the_validated_phases() {
for (backend, phases) in MOE_A2A_BACKENDS {
for phase in phases {
assert!(
A2A_PHASES.contains(phase),
"{backend} declares an unvalidatable phase {phase:?}"
);
}
}
}
#[test]
fn registry_contains_every_python_backend_identity() {
let names = MOE_A2A_BACKENDS
.iter()
.map(|(name, _)| *name)
.collect::<Vec<_>>();
assert_eq!(
names,
vec![
"deepep_ht",
"deepep_ll",
"deepep_v2_context",
"deepep_v2_generation",
"trtllm_deepep_ht",
"trtllm_deepep_ll",
"nvlink_two_sided",
"nvlink_one_sided",
]
);
for backend in names {
validate_a2a_request(backend, "dispatch")
.unwrap_or_else(|error| panic!("{backend} rejected: {error}"));
validate_a2a_request(backend, "combine")
.unwrap_or_else(|error| panic!("{backend} rejected: {error}"));
}
}
#[test]
fn estimation_tiers_raise_empirical_not_implemented() {
for mode in [
DatabaseMode::Sol,
DatabaseMode::SolFull,
DatabaseMode::Empirical,
] {
let (_tmp, db) = synthetic_db(mode);
let result = op("dispatch", "deepep_ht", 1).query(&db, 64);
assert!(
matches!(&result, Err(AicError::EmpiricalNotImplemented(msg))
if msg.contains("silicon data required")),
"{mode:?}: got {result:?}"
);
}
}
#[test]
fn hybrid_uses_silicon_when_covered_and_raises_not_implemented_on_a_miss() {
let (_tmp, db) = synthetic_db(DatabaseMode::Hybrid);
let hit = op("dispatch", "deepep_ht", 1)
.query(&db, 64)
.expect("covered shape");
assert!((hit.latency_ms - 0.640).abs() < 1e-12, "got {hit:?}");
assert_eq!(hit.source, Source::Silicon);
let mut missing = op("dispatch", "deepep_ht", 1);
missing.hidden_size = 9999;
let result = missing.query(&db, 64);
assert!(
matches!(&result, Err(AicError::EmpiricalNotImplemented(msg))
if msg.contains("silicon data required")),
"got {result:?}"
);
}
#[test]
fn registry_validation_precedes_the_mode_gate() {
let (_tmp, db) = synthetic_db(DatabaseMode::Empirical);
let result = op("dispatch", "deepep_xl", 1).query(&db, 64);
assert!(
matches!(result, Err(AicError::InvalidEngineConfig(_))),
"got {result:?}"
);
}
#[test]
fn omitted_optional_fields_take_the_python_ctor_defaults() {
let json = r#"{
"name": "moe_dispatch",
"scale_factor": 1.0,
"phase": "dispatch",
"comm_backend": "deepep_ht",
"hidden_size": 7168,
"topk": 8,
"num_experts": 256,
"moe_ep_size": 16,
"node_num": 2
}"#;
let op: MoeAllToAllOp = serde_json::from_str(json).expect("defaults must fill in");
assert_eq!(op.comm_dtype, "default");
assert_eq!(op.sms, 0);
assert_eq!(op.attention_tp_size, 1);
}
const LFS_POINTER_PREFIX: &[u8] = b"version https://git-lfs";
fn shipped_data_ready(data_root: &Path, basenames: &[&str]) -> bool {
use crate::config::PerfDbSources;
use crate::perf_database::resolve_op_sources;
use std::io::Read;
let mut any_file = false;
for basename in basenames {
for source in resolve_op_sources(&PerfDbSources::default(), basename, data_root) {
let path = source.path();
if !path.exists() {
continue;
}
let mut head = [0u8; LFS_POINTER_PREFIX.len()];
let Ok(mut file) = File::open(path) else {
return false;
};
let Ok(read) = file.read(&mut head) else {
return false;
};
if read >= LFS_POINTER_PREFIX.len() && head == LFS_POINTER_PREFIX {
return false;
}
any_file = true;
}
}
any_file
}
fn oracle_samples(op_kind: &str) -> Vec<serde_json::Value> {
let oracle: serde_json::Value =
serde_json::from_str(include_str!("testdata/op_oracle.json"))
.expect("oracle fixture must parse");
oracle["samples"]
.as_array()
.expect("samples array")
.iter()
.filter(|s| s["op"].as_str() == Some(op_kind))
.cloned()
.collect()
}
#[test]
fn moe_a2a_op_matches_python_oracle() {
let systems = systems_root();
let samples = oracle_samples("moe_a2a");
let mut dbs: BTreeMap<String, PerfDatabase> = BTreeMap::new();
let mut max_rel = 0.0_f64;
let mut checked = 0_usize;
for sample in &samples {
let system = sample["system"].as_str().expect("system");
let backend = sample["backend"].as_str().expect("backend");
let version = sample["version"].as_str().expect("version");
let tuple = format!("{system}/{backend}/{version}");
let data_root = systems.join(sample["data_root"].as_str().expect("data_root"));
if !shipped_data_ready(
&data_root,
&[
"moe_a2a_perf.parquet",
"wideep_deepep_normal_perf.parquet",
"wideep_deepep_ll_perf.parquet",
"trtllm_alltoall_perf.parquet",
],
) {
eprintln!(
"SKIP moe_a2a_op_matches_python_oracle: shipped perf data unavailable at {} \
(run `git lfs pull`)",
data_root.display()
);
return;
}
let db = dbs.entry(tuple.clone()).or_insert_with(|| {
PerfDatabase::load(&systems, system, backend, version)
.unwrap_or_else(|err| panic!("{tuple} must load: {err}"))
});
let u32_of = |field: &str| {
u32::try_from(sample[field].as_u64().expect(field)).expect("fits in u32")
};
let op = MoeAllToAllOp {
name: "oracle".into(),
scale_factor: sample["scale_factor"].as_f64().expect("scale_factor"),
phase: sample["phase"].as_str().expect("phase").to_string(),
comm_backend: sample["comm_backend"]
.as_str()
.expect("comm_backend")
.to_string(),
comm_dtype: sample["comm_dtype"]
.as_str()
.expect("comm_dtype")
.to_string(),
hidden_size: u32_of("hidden_size"),
topk: u32_of("topk"),
num_experts: u32_of("num_experts"),
moe_ep_size: u32_of("moe_ep_size"),
node_num: u32_of("node_num"),
sms: u32_of("sms"),
attention_tp_size: u32_of("attention_tp_size"),
};
let got = op
.query(db, u32_of("x"))
.unwrap_or_else(|err| panic!("oracle sample {sample} must resolve: {err}"));
let want = sample["latency_ms"].as_f64().expect("latency_ms");
assert!(
want > 0.0,
"oracle sample has a non-positive latency: {sample}"
);
let rel = ((got.latency_ms - want) / want).abs();
max_rel = max_rel.max(rel);
assert!(
rel <= 1e-9,
"sample {sample}: rust {} vs python {want} (rel {rel:e})",
got.latency_ms
);
checked += 1;
}
assert!(
checked >= 55,
"oracle unexpectedly small: {checked} samples"
);
eprintln!("moe_a2a op oracle: {checked} samples, max relative error {max_rel:e}");
}
}