use std::collections::BTreeMap;
use std::path::PathBuf;
use std::sync::OnceLock;
use super::axis_curve::LeafAxisCurve;
use super::moe_index::{MoeIndex, MoeShapeKey};
use super::perf_interp::LeafValue;
use super::{SourceResolver, kernel_source_ok};
use crate::common::enums::MoeQuantMode;
use crate::common::error::AicError;
use crate::config::{PerfDbSources, PerfSource};
use crate::perf_database::parquet_loader::PerfReader;
pub struct MoeTable {
data_root: PathBuf,
moe_sources: Vec<PerfSource>,
moe: OnceLock<Result<LoadedMoeGrids, AicError>>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MoeKernel {
Standard,
LowLatency,
}
#[derive(Clone, Debug)]
pub struct MoeSiblingSlice {
pub topk: u32,
pub num_experts: u32,
pub hidden_size: u32,
pub inter_size: u32,
pub points: Vec<(u32, f64)>,
}
struct LoadedMoeGrids {
default: MoeGrids,
low_latency: MoeGrids,
}
struct MoeGrids {
index: MoeIndex<MoeShapeKey, LeafAxisCurve>,
quants_in_load_order: Vec<String>,
}
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord)]
struct MoeKey {
quant: String,
distribution: String,
topk: u32,
num_experts: u32,
hidden_size: u32,
inter_size: u32,
moe_tp_size: u32,
moe_ep_size: u32,
}
impl MoeKey {
fn from_shape(quant: &str, distribution: &str, shape: MoeShapeKey) -> Self {
Self {
quant: quant.to_string(),
distribution: distribution.to_string(),
topk: shape.topk,
num_experts: shape.num_experts,
hidden_size: shape.hidden_size,
inter_size: shape.inter_size,
moe_tp_size: shape.moe_tp_size,
moe_ep_size: shape.moe_ep_size,
}
}
}
impl MoeTable {
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_sources = resolver.sources_for("moe_perf.parquet", &data_root)?;
Ok(Self {
data_root,
moe_sources,
moe: OnceLock::new(),
})
}
#[allow(clippy::too_many_arguments)]
pub fn query(
&self,
num_tokens: u32,
hidden_size: u32,
inter_size: u32,
topk: u32,
num_experts: u32,
moe_tp_size: u32,
moe_ep_size: u32,
quant: MoeQuantMode,
workload_distribution: &str,
sol: &dyn Fn(f64) -> f64,
) -> Result<LeafValue, AicError> {
let loaded = self.load()?;
let grids = &loaded.default;
let quant_name = quant.name();
let shape = MoeShapeKey {
topk,
num_experts,
hidden_size,
inter_size,
moe_tp_size,
moe_ep_size,
};
let (dist, by_tokens) =
grids
.index
.resolve_uniform(quant_name, workload_distribution, &shape);
let by_tokens = by_tokens.ok_or_else(|| {
let key = MoeKey::from_shape(quant_name, dist, shape);
AicError::PerfDatabase(format!(
"MoE data missing for {key:?} at {}",
self.data_root.display()
))
})?;
if by_tokens.is_empty() {
let key = MoeKey::from_shape(quant_name, dist, shape);
return Err(AicError::PerfDatabase(format!(
"MoE data has no token points for {key:?} at {}",
self.data_root.display()
)));
}
if let Some(only) = by_tokens.singleton_underflow(num_tokens) {
let key = MoeKey::from_shape(quant_name, dist, shape);
return Err(AicError::PerfDatabase(format!(
"MoE silicon token underflow has only one measured point; cannot infer \
low-token latency from a singleton. num_tokens={num_tokens}, \
measured_token={only}, key={key:?}"
)));
}
by_tokens.query(num_tokens as f64, sol)
}
#[allow(clippy::too_many_arguments)]
pub fn query_low_latency(
&self,
num_tokens: u32,
hidden_size: u32,
inter_size: u32,
topk: u32,
num_experts: u32,
moe_tp_size: u32,
moe_ep_size: u32,
quant: MoeQuantMode,
workload_distribution: &str,
sol: &dyn Fn(f64) -> f64,
) -> Result<Option<LeafValue>, AicError> {
let loaded = self.load()?;
let grids = &loaded.low_latency;
if grids.index.is_empty() {
return Ok(None);
}
let quant_name = quant.name();
let shape = MoeShapeKey {
topk,
num_experts,
hidden_size,
inter_size,
moe_tp_size,
moe_ep_size,
};
let (dist, by_tokens) =
grids
.index
.resolve_uniform(quant_name, workload_distribution, &shape);
let Some(by_tokens) = by_tokens else {
return Ok(None);
};
if by_tokens.is_empty() {
return Ok(None);
}
if let Some(only) = by_tokens.singleton_underflow(num_tokens) {
let key = MoeKey::from_shape(quant_name, dist, shape);
return Err(AicError::PerfDatabase(format!(
"MoE low-latency token underflow has only one measured point; cannot infer \
low-token latency from a singleton. num_tokens={num_tokens}, \
measured_token={only}, key={key:?}"
)));
}
by_tokens.query(num_tokens as f64, sol).map(Some)
}
pub fn low_latency_available(&self) -> Result<bool, AicError> {
let loaded = self.load()?;
Ok(!loaded.low_latency.index.is_empty())
}
#[allow(clippy::too_many_arguments)]
pub fn slice_points(
&self,
kernel: MoeKernel,
quant_name: &str,
workload_distribution: &str,
topk: u32,
num_experts: u32,
hidden_size: u32,
inter_size: u32,
moe_tp_size: u32,
moe_ep_size: u32,
) -> Result<Vec<(u32, f64)>, AicError> {
let grids = self.grids_for(kernel)?;
let shape = MoeShapeKey {
topk,
num_experts,
hidden_size,
inter_size,
moe_tp_size,
moe_ep_size,
};
let (dist, by_tokens) =
grids
.index
.resolve_uniform(quant_name, workload_distribution, &shape);
let by_tokens = by_tokens.filter(|curve| !curve.is_empty()).ok_or_else(|| {
let key = MoeKey::from_shape(quant_name, dist, shape);
AicError::PerfDatabase(format!(
"MoE data missing for {key:?} ({kernel:?}) at {}",
self.data_root.display()
))
})?;
Ok(by_tokens
.iter()
.map(|(t, leaf)| (t, leaf.latency))
.collect())
}
pub fn sibling_slices(
&self,
kernel: MoeKernel,
quant_name: &str,
workload_distribution: &str,
moe_tp_size: u32,
moe_ep_size: u32,
) -> Result<Vec<MoeSiblingSlice>, AicError> {
let grids = self.grids_for(kernel)?;
let (_, by_shape) = grids
.index
.resolve_uniform_shapes(quant_name, workload_distribution);
let mut slices = Vec::new();
let Some(by_shape) = by_shape else {
return Ok(slices);
};
for (shape, curve) in by_shape {
if shape.moe_tp_size != moe_tp_size
|| shape.moe_ep_size != moe_ep_size
|| curve.is_empty()
{
continue;
}
slices.push(MoeSiblingSlice {
topk: shape.topk,
num_experts: shape.num_experts,
hidden_size: shape.hidden_size,
inter_size: shape.inter_size,
points: curve.iter().map(|(t, leaf)| (t, leaf.latency)).collect(),
});
}
Ok(slices)
}
pub fn available_quants(&self, kernel: MoeKernel) -> Result<Vec<String>, AicError> {
Ok(self.grids_for(kernel)?.quants_in_load_order.clone())
}
fn grids_for(&self, kernel: MoeKernel) -> Result<&MoeGrids, AicError> {
let loaded = self.load()?;
Ok(match kernel {
MoeKernel::Standard => &loaded.default,
MoeKernel::LowLatency => &loaded.low_latency,
})
}
fn load(&self) -> Result<&LoadedMoeGrids, AicError> {
let cell = self.moe.get_or_init(|| load_moe_parquet(&self.moe_sources));
cell.as_ref().map_err(clone_err)
}
}
pub(crate) fn moe_kernel_quant_rewrite(raw_quant: String, kernel_source: &str) -> String {
match (raw_quant.as_str(), kernel_source) {
("w4a8_mxfp4_mxfp8", "sglang_mxfp4_flashinfer_trtllm_moe") => {
"w4a8_mxfp4_mxfp8_trtllm".to_string()
}
("w4a16_mxfp4", "sglang_flashinfer_cutlass_moe") => "w4a16_mxfp4_cutlass".to_string(),
_ => raw_quant,
}
}
fn load_moe_parquet(sources: &[PerfSource]) -> Result<LoadedMoeGrids, AicError> {
let mut default_index: MoeIndex<MoeShapeKey, BTreeMap<u32, LeafValue>> = MoeIndex::default();
let mut low_latency_index: MoeIndex<MoeShapeKey, BTreeMap<u32, LeafValue>> =
MoeIndex::default();
let mut default_quants: Vec<String> = Vec::new();
let mut low_latency_quants: Vec<String> = Vec::new();
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 num_tokens_col = reader.col("num_tokens")?;
let hidden_size_col = reader.col("hidden_size")?;
let inter_size_col = reader.col("inter_size")?;
let topk_col = reader.col("topk")?;
let num_experts_col = reader.col("num_experts")?;
let moe_tp_size_col = reader.col("moe_tp_size")?;
let moe_ep_size_col = reader.col("moe_ep_size")?;
let distribution_col = reader.col("distribution")?;
let latency_col = reader.col("latency")?;
let power_col = reader.col_optional("power");
let kernel_source_col = reader.col_optional("kernel_source");
for row in reader.rows()? {
let row = row?;
if !kernel_source_ok(source.kernel_sources(), kernel_source_col, &row)? {
continue;
}
let kernel_source = row
.str_optional(kernel_source_col)?
.unwrap_or("")
.to_string();
let quant = moe_kernel_quant_rewrite(row.str_owned(moe_dtype_col)?, &kernel_source);
let distribution = row.str_owned(distribution_col)?;
let shape = MoeShapeKey {
topk: row.u32(topk_col)?,
num_experts: row.u32(num_experts_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 (target, target_quants) = if kernel_source == "moe_torch_flow_min_latency" {
(&mut low_latency_index, &mut low_latency_quants)
} else {
(&mut default_index, &mut default_quants)
};
if !target_quants.iter().any(|q| q == &quant) {
target_quants.push(quant.clone());
}
let latency = row.f64(latency_col)?;
let power = row.f64_optional(power_col)?.unwrap_or(0.0);
target
.entry(quant, distribution, shape)
.entry(row.u32(num_tokens_col)?)
.or_insert(LeafValue::with_power(latency, power));
}
}
if !any_source || (default_index.is_empty() && low_latency_index.is_empty()) {
return Err(AicError::PerfDatabase(format!(
"no rows loaded from {} source(s) (first: {})",
sources.len(),
sources
.first()
.map(|s| s.path().display().to_string())
.unwrap_or_default()
)));
}
Ok(LoadedMoeGrids {
default: MoeGrids {
index: default_index.map_values(|curve| LeafAxisCurve::from_map("num_tokens", curve)),
quants_in_load_order: default_quants,
},
low_latency: MoeGrids {
index: low_latency_index
.map_values(|curve| LeafAxisCurve::from_map("num_tokens", curve)),
quants_in_load_order: low_latency_quants,
},
})
}
fn clone_err(err: &AicError) -> AicError {
AicError::PerfDatabase(err.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
const REPO_ROOT_HINT: &str = env!("CARGO_MANIFEST_DIR");
fn b200_vllm_data_root() -> PathBuf {
PathBuf::from(REPO_ROOT_HINT)
.join("../..")
.join("python/aisimulate/src/aiconfigurator_core/systems/data/b200_sxm/vllm/0.19.0")
}
#[test]
fn moe_table_loads_b200_vllm() {
let table = MoeTable::new(b200_vllm_data_root());
let _ = table.load().expect("moe_perf.parquet must load");
}
fn proxy_sol(t: f64) -> f64 {
t
}
#[test]
fn moe_index_resolves_requested_and_uniform_distributions() {
let shape = MoeShapeKey {
topk: 2,
num_experts: 8,
hidden_size: 4096,
inter_size: 2048,
moe_tp_size: 1,
moe_ep_size: 4,
};
let mut index = MoeIndex::default();
*index.entry("fp8".into(), "power_law".into(), shape) = LeafAxisCurve::from_map(
"num_tokens",
BTreeMap::from([(1, LeafValue::latency_only(1.0))]),
);
*index.entry("fp8".into(), "uniform".into(), shape) = LeafAxisCurve::from_map(
"num_tokens",
BTreeMap::from([(1, LeafValue::latency_only(2.0))]),
);
let grids = MoeGrids {
index,
quants_in_load_order: vec!["fp8".to_string()],
};
let (dist, curve) = grids.index.resolve_uniform("fp8", "power_law", &shape);
assert_eq!(dist, "power_law");
assert_eq!(curve.unwrap().get(1).map(|leaf| leaf.latency), Some(1.0));
let (dist, curve) = grids.index.resolve_uniform("fp8", "missing", &shape);
assert_eq!(dist, "uniform");
assert_eq!(curve.unwrap().get(1).map(|leaf| leaf.latency), Some(2.0));
}
#[test]
fn moe_distribution_falls_back_to_uniform() {
let table = MoeTable::new(b200_vllm_data_root());
let result = table.query(
1024,
4096,
2048,
2,
128,
1,
8,
MoeQuantMode::Bfloat16,
"nonexistent_distribution",
&proxy_sol,
);
match result {
Ok(value) => assert!(value.latency > 0.0),
Err(AicError::PerfDatabase(msg)) => {
assert!(
!msg.contains("nonexistent_distribution"),
"expected uniform fallback, not literal distribution name in error: {msg}"
);
}
Err(other) => panic!("unexpected error: {other:?}"),
}
}
#[test]
fn moe_lazy_loads_once() {
let table = MoeTable::new(b200_vllm_data_root());
let r1 = table.load();
let r2 = table.load();
assert_eq!(r1.is_ok(), r2.is_ok());
}
#[test]
fn moe_low_latency_grid_split_routes_by_kernel_source() {
use crate::perf_database::energy_test_fixtures::{Col, write_parquet};
let with_min = tempfile::tempdir().expect("tmpdir");
write_parquet(
&with_min.path().join("moe_perf.parquet"),
&[
Col::Str("moe_dtype", vec!["bfloat16", "bfloat16"]),
Col::I64("num_tokens", vec![1024, 1024]),
Col::I64("hidden_size", vec![4096, 4096]),
Col::I64("inter_size", vec![2048, 2048]),
Col::I64("topk", vec![2, 2]),
Col::I64("num_experts", vec![8, 8]),
Col::I64("moe_tp_size", vec![1, 1]),
Col::I64("moe_ep_size", vec![1, 1]),
Col::Str("distribution", vec!["uniform", "uniform"]),
Col::Str(
"kernel_source",
vec!["moe_torch_flow", "moe_torch_flow_min_latency"],
),
Col::F64("latency", vec![1.0, 0.5]),
],
);
let table = MoeTable::new(with_min.path().to_path_buf());
assert!(
table
.low_latency_available()
.expect("moe_perf.parquet must load")
);
let without = tempfile::tempdir().expect("tmpdir");
write_parquet(
&without.path().join("moe_perf.parquet"),
&[
Col::Str("moe_dtype", vec!["bfloat16"]),
Col::I64("num_tokens", vec![1024]),
Col::I64("hidden_size", vec![4096]),
Col::I64("inter_size", vec![2048]),
Col::I64("topk", vec![2]),
Col::I64("num_experts", vec![8]),
Col::I64("moe_tp_size", vec![1]),
Col::I64("moe_ep_size", vec![1]),
Col::Str("distribution", vec!["uniform"]),
Col::Str("kernel_source", vec!["moe_torch_flow"]),
Col::F64("latency", vec![1.0]),
],
);
let plain = MoeTable::new(without.path().to_path_buf());
assert!(
!plain
.low_latency_available()
.expect("moe_perf.parquet must load")
);
}
#[test]
fn moe_energy_matches_python_oracle() {
use crate::perf_database::energy_test_fixtures::{Col, write_parquet};
let tmp = tempfile::tempdir().expect("tmpdir");
write_parquet(
&tmp.path().join("moe_perf.parquet"),
&[
Col::Str("moe_dtype", vec!["bfloat16", "bfloat16"]),
Col::I64("num_tokens", vec![1024, 2048]),
Col::I64("hidden_size", vec![4096, 4096]),
Col::I64("inter_size", vec![2048, 2048]),
Col::I64("topk", vec![2, 2]),
Col::I64("num_experts", vec![8, 8]),
Col::I64("moe_tp_size", vec![1, 1]),
Col::I64("moe_ep_size", vec![1, 1]),
Col::Str("distribution", vec!["uniform", "uniform"]),
Col::Str("kernel_source", vec!["moe_torch_flow", "moe_torch_flow"]),
Col::F64("latency", vec![1.0, 3.0]),
Col::F64("power", vec![100.0, 200.0]),
],
);
let table = MoeTable::new(tmp.path().to_path_buf());
let v = table
.query(
1536,
4096,
2048,
2,
8,
1,
1,
MoeQuantMode::Bfloat16,
"uniform",
&proxy_sol,
)
.unwrap();
assert!((v.latency - 2.0).abs() < 1e-9, "latency {}", v.latency);
assert!(
(v.energy - 300.0).abs() < 1e-9 * 300.0,
"energy {}",
v.energy
);
}
}