use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicU8, Ordering};
use std::sync::{Arc, Mutex, OnceLock, Weak};
use crate::common::enums::{DatabaseMode, TransferPolicy};
use crate::common::error::AicError;
use crate::common::system_spec::SystemSpec;
use crate::config::{PerfDbSources, PerfSource};
use crate::operators::util_empirical::{DeltaLookupCache, ProvenanceTier, UtilGridCache};
const KNOWN_BACKEND_DIRS: [&str; 5] = ["trtllm", "sglang", "vllm", "nccl", "oneccl"];
#[cfg(test)]
pub(crate) fn resolve_op_sources(
perf_db_sources: &PerfDbSources,
basename: &str,
data_root: &Path,
) -> Vec<PerfSource> {
SourceResolver::fixed(perf_db_sources.clone())
.sources_for(basename, data_root)
.expect("fixed-map resolution is infallible")
}
type SharedTablesMemo = Mutex<HashMap<String, Weak<PerfTables>>>;
fn shared_tables_key(
systems_root: &Path,
system: &str,
backend: &str,
version: &str,
resolver: &SourceResolver,
) -> String {
format!(
"{}\x1f{system}\x1f{backend}\x1f{version}\x1f{}",
systems_root.display(),
resolver.identity_key()
)
}
pub(crate) fn version_dir_is_unusable(version_dir: &Path) -> bool {
if version_dir.join("collection_meta.yaml").is_file() {
return false;
}
version_dir.join("INCOMPLETE.txt").is_file()
}
pub(crate) fn find_in_family_dirs(data_root: &Path, basename: &str) -> Option<PathBuf> {
let version = data_root.file_name()?.to_str()?;
let backend = data_root.parent()?.file_name()?.to_str()?;
let data_dir = data_root.parent()?.parent()?;
for entry in std::fs::read_dir(data_dir).ok()?.flatten() {
let name = entry.file_name();
let name = match name.to_str() {
Some(name) => name,
None => continue,
};
if name.starts_with('.') || KNOWN_BACKEND_DIRS.contains(&name) {
continue;
}
let version_dir = entry.path().join(backend).join(version);
if version_dir_is_unusable(&version_dir) {
continue;
}
let candidate = version_dir.join(basename);
if candidate.is_file() {
return Some(candidate);
}
}
None
}
fn has_family_backend_version(system_data_root: &Path, backend: &str, version: &str) -> bool {
let entries = match std::fs::read_dir(system_data_root) {
Ok(entries) => entries,
Err(_) => return false,
};
for entry in entries.flatten() {
let name = entry.file_name();
let name = match name.to_str() {
Some(name) => name,
None => continue,
};
if KNOWN_BACKEND_DIRS.contains(&name) || !entry.path().is_dir() {
continue;
}
if entry.path().join(backend).join(version).is_dir() {
return true;
}
}
false
}
fn comm_root(system_data_root: &Path, backend_dir: &str, version: &str) -> PathBuf {
let family_root = system_data_root
.join("comm")
.join(backend_dir)
.join(version);
if family_root.is_dir() && !version_dir_is_unusable(&family_root) {
family_root
} else {
system_data_root.join(backend_dir).join(version)
}
}
pub(crate) fn kernel_source_ok(
filter: Option<&[String]>,
ks_col: Option<usize>,
row: &parquet_loader::PerfRow,
) -> Result<bool, AicError> {
match filter {
None => Ok(true),
Some(allow) => match row.str_optional(ks_col)? {
Some(ks) => Ok(allow.iter().any(|a| a == ks)),
None => Ok(false),
},
}
}
pub mod attention;
mod axis_curve;
pub mod communication;
pub mod dsa;
pub mod dsv4;
pub mod dsv4_megamoe;
pub mod fpm_forward;
pub mod gemm;
mod interpolation;
pub mod mhc;
pub mod mla;
pub mod moe;
pub mod moe_a2a;
pub mod moe_expert_compute;
mod moe_index;
pub mod msa;
pub mod parquet_loader;
pub mod perf_interp;
pub mod source_resolution;
pub mod state_space;
pub mod table_view;
pub mod trtllm_alltoall;
pub mod wideep_mla;
pub use attention::AttentionTable;
pub use communication::CommunicationTable;
pub use dsa::DsaTable;
#[allow(unused_imports)]
pub use dsv4::{AttnKind, Dsv4Table};
pub use dsv4_megamoe::Dsv4MegaMoeTable;
pub use fpm_forward::FpmForwardTable;
pub use gemm::GemmTable;
pub use mhc::MhcTable;
pub use mla::MlaTable;
pub use moe::MoeTable;
pub use moe_a2a::MoeA2aTable;
pub use moe_expert_compute::MoeExpertComputeTable;
pub use msa::MsaTable;
#[allow(unused_imports)]
pub use source_resolution::{ResolveCtx, SourceResolver, resolve_one};
pub use state_space::StateSpaceTable;
pub use trtllm_alltoall::TrtllmAlltoallTable;
pub use wideep_mla::WideEpMlaTable;
pub struct PerfTables {
pub system: String,
pub backend: String,
pub version: String,
pub system_spec: SystemSpec,
pub data_root: PathBuf,
pub gemm: GemmTable,
pub attention: AttentionTable,
pub mla: MlaTable,
pub moe: MoeTable,
pub moe_a2a: MoeA2aTable,
pub moe_expert_compute: MoeExpertComputeTable,
pub communication: CommunicationTable,
pub dsa: DsaTable,
pub msa: MsaTable,
pub dsv4: Dsv4Table,
pub dsv4_megamoe: Dsv4MegaMoeTable,
pub mhc: MhcTable,
pub trtllm_alltoall: TrtllmAlltoallTable,
pub wideep_mla: WideEpMlaTable,
pub state_space: StateSpaceTable,
pub fpm_forward: FpmForwardTable,
pub source_resolver: Arc<SourceResolver>,
}
pub struct PerfDatabase {
tables: Arc<PerfTables>,
pub database_mode: DatabaseMode,
pub transfer_policy: TransferPolicy,
pub util_grids: Arc<UtilGridCache>,
pub delta_lookups: Arc<DeltaLookupCache>,
provenance: Arc<AtomicU8>,
}
impl std::ops::Deref for PerfDatabase {
type Target = PerfTables;
fn deref(&self) -> &PerfTables {
&self.tables
}
}
impl PerfDatabase {
#[cfg(test)]
pub(crate) fn set_fpm_forward_for_test(&mut self, table: FpmForwardTable) {
Arc::get_mut(&mut self.tables)
.expect("test fixture db must be uniquely owned")
.fpm_forward = table;
}
pub fn load(
systems_root: &Path,
system: &str,
backend: &str,
version: &str,
) -> Result<Self, AicError> {
Self::load_with_sources(
systems_root,
system,
backend,
version,
&PerfDbSources::default(),
)
}
pub fn load_with_sources(
systems_root: &Path,
system: &str,
backend: &str,
version: &str,
perf_db_sources: &PerfDbSources,
) -> Result<Self, AicError> {
Self::load_with_sources_opts(
systems_root,
system,
backend,
version,
perf_db_sources,
false,
)
}
pub fn load_with_sources_opts(
systems_root: &Path,
system: &str,
backend: &str,
version: &str,
perf_db_sources: &PerfDbSources,
tolerate_missing_data: bool,
) -> Result<Self, AicError> {
Self::load_with_resolver(
systems_root,
system,
backend,
version,
Arc::new(SourceResolver::fixed(perf_db_sources.clone())),
tolerate_missing_data,
)
}
pub fn resolve_ctx(
systems_root: &Path,
system: &str,
backend: &str,
version: &str,
enable_shared_layer: bool,
strict_provenance: bool,
) -> Result<ResolveCtx, AicError> {
let system_yaml = systems_root.join(format!("{system}.yaml"));
let spec = SystemSpec::load(&system_yaml)?;
Ok(ResolveCtx {
systems_root: systems_root.to_path_buf(),
system_data_root: systems_root.join(&spec.data_dir),
backend: backend.to_string(),
version: version.to_string(),
enable_shared_layer,
strict: strict_provenance,
})
}
pub fn load_resolved(
systems_root: &Path,
system: &str,
backend: &str,
version: &str,
enable_shared_layer: bool,
strict_provenance: bool,
tolerate_missing_data: bool,
) -> Result<Self, AicError> {
let ctx = Self::resolve_ctx(
systems_root,
system,
backend,
version,
enable_shared_layer,
strict_provenance,
)?;
Self::load_with_resolver(
systems_root,
system,
backend,
version,
Arc::new(SourceResolver::live(ctx)),
tolerate_missing_data,
)
}
fn load_with_resolver(
systems_root: &Path,
system: &str,
backend: &str,
version: &str,
resolver: Arc<SourceResolver>,
tolerate_missing_data: bool,
) -> Result<Self, AicError> {
let system_yaml = systems_root.join(format!("{system}.yaml"));
let spec = SystemSpec::load(&system_yaml)?;
let system_data_root = systems_root.join(&spec.data_dir);
let data_root = system_data_root.join(backend).join(version);
if !tolerate_missing_data
&& !data_root.is_dir()
&& !has_family_backend_version(&system_data_root, backend, version)
{
return Err(AicError::PerfDatabase(format!(
"perf data directory not found in either the legacy layout ({}) or a family-first layout \
(<family>/{backend}/{version} under {}) (system={system}, backend={backend}, version={version})",
data_root.display(),
system_data_root.display()
)));
}
let nccl_root = spec
.misc
.nccl_version
.as_ref()
.map(|v| comm_root(&system_data_root, "nccl", v));
let oneccl_root = spec
.misc
.oneccl_version
.as_ref()
.map(|v| comm_root(&system_data_root, "oneccl", v));
let tables = PerfTables {
system: system.to_string(),
backend: backend.to_string(),
version: version.to_string(),
gemm: GemmTable::with_sources(data_root.clone(), spec.clone(), &resolver)?,
attention: AttentionTable::with_sources(data_root.clone(), spec.clone(), &resolver)?,
mla: MlaTable::with_sources(data_root.clone(), spec.clone(), &resolver)?,
moe: MoeTable::with_sources(data_root.clone(), &resolver)?,
moe_a2a: MoeA2aTable::with_sources(data_root.clone(), &resolver)?,
moe_expert_compute: MoeExpertComputeTable::with_sources(
data_root.clone(),
spec.clone(),
&resolver,
)?,
communication: CommunicationTable::with_sources(
data_root.clone(),
nccl_root,
oneccl_root,
&resolver,
)?,
dsa: DsaTable::with_sources(data_root.clone(), &resolver)?,
dsv4: Dsv4Table::with_sources(data_root.clone(), &resolver)?,
dsv4_megamoe: Dsv4MegaMoeTable::with_primary(
resolver
.sources_for("dsv4_megamoe_module_perf.parquet", &data_root)?
.into_iter()
.next()
.map(|PerfSource(path, _)| path)
.unwrap_or_else(|| data_root.join("dsv4_megamoe_module_perf.parquet")),
),
msa: MsaTable::with_sources(data_root.clone(), &resolver)?,
mhc: MhcTable::with_sources(data_root.clone(), &resolver)?,
trtllm_alltoall: TrtllmAlltoallTable::with_sources(data_root.clone(), &resolver)?,
wideep_mla: WideEpMlaTable::with_sources(data_root.clone(), spec.clone(), &resolver)?,
state_space: StateSpaceTable::with_sources(
data_root.clone(),
backend,
version,
spec.gpu.sm_version,
&resolver,
)?,
fpm_forward: FpmForwardTable::new(data_root.clone(), system, backend, version),
system_spec: spec,
source_resolver: resolver,
data_root,
};
Ok(Self::from_tables(Arc::new(tables)))
}
pub fn load_resolved_shared(
systems_root: &Path,
system: &str,
backend: &str,
version: &str,
enable_shared_layer: bool,
strict_provenance: bool,
tolerate_missing_data: bool,
) -> Result<Self, AicError> {
static SHARED_TABLES: OnceLock<SharedTablesMemo> = OnceLock::new();
Self::load_resolved_shared_in(
SHARED_TABLES.get_or_init(Default::default),
systems_root,
system,
backend,
version,
enable_shared_layer,
strict_provenance,
tolerate_missing_data,
)
}
#[allow(clippy::too_many_arguments)]
fn load_resolved_shared_in(
memo: &SharedTablesMemo,
systems_root: &Path,
system: &str,
backend: &str,
version: &str,
enable_shared_layer: bool,
strict_provenance: bool,
tolerate_missing_data: bool,
) -> Result<Self, AicError> {
if tolerate_missing_data {
return Self::load_resolved(
systems_root,
system,
backend,
version,
enable_shared_layer,
strict_provenance,
true,
);
}
let ctx = Self::resolve_ctx(
systems_root,
system,
backend,
version,
enable_shared_layer,
strict_provenance,
)?;
let resolver = Arc::new(SourceResolver::live(ctx));
let key = shared_tables_key(systems_root, system, backend, version, &resolver);
if let Some(tables) = memo.lock().unwrap().get(&key).and_then(Weak::upgrade) {
return Ok(Self::from_tables(tables));
}
let db = Self::load_with_resolver(systems_root, system, backend, version, resolver, false)?;
let mut map = memo.lock().unwrap();
map.retain(|_, weak| weak.strong_count() > 0);
map.insert(key, Arc::downgrade(&db.tables));
Ok(db)
}
fn from_tables(tables: Arc<PerfTables>) -> Self {
Self {
tables,
database_mode: DatabaseMode::default(),
transfer_policy: TransferPolicy::ALL,
util_grids: Arc::new(UtilGridCache::new()),
delta_lookups: Arc::new(DeltaLookupCache::new()),
provenance: Arc::new(AtomicU8::new(ProvenanceTier::Silicon as u8)),
}
}
#[cfg(test)]
pub(crate) fn tables_arc(&self) -> &Arc<PerfTables> {
&self.tables
}
pub fn note_provenance(&self, tier: ProvenanceTier) {
self.provenance.fetch_max(tier as u8, Ordering::Relaxed);
}
pub fn reset_provenance(&self) {
self.provenance
.store(ProvenanceTier::Silicon as u8, Ordering::Relaxed);
}
pub fn worst_provenance(&self) -> ProvenanceTier {
ProvenanceTier::from_rank(self.provenance.load(Ordering::Relaxed))
}
pub fn with_mode(
mut self,
database_mode: DatabaseMode,
transfer_policy: TransferPolicy,
) -> Self {
self.database_mode = database_mode;
self.transfer_policy = transfer_policy;
self
}
#[cfg(test)]
pub(crate) fn tables_mut(&mut self) -> &mut PerfTables {
Arc::get_mut(&mut self.tables).expect("tables Arc must be unique for test mutation")
}
pub fn silicon_view(&self) -> PerfDatabase {
PerfDatabase {
tables: Arc::clone(&self.tables),
database_mode: DatabaseMode::Silicon,
transfer_policy: self.transfer_policy,
util_grids: Arc::clone(&self.util_grids),
delta_lookups: Arc::clone(&self.delta_lookups),
provenance: Arc::clone(&self.provenance),
}
}
pub fn sol_full_view(&self) -> PerfDatabase {
PerfDatabase {
tables: Arc::clone(&self.tables),
database_mode: DatabaseMode::SolFull,
transfer_policy: self.transfer_policy,
util_grids: Arc::clone(&self.util_grids),
delta_lookups: Arc::clone(&self.delta_lookups),
provenance: Arc::clone(&self.provenance),
}
}
}
#[cfg(test)]
pub(crate) mod energy_test_fixtures {
use std::fs::File;
use std::path::Path;
use std::sync::Arc;
use parquet::data_type::{BoolType, ByteArray, ByteArrayType, DataType, DoubleType, Int64Type};
use parquet::file::properties::WriterProperties;
use parquet::file::writer::{SerializedFileWriter, SerializedRowGroupWriter};
use parquet::schema::parser::parse_message_type;
use crate::common::system_spec::{GpuSpec, MiscSpec, NodeSpec, SystemSpec};
pub(crate) enum Col {
Str(&'static str, Vec<&'static str>),
I64(&'static str, Vec<i64>),
F64(&'static str, Vec<f64>),
Bool(&'static str, Vec<bool>),
}
fn write_typed<T: DataType>(rg: &mut SerializedRowGroupWriter<'_, File>, values: &[T::T]) {
let mut col = rg.next_column().unwrap().expect("column");
col.typed::<T>().write_batch(values, None, None).unwrap();
col.close().unwrap();
}
pub(crate) fn write_parquet(path: &Path, cols: &[Col]) {
let fields: Vec<String> = cols
.iter()
.map(|c| match c {
Col::Str(name, _) => format!("REQUIRED BYTE_ARRAY {name} (UTF8);"),
Col::I64(name, _) => format!("REQUIRED INT64 {name};"),
Col::F64(name, _) => format!("REQUIRED DOUBLE {name};"),
Col::Bool(name, _) => format!("REQUIRED BOOLEAN {name};"),
})
.collect();
let schema = format!("message energy_fixture {{ {} }}", fields.join(" "));
let schema = Arc::new(parse_message_type(&schema).expect("schema must parse"));
let file = File::create(path).expect("create parquet");
let mut writer =
SerializedFileWriter::new(file, schema, Arc::new(WriterProperties::builder().build()))
.expect("writer");
let mut rg = writer.next_row_group().expect("row group");
for col in cols {
match col {
Col::Str(_, v) => {
let bytes: Vec<ByteArray> = v.iter().map(|s| ByteArray::from(*s)).collect();
write_typed::<ByteArrayType>(&mut rg, &bytes);
}
Col::I64(_, v) => write_typed::<Int64Type>(&mut rg, v),
Col::F64(_, v) => write_typed::<DoubleType>(&mut rg, v),
Col::Bool(_, v) => write_typed::<BoolType>(&mut rg, v),
}
}
rg.close().unwrap();
writer.close().unwrap();
}
pub(crate) fn energy_test_spec() -> SystemSpec {
SystemSpec {
data_dir: "data".into(),
gpu: GpuSpec {
mem_bw: 7.7e12,
mem_bw_empirical_scaling_factor: 0.92,
mem_empirical_constant_latency: 2e-6,
mem_capacity: None,
bfloat16_tc_flops: Some(2.25e15),
int8_tc_flops: Some(4.5e15),
fp8_tc_flops: Some(4.5e15),
fp4_tc_flops: Some(9e15),
power: None,
sm_version: Some(100),
},
node: NodeSpec {
num_gpus_per_node: 8,
intra_node_bw: 900e9,
inter_node_bw: 50e9,
pcie_bw: None,
p2p_latency: 2e-6,
num_gpus_per_rack: None,
inter_rack_bw: None,
},
misc: MiscSpec::default(),
}
}
pub(crate) fn write_energy_systems_root(root: &Path) -> std::path::PathBuf {
let yaml = r#"
data_dir: data
gpu:
mem_bw: 7.7e12
mem_bw_empirical_scaling_factor: 0.92
mem_empirical_constant_latency: 2.0e-6
bfloat16_tc_flops: 2.25e15
int8_tc_flops: 4.5e15
fp8_tc_flops: 4.5e15
fp4_tc_flops: 9.0e15
sm_version: 100
node:
num_gpus_per_node: 8
intra_node_bw: 900.0e9
inter_node_bw: 50.0e9
p2p_latency: 2.0e-6
misc:
nccl_version: test
"#;
std::fs::write(root.join("testsys.yaml"), yaml).expect("write yaml");
let data = root.join("data").join("vllm").join("1.0");
std::fs::create_dir_all(&data).expect("mkdir data");
data
}
}
#[cfg(test)]
mod tests {
use super::*;
const REPO_ROOT_HINT: &str = env!("CARGO_MANIFEST_DIR");
fn systems_root() -> PathBuf {
PathBuf::from(REPO_ROOT_HINT)
.join("../..")
.join("python/aisimulate/src/aiconfigurator_core/systems")
}
#[test]
fn load_b200_sxm_vllm_database() {
let db = PerfDatabase::load(&systems_root(), "b200_sxm", "vllm", "0.24.0")
.expect("b200_sxm/vllm/0.24.0 must load");
assert_eq!(db.system, "b200_sxm");
assert_eq!(db.backend, "vllm");
assert_eq!(db.version, "0.24.0");
let gemm_sources = resolve_op_sources(
&PerfDbSources::default(),
"gemm_perf.parquet",
&db.data_root,
);
assert_eq!(gemm_sources.len(), 1);
assert!(
gemm_sources[0].0.is_file(),
"resolved GEMM parquet must exist: {}",
gemm_sources[0].0.display()
);
}
#[test]
fn shared_load_reuses_tables_and_isolates_view_state() {
let memo = SharedTablesMemo::default();
let db1 = PerfDatabase::load_resolved_shared_in(
&memo,
&systems_root(),
"b200_sxm",
"vllm",
"0.24.0",
false,
false,
false,
)
.expect("shared load must succeed");
let db2 = PerfDatabase::load_resolved_shared_in(
&memo,
&systems_root(),
"b200_sxm",
"vllm",
"0.24.0",
false,
false,
false,
)
.expect("shared load must succeed");
assert!(
Arc::ptr_eq(&db1.tables, &db2.tables),
"same load identity must share the parsed tables"
);
db1.note_provenance(ProvenanceTier::Empirical);
assert_eq!(db2.worst_provenance(), ProvenanceTier::Silicon);
assert!(!Arc::ptr_eq(&db1.util_grids, &db2.util_grids));
assert!(!Arc::ptr_eq(&db1.delta_lookups, &db2.delta_lookups));
}
#[test]
fn missing_data_dir_is_tolerated_only_when_requested() {
let strict = PerfDatabase::load_with_sources_opts(
&systems_root(),
"h100_pcie",
"trtllm",
"estimate",
&PerfDbSources::default(),
false,
);
assert!(
strict
.err()
.map(|e| e.to_string().contains("perf data directory not found"))
.unwrap_or(false),
"strict load of a data-less tuple must raise the missing-directory error"
);
let tolerant = PerfDatabase::load_with_sources_opts(
&systems_root(),
"h100_pcie",
"trtllm",
"estimate",
&PerfDbSources::default(),
true,
)
.expect("tolerant load must succeed from the spec yaml alone");
assert!(
tolerant
.gemm
.query(
crate::common::enums::GemmQuantMode::Bfloat16,
64,
4096,
4096
)
.is_err()
);
}
#[test]
fn shared_load_distinct_policies_load_fresh_tables() {
let db_no_shared = PerfDatabase::load_resolved_shared(
&systems_root(),
"b200_sxm",
"vllm",
"0.24.0",
false,
false,
false,
)
.expect("shared load must succeed");
let db_shared = PerfDatabase::load_resolved_shared(
&systems_root(),
"b200_sxm",
"vllm",
"0.24.0",
true,
false,
false,
)
.expect("shared load must succeed");
assert!(
!Arc::ptr_eq(&db_no_shared.tables, &db_shared.tables),
"a different shared-layer policy is a different load identity"
);
let db_map_a = PerfDatabase::load_with_sources(
&systems_root(),
"b200_sxm",
"vllm",
"0.24.0",
&PerfDbSources::default(),
)
.expect("map load must succeed");
let db_map_b = PerfDatabase::load_with_sources(
&systems_root(),
"b200_sxm",
"vllm",
"0.24.0",
&PerfDbSources::default(),
)
.expect("map load must succeed");
assert!(
!Arc::ptr_eq(&db_map_a.tables, &db_map_b.tables),
"explicit-map loads bypass the shared-tables memo"
);
}
#[test]
fn shared_tables_are_dropped_with_their_last_view() {
let tmp = tempfile::tempdir().unwrap();
energy_test_fixtures::write_energy_systems_root(tmp.path());
let db = PerfDatabase::load_resolved_shared(
tmp.path(),
"testsys",
"vllm",
"1.0",
false,
false,
false,
)
.expect("fixture load must succeed");
let weak = Arc::downgrade(&db.tables);
drop(db);
assert!(
weak.upgrade().is_none(),
"the memo must hold Weak refs only — tables die with their last view"
);
}
#[test]
fn provenance_cell_accumulates_worst_tier_and_is_shared_with_views() {
let db = PerfDatabase::load(&systems_root(), "b200_sxm", "vllm", "0.24.0")
.expect("b200_sxm/vllm/0.24.0 must load");
assert_eq!(db.worst_provenance(), ProvenanceTier::Silicon);
db.note_provenance(ProvenanceTier::XShape);
db.note_provenance(ProvenanceTier::Empirical);
assert_eq!(db.worst_provenance(), ProvenanceTier::XShape);
let view = db.silicon_view();
view.note_provenance(ProvenanceTier::XOp);
assert_eq!(db.worst_provenance(), ProvenanceTier::XOp);
db.reset_provenance();
assert_eq!(db.worst_provenance(), ProvenanceTier::Silicon);
assert_eq!(view.worst_provenance(), ProvenanceTier::Silicon);
}
#[test]
fn load_unknown_version_errors() {
match PerfDatabase::load(&systems_root(), "b200_sxm", "vllm", "99.99.99") {
Err(AicError::PerfDatabase(_)) => {}
Ok(_) => panic!("expected load to fail for missing version"),
Err(other) => panic!("expected PerfDatabase error, got {other:?}"),
}
}
#[test]
fn resolve_op_sources_falls_back_to_family_dir_when_legacy_file_missing() {
let tmp = tempfile::tempdir().unwrap();
let backend = "sglang";
let version = "0.5.14";
let legacy_data_root = tmp.path().join(backend).join(version);
std::fs::create_dir_all(&legacy_data_root).unwrap();
let family_data_root = tmp.path().join("gemm").join(backend).join(version);
std::fs::create_dir_all(&family_data_root).unwrap();
std::fs::write(family_data_root.join("gemm_perf.parquet"), b"stub").unwrap();
let sources = resolve_op_sources(
&PerfDbSources::default(),
"gemm_perf.parquet",
&legacy_data_root,
);
assert_eq!(sources.len(), 1);
assert_eq!(sources[0].0, family_data_root.join("gemm_perf.parquet"));
assert!(sources[0].1.is_none());
}
#[test]
fn resolve_op_sources_skips_known_backend_dirs_when_scanning_families() {
let tmp = tempfile::tempdir().unwrap();
let backend = "nccl";
let version = "2.19";
let legacy_data_root = tmp.path().join(backend).join(version);
std::fs::create_dir_all(&legacy_data_root).unwrap();
let decoy = tmp.path().join("vllm").join(backend).join(version);
std::fs::create_dir_all(&decoy).unwrap();
std::fs::write(decoy.join("nccl_perf.parquet"), b"stub").unwrap();
let sources = resolve_op_sources(
&PerfDbSources::default(),
"nccl_perf.parquet",
&legacy_data_root,
);
assert_eq!(sources.len(), 1);
assert_eq!(sources[0].0, legacy_data_root.join("nccl_perf.parquet"));
}
fn write_synthetic_system_yaml(systems_root: &Path, system: &str) {
let yaml = "\
data_dir: data
gpu:
mem_bw: 1000000000000
node:
num_gpus_per_node: 8
inter_node_bw: 100000000000
intra_node_bw: 900000000000
";
std::fs::write(systems_root.join(format!("{system}.yaml")), yaml).unwrap();
}
#[test]
fn load_succeeds_on_family_only_layout() {
let tmp = tempfile::tempdir().unwrap();
let systems_root = tmp.path();
let (system, backend, version) = ("synth_family", "vllm", "1.2.3");
write_synthetic_system_yaml(systems_root, system);
let family_dir = systems_root
.join("data")
.join("gemm")
.join(backend)
.join(version);
std::fs::create_dir_all(&family_dir).unwrap();
std::fs::write(family_dir.join("gemm_perf.parquet"), b"stub").unwrap();
let legacy_data_root = systems_root.join("data").join(backend).join(version);
assert!(
!legacy_data_root.is_dir(),
"fixture must not have a legacy dir"
);
let db = PerfDatabase::load(systems_root, system, backend, version)
.expect("family-only layout must load");
assert_eq!(db.system, system);
assert_eq!(db.data_root, legacy_data_root);
}
#[test]
fn load_succeeds_on_legacy_only_layout() {
let tmp = tempfile::tempdir().unwrap();
let systems_root = tmp.path();
let (system, backend, version) = ("synth_legacy", "vllm", "1.2.3");
write_synthetic_system_yaml(systems_root, system);
let legacy_dir = systems_root.join("data").join(backend).join(version);
std::fs::create_dir_all(&legacy_dir).unwrap();
std::fs::write(legacy_dir.join("gemm_perf.parquet"), b"stub").unwrap();
let db = PerfDatabase::load(systems_root, system, backend, version)
.expect("legacy-only layout must load");
assert_eq!(db.data_root, legacy_dir);
}
#[test]
fn load_errors_mentioning_both_layouts_on_total_miss() {
let tmp = tempfile::tempdir().unwrap();
let systems_root = tmp.path();
let (system, backend, version) = ("synth_missing", "vllm", "9.9.9");
write_synthetic_system_yaml(systems_root, system);
std::fs::create_dir_all(systems_root.join("data")).unwrap();
match PerfDatabase::load(systems_root, system, backend, version) {
Err(AicError::PerfDatabase(msg)) => {
assert!(
msg.contains("legacy"),
"error should mention legacy layout: {msg}"
);
assert!(
msg.contains("family"),
"error should mention family layout: {msg}"
);
}
Ok(_) => panic!("expected load to fail for a totally missing tuple"),
Err(other) => panic!("expected PerfDatabase error, got {other:?}"),
}
}
#[test]
fn comm_root_prefers_family_comm_dir_over_legacy_nccl() {
let tmp = tempfile::tempdir().unwrap();
let system_data_root = tmp.path();
let version = "2.27.3";
let legacy = system_data_root.join("nccl").join(version);
std::fs::create_dir_all(&legacy).unwrap();
assert_eq!(comm_root(system_data_root, "nccl", version), legacy);
let family = system_data_root.join("comm").join("nccl").join(version);
std::fs::create_dir_all(&family).unwrap();
assert_eq!(comm_root(system_data_root, "nccl", version), family);
}
}