use crate::common::enums::{DatabaseMode, FmhaQuantMode, GemmQuantMode, KvCacheQuantMode};
use crate::common::error::AicError;
use crate::operators::base::{PerformanceResult, Source};
use crate::operators::util_empirical::{self, UtilGrid};
use crate::perf_database::PerfDatabase;
use crate::perf_database::attention::generation_attn_flops;
use crate::perf_database::gemm::quant_tc_flops;
use crate::perf_database::mla::{
context_mla_sol_ms, context_mla_sol_prefix, context_mla_sol_prefix_ms,
generation_mla_module_sol, generation_mla_module_sol_ms, generation_mla_sol,
generation_mla_sol_ms, mla_bmm_sol, mla_bmm_sol_ms,
};
use serde::{Deserialize, Serialize};
fn prefix_correction(full_s: u32, prefix: u32) -> f64 {
if full_s == 0 {
return 0.0;
}
let f = full_s as f64;
let p = prefix as f64;
(f * f - p * p) / (f * f)
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ContextMlaOp {
pub name: String,
pub scale_factor: f64,
pub num_heads: u32,
pub kv_cache_dtype: KvCacheQuantMode,
pub fmha_quant_mode: FmhaQuantMode,
#[serde(default = "crate::operators::gemm::default_seq_split")]
pub cp_size: u32,
}
impl ContextMlaOp {
pub fn new(
name: impl Into<String>,
num_heads: u32,
kv_cache_dtype: KvCacheQuantMode,
fmha_quant_mode: FmhaQuantMode,
) -> Self {
Self {
name: name.into(),
scale_factor: 1.0,
num_heads,
kv_cache_dtype,
fmha_quant_mode,
cp_size: 1,
}
}
pub fn query(
&self,
db: &PerfDatabase,
batch_size: u32,
isl: u32,
prefix: u32,
) -> Result<PerformanceResult, AicError> {
let ctx = |s: u32, pfx: u32| -> Result<PerformanceResult, AicError> {
query_context_mla_table(
db,
batch_size,
s,
pfx,
self.num_heads,
self.kv_cache_dtype,
self.fmha_quant_mode,
)
};
let result = if self.cp_size > 1 {
let c = isl.div_ceil(2 * self.cp_size).max(1);
ctx(c, prefix)?.plus(ctx(c, prefix + isl - c)?)
} else {
ctx(isl, prefix)?
};
Ok(result.clamp_non_negative().scaled(self.scale_factor))
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct GenerationMlaOp {
pub name: String,
pub scale_factor: f64,
pub num_heads: u32,
pub kv_cache_dtype: KvCacheQuantMode,
}
impl GenerationMlaOp {
pub fn new(name: impl Into<String>, num_heads: u32, kv_cache_dtype: KvCacheQuantMode) -> Self {
Self {
name: name.into(),
scale_factor: 1.0,
num_heads,
kv_cache_dtype,
}
}
pub fn query(
&self,
db: &PerfDatabase,
batch_size: u32,
s: u32,
) -> Result<PerformanceResult, AicError> {
let result =
query_generation_mla_table(db, batch_size, s, self.num_heads, self.kv_cache_dtype)?;
Ok(result.clamp_non_negative().scaled(self.scale_factor))
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct MlaModuleOp {
pub name: String,
pub scale_factor: f64,
pub num_heads: u32,
pub kv_cache_dtype: KvCacheQuantMode,
pub fmha_quant_mode: FmhaQuantMode,
pub gemm_quant_mode: GemmQuantMode,
#[serde(default)]
pub native_num_heads: Option<u32>,
}
impl MlaModuleOp {
pub fn new(
name: impl Into<String>,
num_heads: u32,
kv_cache_dtype: KvCacheQuantMode,
fmha_quant_mode: FmhaQuantMode,
gemm_quant_mode: GemmQuantMode,
) -> Self {
Self {
name: name.into(),
scale_factor: 1.0,
num_heads,
kv_cache_dtype,
fmha_quant_mode,
gemm_quant_mode,
native_num_heads: None,
}
}
pub fn query_context(
&self,
db: &PerfDatabase,
batch_size: u32,
isl: u32,
prefix: u32,
) -> Result<PerformanceResult, AicError> {
let result = query_context_mla_module_table(
db,
batch_size,
isl,
prefix,
self.num_heads,
self.kv_cache_dtype,
self.fmha_quant_mode,
self.gemm_quant_mode,
self.native_num_heads,
)?;
Ok(result.clamp_non_negative().scaled(self.scale_factor))
}
pub fn query_generation(
&self,
db: &PerfDatabase,
batch_size: u32,
s: u32,
) -> Result<PerformanceResult, AicError> {
let result = query_generation_mla_module_table(
db,
batch_size,
s,
self.num_heads,
self.kv_cache_dtype,
self.gemm_quant_mode,
self.native_num_heads,
)?;
Ok(result.clamp_non_negative().scaled(self.scale_factor))
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct MlaBmmOp {
pub name: String,
pub scale_factor: f64,
pub num_heads: u32,
pub quant_mode: GemmQuantMode,
pub is_pre: bool,
}
impl MlaBmmOp {
pub fn new(
name: impl Into<String>,
num_heads: u32,
quant_mode: GemmQuantMode,
is_pre: bool,
) -> Self {
Self {
name: name.into(),
scale_factor: 1.0,
num_heads,
quant_mode,
is_pre,
}
}
pub fn query(&self, db: &PerfDatabase, num_tokens: u32) -> Result<PerformanceResult, AicError> {
let result =
query_mla_bmm_table(db, num_tokens, self.num_heads, self.quant_mode, self.is_pre)?;
Ok(result.clamp_non_negative().scaled(self.scale_factor))
}
}
fn query_context_mla_table(
db: &PerfDatabase,
b: u32,
s: u32,
prefix: u32,
num_heads: u32,
kv_quant: KvCacheQuantMode,
fmha_quant: FmhaQuantMode,
) -> Result<PerformanceResult, AicError> {
let silicon = || -> Result<PerformanceResult, AicError> {
let full_s = s + prefix;
let raw = db
.mla
.query_context(b, full_s, num_heads, kv_quant, fmha_quant)?;
let correction = prefix_correction(full_s, prefix);
Ok(PerformanceResult::with_energy(
raw.latency * correction,
raw.energy * correction,
Source::Silicon,
))
};
match db.database_mode {
DatabaseMode::Sol | DatabaseMode::SolFull => {
let attn_flops = quant_tc_flops(&db.system_spec, fmha_quant.mapping())?;
Ok(PerformanceResult::sol(context_mla_sol_prefix(
&db.system_spec,
kv_quant,
num_heads as f64,
s as f64,
prefix as f64,
b as f64,
attn_flops,
)))
}
DatabaseMode::Empirical => Ok(PerformanceResult::new(
context_mla_empirical(db, b, s, prefix, num_heads, kv_quant, fmha_quant)?,
Source::Empirical,
)),
DatabaseMode::Hybrid => match silicon() {
Ok(result) => Ok(result),
Err(err) if err.is_missing_perf_data() => Ok(PerformanceResult::new(
context_mla_empirical(db, b, s, prefix, num_heads, kv_quant, fmha_quant)?,
Source::Empirical,
)),
Err(err) => Err(err),
},
_ => silicon(),
}
}
fn context_mla_empirical(
db: &PerfDatabase,
b: u32,
s: u32,
prefix: u32,
num_heads: u32,
kv_quant: KvCacheQuantMode,
fmha_quant: FmhaQuantMode,
) -> Result<f64, AicError> {
let spec = &db.system_spec;
let attn_flops = quant_tc_flops(spec, fmha_quant.mapping())?;
let sol = |c: &[f64]| context_mla_sol_ms(spec, kv_quant, c[0], c[1], c[2], attn_flops);
let key = format!("ctx_mla:{}:{}", fmha_quant.name(), kv_quant.name());
let grid = db.util_grids.get_or_try_build(&key, || {
match db.mla.context_points(kv_quant, fmha_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 sol_query = context_mla_sol_prefix_ms(
spec,
kv_quant,
num_heads as f64,
s as f64,
prefix as f64,
b as f64,
attn_flops,
);
let query = [num_heads as f64, (s + prefix) as f64, b as f64];
let (latency, _) = util_empirical::estimate(sol_query, &query, grid.as_deref(), 1.0)?;
db.note_provenance(util_empirical::ProvenanceTier::Empirical);
Ok(latency)
}
fn query_generation_mla_table(
db: &PerfDatabase,
b: u32,
s: u32,
num_heads: u32,
kv_quant: KvCacheQuantMode,
) -> 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 attn_flops = generation_attn_flops(&db.system_spec, kv_quant)?;
Ok(PerformanceResult::sol(generation_mla_sol(
&db.system_spec,
kv_quant,
num_heads as f64,
b as f64,
s as f64,
attn_flops,
)))
}
DatabaseMode::Empirical => Ok(PerformanceResult::new(
generation_mla_empirical(db, b, s, num_heads, kv_quant)?,
Source::Empirical,
)),
DatabaseMode::Hybrid => match db.mla.query_generation(b, s, num_heads, kv_quant) {
Ok(value) => Ok(silicon(value)),
Err(err) if err.is_missing_perf_data() => Ok(PerformanceResult::new(
generation_mla_empirical(db, b, s, num_heads, kv_quant)?,
Source::Empirical,
)),
Err(err) => Err(err),
},
_ => Ok(silicon(db.mla.query_generation(b, s, num_heads, kv_quant)?)),
}
}
fn generation_mla_empirical(
db: &PerfDatabase,
b: u32,
s: u32,
num_heads: u32,
kv_quant: KvCacheQuantMode,
) -> Result<f64, AicError> {
let spec = &db.system_spec;
let attn_flops = generation_attn_flops(spec, kv_quant)?;
let sol = |c: &[f64]| generation_mla_sol_ms(spec, kv_quant, c[0], c[1], c[2], attn_flops);
let key = format!("gen_mla:{}", kv_quant.name());
let grid =
db.util_grids
.get_or_try_build(&key, || match db.mla.generation_points(kv_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 = [num_heads as f64, b as f64, s as f64];
let (latency, _) = util_empirical::estimate(sol(&query), &query, grid.as_deref(), 1.0)?;
db.note_provenance(util_empirical::ProvenanceTier::Empirical);
Ok(latency)
}
fn resolve_bmm_slice_heads(
db: &PerfDatabase,
num_heads: u32,
quant: GemmQuantMode,
is_pre: bool,
) -> Result<(u32, f64), AicError> {
let pow2 = num_heads.next_power_of_two();
if pow2 == num_heads {
return Ok((num_heads, 1.0));
}
let has_exact = match db.mla.bmm_selected_quant(quant) {
Ok(selected) => db.mla.bmm_has_heads(selected, is_pre, num_heads)?,
Err(err) if err.is_missing_perf_data() => false,
Err(err) => return Err(err),
};
if has_exact {
Ok((num_heads, 1.0))
} else {
Ok((pow2, f64::from(num_heads) / f64::from(pow2)))
}
}
fn query_mla_bmm_table(
db: &PerfDatabase,
num_tokens: u32,
num_heads: u32,
quant: GemmQuantMode,
is_pre: bool,
) -> Result<PerformanceResult, AicError> {
if matches!(db.database_mode, DatabaseMode::Sol | DatabaseMode::SolFull) {
let spec = &db.system_spec;
let bmm_flops = quant_tc_flops(spec, quant.mapping())?;
return Ok(PerformanceResult::sol(mla_bmm_sol(
spec,
quant,
num_heads as f64,
num_tokens as f64,
bmm_flops,
)));
}
let (num_heads, head_scale) = resolve_bmm_slice_heads(db, num_heads, quant, is_pre)?;
let silicon = |v: crate::perf_database::perf_interp::LeafValue| {
PerformanceResult::with_energy(
v.latency * head_scale,
v.energy * head_scale,
Source::Silicon,
)
};
match db.database_mode {
DatabaseMode::Empirical => Ok(PerformanceResult::new(
mla_bmm_empirical(db, num_tokens, num_heads, quant, is_pre)? * head_scale,
Source::Empirical,
)),
DatabaseMode::Hybrid => match db.mla.query_bmm(num_tokens, num_heads, quant, is_pre) {
Ok(value) => Ok(silicon(value)),
Err(err) if err.is_missing_perf_data() => Ok(PerformanceResult::new(
mla_bmm_empirical(db, num_tokens, num_heads, quant, is_pre)? * head_scale,
Source::Empirical,
)),
Err(err) => Err(err),
},
_ => Ok(silicon(
db.mla.query_bmm(num_tokens, num_heads, quant, is_pre)?,
)),
}
}
fn mla_bmm_empirical(
db: &PerfDatabase,
num_tokens: u32,
num_heads: u32,
quant: GemmQuantMode,
is_pre: bool,
) -> Result<f64, AicError> {
let spec = &db.system_spec;
let bmm_flops = quant_tc_flops(spec, quant.mapping())?;
let sol = |c: &[f64]| mla_bmm_sol_ms(spec, quant, num_heads as f64, c[0], bmm_flops);
let op_name = if is_pre {
"mla_gen_pre"
} else {
"mla_gen_post"
};
let grid = match db.mla.bmm_selected_quant(quant) {
Ok(selected) => {
let key = format!(
"mla_bmm:{}:{}:{}:{}",
quant.name(),
selected.name(),
op_name,
num_heads
);
db.util_grids.get_or_try_build(&key, || {
match db.mla.bmm_points(selected, is_pre, num_heads) {
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),
}
})?
}
Err(err) if err.is_missing_perf_data() => None,
Err(err) => return Err(err),
};
let query = [num_tokens as f64];
let (latency, _) = util_empirical::estimate(sol(&query), &query, grid.as_deref(), 1.0)?;
db.note_provenance(util_empirical::ProvenanceTier::Empirical);
Ok(latency)
}
#[allow(clippy::too_many_arguments)]
fn query_context_mla_module_table(
db: &PerfDatabase,
b: u32,
s: u32,
prefix: u32,
num_heads: u32,
kv_quant: KvCacheQuantMode,
fmha_quant: FmhaQuantMode,
gemm_quant: GemmQuantMode,
native_heads: Option<u32>,
) -> Result<PerformanceResult, AicError> {
let silicon = || -> Result<PerformanceResult, AicError> {
let full_s = s + prefix;
let raw = db.mla.query_context_module(
b,
full_s,
num_heads,
kv_quant,
fmha_quant,
gemm_quant,
native_heads,
)?;
let correction = prefix_correction(full_s, prefix);
Ok(PerformanceResult::with_energy(
raw.latency * correction,
raw.energy * correction,
Source::Silicon,
))
};
match db.database_mode {
DatabaseMode::Sol | DatabaseMode::SolFull => {
let attn_flops = quant_tc_flops(&db.system_spec, fmha_quant.mapping())?;
Ok(PerformanceResult::sol(context_mla_sol_prefix(
&db.system_spec,
kv_quant,
num_heads as f64,
s as f64,
prefix as f64,
b as f64,
attn_flops,
)))
}
DatabaseMode::Empirical => Ok(PerformanceResult::new(
context_mla_module_empirical(
db,
b,
s,
prefix,
num_heads,
kv_quant,
fmha_quant,
gemm_quant,
native_heads,
)?,
Source::Empirical,
)),
DatabaseMode::Hybrid => match silicon() {
Ok(result) => Ok(result),
Err(err) if err.is_missing_perf_data() => Ok(PerformanceResult::new(
context_mla_module_empirical(
db,
b,
s,
prefix,
num_heads,
kv_quant,
fmha_quant,
gemm_quant,
native_heads,
)?,
Source::Empirical,
)),
Err(err) => Err(err),
},
_ => silicon(),
}
}
#[allow(clippy::too_many_arguments)]
fn context_mla_module_empirical(
db: &PerfDatabase,
b: u32,
s: u32,
prefix: u32,
num_heads: u32,
kv_quant: KvCacheQuantMode,
fmha_quant: FmhaQuantMode,
gemm_quant: GemmQuantMode,
native_heads: Option<u32>,
) -> Result<f64, AicError> {
let spec = &db.system_spec;
let attn_flops = quant_tc_flops(spec, fmha_quant.mapping())?;
let sol = |c: &[f64]| context_mla_sol_ms(spec, kv_quant, c[0], c[1], c[2], attn_flops);
let key = format!(
"ctx_mla_mod:{}:{}:{}:{:?}",
fmha_quant.name(),
kv_quant.name(),
gemm_quant.name(),
native_heads
);
let grid = db.util_grids.get_or_try_build(&key, || {
match db
.mla
.context_module_points(kv_quant, fmha_quant, gemm_quant, native_heads)
{
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 sol_query = context_mla_sol_prefix_ms(
spec,
kv_quant,
num_heads as f64,
s as f64,
prefix as f64,
b as f64,
attn_flops,
);
let query = [num_heads as f64, (s + prefix) as f64, b as f64];
let (latency, _) = util_empirical::estimate(sol_query, &query, grid.as_deref(), 1.0)?;
db.note_provenance(util_empirical::ProvenanceTier::Empirical);
Ok(latency)
}
fn query_generation_mla_module_table(
db: &PerfDatabase,
b: u32,
s: u32,
num_heads: u32,
kv_quant: KvCacheQuantMode,
gemm_quant: GemmQuantMode,
native_heads: Option<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 spec = &db.system_spec;
let attn_flops = generation_attn_flops(spec, kv_quant)?;
let bmm_flops = quant_tc_flops(spec, gemm_quant.mapping())?;
Ok(PerformanceResult::sol(generation_mla_module_sol(
spec,
kv_quant,
gemm_quant,
num_heads as f64,
b as f64,
s as f64,
attn_flops,
bmm_flops,
)))
}
DatabaseMode::Empirical => Ok(PerformanceResult::new(
generation_mla_module_empirical(
db,
b,
s,
num_heads,
kv_quant,
gemm_quant,
native_heads,
)?,
Source::Empirical,
)),
DatabaseMode::Hybrid => {
match db.mla.query_generation_module(
b,
s,
num_heads,
kv_quant,
gemm_quant,
native_heads,
) {
Ok(value) => Ok(silicon(value)),
Err(err) if err.is_missing_perf_data() => Ok(PerformanceResult::new(
generation_mla_module_empirical(
db,
b,
s,
num_heads,
kv_quant,
gemm_quant,
native_heads,
)?,
Source::Empirical,
)),
Err(err) => Err(err),
}
}
_ => Ok(silicon(db.mla.query_generation_module(
b,
s,
num_heads,
kv_quant,
gemm_quant,
native_heads,
)?)),
}
}
fn generation_mla_module_empirical(
db: &PerfDatabase,
b: u32,
s: u32,
num_heads: u32,
kv_quant: KvCacheQuantMode,
gemm_quant: GemmQuantMode,
native_heads: Option<u32>,
) -> Result<f64, AicError> {
let spec = &db.system_spec;
let attn_flops = generation_attn_flops(spec, kv_quant)?;
let bmm_flops = quant_tc_flops(spec, gemm_quant.mapping())?;
let sol = |c: &[f64]| {
generation_mla_module_sol_ms(
spec, kv_quant, gemm_quant, c[0], c[1], c[2], attn_flops, bmm_flops,
)
};
let key = format!(
"gen_mla_mod:{}:{}:{:?}",
kv_quant.name(),
gemm_quant.name(),
native_heads
);
let grid = db.util_grids.get_or_try_build(&key, || {
match db
.mla
.generation_module_points(kv_quant, gemm_quant, native_heads)
{
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 = [num_heads as f64, b as f64, s 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 mla_op_context_absent_on_vllm_b200() {
let db = b200_vllm_db();
let op = ContextMlaOp::new(
"ctx_op",
128,
KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Bfloat16,
);
let err = op.query(&db, 1, 1024, 0).unwrap_err();
match err {
AicError::Io { .. } | AicError::PerfDatabase(_) => {}
other => panic!("unexpected error: {other:?}"),
}
}
fn gb200_trtllm_db() -> PerfDatabase {
let systems_root = PathBuf::from(REPO_ROOT_HINT)
.join("../..")
.join("python/aisimulate/src/aiconfigurator_core/systems");
PerfDatabase::load(&systems_root, "gb200", "trtllm", "1.3.0rc20").expect("db must load")
}
fn assert_close(got: f64, expected: f64, what: &str) {
assert!(
(got - expected).abs() < 1e-9,
"{what}: expected {expected}, got {got}"
);
}
#[test]
fn context_mla_empirical_regime_routing() {
let mut db = gb200_trtllm_db();
db.database_mode = DatabaseMode::Empirical;
let cases: &[(u32, u32, u32, u32)] =
&[(4, 5000, 0, 128), (2, 3000, 1024, 16), (4, 4096, 0, 128)];
for &(b, s, prefix, n) in cases {
let __r = query_context_mla_table(
&db,
b,
s,
prefix,
n,
KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Bfloat16,
)
.expect("empirical query");
assert!(__r.latency_ms.is_finite() && __r.latency_ms > 0.0);
assert_eq!(
__r.source,
Source::Empirical,
"(b={b}, s={s}, pfx={prefix}, n={n})"
);
}
}
#[test]
fn context_mla_hybrid_missing_slice_raises_empirical_not_implemented() {
let mut db = gb200_trtllm_db();
db.database_mode = DatabaseMode::Hybrid;
let result = query_context_mla_table(
&db,
4,
4096,
0,
128,
KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Fp8,
);
assert!(
matches!(result, Err(AicError::EmpiricalNotImplemented(_))),
"got {result:?}"
);
}
#[test]
fn generation_mla_empirical_regime_routing() {
let mut db = gb200_trtllm_db();
db.database_mode = DatabaseMode::Empirical;
for &(b, s, n) in &[(7u32, 9000u32, 128u32), (1, 4096, 128)] {
let __r = query_generation_mla_table(&db, b, s, n, KvCacheQuantMode::Bfloat16)
.expect("empirical query");
assert!(__r.latency_ms.is_finite() && __r.latency_ms > 0.0);
assert_eq!(__r.source, Source::Empirical, "(b={b}, s={s}, n={n})");
}
db.database_mode = DatabaseMode::Hybrid;
let result = query_generation_mla_table(&db, 1, 4096, 128, KvCacheQuantMode::Int8);
assert!(
matches!(result, Err(AicError::EmpiricalNotImplemented(_))),
"got {result:?}"
);
}
#[test]
fn mla_bmm_empirical_regime_routing() {
let mut db = gb200_trtllm_db();
db.database_mode = DatabaseMode::Empirical;
let cases: &[(u32, u32, GemmQuantMode, bool)] = &[
(100, 128, GemmQuantMode::Bfloat16, true),
(256, 128, GemmQuantMode::Bfloat16, true),
(20000, 128, GemmQuantMode::Fp8, true),
(777, 64, GemmQuantMode::Fp8, false),
];
for &(t, n, quant, is_pre) in cases {
let __r = query_mla_bmm_table(&db, t, n, quant, is_pre).expect("empirical query");
assert!(__r.latency_ms.is_finite() && __r.latency_ms > 0.0);
assert_eq!(
__r.source,
Source::Empirical,
"mla_bmm(t={t}, n={n}, {quant:?}, pre={is_pre})"
);
}
db.database_mode = DatabaseMode::Hybrid;
let result = query_mla_bmm_table(&db, 64, 130, GemmQuantMode::Bfloat16, true);
assert!(
matches!(result, Err(AicError::EmpiricalNotImplemented(_))),
"got {result:?}"
);
}
#[test]
fn mla_bmm_non_pow2_heads_reroute_to_next_pow2_slice() {
let db = gb200_trtllm_db();
let __r = query_mla_bmm_table(&db, 64, 7, GemmQuantMode::Bfloat16, true).expect("reroute");
let (lat7, source) = (__r.latency_ms, __r.source);
let __r = query_mla_bmm_table(&db, 64, 8, GemmQuantMode::Bfloat16, true).expect("exact");
let lat8 = __r.latency_ms;
assert_eq!(source, Source::Silicon);
assert_close(lat7, lat8 * 7.0 / 8.0, "mla_bmm 7 -> 8-head slice reroute");
}
#[test]
fn context_mla_module_empirical_regime_routing() {
let mut db = b200_vllm_db();
db.database_mode = DatabaseMode::Empirical;
type Case = (
u32,
u32,
u32,
u32,
FmhaQuantMode,
KvCacheQuantMode,
GemmQuantMode,
);
let cases: &[Case] = &[
(
2,
5000,
0,
128,
FmhaQuantMode::Bfloat16,
KvCacheQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
),
(
1,
2000,
2048,
16,
FmhaQuantMode::Bfloat16,
KvCacheQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
),
(
1,
1,
0,
128,
FmhaQuantMode::Bfloat16,
KvCacheQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
),
(
2,
5000,
0,
128,
FmhaQuantMode::Fp8,
KvCacheQuantMode::Fp8,
GemmQuantMode::Fp8Block,
),
];
for &(b, s, prefix, n, fmha, kv, gemm) in cases {
let __r = query_context_mla_module_table(&db, b, s, prefix, n, kv, fmha, gemm, None)
.expect("empirical query");
assert!(__r.latency_ms.is_finite() && __r.latency_ms > 0.0);
assert_eq!(__r.source, Source::Empirical, "(b={b}, s={s}, n={n})");
}
db.database_mode = DatabaseMode::Hybrid;
let result = query_context_mla_module_table(
&db,
2,
5000,
0,
128,
KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Bfloat16,
GemmQuantMode::Fp8,
None,
);
assert!(
matches!(result, Err(AicError::EmpiricalNotImplemented(_))),
"got {result:?}"
);
}
#[test]
fn generation_mla_module_empirical_regime_routing() {
let mut db = b200_vllm_db();
db.database_mode = DatabaseMode::Empirical;
let cases: &[(u32, u32, u32, KvCacheQuantMode, GemmQuantMode)] = &[
(
8,
3000,
128,
KvCacheQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
),
(
1,
4097,
128,
KvCacheQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
),
(8, 3000, 16, KvCacheQuantMode::Fp8, GemmQuantMode::Fp8Block),
];
for &(b, s, n, kv, gemm) in cases {
let __r = query_generation_mla_module_table(&db, b, s, n, kv, gemm, None)
.expect("empirical query");
assert!(__r.latency_ms.is_finite() && __r.latency_ms > 0.0);
assert_eq!(__r.source, Source::Empirical, "(b={b}, s={s}, n={n})");
}
db.database_mode = DatabaseMode::Hybrid;
let result = query_generation_mla_module_table(
&db,
8,
3000,
128,
KvCacheQuantMode::Bfloat16,
GemmQuantMode::Fp8,
None,
);
assert!(
matches!(result, Err(AicError::EmpiricalNotImplemented(_))),
"got {result:?}"
);
}
#[test]
fn mla_sol_mode_returns_roofline_with_sol_source() {
let mut db = b200_vllm_db();
db.database_mode = DatabaseMode::Sol;
let spec = db.system_spec.clone();
let ctx = ContextMlaOp::new(
"ctx_mla",
128,
KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Bfloat16,
);
let result = ctx.query(&db, 2, 1024, 512).expect("ctx sol");
let attn_flops = quant_tc_flops(&spec, FmhaQuantMode::Bfloat16.mapping()).unwrap();
let expected = context_mla_sol_prefix_ms(
&spec,
KvCacheQuantMode::Bfloat16,
128.0,
1024.0,
512.0,
2.0,
attn_flops,
);
assert_eq!(result.latency_ms, expected);
assert_eq!(result.source, Source::Sol);
assert_eq!(result.energy_wms, 0.0);
let generation = GenerationMlaOp::new("gen_mla", 128, KvCacheQuantMode::Bfloat16);
let result = generation.query(&db, 8, 1024).expect("gen sol");
let gen_flops = generation_attn_flops(&spec, KvCacheQuantMode::Bfloat16).unwrap();
let expected = generation_mla_sol_ms(
&spec,
KvCacheQuantMode::Bfloat16,
128.0,
8.0,
1024.0,
gen_flops,
);
assert_eq!(result.latency_ms, expected);
assert_eq!(result.source, Source::Sol);
let bmm = MlaBmmOp::new("bmm", 96, GemmQuantMode::Bfloat16, true);
let result = bmm.query(&db, 64).expect("bmm sol");
let bmm_flops = quant_tc_flops(&spec, GemmQuantMode::Bfloat16.mapping()).unwrap();
let expected = mla_bmm_sol_ms(&spec, GemmQuantMode::Bfloat16, 96.0, 64.0, bmm_flops);
assert_eq!(result.latency_ms, expected);
assert_eq!(result.source, Source::Sol);
}
}