use std::collections::{BTreeMap, BTreeSet};
use std::path::PathBuf;
use std::sync::OnceLock;
use super::axis_curve::AxisCurve;
use super::{SourceResolver, kernel_source_ok};
use crate::common::enums::MoeQuantMode;
use crate::common::error::AicError;
use crate::common::system_spec::{SystemSpec, quant_tc_flops};
use crate::config::{PerfDbSources, PerfSource};
use crate::perf_database::parquet_loader::PerfReader;
fn token_axis_curve(points: &std::collections::BTreeMap<u32, f64>) -> AxisCurve {
AxisCurve::from_sorted_iter(
"num_tokens",
points
.iter()
.map(|(&coordinate, &value)| (coordinate, value)),
)
}
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub struct MoeExpertComputeKey {
pub kernel_source: String,
pub quant: String,
pub distribution: String,
pub inference_phase: String,
pub topk: u32,
pub num_experts: u32,
pub num_slots: u32,
pub hidden_size: u32,
pub inter_size: u32,
pub moe_tp_size: u32,
pub moe_ep_size: u32,
}
struct MoeEpGrids {
by_keys: BTreeMap<MoeExpertComputeKey, BTreeMap<u32, f64>>,
dist_order: BTreeMap<(String, String), Vec<String>>,
}
impl MoeEpGrids {
fn new() -> Self {
Self {
by_keys: BTreeMap::new(),
dist_order: BTreeMap::new(),
}
}
fn note_distribution(&mut self, key: &MoeExpertComputeKey) {
let order = self
.dist_order
.entry((key.kernel_source.clone(), key.quant.clone()))
.or_default();
if !order.iter().any(|d| d == &key.distribution) {
order.push(key.distribution.clone());
}
}
fn store_overwrite(&mut self, key: MoeExpertComputeKey, num_tokens: u32, latency_ms: f64) {
self.note_distribution(&key);
self.by_keys
.entry(key)
.or_default()
.insert(num_tokens, latency_ms);
}
fn store_keep_first(&mut self, key: MoeExpertComputeKey, num_tokens: u32, latency_ms: f64) {
self.note_distribution(&key);
self.by_keys
.entry(key)
.or_default()
.entry(num_tokens)
.or_insert(latency_ms);
}
}
pub struct MoeExpertComputeTable {
data_root: PathBuf,
spec: SystemSpec,
moe_ep_sources: Vec<PerfSource>,
legacy_context_sources: Vec<PerfSource>,
legacy_generation_sources: Vec<PerfSource>,
legacy_trtllm_wideep_sources: Vec<PerfSource>,
grids: OnceLock<Result<MoeEpGrids, AicError>>,
}
impl MoeExpertComputeTable {
pub fn new(data_root: PathBuf, spec: SystemSpec) -> Self {
Self::with_sources(
data_root,
spec,
&SourceResolver::fixed(PerfDbSources::default()),
)
.expect("fixed-map resolution is infallible")
}
pub fn with_sources(
data_root: PathBuf,
spec: SystemSpec,
resolver: &SourceResolver,
) -> Result<Self, AicError> {
let moe_ep_sources = resolver.sources_for("moe_expert_compute_perf.parquet", &data_root)?;
let legacy_context_sources =
resolver.sources_for("wideep_context_moe_perf.parquet", &data_root)?;
let legacy_generation_sources =
resolver.sources_for("wideep_generation_moe_perf.parquet", &data_root)?;
let legacy_trtllm_wideep_sources =
resolver.sources_for("wideep_moe_perf.parquet", &data_root)?;
Ok(Self {
data_root,
spec,
moe_ep_sources,
legacy_context_sources,
legacy_generation_sources,
legacy_trtllm_wideep_sources,
grids: OnceLock::new(),
})
}
#[allow(clippy::too_many_arguments)]
pub fn query(
&self,
kernel_source: &str,
quant: MoeQuantMode,
workload_distribution: &str,
inference_phase: &str,
topk: u32,
num_experts: u32,
num_slots: u32,
hidden_size: u32,
inter_size: u32,
moe_tp_size: u32,
moe_ep_size: u32,
num_tokens: u32,
is_gated: bool,
) -> Result<f64, AicError> {
let tc_flops = quant_tc_flops(&self.spec, MoeQuantMode::Bfloat16.mapping())?;
let grids = self.load()?;
let quant_name = quant.name();
let slice_key = (kernel_source.to_string(), quant_name.to_string());
let Some(dist_order) = grids.dist_order.get(&slice_key) else {
let kernel_seen = grids
.dist_order
.keys()
.any(|(kernel, _)| kernel == kernel_source);
return Err(AicError::PerfDatabase(if kernel_seen {
format!(
"moe_expert_compute data missing for kernel_source={kernel_source:?} \
quant={quant_name:?} at {}",
self.data_root.display()
)
} else {
format!(
"moe_expert_compute data missing for kernel_source={kernel_source:?} at {}",
self.data_root.display()
)
}));
};
let bound_key = |distribution: &str, fill: u32| MoeExpertComputeKey {
kernel_source: kernel_source.to_string(),
quant: quant_name.to_string(),
distribution: distribution.to_string(),
inference_phase: inference_phase.to_string(),
topk: fill,
num_experts: fill,
num_slots: fill,
hidden_size: fill,
inter_size: fill,
moe_tp_size: fill,
moe_ep_size: fill,
};
let carries_phase = |distribution: &str| {
grids
.by_keys
.range(bound_key(distribution, 0)..=bound_key(distribution, u32::MAX))
.next()
.is_some()
};
let used_distribution: &str = if carries_phase(workload_distribution) {
workload_distribution
} else if carries_phase("uniform") {
"uniform"
} else if let Some(first) = dist_order.iter().find(|dist| carries_phase(dist)) {
first
} else {
return Err(AicError::PerfDatabase(format!(
"moe_expert_compute workload_distribution {workload_distribution:?} is not available for \
{kernel_source}/{quant_name} at {}; no collected distribution carries \
{inference_phase:?} data",
self.data_root.display()
)));
};
let key = MoeExpertComputeKey {
kernel_source: kernel_source.to_string(),
quant: quant_name.to_string(),
distribution: used_distribution.to_string(),
inference_phase: inference_phase.to_string(),
topk,
num_experts,
num_slots,
hidden_size,
inter_size,
moe_tp_size,
moe_ep_size,
};
let curve = grids
.by_keys
.get(&key)
.filter(|curve| !curve.is_empty())
.ok_or_else(|| {
AicError::PerfDatabase(format!(
"moe_expert_compute data missing for {key:?} at {}",
self.data_root.display()
))
})?;
if let Some(only) = token_axis_curve(curve).singleton_underflow(num_tokens) {
return Err(AicError::PerfDatabase(format!(
"MoE EP silicon token underflow has only one measured point; cannot infer \
low-token latency from a singleton. measured_token={only}, requested \
num_tokens={num_tokens} for {key:?} at {}",
self.data_root.display()
)));
}
let sol = |tokens: f64| {
ep_sol_latency_ms(
&self.spec,
quant,
topk,
num_slots,
hidden_size,
inter_size,
moe_tp_size,
moe_ep_size,
tokens.round() as u32,
is_gated,
tc_flops,
)
};
token_axis_curve(curve).query(f64::from(num_tokens), &sol)
}
pub fn available_kernels(&self) -> Result<Vec<String>, AicError> {
let grids = match self.load() {
Ok(grids) => grids,
Err(err) if err.is_missing_perf_data() => return Ok(Vec::new()),
Err(err) => return Err(err),
};
let mut names: Vec<String> = Vec::new();
for key in grids.by_keys.keys() {
if names.last().map(String::as_str) != Some(key.kernel_source.as_str()) {
names.push(key.kernel_source.clone());
}
}
Ok(names)
}
fn load(&self) -> Result<&MoeEpGrids, AicError> {
let cell = self.grids.get_or_init(|| {
load_moe_ep_grids(
&self.moe_ep_sources,
&self.legacy_context_sources,
&self.legacy_generation_sources,
&self.legacy_trtllm_wideep_sources,
)
});
cell.as_ref().map_err(clone_err)
}
}
pub(crate) const SGLANG_ADAPTED_KERNEL_SOURCE: &str = "deepep_moe";
pub(crate) const LEGACY_TRTLLM_DEFAULT_KERNEL_SOURCE: &str = "moe_torch_flow";
#[allow(clippy::too_many_arguments)]
fn ep_sol_latency_ms(
spec: &SystemSpec,
quant: MoeQuantMode,
topk: u32,
num_slots: u32,
hidden_size: u32,
inter_size: u32,
moe_tp_size: u32,
moe_ep_size: u32,
tokens: u32,
is_gated: bool,
tc_flops: f64,
) -> f64 {
let total_tokens = tokens as u64 * topk as u64;
let moe_expert_compute = (moe_ep_size as u64).max(1);
let moe_tp = (moe_tp_size as u64).max(1);
let h = hidden_size as u64;
let inter = inter_size as u64;
let slots = num_slots as u64;
let num_gemms: u64 = if is_gated { 3 } else { 2 };
let ops = total_tokens * h * inter * num_gemms * 2 / moe_expert_compute / moe_tp;
let mem_bytes_int = total_tokens / moe_expert_compute * h * 2 + total_tokens / moe_expert_compute * inter * num_gemms / moe_tp + h * inter * num_gemms / moe_tp
* std::cmp::min(slots / moe_expert_compute, total_tokens / moe_expert_compute); let mem_bytes = (mem_bytes_int as f64) * quant.mapping().memory;
let sol_math = (ops as f64) / (tc_flops * quant.mapping().compute) * 1000.0;
let sol_mem = mem_bytes / spec.gpu.mem_bw * 1000.0;
sol_math.max(sol_mem)
}
fn load_moe_ep_grids(
moe_ep_sources: &[PerfSource],
context_sources: &[PerfSource],
generation_sources: &[PerfSource],
trtllm_wideep_sources: &[PerfSource],
) -> Result<MoeEpGrids, AicError> {
let mut grids = MoeEpGrids::new();
let mut any_source = adapt_legacy_sglang_wideep_moe(context_sources, "context", &mut grids)?;
any_source |= adapt_legacy_sglang_wideep_moe(generation_sources, "generation", &mut grids)?;
any_source |= adapt_legacy_trtllm_wideep_moe(trtllm_wideep_sources, &mut grids)?;
any_source |= load_new_schema(moe_ep_sources, &mut grids)?;
if !any_source || grids.by_keys.is_empty() {
return Err(AicError::PerfDatabase(format!(
"no MoE EP rows loaded from {} source(s) (moe_expert_compute + 3 legacy wideep tables; first: {})",
moe_ep_sources.len()
+ context_sources.len()
+ generation_sources.len()
+ trtllm_wideep_sources.len(),
moe_ep_sources
.first()
.map(|s| s.path().display().to_string())
.unwrap_or_else(|| "<no moe_expert_compute sources>".to_string())
)));
}
Ok(grids)
}
fn adapt_legacy_sglang_wideep_moe(
sources: &[PerfSource],
inference_phase: &str,
grids: &mut MoeEpGrids,
) -> 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 moe_dtype_col = reader.col("moe_dtype")?;
let distribution_col = reader.col("distribution")?;
let topk_col = reader.col("topk")?;
let num_experts_col = reader.col("num_experts")?;
let hidden_size_col = reader.col("hidden_size")?;
let inter_size_col = reader.col_optional("inter_size");
let moe_tp_size_col = reader.col_optional("moe_tp_size");
let moe_ep_size_col = reader.col("moe_ep_size")?;
let num_tokens_col = reader.col("num_tokens")?;
let latency_col = reader.col("latency")?;
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 num_experts = row.u32(num_experts_col)?;
let key = MoeExpertComputeKey {
kernel_source: SGLANG_ADAPTED_KERNEL_SOURCE.to_string(),
quant: row.str_owned(moe_dtype_col)?,
distribution: row.str_owned(distribution_col)?,
inference_phase: inference_phase.to_string(),
topk: row.u32(topk_col)?,
num_experts,
num_slots: num_experts,
hidden_size: row.u32(hidden_size_col)?,
inter_size: row.u32_optional(inter_size_col)?.unwrap_or(0),
moe_tp_size: row.u32_optional(moe_tp_size_col)?.unwrap_or(1),
moe_ep_size: row.u32(moe_ep_size_col)?,
};
grids.store_keep_first(key, row.u32(num_tokens_col)?, row.f64(latency_col)?);
}
}
Ok(any_source)
}
fn adapt_legacy_trtllm_wideep_moe(
sources: &[PerfSource],
grids: &mut MoeEpGrids,
) -> 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 ks_col = reader.col_optional("kernel_source");
let moe_dtype_col = reader.col("moe_dtype")?;
let distribution_col = reader.col("distribution")?;
let topk_col = reader.col("topk")?;
let num_experts_col = reader.col("num_experts")?;
let num_slots_col = reader.col("num_slots")?;
let hidden_size_col = reader.col("hidden_size")?;
let inter_size_col = reader.col("inter_size")?;
let moe_tp_size_col = reader.col("moe_tp_size")?;
let moe_ep_size_col = reader.col("moe_ep_size")?;
let num_tokens_col = reader.col("num_tokens")?;
let latency_col = reader.col("latency")?;
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.to_string(),
Some(_) => row.str_optional(ks_col)?.unwrap_or("").to_string(),
};
let quant = row.str_owned(moe_dtype_col)?;
let distribution = row.str_owned(distribution_col)?;
let topk = row.u32(topk_col)?;
let num_experts = row.u32(num_experts_col)?;
let num_slots = row.u32(num_slots_col)?;
let hidden_size = row.u32(hidden_size_col)?;
let inter_size = row.u32(inter_size_col)?;
let moe_tp_size = row.u32(moe_tp_size_col)?;
let moe_ep_size = row.u32(moe_ep_size_col)?;
let num_tokens = row.u32(num_tokens_col)?;
let latency_ms = row.f64(latency_col)?;
for inference_phase in ["context", "generation"] {
let key = MoeExpertComputeKey {
kernel_source: kernel_source.clone(),
quant: quant.clone(),
distribution: distribution.clone(),
inference_phase: inference_phase.to_string(),
topk,
num_experts,
num_slots,
hidden_size,
inter_size,
moe_tp_size,
moe_ep_size,
};
grids.store_keep_first(key, num_tokens, latency_ms);
}
}
}
Ok(any_source)
}
fn load_new_schema(sources: &[PerfSource], grids: &mut MoeEpGrids) -> Result<bool, AicError> {
let mut any_source = false;
let mut seen: BTreeSet<(MoeExpertComputeKey, u32)> = BTreeSet::new();
for source in sources {
let path = source.path();
if !path.exists() {
continue;
}
any_source = true;
let reader = PerfReader::open(path)?;
let ks_col = reader.col("kernel_source")?;
let moe_dtype_col = reader.col("moe_dtype")?;
let distribution_col = reader.col("distribution")?;
let inference_phase_col = reader.col("inference_phase")?;
let topk_col = reader.col("topk")?;
let num_experts_col = reader.col("num_experts")?;
let num_slots_col = reader.col("num_slots")?;
let hidden_size_col = reader.col("hidden_size")?;
let inter_size_col = reader.col("inter_size")?;
let moe_tp_size_col = reader.col("moe_tp_size")?;
let moe_ep_size_col = reader.col("moe_ep_size")?;
let num_tokens_col = reader.col("num_tokens")?;
let latency_col = reader.col("latency")?;
for row in reader.rows()? {
let row = row?;
if !kernel_source_ok(source.kernel_sources(), Some(ks_col), &row)? {
continue;
}
let key = MoeExpertComputeKey {
kernel_source: row.str_owned(ks_col)?,
quant: row.str_owned(moe_dtype_col)?,
distribution: row.str_owned(distribution_col)?,
inference_phase: row.str_owned(inference_phase_col)?,
topk: row.u32(topk_col)?,
num_experts: row.u32(num_experts_col)?,
num_slots: row.u32(num_slots_col)?,
hidden_size: row.u32(hidden_size_col)?,
inter_size: row.u32(inter_size_col)?,
moe_tp_size: row.u32(moe_tp_size_col)?,
moe_ep_size: row.u32(moe_ep_size_col)?,
};
let num_tokens = row.u32(num_tokens_col)?;
let latency_ms = row.f64(latency_col)?;
grids.note_distribution(&key);
if seen.insert((key.clone(), num_tokens)) {
grids
.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 crate::common::system_spec::{GpuSpec, MiscSpec, NodeSpec};
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 test_spec() -> SystemSpec {
SystemSpec {
data_dir: PathBuf::from("data/synthetic"),
gpu: GpuSpec {
mem_bw: 1e9,
mem_bw_empirical_scaling_factor: 1.0,
mem_empirical_constant_latency: 0.0,
mem_capacity: None,
bfloat16_tc_flops: Some(1e12),
int8_tc_flops: None,
fp8_tc_flops: None,
fp4_tc_flops: None,
power: None,
sm_version: None,
},
node: NodeSpec {
num_gpus_per_node: 8,
inter_node_bw: 100.0,
intra_node_bw: 900.0,
pcie_bw: None,
p2p_latency: 0.0,
num_gpus_per_rack: None,
inter_rack_bw: None,
},
misc: MiscSpec::default(),
}
}
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 EpRow {
kernel_source: &'static str,
moe_dtype: &'static str,
distribution: &'static str,
inference_phase: &'static str,
num_slots: i64,
moe_ep_size: i64,
num_tokens: i64,
latency_ms: f64,
}
#[allow(clippy::too_many_arguments)]
fn ep_row(
kernel_source: &'static str,
moe_dtype: &'static str,
distribution: &'static str,
inference_phase: &'static str,
num_slots: i64,
moe_ep_size: i64,
num_tokens: i64,
latency_ms: f64,
) -> EpRow {
EpRow {
kernel_source,
moe_dtype,
distribution,
inference_phase,
num_slots,
moe_ep_size,
num_tokens,
latency_ms,
}
}
fn write_moe_ep_parquet(path: &Path, rows: &[EpRow]) {
let schema = Arc::new(
parse_message_type(
"message ep {
REQUIRED BYTE_ARRAY kernel_source (UTF8);
REQUIRED BYTE_ARRAY moe_dtype (UTF8);
REQUIRED BYTE_ARRAY distribution (UTF8);
REQUIRED BYTE_ARRAY inference_phase (UTF8);
REQUIRED INT64 topk;
REQUIRED INT64 num_experts;
REQUIRED INT64 num_slots;
REQUIRED INT64 hidden_size;
REQUIRED INT64 inter_size;
REQUIRED INT64 moe_tp_size;
REQUIRED INT64 moe_ep_size;
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.kernel_source))
.collect::<Vec<_>>(),
);
write_column::<ByteArrayType>(
&mut rg,
&rows
.iter()
.map(|r| ByteArray::from(r.moe_dtype))
.collect::<Vec<_>>(),
);
write_column::<ByteArrayType>(
&mut rg,
&rows
.iter()
.map(|r| ByteArray::from(r.distribution))
.collect::<Vec<_>>(),
);
write_column::<ByteArrayType>(
&mut rg,
&rows
.iter()
.map(|r| ByteArray::from(r.inference_phase))
.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.num_slots).collect::<Vec<_>>(),
);
write_column::<Int64Type>(&mut rg, &vec![7168_i64; n]);
write_column::<Int64Type>(&mut rg, &vec![2048_i64; n]);
write_column::<Int64Type>(&mut rg, &vec![1_i64; n]);
write_column::<Int64Type>(
&mut rg,
&rows.iter().map(|r| r.moe_ep_size).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_ms).collect::<Vec<_>>(),
);
rg.close().unwrap();
writer.close().unwrap();
}
fn write_legacy_sglang_parquet(
path: &Path,
rows: &[(&'static str, i64, i64, f64)],
with_inter_and_tp: bool,
) {
let inter_decl = if with_inter_and_tp {
"REQUIRED INT64 inter_size;"
} else {
""
};
let tp_decl = if with_inter_and_tp {
"REQUIRED INT64 moe_tp_size;"
} else {
""
};
let schema = Arc::new(
parse_message_type(&format!(
"message sglang_wideep {{
REQUIRED BYTE_ARRAY kernel_source (UTF8);
REQUIRED BYTE_ARRAY moe_dtype (UTF8);
REQUIRED INT64 num_tokens;
REQUIRED INT64 hidden_size;
{inter_decl}
REQUIRED INT64 topk;
REQUIRED INT64 num_experts;
{tp_decl}
REQUIRED INT64 moe_ep_size;
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, &vec![ByteArray::from("deepepmoe"); n]);
write_column::<ByteArrayType>(&mut rg, &vec![ByteArray::from("fp8_block"); n]);
write_column::<Int64Type>(&mut rg, &rows.iter().map(|r| r.2).collect::<Vec<_>>());
write_column::<Int64Type>(&mut rg, &vec![7168_i64; n]);
if with_inter_and_tp {
write_column::<Int64Type>(&mut rg, &vec![2048_i64; n]);
}
write_column::<Int64Type>(&mut rg, &vec![8_i64; n]);
write_column::<Int64Type>(&mut rg, &vec![256_i64; n]);
if with_inter_and_tp {
write_column::<Int64Type>(&mut rg, &vec![1_i64; n]);
}
write_column::<Int64Type>(&mut rg, &rows.iter().map(|r| r.1).collect::<Vec<_>>());
write_column::<ByteArrayType>(
&mut rg,
&rows
.iter()
.map(|r| ByteArray::from(r.0))
.collect::<Vec<_>>(),
);
write_column::<DoubleType>(&mut rg, &rows.iter().map(|r| r.3).collect::<Vec<_>>());
rg.close().unwrap();
writer.close().unwrap();
}
fn write_legacy_trtllm_wideep_parquet(
path: &Path,
rows: &[(Option<&'static str>, &'static str, i64, i64, i64, f64)],
with_kernel_source_column: bool,
) {
let ks_decl = if with_kernel_source_column {
"OPTIONAL BYTE_ARRAY kernel_source (UTF8);"
} else {
""
};
let schema = Arc::new(
parse_message_type(&format!(
"message trtllm_wideep {{
{ks_decl}
REQUIRED BYTE_ARRAY moe_dtype (UTF8);
REQUIRED INT64 num_tokens;
REQUIRED INT64 hidden_size;
REQUIRED INT64 inter_size;
REQUIRED INT64 topk;
REQUIRED INT64 num_experts;
REQUIRED INT64 num_slots;
REQUIRED INT64 moe_tp_size;
REQUIRED INT64 moe_ep_size;
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();
if with_kernel_source_column {
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("nvfp4"); n]);
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![2048_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.2).collect::<Vec<_>>());
write_column::<Int64Type>(&mut rg, &vec![1_i64; n]);
write_column::<Int64Type>(&mut rg, &rows.iter().map(|r| r.3).collect::<Vec<_>>());
write_column::<ByteArrayType>(
&mut rg,
&rows
.iter()
.map(|r| ByteArray::from(r.1))
.collect::<Vec<_>>(),
);
write_column::<DoubleType>(&mut rg, &rows.iter().map(|r| r.5).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}"
);
}
#[allow(clippy::too_many_arguments)]
fn q(
table: &MoeExpertComputeTable,
kernel_source: &str,
quant: MoeQuantMode,
distribution: &str,
phase: &str,
num_slots: u32,
moe_ep_size: u32,
num_tokens: u32,
) -> Result<f64, AicError> {
table.query(
kernel_source,
quant,
distribution,
phase,
8,
256,
num_slots,
7168,
2048,
1,
moe_ep_size,
num_tokens,
true,
)
}
#[test]
fn new_schema_stores_ms_raw_and_keys_all_axes() {
let tmp = tempfile::tempdir().unwrap();
write_moe_ep_parquet(
&tmp.path().join("moe_expert_compute_perf.parquet"),
&[
ep_row(
"deepep_moe",
"fp8_block",
"uniform",
"context",
256,
16,
128,
0.25,
),
ep_row(
"deepep_moe",
"fp8_block",
"uniform",
"generation",
256,
16,
128,
0.5,
),
ep_row(
"deepgemm",
"fp8_block",
"uniform",
"context",
288,
16,
128,
0.75,
),
],
);
let table = MoeExpertComputeTable::new(tmp.path().to_path_buf(), test_spec());
approx(
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"context",
256,
16,
128,
)
.unwrap(),
0.25,
);
approx(
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"generation",
256,
16,
128,
)
.unwrap(),
0.5,
);
approx(
q(
&table,
"deepgemm",
MoeQuantMode::Fp8Block,
"uniform",
"context",
288,
16,
128,
)
.unwrap(),
0.75,
);
assert!(
q(
&table,
"deepgemm",
MoeQuantMode::Fp8Block,
"uniform",
"context",
256,
16,
128
)
.is_err()
);
assert!(
q(
&table,
"deepgemm",
MoeQuantMode::Fp8Block,
"uniform",
"generation",
288,
16,
128
)
.is_err()
);
assert!(
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"prefill",
256,
16,
128
)
.is_err()
);
}
#[test]
fn legacy_sglang_context_and_generation_adapters() {
let tmp = tempfile::tempdir().unwrap();
write_legacy_sglang_parquet(
&tmp.path().join("wideep_context_moe_perf.parquet"),
&[("uniform", 2, 32, 0.375)],
true,
);
write_legacy_sglang_parquet(
&tmp.path().join("wideep_generation_moe_perf.parquet"),
&[("uniform", 4, 2, 0.125)],
true,
);
let table = MoeExpertComputeTable::new(tmp.path().to_path_buf(), test_spec());
approx(
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"context",
256,
2,
32,
)
.unwrap(),
0.375,
);
approx(
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"generation",
256,
4,
2,
)
.unwrap(),
0.125,
);
assert!(
q(
&table,
"deepepmoe",
MoeQuantMode::Fp8Block,
"uniform",
"context",
256,
2,
32
)
.is_err()
);
assert!(
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"context",
288,
2,
32
)
.is_err()
);
assert!(
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"generation",
256,
2,
32
)
.is_err()
);
assert!(
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"context",
256,
4,
2
)
.is_err()
);
}
#[test]
fn legacy_sglang_missing_inter_and_tp_columns_default() {
let tmp = tempfile::tempdir().unwrap();
write_legacy_sglang_parquet(
&tmp.path().join("wideep_context_moe_perf.parquet"),
&[("uniform", 2, 32, 0.25)],
false,
);
let table = MoeExpertComputeTable::new(tmp.path().to_path_buf(), test_spec());
approx(
table
.query(
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"context",
8,
256,
256,
7168,
0, 1, 2,
32,
true,
)
.unwrap(),
0.25,
);
assert!(
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"context",
256,
2,
32
)
.is_err()
);
}
#[test]
fn legacy_sglang_duplicate_rows_first_wins() {
let tmp = tempfile::tempdir().unwrap();
write_legacy_sglang_parquet(
&tmp.path().join("wideep_context_moe_perf.parquet"),
&[("uniform", 2, 32, 0.1), ("uniform", 2, 32, 0.9)],
true,
);
let table = MoeExpertComputeTable::new(tmp.path().to_path_buf(), test_spec());
approx(
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"context",
256,
2,
32,
)
.unwrap(),
0.1,
);
}
#[test]
fn legacy_trtllm_registers_both_phases_and_passes_num_slots() {
let tmp = tempfile::tempdir().unwrap();
write_legacy_trtllm_wideep_parquet(
&tmp.path().join("wideep_moe_perf.parquet"),
&[(
Some("wideep_compute_cutlass"),
"power_law_1.01_eplb",
288,
2,
1,
0.0611904,
)],
true,
);
let table = MoeExpertComputeTable::new(tmp.path().to_path_buf(), test_spec());
for phase in ["context", "generation"] {
approx(
q(
&table,
"wideep_compute_cutlass",
MoeQuantMode::Nvfp4,
"power_law_1.01_eplb",
phase,
288,
2,
1,
)
.unwrap(),
0.0611904,
);
}
assert!(
q(
&table,
"wideep_compute_cutlass",
MoeQuantMode::Nvfp4,
"power_law_1.01_eplb",
"context",
256,
2,
1,
)
.is_err()
);
let tmp2 = tempfile::tempdir().unwrap();
write_legacy_sglang_parquet(
&tmp2.path().join("wideep_context_moe_perf.parquet"),
&[("uniform", 2, 32, 0.1)],
true,
);
write_legacy_trtllm_wideep_parquet(
&tmp2.path().join("wideep_moe_perf.parquet"),
&[
(Some("deepep_moe"), "uniform", 256, 2, 32, 0.7),
(Some("deepep_moe"), "uniform", 256, 2, 32, 0.8),
],
true,
);
let table2 = MoeExpertComputeTable::new(tmp2.path().to_path_buf(), test_spec());
approx(
q(
&table2,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"context",
256,
2,
32,
)
.unwrap(),
0.1,
);
for phase in ["context", "generation"] {
approx(
q(
&table2,
"deepep_moe",
MoeQuantMode::Nvfp4,
"uniform",
phase,
256,
2,
32,
)
.unwrap(),
0.7,
);
}
}
#[test]
fn legacy_trtllm_kernel_source_column_absent_defaults_null_is_empty() {
let tmp = tempfile::tempdir().unwrap();
write_legacy_trtllm_wideep_parquet(
&tmp.path().join("wideep_moe_perf.parquet"),
&[(None, "uniform", 256, 2, 1, 1.5)],
false,
);
let table = MoeExpertComputeTable::new(tmp.path().to_path_buf(), test_spec());
approx(
q(
&table,
"moe_torch_flow",
MoeQuantMode::Nvfp4,
"uniform",
"context",
256,
2,
1,
)
.unwrap(),
1.5,
);
let tmp2 = tempfile::tempdir().unwrap();
write_legacy_trtllm_wideep_parquet(
&tmp2.path().join("wideep_moe_perf.parquet"),
&[
(Some("wideep_compute_cutlass"), "uniform", 256, 2, 1, 2.5),
(None, "uniform", 256, 4, 1, 9.0),
],
true,
);
let table2 = MoeExpertComputeTable::new(tmp2.path().to_path_buf(), test_spec());
approx(
q(
&table2,
"wideep_compute_cutlass",
MoeQuantMode::Nvfp4,
"uniform",
"context",
256,
2,
1,
)
.unwrap(),
2.5,
);
approx(
q(
&table2,
"",
MoeQuantMode::Nvfp4,
"uniform",
"context",
256,
4,
1,
)
.unwrap(),
9.0,
);
assert!(
q(
&table2,
"moe_torch_flow",
MoeQuantMode::Nvfp4,
"uniform",
"context",
256,
4,
1
)
.is_err()
);
}
#[test]
fn new_schema_overwrites_legacy_and_repeats_keep_first() {
let tmp = tempfile::tempdir().unwrap();
write_legacy_sglang_parquet(
&tmp.path().join("wideep_context_moe_perf.parquet"),
&[("uniform", 2, 64, 0.1), ("uniform", 2, 128, 0.2)],
true,
);
write_moe_ep_parquet(
&tmp.path().join("moe_expert_compute_perf.parquet"),
&[
ep_row(
"deepep_moe",
"fp8_block",
"uniform",
"context",
256,
2,
64,
0.7,
),
ep_row(
"deepep_moe",
"fp8_block",
"uniform",
"context",
256,
2,
64,
0.9,
),
],
);
let table = MoeExpertComputeTable::new(tmp.path().to_path_buf(), test_spec());
approx(
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"context",
256,
2,
64,
)
.unwrap(),
0.7,
);
approx(
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"context",
256,
2,
128,
)
.unwrap(),
0.2,
);
}
#[test]
fn distribution_chain_requested_uniform_first_available_and_phase_scoping() {
let tmp = tempfile::tempdir().unwrap();
write_legacy_sglang_parquet(
&tmp.path().join("wideep_context_moe_perf.parquet"),
&[("zeta_dist", 2, 64, 0.111), ("alpha_dist", 2, 64, 0.222)],
true,
);
write_legacy_sglang_parquet(
&tmp.path().join("wideep_generation_moe_perf.parquet"),
&[("uniform", 2, 64, 0.333), ("gen_only", 2, 64, 0.444)],
true,
);
let table = MoeExpertComputeTable::new(tmp.path().to_path_buf(), test_spec());
let ctx = |dist: &str| {
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
dist,
"context",
256,
2,
64,
)
};
let generation = |dist: &str| {
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
dist,
"generation",
256,
2,
64,
)
};
approx(ctx("zeta_dist").unwrap(), 0.111);
approx(ctx("alpha_dist").unwrap(), 0.222);
approx(ctx("power_law").unwrap(), 0.111);
approx(ctx("gen_only").unwrap(), 0.111);
approx(generation("power_law").unwrap(), 0.333);
approx(generation("gen_only").unwrap(), 0.444);
let tmp2 = tempfile::tempdir().unwrap();
write_legacy_sglang_parquet(
&tmp2.path().join("wideep_generation_moe_perf.parquet"),
&[("uniform", 2, 64, 0.5)],
true,
);
let table2 = MoeExpertComputeTable::new(tmp2.path().to_path_buf(), test_spec());
assert!(
q(
&table2,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"context",
256,
2,
64
)
.is_err()
);
}
#[test]
fn distribution_fallback_does_not_rescue_a_missing_shape() {
let tmp = tempfile::tempdir().unwrap();
write_legacy_sglang_parquet(
&tmp.path().join("wideep_context_moe_perf.parquet"),
&[("uniform", 2, 64, 0.1), ("power_law_0.8", 4, 64, 0.2)],
true,
);
let table = MoeExpertComputeTable::new(tmp.path().to_path_buf(), test_spec());
assert!(
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"context",
256,
4,
64
)
.is_err()
);
approx(
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
"power_law_0.8",
"context",
256,
4,
64,
)
.unwrap(),
0.2,
);
}
#[test]
fn kernel_and_quant_are_exact_typed_misses() {
let tmp = tempfile::tempdir().unwrap();
write_legacy_sglang_parquet(
&tmp.path().join("wideep_context_moe_perf.parquet"),
&[("uniform", 2, 64, 0.1)],
true,
);
let table = MoeExpertComputeTable::new(tmp.path().to_path_buf(), test_spec());
match q(
&table,
"deepep_moe",
MoeQuantMode::Nvfp4,
"uniform",
"context",
256,
2,
64,
)
.unwrap_err()
{
AicError::PerfDatabase(msg) => assert!(msg.contains("quant"), "got: {msg}"),
other => panic!("unexpected error: {other:?}"),
}
assert!(
q(
&table,
"deepep_moe",
MoeQuantMode::Fp8,
"uniform",
"context",
256,
2,
64
)
.is_err()
);
match q(
&table,
"deepgemm",
MoeQuantMode::Fp8Block,
"uniform",
"context",
256,
2,
64,
)
.unwrap_err()
{
AicError::PerfDatabase(msg) => {
assert!(msg.contains("kernel_source"), "got: {msg}")
}
other => panic!("unexpected error: {other:?}"),
}
}
fn hand_sol_ms(tokens: u64) -> f64 {
let mem_int = 81920 * tokens + 44040192 * (4 * tokens).min(128);
(mem_int as f64) * 2.0 / 1e9 * 1000.0
}
#[test]
fn query_rejects_missing_or_invalid_bfloat16_flops() {
let tmp = tempfile::tempdir().unwrap();
write_moe_ep_parquet(
&tmp.path().join("moe_expert_compute_perf.parquet"),
&[ep_row(
"deepep_moe",
"bfloat16",
"uniform",
"context",
256,
2,
32,
0.5,
)],
);
for value in [None, Some(0.0), Some(f64::NAN)] {
let mut spec = test_spec();
spec.gpu.bfloat16_tc_flops = value;
let table = MoeExpertComputeTable::new(tmp.path().to_path_buf(), spec);
match q(
&table,
"deepep_moe",
MoeQuantMode::Bfloat16,
"uniform",
"context",
256,
2,
32,
) {
Err(AicError::MissingSystemFlops(message)) => {
assert!(message.contains("bfloat16_tc_flops"), "got: {message}");
}
other => panic!("expected MissingSystemFlops, got {other:?}"),
}
}
}
#[test]
fn token_curve_lerp_and_roofline_util_hold() {
let tmp = tempfile::tempdir().unwrap();
write_moe_ep_parquet(
&tmp.path().join("moe_expert_compute_perf.parquet"),
&[
ep_row(
"deepep_moe",
"bfloat16",
"uniform",
"context",
256,
2,
32,
0.5,
),
ep_row(
"deepep_moe",
"bfloat16",
"uniform",
"context",
256,
2,
64,
0.8,
),
],
);
let table = MoeExpertComputeTable::new(tmp.path().to_path_buf(), test_spec());
let at = |tokens: u32| {
q(
&table,
"deepep_moe",
MoeQuantMode::Bfloat16,
"uniform",
"context",
256,
2,
tokens,
)
.unwrap()
};
approx(at(32), 0.5);
approx(at(64), 0.8);
approx(at(48), 0.65);
approx(at(100), hand_sol_ms(100) / (hand_sol_ms(64) / 0.8));
approx(at(16), hand_sol_ms(16) / (hand_sol_ms(32) / 0.5));
}
#[test]
fn singleton_token_curve_underflow_is_a_typed_miss() {
let tmp = tempfile::tempdir().unwrap();
write_moe_ep_parquet(
&tmp.path().join("moe_expert_compute_perf.parquet"),
&[ep_row(
"deepep_moe",
"bfloat16",
"uniform",
"context",
256,
2,
64,
0.8,
)],
);
let table = MoeExpertComputeTable::new(tmp.path().to_path_buf(), test_spec());
let at = |tokens: u32| {
q(
&table,
"deepep_moe",
MoeQuantMode::Bfloat16,
"uniform",
"context",
256,
2,
tokens,
)
};
match at(32).unwrap_err() {
AicError::PerfDatabase(msg) => {
assert!(msg.contains("singleton"), "got: {msg}")
}
other => panic!("unexpected error: {other:?}"),
}
approx(at(64).unwrap(), 0.8);
approx(at(128).unwrap(), hand_sol_ms(128) / (hand_sol_ms(64) / 0.8));
}
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_expert_compute_perf.parquet",
"wideep_context_moe_perf.parquet",
"wideep_generation_moe_perf.parquet",
"wideep_moe_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_ep_matches_python_oracle() {
let oracle: serde_json::Value =
serde_json::from_str(include_str!("testdata/moe_expert_compute_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, MoeExpertComputeTable> = 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_ep_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(|| {
let system = sample["system"].as_str().expect("system");
let spec = SystemSpec::load(&systems.join(format!("{system}.yaml")))
.expect("system yaml must load");
MoeExpertComputeTable::new(data_root.clone(), spec)
});
let u32_of = |field: &str| {
u32::try_from(sample[field].as_u64().expect(field)).expect("fits in u32")
};
let quant: MoeQuantMode = serde_json::from_value(sample["quant"].clone())
.expect("quant must map to a MoeQuantMode");
let got = table
.query(
sample["kernel_source"].as_str().expect("kernel_source"),
quant,
sample["distribution"].as_str().expect("distribution"),
sample["inference_phase"].as_str().expect("inference_phase"),
u32_of("topk"),
u32_of("num_experts"),
u32_of("num_slots"),
u32_of("hidden_size"),
u32_of("inter_size"),
u32_of("moe_tp_size"),
u32_of("moe_ep_size"),
u32_of("num_tokens"),
true, )
.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_expert_compute oracle: {checked} samples, max relative error {max_rel:e}");
}
#[test]
fn missing_sources_are_a_typed_miss() {
let tmp = tempfile::tempdir().unwrap();
let table = MoeExpertComputeTable::new(tmp.path().to_path_buf(), test_spec());
match q(
&table,
"deepep_moe",
MoeQuantMode::Fp8Block,
"uniform",
"context",
256,
2,
64,
)
.unwrap_err()
{
AicError::PerfDatabase(_) | AicError::Io { .. } => {}
other => panic!("unexpected error: {other:?}"),
}
}
}