use std::collections::{BTreeMap, BTreeSet};
use std::path::PathBuf;
use std::sync::OnceLock;
use super::axis_curve::AxisCurve;
use super::perf_interp::{self, Node, OpInterpConfig};
use super::{SourceResolver, kernel_source_ok};
use crate::common::error::AicError;
use crate::config::{PerfDbSources, PerfSource};
use crate::perf_database::parquet_loader::{PerfReader, PerfRow};
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub struct MoeA2aKey {
pub comm_backend: String,
pub phase: String,
pub comm_dtype: String,
pub ep_size: u32,
pub node_num: u32,
pub hidden_size: u32,
pub topk: u32,
pub num_experts: u32,
pub sms: u32,
}
struct MoeA2aGrids {
by_keys: BTreeMap<MoeA2aKey, BTreeMap<u32, f64>>,
dtypes_by_phase: BTreeMap<(String, String), BTreeSet<String>>,
}
type A2aGrid = BTreeMap<MoeA2aKey, BTreeMap<u32, f64>>;
pub struct MoeA2aTable {
data_root: PathBuf,
moe_a2a_sources: Vec<PerfSource>,
legacy_normal_sources: Vec<PerfSource>,
legacy_ll_sources: Vec<PerfSource>,
legacy_trtllm_alltoall_sources: Vec<PerfSource>,
grids: OnceLock<Result<MoeA2aGrids, AicError>>,
}
impl MoeA2aTable {
pub fn new(data_root: PathBuf) -> Self {
Self::with_sources(data_root, &SourceResolver::fixed(PerfDbSources::default()))
.expect("fixed-map resolution is infallible")
}
pub fn with_sources(data_root: PathBuf, resolver: &SourceResolver) -> Result<Self, AicError> {
let moe_a2a_sources = resolver.sources_for("moe_a2a_perf.parquet", &data_root)?;
let legacy_normal_sources =
resolver.sources_for("wideep_deepep_normal_perf.parquet", &data_root)?;
let legacy_ll_sources =
resolver.sources_for("wideep_deepep_ll_perf.parquet", &data_root)?;
let legacy_trtllm_alltoall_sources =
resolver.sources_for("trtllm_alltoall_perf.parquet", &data_root)?;
Ok(Self {
data_root,
moe_a2a_sources,
legacy_normal_sources,
legacy_ll_sources,
legacy_trtllm_alltoall_sources,
grids: OnceLock::new(),
})
}
#[allow(clippy::too_many_arguments)]
pub fn query(
&self,
comm_backend: &str,
phase: &str,
comm_dtype: &str,
ep_size: u32,
node_num: u32,
hidden_size: u32,
topk: u32,
num_experts: u32,
num_tokens: u32,
sms: u32,
) -> Result<f64, AicError> {
let grids = self.load()?;
let resolve_dtype = |phase_name: &str| -> Result<String, AicError> {
let phase_slice = (comm_backend.to_string(), phase_name.to_string());
let dtypes = grids.dtypes_by_phase.get(&phase_slice).ok_or_else(|| {
AicError::PerfDatabase(format!(
"moe_a2a data missing for comm_backend={comm_backend:?} \
phase={phase_name:?} at {}",
self.data_root.display()
))
})?;
if dtypes.contains(comm_dtype) {
Ok(comm_dtype.to_string())
} else if comm_dtype == "fp8_block" && dtypes.contains("fp8") {
Ok("fp8".to_string())
} else if dtypes.len() == 1 && dtypes.contains("default") {
Ok("default".to_string())
} else {
Err(AicError::PerfDatabase(format!(
"moe_a2a comm_dtype {comm_dtype:?} is not available for \
{comm_backend}/{phase_name} at {}; collected dtypes: {dtypes:?}",
self.data_root.display()
)))
}
};
let used_dtype = resolve_dtype(phase)?;
let collect_by_sms = |phase_name: &str, dtype: &str| {
let key_at = |sms: u32| MoeA2aKey {
comm_backend: comm_backend.to_string(),
phase: phase_name.to_string(),
comm_dtype: dtype.to_string(),
ep_size,
node_num,
hidden_size,
topk,
num_experts,
sms,
};
grids
.by_keys
.range(key_at(0)..=key_at(u32::MAX))
.map(|(key, curve)| (key.sms, curve))
.collect::<BTreeMap<_, _>>()
};
let by_sms = collect_by_sms(phase, &used_dtype);
if by_sms.is_empty() {
return Err(AicError::PerfDatabase(format!(
"moe_a2a data missing for {comm_backend}/{phase}, dtype={used_dtype}, \
ep={ep_size}, nodes={node_num}, hidden={hidden_size}, topk={topk}, \
experts={num_experts} at {}",
self.data_root.display(),
)));
}
let use_token_curve = by_sms.contains_key(&sms)
&& (comm_backend != "deepep_ht" || (node_num == 1 && sms == 20));
if use_token_curve {
let curve = by_sms.get(&sms).expect("contains_key checked");
return token_axis_curve(curve).query(num_tokens as f64, &|t| t);
}
let latency = query_sms_grid(&by_sms, sms, num_tokens)?;
if comm_backend == "deepep_ht" && matches!(phase, "dispatch" | "combine") {
let other_phase = if phase == "dispatch" {
"combine"
} else {
"dispatch"
};
let other_dtype = resolve_dtype(other_phase)?;
let other_by_sms = collect_by_sms(other_phase, &other_dtype);
let mut combined = BTreeMap::<u32, BTreeMap<u32, f64>>::new();
for (&sm, curve) in &by_sms {
let Some(other_curve) = other_by_sms.get(&sm) else {
continue;
};
for (&tokens, &value) in curve.iter() {
if let Some(&other_value) = other_curve.get(&tokens) {
combined
.entry(sm)
.or_default()
.insert(tokens, value + other_value);
}
}
}
if !combined.is_empty() {
let other_latency = query_sms_grid(&other_by_sms, sms, num_tokens)?;
let combined_refs = combined
.iter()
.map(|(&sm, curve)| (sm, curve))
.collect::<BTreeMap<_, _>>();
let combined_latency = query_sms_grid(&combined_refs, sms, num_tokens)?;
let phase_sum = latency + other_latency;
if phase_sum > 0.0 {
return Ok(latency * combined_latency / phase_sum);
}
}
}
Ok(latency)
}
fn load(&self) -> Result<&MoeA2aGrids, AicError> {
let cell = self.grids.get_or_init(|| {
load_moe_a2a_grids(
&self.moe_a2a_sources,
&self.legacy_normal_sources,
&self.legacy_ll_sources,
&self.legacy_trtllm_alltoall_sources,
)
});
cell.as_ref().map_err(clone_err)
}
}
fn query_sms_grid(
by_sms: &BTreeMap<u32, &BTreeMap<u32, f64>>,
sms: u32,
num_tokens: u32,
) -> Result<f64, AicError> {
let mut node = Node::branch();
for (&sm, curve) in by_sms {
for (&tokens, &latency) in curve.iter() {
node.insert(&[sm, tokens], latency);
}
}
let sol = |c: &[f64]| c[1];
let cfg = OpInterpConfig::grid(&["sms", "num_tokens"], &sol);
perf_interp::query(&cfg, &node, &[f64::from(sms), f64::from(num_tokens)])
}
pub(crate) const LEGACY_DEEPEP_DTYPE: &str = "default";
pub(crate) fn legacy_deepep_ep_size(node_num: u32) -> u32 {
node_num.saturating_mul(8)
}
pub(crate) const LEGACY_TRTLLM_DEFAULT_KERNEL_SOURCE: &str = "NVLinkTwoSided";
fn load_moe_a2a_grids(
a2a_sources: &[PerfSource],
normal_sources: &[PerfSource],
ll_sources: &[PerfSource],
trtllm_sources: &[PerfSource],
) -> Result<MoeA2aGrids, AicError> {
let mut by_keys: A2aGrid = BTreeMap::new();
let mut any_source = adapt_legacy_deepep_normal(normal_sources, &mut by_keys)?;
any_source |= adapt_legacy_deepep_ll(ll_sources, &mut by_keys)?;
any_source |= adapt_legacy_trtllm_alltoall(trtllm_sources, &mut by_keys)?;
any_source |= load_new_schema(a2a_sources, &mut by_keys)?;
if !any_source || by_keys.is_empty() {
return Err(AicError::PerfDatabase(format!(
"no MoE all-to-all rows loaded from {} source(s) (moe_a2a + 3 legacy tables; \
first: {})",
a2a_sources.len() + normal_sources.len() + ll_sources.len() + trtllm_sources.len(),
a2a_sources
.first()
.map(|s| s.path().display().to_string())
.unwrap_or_else(|| "<no moe_a2a sources>".to_string())
)));
}
let mut dtypes_by_phase: BTreeMap<(String, String), BTreeSet<String>> = BTreeMap::new();
for key in by_keys.keys() {
dtypes_by_phase
.entry((key.comm_backend.clone(), key.phase.clone()))
.or_default()
.insert(key.comm_dtype.clone());
}
Ok(MoeA2aGrids {
by_keys,
dtypes_by_phase,
})
}
fn token_axis_curve(points: &std::collections::BTreeMap<u32, f64>) -> AxisCurve {
AxisCurve::from_sorted_iter(
"num_tokens",
points
.iter()
.map(|(&coordinate, &value)| (coordinate, value)),
)
}
fn store_first_wins(by_keys: &mut A2aGrid, key: MoeA2aKey, num_tokens: u32, latency_ms: f64) {
by_keys
.entry(key)
.or_default()
.entry(num_tokens)
.or_insert(latency_ms);
}
fn legacy_deepep_key(
comm_backend: &str,
phase: &str,
node_num: u32,
hidden_size: u32,
topk: u32,
num_experts: u32,
sms: u32,
) -> MoeA2aKey {
MoeA2aKey {
comm_backend: comm_backend.to_string(),
phase: phase.to_string(),
comm_dtype: LEGACY_DEEPEP_DTYPE.to_string(),
ep_size: legacy_deepep_ep_size(node_num),
node_num,
hidden_size,
topk,
num_experts,
sms,
}
}
fn adapt_legacy_deepep_normal(
sources: &[PerfSource],
by_keys: &mut A2aGrid,
) -> Result<bool, AicError> {
let mut any_source = false;
for source in sources {
let path = source.path();
if !path.exists() {
continue;
}
any_source = true;
let reader = PerfReader::open(path)?;
let node_num_col = reader.col("node_num")?;
let hidden_size_col = reader.col("hidden_size")?;
let num_token_col = reader.col("num_token")?;
let num_topk_col = reader.col("num_topk")?;
let num_experts_col = reader.col("num_experts")?;
let dispatch_sms_col = reader.col("dispatch_sms")?;
let dispatch_transmit_us_col = reader.col_optional("dispatch_transmit_us");
let dispatch_notify_us_col = reader.col_optional("dispatch_notify_us");
let combine_transmit_us_col = reader.col_optional("combine_transmit_us");
let combine_notify_us_col = reader.col_optional("combine_notify_us");
let ks_col = reader.col_optional("kernel_source");
for row in reader.rows()? {
let row = row?;
if !kernel_source_ok(source.kernel_sources(), ks_col, &row)? {
continue;
}
let node_num = row.u32(node_num_col)?;
let hidden_size = row.u32(hidden_size_col)?;
let topk = row.u32(num_topk_col)?;
let num_experts = row.u32(num_experts_col)?;
let sms = row.u32(dispatch_sms_col)?;
let num_tokens = row.u32(num_token_col)?;
let dispatch_us = row.f64_optional(dispatch_transmit_us_col)?.unwrap_or(0.0)
+ row.f64_optional(dispatch_notify_us_col)?.unwrap_or(0.0);
let combine_us = row.f64_optional(combine_transmit_us_col)?.unwrap_or(0.0)
+ row.f64_optional(combine_notify_us_col)?.unwrap_or(0.0);
for (phase, latency_us) in [("dispatch", dispatch_us), ("combine", combine_us)] {
store_first_wins(
by_keys,
legacy_deepep_key(
"deepep_ht",
phase,
node_num,
hidden_size,
topk,
num_experts,
sms,
),
num_tokens,
latency_us / 1000.0,
);
}
}
}
Ok(any_source)
}
fn adapt_legacy_deepep_ll(sources: &[PerfSource], by_keys: &mut A2aGrid) -> Result<bool, AicError> {
let mut any_source = false;
for source in sources {
let path = source.path();
if !path.exists() {
continue;
}
any_source = true;
let reader = PerfReader::open(path)?;
let node_num_col = reader.col("node_num")?;
let hidden_size_col = reader.col("hidden_size")?;
let num_token_col = reader.col("num_token")?;
let num_topk_col = reader.col("num_topk")?;
let num_experts_col = reader.col("num_experts")?;
let dispatch_avg_t_us_col = reader.col_optional("dispatch_avg_t_us");
let combine_avg_t_us_col = reader.col_optional("combine_avg_t_us");
let ks_col = reader.col_optional("kernel_source");
for row in reader.rows()? {
let row = row?;
if !kernel_source_ok(source.kernel_sources(), ks_col, &row)? {
continue;
}
let node_num = row.u32(node_num_col)?;
let hidden_size = row.u32(hidden_size_col)?;
let topk = row.u32(num_topk_col)?;
let num_experts = row.u32(num_experts_col)?;
let num_tokens = row.u32(num_token_col)?;
let dispatch_us = row.f64_optional(dispatch_avg_t_us_col)?.unwrap_or(0.0);
let combine_us = row.f64_optional(combine_avg_t_us_col)?.unwrap_or(0.0);
for (phase, latency_us) in [("dispatch", dispatch_us), ("combine", combine_us)] {
store_first_wins(
by_keys,
legacy_deepep_key(
"deepep_ll",
phase,
node_num,
hidden_size,
topk,
num_experts,
0,
),
num_tokens,
latency_us / 1000.0,
);
}
}
}
Ok(any_source)
}
pub(crate) fn legacy_trtllm_backend(kernel_source: &str) -> Option<&'static str> {
match kernel_source {
"NVLinkTwoSided" => Some("nvlink_two_sided"),
"NVLinkOneSided" => Some("nvlink_one_sided"),
_ => None,
}
}
pub(crate) fn legacy_trtllm_phase_dtype(
op_name: &str,
) -> Option<(&'static str, Option<&'static str>)> {
match op_name {
"alltoall_prepare" => Some(("prepare", None)),
"alltoall_dispatch" => Some(("dispatch", None)),
"alltoall_combine" => Some(("combine", None)),
"alltoall_combine_low_precision" => Some(("combine", Some("fp4"))),
_ => None,
}
}
fn adapt_legacy_trtllm_alltoall(
sources: &[PerfSource],
by_keys: &mut A2aGrid,
) -> Result<bool, AicError> {
let mut any_source = false;
for source in sources {
let path = source.path();
if !path.exists() {
continue;
}
any_source = true;
let reader = PerfReader::open(path)?;
let op_name_col = reader.col("op_name")?;
let moe_dtype_col = reader.col("moe_dtype")?;
let num_tokens_col = reader.col("num_tokens")?;
let hidden_size_col = reader.col("hidden_size")?;
let topk_col = reader.col("topk")?;
let num_experts_col = reader.col("num_experts")?;
let moe_ep_size_col = reader.col("moe_ep_size")?;
let latency_col = reader.col("latency")?;
let ks_col = reader.col_optional("kernel_source");
let num_nodes_col = reader.col_optional("num_nodes");
for row in reader.rows()? {
let row = row?;
if !kernel_source_ok(source.kernel_sources(), ks_col, &row)? {
continue;
}
let kernel_source = match ks_col {
None => LEGACY_TRTLLM_DEFAULT_KERNEL_SOURCE,
Some(_) => row.str_optional(ks_col)?.unwrap_or(""),
};
let Some(comm_backend) = legacy_trtllm_backend(kernel_source) else {
continue;
};
let Some((phase, dtype_override)) = legacy_trtllm_phase_dtype(row.str(op_name_col)?)
else {
continue;
};
let comm_dtype = match dtype_override {
Some(dtype) => dtype.to_string(),
None => row.str_owned(moe_dtype_col)?,
};
let ep_size = row.u32(moe_ep_size_col)?;
let node_num = match num_nodes_col {
Some(_) => row.u32_optional(num_nodes_col)?.ok_or_else(|| {
AicError::PerfDatabase(format!(
"legacy trtllm alltoall row has a null num_nodes cell at {}",
path.display()
))
})?,
None => crate::perf_database::trtllm_alltoall::legacy_num_nodes_fallback(ep_size),
};
let key = MoeA2aKey {
comm_backend: comm_backend.to_string(),
phase: phase.to_string(),
comm_dtype,
ep_size,
node_num,
hidden_size: row.u32(hidden_size_col)?,
topk: row.u32(topk_col)?,
num_experts: row.u32(num_experts_col)?,
sms: 0,
};
store_first_wins(
by_keys,
key,
row.u32(num_tokens_col)?,
row.f64(latency_col)?,
);
}
}
Ok(any_source)
}
pub(crate) fn normalize_sms(row: &PerfRow, col: Option<usize>) -> Result<u32, AicError> {
if let Some(value) = row.u32_optional(col)? {
return Ok(value);
}
match row.f64_optional(col)? {
Some(value) if value.is_finite() => Ok(value.max(0.0) as u32),
_ => Ok(0),
}
}
fn load_new_schema(sources: &[PerfSource], by_keys: &mut A2aGrid) -> Result<bool, AicError> {
let mut any_source = false;
let mut seen: BTreeSet<(MoeA2aKey, u32)> = BTreeSet::new();
for source in sources {
let path = source.path();
if !path.exists() {
continue;
}
any_source = true;
let reader = PerfReader::open(path)?;
let comm_backend_col = reader.col("comm_backend")?;
let phase_col = reader.col("phase")?;
let comm_dtype_col = reader.col("comm_dtype")?;
let ep_size_col = reader.col("ep_size")?;
let node_num_col = reader.col("node_num")?;
let hidden_size_col = reader.col("hidden_size")?;
let topk_col = reader.col("topk")?;
let num_experts_col = reader.col("num_experts")?;
let num_tokens_col = reader.col("num_tokens")?;
let latency_col = reader.col("latency")?;
let sms_col = reader.col_optional("sms");
let ks_col = reader.col_optional("kernel_source");
for row in reader.rows()? {
let row = row?;
if !kernel_source_ok(source.kernel_sources(), ks_col, &row)? {
continue;
}
let key = MoeA2aKey {
comm_backend: row.str_owned(comm_backend_col)?,
phase: row.str_owned(phase_col)?,
comm_dtype: row.str_owned(comm_dtype_col)?,
ep_size: row.u32(ep_size_col)?,
node_num: row.u32(node_num_col)?,
hidden_size: row.u32(hidden_size_col)?,
topk: row.u32(topk_col)?,
num_experts: row.u32(num_experts_col)?,
sms: normalize_sms(&row, sms_col)?,
};
let num_tokens = row.u32(num_tokens_col)?;
let latency_ms = row.f64(latency_col)? / 1000.0;
if seen.insert((key.clone(), num_tokens)) {
by_keys
.entry(key)
.or_default()
.insert(num_tokens, latency_ms);
}
}
}
Ok(any_source)
}
fn clone_err(err: &AicError) -> AicError {
AicError::PerfDatabase(err.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
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::fs::File;
use std::path::Path;
use std::sync::Arc;
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();
}
#[derive(Clone)]
struct A2aRow {
comm_backend: &'static str,
phase: &'static str,
comm_dtype: &'static str,
ep_size: i64,
node_num: i64,
sms: Option<i64>,
num_tokens: i64,
latency_us: f64,
}
fn a2a_row(
comm_backend: &'static str,
phase: &'static str,
comm_dtype: &'static str,
ep_size: i64,
node_num: i64,
sms: Option<i64>,
num_tokens: i64,
latency_us: f64,
) -> A2aRow {
A2aRow {
comm_backend,
phase,
comm_dtype,
ep_size,
node_num,
sms,
num_tokens,
latency_us,
}
}
fn write_a2a_parquet(path: &Path, rows: &[A2aRow], with_sms_column: bool) {
let sms_decl = if with_sms_column {
"OPTIONAL INT64 sms;"
} else {
""
};
let schema = Arc::new(
parse_message_type(&format!(
"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;
{sms_decl}
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,
&rows
.iter()
.map(|r| ByteArray::from(r.comm_dtype))
.collect::<Vec<_>>(),
);
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]);
if with_sms_column {
let values: Vec<i64> = rows.iter().filter_map(|r| r.sms).collect();
let def_levels: Vec<i16> = rows
.iter()
.map(|r| if r.sms.is_some() { 1 } else { 0 })
.collect();
let mut col = rg.next_column().unwrap().unwrap();
col.typed::<Int64Type>()
.write_batch(&values, Some(&def_levels), None)
.unwrap();
col.close().unwrap();
}
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 write_deepep_normal_parquet(path: &Path, rows: &[(i64, i64, i64, f64, f64, f64, f64)]) {
let schema = Arc::new(
parse_message_type(
"message normal {
REQUIRED INT64 node_num;
REQUIRED INT64 hidden_size;
REQUIRED INT64 num_token;
REQUIRED INT64 num_topk;
REQUIRED INT64 num_experts;
REQUIRED INT64 dispatch_sms;
REQUIRED DOUBLE dispatch_transmit_us;
REQUIRED DOUBLE dispatch_notify_us;
REQUIRED DOUBLE combine_transmit_us;
REQUIRED DOUBLE combine_notify_us;
}",
)
.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::<Int64Type>(&mut rg, &rows.iter().map(|r| r.0).collect::<Vec<_>>());
write_column::<Int64Type>(&mut rg, &vec![7168_i64; n]);
write_column::<Int64Type>(&mut rg, &rows.iter().map(|r| r.2).collect::<Vec<_>>());
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.1).collect::<Vec<_>>());
write_column::<DoubleType>(&mut rg, &rows.iter().map(|r| r.3).collect::<Vec<_>>());
write_column::<DoubleType>(&mut rg, &rows.iter().map(|r| r.4).collect::<Vec<_>>());
write_column::<DoubleType>(&mut rg, &rows.iter().map(|r| r.5).collect::<Vec<_>>());
write_column::<DoubleType>(&mut rg, &rows.iter().map(|r| r.6).collect::<Vec<_>>());
rg.close().unwrap();
writer.close().unwrap();
}
fn write_deepep_ll_parquet(path: &Path, rows: &[(i64, i64, f64, f64)]) {
let schema = Arc::new(
parse_message_type(
"message ll {
REQUIRED INT64 node_num;
REQUIRED INT64 hidden_size;
REQUIRED INT64 num_token;
REQUIRED INT64 num_topk;
REQUIRED INT64 num_experts;
REQUIRED DOUBLE combine_avg_t_us;
REQUIRED DOUBLE dispatch_avg_t_us;
}",
)
.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::<Int64Type>(&mut rg, &rows.iter().map(|r| r.0).collect::<Vec<_>>());
write_column::<Int64Type>(&mut rg, &vec![7168_i64; n]);
write_column::<Int64Type>(&mut rg, &rows.iter().map(|r| r.1).collect::<Vec<_>>());
write_column::<Int64Type>(&mut rg, &vec![8_i64; n]);
write_column::<Int64Type>(&mut rg, &vec![256_i64; n]);
write_column::<DoubleType>(&mut rg, &rows.iter().map(|r| r.3).collect::<Vec<_>>());
write_column::<DoubleType>(&mut rg, &rows.iter().map(|r| r.2).collect::<Vec<_>>());
rg.close().unwrap();
writer.close().unwrap();
}
fn write_trtllm_alltoall_parquet(
path: &Path,
rows: &[(&'static str, &'static str, &'static str, i64, i64, f64)],
num_nodes: Option<i64>,
) {
let num_nodes_decl = if num_nodes.is_some() {
"REQUIRED INT64 num_nodes;"
} else {
""
};
let schema = Arc::new(
parse_message_type(&format!(
"message alltoall {{
REQUIRED BYTE_ARRAY op_name (UTF8);
REQUIRED BYTE_ARRAY kernel_source (UTF8);
REQUIRED BYTE_ARRAY moe_dtype (UTF8);
REQUIRED INT64 num_tokens;
REQUIRED INT64 hidden_size;
REQUIRED INT64 topk;
REQUIRED INT64 num_experts;
REQUIRED INT64 moe_ep_size;
{num_nodes_decl}
REQUIRED BYTE_ARRAY distribution (UTF8);
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.1))
.collect::<Vec<_>>(),
);
write_column::<ByteArrayType>(
&mut rg,
&rows
.iter()
.map(|r| ByteArray::from(r.0))
.collect::<Vec<_>>(),
);
write_column::<ByteArrayType>(
&mut rg,
&rows
.iter()
.map(|r| ByteArray::from(r.2))
.collect::<Vec<_>>(),
);
write_column::<Int64Type>(&mut rg, &rows.iter().map(|r| r.4).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.3).collect::<Vec<_>>());
if let Some(nodes) = num_nodes {
write_column::<Int64Type>(&mut rg, &vec![nodes; n]);
}
write_column::<ByteArrayType>(&mut rg, &vec![ByteArray::from("balanced"); n]);
write_column::<DoubleType>(&mut rg, &rows.iter().map(|r| r.5).collect::<Vec<_>>());
rg.close().unwrap();
writer.close().unwrap();
}
fn write_trtllm_alltoall_nullable_ks_parquet(
path: &Path,
rows: &[(Option<&'static str>, i64, f64)],
) {
let schema = Arc::new(
parse_message_type(
"message alltoall {
REQUIRED BYTE_ARRAY op_name (UTF8);
OPTIONAL BYTE_ARRAY kernel_source (UTF8);
REQUIRED BYTE_ARRAY moe_dtype (UTF8);
REQUIRED INT64 num_tokens;
REQUIRED INT64 hidden_size;
REQUIRED INT64 topk;
REQUIRED INT64 num_experts;
REQUIRED INT64 moe_ep_size;
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, &vec![ByteArray::from("alltoall_dispatch"); n]);
let values: Vec<ByteArray> = rows
.iter()
.filter_map(|r| r.0.map(ByteArray::from))
.collect();
let def_levels: Vec<i16> = rows.iter().map(|r| i16::from(r.0.is_some())).collect();
let mut col = rg.next_column().unwrap().unwrap();
col.typed::<ByteArrayType>()
.write_batch(&values, Some(&def_levels), None)
.unwrap();
col.close().unwrap();
write_column::<ByteArrayType>(&mut rg, &vec![ByteArray::from("fp8"); n]);
write_column::<Int64Type>(&mut rg, &vec![64_i64; n]);
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.1).collect::<Vec<_>>());
write_column::<DoubleType>(&mut rg, &rows.iter().map(|r| r.2).collect::<Vec<_>>());
rg.close().unwrap();
writer.close().unwrap();
}
fn approx(got: f64, want: f64) {
assert!(
(got - want).abs() <= 1e-12 * want.abs().max(1.0),
"got {got}, want {want}"
);
}
#[test]
fn new_schema_converts_us_to_ms_and_normalizes_sms() {
let tmp = tempfile::tempdir().unwrap();
write_a2a_parquet(
&tmp.path().join("moe_a2a_perf.parquet"),
&[
a2a_row("deepep_ht", "dispatch", "fp8", 16, 2, Some(20), 64, 250.0),
a2a_row("deepep_ht", "combine", "fp8", 16, 2, Some(20), 64, 250.0),
a2a_row("deepep_ll", "dispatch", "fp8", 16, 2, None, 64, 125.0),
],
true,
);
let table = MoeA2aTable::new(tmp.path().to_path_buf());
approx(
table
.query("deepep_ht", "dispatch", "fp8", 16, 2, 7168, 8, 256, 64, 20)
.unwrap(),
0.25,
);
approx(
table
.query("deepep_ll", "dispatch", "fp8", 16, 2, 7168, 8, 256, 64, 0)
.unwrap(),
0.125,
);
let tmp2 = tempfile::tempdir().unwrap();
write_a2a_parquet(
&tmp2.path().join("moe_a2a_perf.parquet"),
&[
a2a_row("deepep_ht", "dispatch", "fp8", 16, 2, Some(20), 64, 250.0),
a2a_row("deepep_ht", "combine", "fp8", 16, 2, Some(20), 64, 250.0),
],
false,
);
let table2 = MoeA2aTable::new(tmp2.path().to_path_buf());
approx(
table2
.query("deepep_ht", "dispatch", "fp8", 16, 2, 7168, 8, 256, 64, 0)
.unwrap(),
0.25,
);
}
#[test]
fn legacy_deepep_normal_sums_component_columns() {
let tmp = tempfile::tempdir().unwrap();
write_deepep_normal_parquet(
&tmp.path().join("wideep_deepep_normal_perf.parquet"),
&[(2, 20, 64, 100.0, 25.0, 300.0, 75.0)],
);
let table = MoeA2aTable::new(tmp.path().to_path_buf());
approx(
table
.query(
"deepep_ht",
"dispatch",
"default",
16,
2,
7168,
8,
256,
64,
20,
)
.unwrap(),
0.125,
);
approx(
table
.query(
"deepep_ht",
"combine",
"default",
16,
2,
7168,
8,
256,
64,
20,
)
.unwrap(),
0.375,
);
assert!(
table
.query(
"deepep_ht",
"dispatch",
"default",
2,
2,
7168,
8,
256,
64,
20
)
.is_err()
);
}
#[test]
fn legacy_deepep_ll_uses_average_columns_at_sms_zero() {
let tmp = tempfile::tempdir().unwrap();
write_deepep_ll_parquet(
&tmp.path().join("wideep_deepep_ll_perf.parquet"),
&[(4, 64, 90.0, 210.0)],
);
let table = MoeA2aTable::new(tmp.path().to_path_buf());
approx(
table
.query(
"deepep_ll",
"dispatch",
"default",
32,
4,
7168,
8,
256,
64,
0,
)
.unwrap(),
0.09,
);
approx(
table
.query(
"deepep_ll",
"combine",
"default",
32,
4,
7168,
8,
256,
64,
0,
)
.unwrap(),
0.21,
);
}
#[test]
fn legacy_trtllm_alltoall_maps_phases_dtypes_and_node_num() {
let tmp = tempfile::tempdir().unwrap();
write_trtllm_alltoall_parquet(
&tmp.path().join("trtllm_alltoall_perf.parquet"),
&[
("NVLinkTwoSided", "alltoall_prepare", "nvfp4", 16, 64, 0.5),
("NVLinkTwoSided", "alltoall_dispatch", "nvfp4", 16, 64, 1.5),
("NVLinkTwoSided", "alltoall_combine", "nvfp4", 16, 64, 2.5),
(
"NVLinkTwoSided",
"alltoall_combine_low_precision",
"nvfp4",
16,
64,
3.5,
),
("NVLinkOneSided", "alltoall_dispatch", "fp8", 2, 64, 4.5),
(
"MnnvlThreeSided",
"alltoall_dispatch",
"bfloat16",
8,
64,
9.0,
),
(
"NVLinkTwoSided",
"alltoall_something",
"bfloat16",
8,
64,
9.0,
),
],
None,
);
let table = MoeA2aTable::new(tmp.path().to_path_buf());
let q = |backend: &str, phase: &str, dtype: &str, ep: u32, node: u32| {
table.query(backend, phase, dtype, ep, node, 7168, 8, 256, 64, 0)
};
approx(
q("nvlink_two_sided", "prepare", "nvfp4", 16, 4).unwrap(),
0.5,
);
approx(
q("nvlink_two_sided", "dispatch", "nvfp4", 16, 4).unwrap(),
1.5,
);
approx(
q("nvlink_two_sided", "combine", "nvfp4", 16, 4).unwrap(),
2.5,
);
approx(q("nvlink_two_sided", "combine", "fp4", 16, 4).unwrap(), 3.5);
approx(q("nvlink_one_sided", "dispatch", "fp8", 2, 1).unwrap(), 4.5);
for phase in ["prepare", "dispatch", "combine"] {
assert!(
q("nvlink_two_sided", phase, "bfloat16", 8, 2).is_err(),
"an unmapped row leaked into nvlink_two_sided/{phase}"
);
assert!(
q("nvlink_one_sided", phase, "bfloat16", 8, 2).is_err(),
"an unmapped row leaked into nvlink_one_sided/{phase}"
);
}
assert!(q("nvlink_two_sided", "alltoall_something", "bfloat16", 8, 2).is_err());
}
#[test]
fn legacy_trtllm_alltoall_null_kernel_source_row_is_dropped() {
let tmp = tempfile::tempdir().unwrap();
write_trtllm_alltoall_nullable_ks_parquet(
&tmp.path().join("trtllm_alltoall_perf.parquet"),
&[(Some("NVLinkTwoSided"), 16, 1.5), (None, 32, 9.0)],
);
let table = MoeA2aTable::new(tmp.path().to_path_buf());
let q = |backend: &str, ep: u32, node: u32| {
table.query(backend, "dispatch", "fp8", ep, node, 7168, 8, 256, 64, 0)
};
approx(q("nvlink_two_sided", 16, 4).unwrap(), 1.5);
assert!(
q("nvlink_two_sided", 32, 8).is_err(),
"a null kernel_source cell must not default to NVLinkTwoSided"
);
assert!(q("nvlink_one_sided", 32, 8).is_err());
}
#[test]
fn legacy_trtllm_alltoall_num_nodes_column_wins() {
let tmp = tempfile::tempdir().unwrap();
write_trtllm_alltoall_parquet(
&tmp.path().join("trtllm_alltoall_perf.parquet"),
&[("NVLinkTwoSided", "alltoall_dispatch", "fp8", 16, 64, 1.25)],
Some(2),
);
let table = MoeA2aTable::new(tmp.path().to_path_buf());
approx(
table
.query(
"nvlink_two_sided",
"dispatch",
"fp8",
16,
2,
7168,
8,
256,
64,
0,
)
.unwrap(),
1.25,
);
assert!(
table
.query(
"nvlink_two_sided",
"dispatch",
"fp8",
16,
4,
7168,
8,
256,
64,
0
)
.is_err()
);
}
#[test]
fn dtype_chain_exact_then_fp8_block_alias_then_sole_then_miss() {
let tmp = tempfile::tempdir().unwrap();
write_a2a_parquet(
&tmp.path().join("moe_a2a_perf.parquet"),
&[
a2a_row("deepep_ht", "dispatch", "fp8", 16, 2, Some(20), 64, 100.0),
a2a_row("deepep_ht", "dispatch", "nvfp4", 16, 2, Some(20), 64, 200.0),
a2a_row("deepep_ht", "combine", "fp8", 16, 2, Some(20), 64, 100.0),
a2a_row("deepep_ht", "combine", "nvfp4", 16, 2, Some(20), 64, 200.0),
],
true,
);
let table = MoeA2aTable::new(tmp.path().to_path_buf());
let q =
|dtype: &str| table.query("deepep_ht", "dispatch", dtype, 16, 2, 7168, 8, 256, 64, 20);
approx(q("fp8").unwrap(), 0.1);
approx(q("nvfp4").unwrap(), 0.2);
approx(q("fp8_block").unwrap(), 0.1);
assert!(q("bfloat16").is_err());
let tmp2 = tempfile::tempdir().unwrap();
write_deepep_ll_parquet(
&tmp2.path().join("wideep_deepep_ll_perf.parquet"),
&[(2, 64, 100.0, 300.0)],
);
let table2 = MoeA2aTable::new(tmp2.path().to_path_buf());
approx(
table2
.query("deepep_ll", "dispatch", "nvfp4", 16, 2, 7168, 8, 256, 64, 0)
.unwrap(),
0.1,
);
assert!(
table2
.query(
"deepep_ll",
"prepare",
"default",
16,
2,
7168,
8,
256,
64,
0
)
.is_err()
);
}
#[test]
fn dtype_chain_exact_fp8_block_beats_the_fp8_alias() {
let tmp = tempfile::tempdir().unwrap();
write_a2a_parquet(
&tmp.path().join("moe_a2a_perf.parquet"),
&[
a2a_row("deepep_ht", "dispatch", "fp8", 16, 2, Some(20), 64, 100.0),
a2a_row("deepep_ht", "combine", "fp8", 16, 2, Some(20), 64, 100.0),
a2a_row(
"deepep_ht",
"dispatch",
"fp8_block",
16,
2,
Some(20),
64,
700.0,
),
a2a_row(
"deepep_ht",
"combine",
"fp8_block",
16,
2,
Some(20),
64,
700.0,
),
],
true,
);
let table = MoeA2aTable::new(tmp.path().to_path_buf());
approx(
table
.query(
"deepep_ht",
"dispatch",
"fp8_block",
16,
2,
7168,
8,
256,
64,
20,
)
.unwrap(),
0.7,
);
}
#[test]
fn sms_exact_is_1d_and_off_grid_is_2d() {
let tmp = tempfile::tempdir().unwrap();
write_a2a_parquet(
&tmp.path().join("moe_a2a_perf.parquet"),
&[
a2a_row("deepep_ht", "dispatch", "fp8", 16, 2, Some(16), 64, 100.0),
a2a_row("deepep_ht", "dispatch", "fp8", 16, 2, Some(16), 128, 200.0),
a2a_row("deepep_ht", "dispatch", "fp8", 16, 2, Some(32), 64, 500.0),
a2a_row("deepep_ht", "dispatch", "fp8", 16, 2, Some(32), 128, 900.0),
a2a_row("deepep_ht", "combine", "fp8", 16, 2, Some(16), 64, 100.0),
a2a_row("deepep_ht", "combine", "fp8", 16, 2, Some(16), 128, 200.0),
a2a_row("deepep_ht", "combine", "fp8", 16, 2, Some(32), 64, 500.0),
a2a_row("deepep_ht", "combine", "fp8", 16, 2, Some(32), 128, 900.0),
],
true,
);
let table = MoeA2aTable::new(tmp.path().to_path_buf());
let q = |sms: u32, tokens: u32| {
table
.query(
"deepep_ht",
"dispatch",
"fp8",
16,
2,
7168,
8,
256,
tokens,
sms,
)
.unwrap()
};
approx(q(16, 64), 0.1);
approx(q(16, 96), 0.15);
approx(q(32, 128), 0.9);
approx(q(24, 64), 0.3);
approx(q(24, 128), 0.55);
approx(q(24, 96), 0.425);
approx(q(8, 96), 0.182_754_372_856_303_42);
approx(q(40, 256), 0.856_666_877_173_728_9);
}
#[test]
fn new_schema_overwrites_legacy_and_repeats_keep_first() {
let tmp = tempfile::tempdir().unwrap();
write_deepep_ll_parquet(
&tmp.path().join("wideep_deepep_ll_perf.parquet"),
&[(2, 64, 100.0, 300.0)],
);
write_a2a_parquet(
&tmp.path().join("moe_a2a_perf.parquet"),
&[
a2a_row("deepep_ll", "dispatch", "default", 16, 2, None, 64, 700.0),
a2a_row("deepep_ll", "dispatch", "default", 16, 2, None, 64, 900.0),
],
true,
);
let table = MoeA2aTable::new(tmp.path().to_path_buf());
approx(
table
.query(
"deepep_ll",
"dispatch",
"default",
16,
2,
7168,
8,
256,
64,
0,
)
.unwrap(),
0.7,
);
approx(
table
.query(
"deepep_ll",
"combine",
"default",
16,
2,
7168,
8,
256,
64,
0,
)
.unwrap(),
0.3,
);
}
#[test]
fn legacy_duplicate_rows_keep_first() {
let tmp = tempfile::tempdir().unwrap();
write_deepep_normal_parquet(
&tmp.path().join("wideep_deepep_normal_perf.parquet"),
&[
(2, 20, 64, 100.0, 0.0, 0.0, 0.0),
(2, 20, 64, 999.0, 0.0, 0.0, 0.0),
],
);
let table = MoeA2aTable::new(tmp.path().to_path_buf());
approx(
table
.query(
"deepep_ht",
"dispatch",
"default",
16,
2,
7168,
8,
256,
64,
20,
)
.unwrap(),
0.1,
);
}
const LFS_POINTER_PREFIX: &[u8] = b"version https://git-lfs";
fn shipped_data_ready(data_root: &Path) -> bool {
use std::io::Read;
let mut any_file = false;
for basename in [
"moe_a2a_perf.parquet",
"wideep_deepep_normal_perf.parquet",
"wideep_deepep_ll_perf.parquet",
"trtllm_alltoall_perf.parquet",
] {
for source in crate::perf_database::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
}
#[test]
fn moe_a2a_matches_python_oracle() {
let oracle: serde_json::Value =
serde_json::from_str(include_str!("testdata/moe_a2a_oracle.json"))
.expect("oracle fixture must parse");
let systems = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("../../python/aisimulate/src/aiconfigurator_core/systems");
let samples = oracle["samples"].as_array().expect("samples array");
let mut tables: BTreeMap<String, MoeA2aTable> = BTreeMap::new();
let mut max_rel = 0.0_f64;
let mut checked = 0_usize;
for sample in samples {
let rel_root = sample["data_root"].as_str().expect("data_root");
let data_root = systems.join(rel_root);
if !shipped_data_ready(&data_root) {
eprintln!(
"SKIP moe_a2a_matches_python_oracle: shipped perf data unavailable at {} \
(run `git lfs pull`)",
data_root.display()
);
return;
}
let table = tables
.entry(rel_root.to_string())
.or_insert_with(|| MoeA2aTable::new(data_root.clone()));
let u32_of = |field: &str| {
u32::try_from(sample[field].as_u64().expect(field)).expect("fits in u32")
};
let got = table
.query(
sample["comm_backend"].as_str().expect("comm_backend"),
sample["phase"].as_str().expect("phase"),
sample["comm_dtype"].as_str().expect("comm_dtype"),
u32_of("ep_size"),
u32_of("node_num"),
u32_of("hidden_size"),
u32_of("topk"),
u32_of("num_experts"),
u32_of("num_tokens"),
u32_of("sms"),
)
.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 - want) / want).abs();
max_rel = max_rel.max(rel);
assert!(
rel <= 1e-9,
"sample {sample}: rust {got} vs python {want} (rel {rel:e})"
);
checked += 1;
}
assert!(
checked >= 150,
"oracle unexpectedly small: {checked} samples"
);
eprintln!("moe_a2a oracle: {checked} samples, max relative error {max_rel:e}");
}
#[test]
fn missing_sources_are_a_typed_miss() {
let tmp = tempfile::tempdir().unwrap();
let table = MoeA2aTable::new(tmp.path().to_path_buf());
match table
.query(
"deepep_ht",
"dispatch",
"default",
16,
2,
7168,
8,
256,
64,
20,
)
.unwrap_err()
{
AicError::PerfDatabase(_) | AicError::Io { .. } => {}
other => panic!("unexpected error: {other:?}"),
}
}
}