use crate::common::enums::{DatabaseMode, GemmQuantMode, TransferKind};
use crate::common::error::AicError;
use crate::operators::base::{PerformanceResult, SolComponents, Source, subtract_sol};
use crate::operators::moe::policy_fingerprint;
use crate::operators::util_empirical::{self, UtilGrid, ZeroAwareDeltaLookup};
use crate::perf_database::PerfDatabase;
use crate::perf_database::gemm::{
gemm_quant_by_name, gemm_sol_latency_ms_with_flops, gemm_sol_with_flops,
normalize_fp8_static_quant, quant_tc_flops,
};
use serde::{Deserialize, Serialize};
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct GemmOp {
pub name: String,
pub scale_factor: f64,
pub n: u32,
pub k: u32,
pub quant_mode: GemmQuantMode,
pub scale_num_tokens: u32,
pub low_precision_input: bool,
#[serde(default = "default_seq_split")]
pub seq_split: u32,
#[serde(default)]
pub below_grid_sol: bool,
}
pub(crate) fn default_seq_split() -> u32 {
1
}
impl GemmOp {
pub fn new(name: impl Into<String>, n: u32, k: u32, quant_mode: GemmQuantMode) -> Self {
Self {
name: name.into(),
scale_factor: 1.0,
n,
k,
quant_mode,
scale_num_tokens: 1,
low_precision_input: false,
seq_split: 1,
below_grid_sol: false,
}
}
pub fn query(
&self,
db: &PerfDatabase,
x: u32,
quant_override: Option<GemmQuantMode>,
) -> Result<PerformanceResult, AicError> {
let m = x / self.scale_num_tokens.max(1);
let m = m.div_ceil(self.seq_split.max(1));
let quant = quant_override.unwrap_or(self.quant_mode);
let base = match query_gemm_table(db, quant, m, self.n, self.k) {
Err(err)
if self.below_grid_sol
&& db.database_mode == DatabaseMode::Silicon
&& err.is_missing_perf_data()
&& db.gemm.has_quant(quant)? =>
{
let tc_flops = quant_tc_flops(&db.system_spec, quant.mapping())?;
PerformanceResult::sol(gemm_sol_with_flops(
&db.system_spec,
quant,
tc_flops,
m as f64,
self.n as f64,
self.k as f64,
))
}
other => other?,
};
let mut latency = base.latency_ms;
let mut energy = base.energy_wms;
let mut source = base.source;
let mut sol = base.sol;
let mut latency_floor = 0.0_f64;
if quant == GemmQuantMode::Fp8Static {
let cs = query_compute_scale_table(db, quant, m, self.k)?;
latency -= cs.latency_ms;
energy -= cs.energy_wms;
sol = subtract_sol(sol, cs.sol);
if self.low_precision_input {
let sm = query_scale_matrix_table(db, quant, m, self.k)?;
latency -= sm.latency_ms;
energy -= sm.energy_wms;
sol = subtract_sol(sol, sm.sol);
}
let tc_flops = quant_tc_flops(&db.system_spec, quant.mapping())?;
let floor_components = gemm_sol_with_flops(
&db.system_spec,
quant,
tc_flops,
m as f64,
self.n as f64,
self.k as f64,
);
latency_floor = floor_components.time_ms();
if sol.is_some() && latency < latency_floor {
sol = Some(floor_components);
}
source = Source::Estimated;
}
let mut result = PerformanceResult::with_energy(latency.max(latency_floor), energy, source);
if let Some(components) = sol {
result = result.with_sol(components);
}
Ok(result.clamp_non_negative().scaled(self.scale_factor))
}
pub fn weights_bytes(&self) -> f64 {
(self.n as f64) * (self.k as f64) * self.quant_mode.mapping().memory * self.scale_factor
}
}
fn query_gemm_table(
db: &PerfDatabase,
quant: GemmQuantMode,
m: u32,
n: u32,
k: u32,
) -> Result<PerformanceResult, AicError> {
let silicon = |v: crate::perf_database::perf_interp::LeafValue| {
PerformanceResult::with_energy(v.latency, v.energy, Source::Silicon)
};
match db.database_mode {
DatabaseMode::Sol | DatabaseMode::SolFull => {
let tc_flops = quant_tc_flops(&db.system_spec, quant.mapping())?;
Ok(PerformanceResult::sol(gemm_sol_with_flops(
&db.system_spec,
quant,
tc_flops,
m as f64,
n as f64,
k as f64,
)))
}
DatabaseMode::Empirical => Ok(PerformanceResult::new(
gemm_empirical(db, quant, m, n, k)?,
Source::Empirical,
)),
DatabaseMode::Hybrid => match db.gemm.query(quant, m, n, k) {
Ok(value) => Ok(silicon(value)),
Err(err) if err.is_missing_perf_data() => Ok(PerformanceResult::new(
gemm_empirical(db, quant, m, n, k)?,
Source::Empirical,
)),
Err(err) => Err(err),
},
_ => Ok(silicon(db.gemm.query(quant, m, n, k)?)),
}
}
pub(crate) const GEMM_QUANT_UTIL_LEVEL: &[(f64, f64, f64)] = &[
(2.0, 1.0, 0.70), (1.0, 1.0, 0.55), (0.5625, 1.0, 0.45), (0.5, 1.0, 0.45), (1.0, 2.0, 0.45), (0.5, 2.0, 0.35), (1.0, 4.0, 0.30), (0.5, 4.0, 0.30), (0.5625, 4.0, 0.30), ];
const GEMM_QUANT_UTIL_DEFAULT: f64 = 0.45;
fn gemm_quant_util_level(quant: GemmQuantMode) -> f64 {
let mapping = quant.mapping();
GEMM_QUANT_UTIL_LEVEL
.iter()
.find(|(memory, compute, _)| *memory == mapping.memory && *compute == mapping.compute)
.map(|(_, _, level)| *level)
.unwrap_or(GEMM_QUANT_UTIL_DEFAULT)
}
fn xprofile_gemm_quants(
query: GemmQuantMode,
table_quants: &[GemmQuantMode],
) -> Vec<GemmQuantMode> {
let qp = query.mapping();
let mut refs: Vec<GemmQuantMode> = table_quants
.iter()
.copied()
.filter(|q| {
let mapping = q.mapping();
*q != query && !(mapping.memory == qp.memory && mapping.compute == qp.compute)
})
.collect();
let dist = |q: GemmQuantMode| {
let mapping = q.mapping();
(
(mapping.compute - qp.compute).abs(),
(mapping.memory - qp.memory).abs(),
)
};
refs.sort_by(|a, b| {
dist(*a)
.partial_cmp(&dist(*b))
.expect("finite profile distances")
});
refs
}
fn gemm_table_quants(db: &PerfDatabase) -> Result<Vec<GemmQuantMode>, AicError> {
match db.gemm.available_quants() {
Ok(names) => Ok(names.iter().filter_map(|n| gemm_quant_by_name(n)).collect()),
Err(err) if err.is_missing_perf_data() => Ok(Vec::new()),
Err(err) => Err(err),
}
}
fn gemm_reference_grid(
db: &PerfDatabase,
source_quant: GemmQuantMode,
sol_quant: GemmQuantMode,
provenance: &'static str,
key: &str,
) -> Result<Option<std::sync::Arc<UtilGrid>>, AicError> {
let spec = &db.system_spec;
db.util_grids.get_or_try_build(key, || {
let tc_flops = quant_tc_flops(spec, sol_quant.mapping())?;
let sol =
|c: &[f64]| gemm_sol_latency_ms_with_flops(spec, sol_quant, tc_flops, c[0], c[1], c[2]);
match db.gemm.gemm_points(source_quant) {
Ok(points) => {
let mut grid = UtilGrid::new(util_empirical::build_samples(points, sol));
grid.reference_provenance = Some(provenance);
Ok(Some(grid))
}
Err(err) if err.is_missing_perf_data() => Ok(None),
Err(err) => Err(err),
}
})
}
fn gemm_empirical(
db: &PerfDatabase,
quant: GemmQuantMode,
m: u32,
n: u32,
k: u32,
) -> Result<f64, AicError> {
let spec = &db.system_spec;
let tc_flops = quant_tc_flops(spec, quant.mapping())?;
let sol = |c: &[f64]| gemm_sol_latency_ms_with_flops(spec, quant, tc_flops, c[0], c[1], c[2]);
let tqm = normalize_fp8_static_quant(quant);
let key = format!("gemm:{}", tqm.name());
let mut grid = db.util_grids.get_or_try_build(&key, || {
match db.gemm.gemm_points(quant) {
Ok(points) => Ok(Some(UtilGrid::new(util_empirical::build_samples(
points, sol,
)))),
Err(err) if err.is_missing_perf_data() => Ok(None),
Err(err) => Err(err),
}
})?;
let mut util_scale = 1.0;
if grid.as_deref().is_none_or(UtilGrid::is_empty) {
let policy = db.transfer_policy;
let table_quants = gemm_table_quants(db)?;
let fingerprint = policy_fingerprint(policy);
if policy.contains(TransferKind::XQuant) {
let qp = tqm.mapping();
if let Some(ref_q) = table_quants.iter().copied().find(|q| {
let mapping = q.mapping();
*q != tqm && mapping.memory == qp.memory && mapping.compute == qp.compute
}) {
let key = format!(
"gemm_xquant:{}:policy={}:ref={}",
tqm.name(),
fingerprint,
ref_q.name()
);
if let Some(reference) = gemm_reference_grid(db, ref_q, tqm, "xquant", &key)? {
grid = Some(reference);
}
}
}
if grid.as_deref().is_none_or(UtilGrid::is_empty) && policy.contains(TransferKind::XProfile)
{
for ref_q in xprofile_gemm_quants(tqm, &table_quants) {
let key = format!(
"gemm_xprofile:{}:policy={}:ref={}",
tqm.name(),
fingerprint,
ref_q.name()
);
if let Some(reference) = gemm_reference_grid(db, ref_q, ref_q, "xprofile", &key)? {
if !reference.is_empty() {
grid = Some(reference);
util_scale = gemm_quant_util_level(tqm) / gemm_quant_util_level(ref_q);
break;
}
}
}
}
}
let query = [m as f64, n as f64, k as f64];
let (latency, _) = util_empirical::estimate(sol(&query), &query, grid.as_deref(), util_scale)?;
db.note_provenance(
grid.as_deref()
.and_then(|g| g.reference_provenance)
.and_then(util_empirical::ProvenanceTier::from_tag)
.unwrap_or(util_empirical::ProvenanceTier::Empirical),
);
Ok(latency)
}
fn query_compute_scale_table(
db: &PerfDatabase,
quant: GemmQuantMode,
m: u32,
k: u32,
) -> Result<PerformanceResult, AicError> {
let silicon = |v: crate::perf_database::perf_interp::LeafValue| {
PerformanceResult::with_energy(v.latency, v.energy, Source::Silicon)
};
match db.database_mode {
DatabaseMode::Sol | DatabaseMode::SolFull => {
Ok(PerformanceResult::sol(SolComponents::new(
0.0,
2.0 * m as f64 * k as f64 / db.system_spec.gpu.mem_bw * 1000.0,
)))
}
DatabaseMode::Empirical => Ok(PerformanceResult::new(
compute_scale_empirical(db, quant, m, k)?,
Source::Empirical,
)),
DatabaseMode::Hybrid => match db.gemm.query_compute_scale(quant, m, k) {
Ok(value) => Ok(silicon(value)),
Err(err) if err.is_missing_perf_data() => Ok(PerformanceResult::new(
compute_scale_empirical(db, quant, m, k)?,
Source::Empirical,
)),
Err(err) => Err(err),
},
_ => Ok(silicon(db.gemm.query_compute_scale(quant, m, k)?)),
}
}
fn compute_scale_empirical(
db: &PerfDatabase,
quant: GemmQuantMode,
m: u32,
k: u32,
) -> Result<f64, AicError> {
let key = format!("compute_scale:{}", normalize_fp8_static_quant(quant).name());
let lookup =
db.delta_lookups
.get_or_try_build(&key, || match db.gemm.compute_scale_points(quant) {
Ok(points) => Ok(ZeroAwareDeltaLookup::new(points)),
Err(err) if err.is_missing_perf_data() => Err(AicError::EmpiricalNotImplemented(
format!("No empirical compute_scale data is available for m={m}, k={k}."),
)),
Err(err) => Err(err),
})?;
let spec = &db.system_spec;
let latency = lookup.estimate(&[m as f64, k as f64], |c| {
2.0 * c[0] * c[1] / spec.gpu.mem_bw * 1000.0
})?;
db.note_provenance(util_empirical::ProvenanceTier::Empirical);
Ok(latency)
}
fn query_scale_matrix_table(
db: &PerfDatabase,
quant: GemmQuantMode,
m: u32,
k: u32,
) -> Result<PerformanceResult, AicError> {
let silicon = |v: crate::perf_database::perf_interp::LeafValue| {
PerformanceResult::with_energy(v.latency, v.energy, Source::Silicon)
};
match db.database_mode {
DatabaseMode::Sol | DatabaseMode::SolFull => {
Ok(PerformanceResult::sol(SolComponents::new(
0.0,
3.0 * m as f64 * k as f64 / db.system_spec.gpu.mem_bw * 1000.0,
)))
}
DatabaseMode::Empirical => Ok(PerformanceResult::new(
scale_matrix_empirical(db, quant, m, k)?,
Source::Empirical,
)),
DatabaseMode::Hybrid => match db.gemm.query_scale_matrix(quant, m, k) {
Ok(value) => Ok(silicon(value)),
Err(err) if err.is_missing_perf_data() => Ok(PerformanceResult::new(
scale_matrix_empirical(db, quant, m, k)?,
Source::Empirical,
)),
Err(err) => Err(err),
},
_ => Ok(silicon(db.gemm.query_scale_matrix(quant, m, k)?)),
}
}
fn scale_matrix_empirical(
db: &PerfDatabase,
quant: GemmQuantMode,
m: u32,
k: u32,
) -> Result<f64, AicError> {
let spec = &db.system_spec;
let sol = |c: &[f64]| 3.0 * c[0] * c[1] / spec.gpu.mem_bw * 1000.0;
let key = format!("scale_matrix:{}", normalize_fp8_static_quant(quant).name());
let grid =
db.util_grids
.get_or_try_build(&key, || match db.gemm.scale_matrix_points(quant) {
Ok(points) => Ok(Some(UtilGrid::new(util_empirical::build_samples(
points, sol,
)))),
Err(err) if err.is_missing_perf_data() => Ok(None),
Err(err) => Err(err),
})?;
let query = [m as f64, k as f64];
let (latency, _) = util_empirical::estimate(sol(&query), &query, grid.as_deref(), 1.0)?;
db.note_provenance(util_empirical::ProvenanceTier::Empirical);
Ok(latency)
}
#[cfg(test)]
mod tests {
use super::*;
use std::path::PathBuf;
const REPO_ROOT_HINT: &str = env!("CARGO_MANIFEST_DIR");
fn b200_vllm_db() -> PerfDatabase {
let systems_root = PathBuf::from(REPO_ROOT_HINT)
.join("../..")
.join("python/aisimulate/src/aiconfigurator_core/systems");
PerfDatabase::load(&systems_root, "b200_sxm", "vllm", "0.24.0").expect("db must load")
}
#[test]
fn gemm_op_scale_factor_multiplies_latency() {
let db = b200_vllm_db();
let base = GemmOp::new("base", 65536, 16384, GemmQuantMode::Bfloat16)
.query(&db, 32768, None)
.expect("query must succeed");
assert_eq!(base.source, Source::Silicon);
let op = GemmOp {
name: "scaled".to_string(),
scale_factor: 0.5,
n: 65536,
k: 16384,
quant_mode: GemmQuantMode::Bfloat16,
scale_num_tokens: 1,
low_precision_input: false,
seq_split: 1,
below_grid_sol: false,
};
let result = op.query(&db, 32768, None).expect("query must succeed");
assert!(
(result.latency_ms - base.latency_ms * 0.5).abs() < 1e-12,
"scale_factor must halve the base latency: got {} vs base {}",
result.latency_ms,
base.latency_ms
);
}
#[test]
fn gemm_op_scale_num_tokens_divides_x() {
let db = b200_vllm_db();
let base = GemmOp::new("base", 65536, 16384, GemmQuantMode::Bfloat16)
.query(&db, 32768, None)
.expect("query must succeed");
let op = GemmOp {
name: "halved".to_string(),
scale_factor: 1.0,
n: 65536,
k: 16384,
quant_mode: GemmQuantMode::Bfloat16,
scale_num_tokens: 2,
low_precision_input: false,
seq_split: 1,
below_grid_sol: false,
};
let result = op.query(&db, 65536, None).expect("query must succeed");
assert!(
(result.latency_ms - base.latency_ms).abs() < 1e-12,
"scale_num_tokens must divide x: got {} vs base {}",
result.latency_ms,
base.latency_ms
);
}
#[test]
fn gemm_op_below_grid_sol_flag_degrades_shape_miss_to_sol() {
let db = b200_vllm_db();
let strict = GemmOp::new("gate", 1, 2048, GemmQuantMode::Bfloat16);
assert!(strict.query(&db, 8, None).is_err());
let op = GemmOp {
below_grid_sol: true,
..strict
};
let result = op
.query(&db, 8, None)
.expect("below-grid opt-in must degrade to SOL");
assert_eq!(result.source, Source::Sol);
assert!(
(result.latency_ms - 4.78961038961039e-06).abs() < 1e-15,
"expected the Python SOL oracle, got {}",
result.latency_ms
);
}
#[test]
fn gemm_op_quant_override_routes_to_different_quant() {
let db = b200_vllm_db();
let op = GemmOp::new("default-bf16", 65536, 16384, GemmQuantMode::Bfloat16);
let default = op.query(&db, 32768, None).expect("default query");
let overridden = op
.query(&db, 32768, Some(GemmQuantMode::Nvfp4))
.expect("override query must succeed");
let native = GemmOp::new("native-nvfp4", 65536, 16384, GemmQuantMode::Nvfp4)
.query(&db, 32768, None)
.expect("native query");
assert!(
(overridden.latency_ms - native.latency_ms).abs() < 1e-12,
"override must equal the native nvfp4 lookup: {} vs {}",
overridden.latency_ms,
native.latency_ms
);
assert!(
(overridden.latency_ms - default.latency_ms).abs() > 1e-9,
"override must change the lookup away from bf16"
);
}
#[test]
fn gemm_empirical_regime_routing() {
let mut db = b200_vllm_db();
db.database_mode = crate::common::enums::DatabaseMode::Empirical;
let cases = [
(3000u32, 65536u32, 16384u32, GemmQuantMode::Bfloat16),
(777, 4000, 5000, GemmQuantMode::Bfloat16),
(32768, 65536, 16384, GemmQuantMode::Nvfp4),
(1, 129, 130, GemmQuantMode::Fp8),
];
for (m, n, k, quant) in cases {
let result = query_gemm_table(&db, quant, m, n, k).expect("empirical query");
assert!(result.latency_ms.is_finite() && result.latency_ms > 0.0);
assert_eq!(
result.source,
Source::Empirical,
"({m}, {n}, {k}, {quant:?})"
);
assert_eq!(
result.energy_wms, 0.0,
"empirical fallback carries no energy"
);
}
}
#[test]
fn w4a16_nvfp4_uses_xprofile_only_in_hybrid() {
let systems_root = PathBuf::from(REPO_ROOT_HINT)
.join("../..")
.join("python/aisimulate/src/aiconfigurator_core/systems");
let mut db = PerfDatabase::load(&systems_root, "h200_sxm", "trtllm", "1.3.0rc20")
.expect("database must load");
let silicon = query_gemm_table(&db, GemmQuantMode::W4a16Nvfp4, 8192, 1536, 32);
assert!(matches!(silicon, Err(AicError::PerfDatabase(_))));
db.database_mode = crate::common::enums::DatabaseMode::Hybrid;
let hybrid = query_gemm_table(&db, GemmQuantMode::W4a16Nvfp4, 8192, 1536, 32)
.expect("XPROFILE transfer should make W4A16-NVFP4 HYBRID-estimable");
assert_eq!(hybrid.source, Source::Empirical);
assert!(hybrid.latency_ms.is_finite() && hybrid.latency_ms > 0.0);
}
#[test]
fn gemm_hybrid_missing_quant_raises_when_policy_forbids_transfer() {
let balanced = crate::common::enums::TransferPolicy {
xshape: true,
xquant: true,
xprofile: false,
xop: false,
};
let db = b200_vllm_db().with_mode(crate::common::enums::DatabaseMode::Hybrid, balanced);
let result = query_gemm_table(&db, GemmQuantMode::Int4Wo, 64, 64, 64);
assert!(
matches!(result, Err(AicError::EmpiricalNotImplemented(_))),
"got {result:?}"
);
}
#[test]
fn gemm_hybrid_missing_quant_borrows_xprofile_under_default_policy() {
use crate::perf_database::energy_test_fixtures::{
Col, write_energy_systems_root, write_parquet,
};
let tmp = tempfile::tempdir().expect("tmpdir");
let data = write_energy_systems_root(tmp.path());
write_parquet(
&data.join("gemm_perf.parquet"),
&[
Col::Str("gemm_dtype", vec!["fp8", "fp8", "bfloat16", "bfloat16"]),
Col::I64("m", vec![64, 128, 64, 128]),
Col::I64("n", vec![64, 64, 64, 64]),
Col::I64("k", vec![64, 64, 64, 64]),
Col::F64("latency", vec![1000.0, 2000.0, 1.0, 2.0]),
],
);
let mut db = PerfDatabase::load(tmp.path(), "testsys", "vllm", "1.0").expect("db");
db.database_mode = crate::common::enums::DatabaseMode::Hybrid;
let result =
query_gemm_table(&db, GemmQuantMode::Int4Wo, 64, 64, 64).expect("xprofile borrow");
assert_eq!(result.source, Source::Empirical);
assert_eq!(
db.worst_provenance(),
util_empirical::ProvenanceTier::XProfile
);
assert!(
result.latency_ms < 100.0,
"compute-first must borrow the bfloat16 donor (~O(1) ms), got {} \
(file-order fp8 donor would be ~O(1000))",
result.latency_ms
);
}
fn h200_vllm_db() -> PerfDatabase {
let systems_root = PathBuf::from(REPO_ROOT_HINT)
.join("../..")
.join("python/aisimulate/src/aiconfigurator_core/systems");
PerfDatabase::load(&systems_root, "h200_sxm", "vllm", "0.24.0").expect("db must load")
}
#[test]
fn gemm_quant_transfer_ladder_matches_python_oracles() {
let mut db = h200_vllm_db();
db.database_mode = crate::common::enums::DatabaseMode::Hybrid;
let cases = [
(
GemmQuantMode::Sq,
512u32,
util_empirical::ProvenanceTier::XQuant,
),
(
GemmQuantMode::Sq,
8192,
util_empirical::ProvenanceTier::XQuant,
),
(
GemmQuantMode::Int4Wo,
512,
util_empirical::ProvenanceTier::XProfile,
),
(
GemmQuantMode::Int4Wo,
8192,
util_empirical::ProvenanceTier::XProfile,
),
];
for (quant, m, tier) in cases {
db.reset_provenance();
let result = query_gemm_table(&db, quant, m, 4096, 4096).expect("ladder query");
let latency = result.latency_ms;
assert!(
latency.is_finite() && latency > 0.0,
"({quant:?}, m={m}): ladder estimate must be positive, got {latency}"
);
assert_eq!(
result.source,
Source::Empirical,
"({quant:?}, m={m}): wrong source"
);
assert_eq!(
db.worst_provenance(),
tier,
"({quant:?}, m={m}): wrong tier"
);
}
assert!(matches!(
query_gemm_table(&db, GemmQuantMode::Nvfp4, 512, 4096, 4096),
Err(AicError::MissingSystemFlops(_))
));
}
#[test]
fn gemm_op_weights_bytes_matches_python_formula() {
let op = GemmOp::new("w", 1024, 4096, GemmQuantMode::Bfloat16);
assert_eq!(op.weights_bytes(), 1024.0 * 4096.0 * 2.0);
let fp8_op = GemmOp::new("w-fp8", 1024, 4096, GemmQuantMode::Fp8);
assert_eq!(fp8_op.weights_bytes(), 1024.0 * 4096.0 * 1.0);
}
#[test]
fn gemm_fp8_static_energy_composition_matches_python_oracle() {
use crate::perf_database::energy_test_fixtures::{
Col, write_energy_systems_root, write_parquet,
};
let tmp = tempfile::tempdir().expect("tmpdir");
let data = write_energy_systems_root(tmp.path());
write_parquet(
&data.join("gemm_perf.parquet"),
&[
Col::Str("gemm_dtype", vec!["fp8", "fp8"]),
Col::I64("m", vec![128, 256]),
Col::I64("n", vec![1024, 1024]),
Col::I64("k", vec![1024, 1024]),
Col::F64("latency", vec![1.0, 3.0]),
Col::F64("power", vec![100.0, 200.0]),
],
);
write_parquet(
&data.join("computescale_perf.parquet"),
&[
Col::Str("quant_dtype", vec!["fp8", "fp8"]),
Col::I64("m", vec![128, 256]),
Col::I64("k", vec![1024, 1024]),
Col::F64("latency", vec![0.25, 0.75]),
Col::F64("power", vec![40.0, 80.0]),
],
);
write_parquet(
&data.join("scale_matrix_perf.parquet"),
&[
Col::Str("quant_dtype", vec!["fp8", "fp8"]),
Col::I64("m", vec![128, 256]),
Col::I64("k", vec![1024, 1024]),
Col::F64("latency", vec![0.125, 0.375]),
Col::F64("power", vec![20.0, 60.0]),
],
);
let db = PerfDatabase::load(tmp.path(), "testsys", "vllm", "1.0").expect("db must load");
let op = GemmOp {
name: "g".to_string(),
scale_factor: 2.0,
n: 1024,
k: 1024,
quant_mode: GemmQuantMode::Fp8Static,
scale_num_tokens: 1,
low_precision_input: true,
seq_split: 1,
below_grid_sol: false,
};
let r = op.query(&db, 192, None).expect("fp8_static query");
assert!(
(r.latency_ms - 2.5).abs() < 1e-9,
"latency {}",
r.latency_ms
);
assert!(
(r.energy_wms - 520.0).abs() < 1e-9 * 520.0,
"energy {}",
r.energy_wms
);
assert_eq!(r.source, Source::Estimated);
let op1 = GemmOp {
name: "g1".to_string(),
scale_factor: 1.0,
n: 1024,
k: 1024,
quant_mode: GemmQuantMode::Fp8Static,
scale_num_tokens: 1,
low_precision_input: false,
seq_split: 1,
below_grid_sol: false,
};
let r1 = op1.query(&db, 192, None).expect("fp8_static query");
assert!(
(r1.latency_ms - 1.5).abs() < 1e-9,
"latency {}",
r1.latency_ms
);
assert!(
(r1.energy_wms - 270.0).abs() < 1e-9 * 270.0,
"energy {}",
r1.energy_wms
);
}
#[test]
fn gemm_sol_mode_returns_roofline_with_sol_source() {
let mut db = b200_vllm_db();
db.database_mode = DatabaseMode::Sol;
let quant = GemmQuantMode::Bfloat16;
let op = GemmOp::new("gemm", 4096, 4096, quant);
let result = op.query(&db, 512, None).expect("sol query");
let tc_flops = quant_tc_flops(&db.system_spec, quant.mapping()).expect("flops");
let expected =
gemm_sol_latency_ms_with_flops(&db.system_spec, quant, tc_flops, 512.0, 4096.0, 4096.0);
assert_eq!(result.latency_ms, expected);
assert_eq!(result.source, Source::Sol);
assert_eq!(result.energy_wms, 0.0);
db.database_mode = DatabaseMode::SolFull;
let alias = op.query(&db, 512, None).expect("sol_full query");
assert_eq!(alias.latency_ms, expected);
assert_eq!(alias.source, Source::Sol);
let mem_bw = db.system_spec.gpu.mem_bw;
let cs = query_compute_scale_table(&db, quant, 512, 4096).expect("cs");
assert_eq!(cs.latency_ms, 2.0 * 512.0 * 4096.0 / mem_bw * 1000.0);
assert_eq!(cs.source, Source::Sol);
let sm = query_scale_matrix_table(&db, quant, 512, 4096).expect("sm");
assert_eq!(sm.latency_ms, 3.0 * 512.0 * 4096.0 / mem_bw * 1000.0);
assert_eq!(sm.source, Source::Sol);
}
}