use serde::{Deserialize, Serialize};
use crate::common::enums::{
DatabaseMode, FmhaQuantMode, GemmQuantMode, KvCacheQuantMode, TransferKind,
};
use crate::common::error::AicError;
use crate::common::system_spec::SystemSpec;
use crate::operators::base::{PerformanceResult, Source};
use crate::operators::dsa::DsaModuleOp;
use crate::perf_database::PerfDatabase;
use crate::perf_database::dsa::{
dsa_context_sol_flops, dsa_context_sol_ms, dsa_dims, dsa_generation_sol_flops,
dsa_generation_sol_ms,
};
use crate::perf_database::gemm::quant_tc_flops;
const MSA_ARCHITECTURE: &str = "MiniMaxM3ForCausalLM";
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct MsaModuleOp {
pub name: String,
pub scale_factor: f64,
pub num_heads: u32,
pub num_kv_heads: u32,
pub hidden_size: u32,
pub head_dim: u32,
pub v_head_dim: u32,
pub index_n_heads: u32,
pub index_head_dim: u32,
pub index_topk: u32,
pub block_size: u32,
pub kv_cache_dtype: KvCacheQuantMode,
pub fmha_quant_mode: FmhaQuantMode,
pub gemm_quant_mode: GemmQuantMode,
pub dsa_architecture: String,
pub dsa_scale_k: f64,
}
impl MsaModuleOp {
pub fn query_context(
&self,
db: &PerfDatabase,
batch_size: u32,
s: u32,
prefix: u32,
) -> Result<PerformanceResult, AicError> {
let sol = self.sol_ms(db, batch_size, s, prefix, true)?;
match db.database_mode {
DatabaseMode::Sol | DatabaseMode::SolFull => {
Ok(PerformanceResult::new(sol * self.scale_factor, Source::Sol))
}
DatabaseMode::Silicon => {
self.silicon_context(db, batch_size, s, prefix)
.map(|latency| {
PerformanceResult::new(latency * self.scale_factor, Source::Silicon)
})
}
DatabaseMode::Empirical => self.xop_context(db, batch_size, s, prefix, sol),
DatabaseMode::Hybrid => match self.silicon_context(db, batch_size, s, prefix) {
Ok(latency) => Ok(PerformanceResult::new(
latency * self.scale_factor,
Source::Silicon,
)),
Err(err) if err.is_missing_perf_data() => {
self.xop_context(db, batch_size, s, prefix, sol)
}
Err(err) => Err(err),
},
}
}
fn xop_context(
&self,
db: &PerfDatabase,
batch_size: u32,
s: u32,
prefix: u32,
sol: f64,
) -> Result<PerformanceResult, AicError> {
if !db.transfer_policy.contains(TransferKind::XOp) {
return Err(AicError::EmpiricalNotImplemented(
"MSA context: cross-op transfer (xop) is disabled by the transfer policy \
and no MSA silicon data is available for this workload."
.to_string(),
));
}
let util = self.dsa_context_util(db, batch_size, s, prefix);
match util {
Some(util) if util > 0.0 => {
let latency = sol / (util * self.dsa_scale_k);
db.note_provenance(crate::operators::util_empirical::ProvenanceTier::XOp);
Ok(PerformanceResult::new(
latency * self.scale_factor,
Source::Empirical,
))
}
_ => Err(AicError::EmpiricalNotImplemented(format!(
"MSA context: no DSA util to transfer from (arch={}, b={batch_size}, \
s={s}); collect MSA/DSA data or set msa_dsa_scale_k against an available \
quant.",
self.dsa_architecture
))),
}
}
pub fn query_generation(
&self,
db: &PerfDatabase,
batch_size: u32,
s: u32,
) -> Result<PerformanceResult, AicError> {
let sol = self.sol_ms(db, batch_size, s, 0, false)?;
match db.database_mode {
DatabaseMode::Sol | DatabaseMode::SolFull => {
Ok(PerformanceResult::new(sol * self.scale_factor, Source::Sol))
}
DatabaseMode::Silicon => self.silicon_generation(db, batch_size, s).map(|latency| {
PerformanceResult::new(latency * self.scale_factor, Source::Silicon)
}),
DatabaseMode::Empirical => self.xop_generation(db, batch_size, s, sol),
DatabaseMode::Hybrid => match self.silicon_generation(db, batch_size, s) {
Ok(latency) => Ok(PerformanceResult::new(
latency * self.scale_factor,
Source::Silicon,
)),
Err(err) if err.is_missing_perf_data() => {
self.xop_generation(db, batch_size, s, sol)
}
Err(err) => Err(err),
},
}
}
fn xop_generation(
&self,
db: &PerfDatabase,
batch_size: u32,
s: u32,
sol: f64,
) -> Result<PerformanceResult, AicError> {
if !db.transfer_policy.contains(TransferKind::XOp) {
return Err(AicError::EmpiricalNotImplemented(
"MSA generation: cross-op transfer (xop) is disabled by the transfer \
policy and no MSA silicon data is available for this workload."
.to_string(),
));
}
let util = self.dsa_generation_util(db, batch_size, s);
match util {
Some(util) if util > 0.0 => {
let latency = sol / (util * self.dsa_scale_k);
db.note_provenance(crate::operators::util_empirical::ProvenanceTier::XOp);
Ok(PerformanceResult::new(
latency * self.scale_factor,
Source::Empirical,
))
}
_ => Err(AicError::EmpiricalNotImplemented(format!(
"MSA generation: no DSA util to transfer from (arch={}, b={batch_size}, \
s={s}); collect MSA/DSA data or set msa_dsa_scale_k against an available \
quant.",
self.dsa_architecture
))),
}
}
fn silicon_context(
&self,
db: &PerfDatabase,
b: u32,
s: u32,
prefix: u32,
) -> Result<f64, AicError> {
let flops = msa_sol_flops(&db.system_spec, self.gemm_quant_mode, self.fmha_quant_mode)?;
let spec = &db.system_spec;
let sol = move |c: &[f64]| {
msa_attention_sol_ms_with(
spec,
true,
c[3] as i128, c[2] as i128, c[1] as i128, c[0] as i128, self.num_kv_heads as i128,
self.hidden_size as i128,
self.head_dim as i128,
self.v_head_dim as i128,
self.index_n_heads as i128,
self.index_head_dim as i128,
self.index_topk as i128,
self.block_size as i128,
self.kv_cache_dtype,
self.fmha_quant_mode,
self.gemm_quant_mode,
flops,
)
};
db.msa.query_context(
b,
s,
prefix,
self.num_heads,
self.kv_cache_dtype,
self.fmha_quant_mode,
self.gemm_quant_mode,
MSA_ARCHITECTURE,
&sol,
)
}
fn silicon_generation(&self, db: &PerfDatabase, b: u32, s: u32) -> Result<f64, AicError> {
let flops = msa_sol_flops(&db.system_spec, self.gemm_quant_mode, self.fmha_quant_mode)?;
let spec = &db.system_spec;
let sol = move |c: &[f64]| {
msa_attention_sol_ms_with(
spec,
false,
c[1] as i128, c[2] as i128, 0, c[0] as i128, self.num_kv_heads as i128,
self.hidden_size as i128,
self.head_dim as i128,
self.v_head_dim as i128,
self.index_n_heads as i128,
self.index_head_dim as i128,
self.index_topk as i128,
self.block_size as i128,
self.kv_cache_dtype,
self.fmha_quant_mode,
self.gemm_quant_mode,
flops,
)
};
db.msa.query_generation(
b,
s,
self.num_heads,
self.kv_cache_dtype,
self.gemm_quant_mode,
MSA_ARCHITECTURE,
&sol,
)
}
fn sol_ms(
&self,
db: &PerfDatabase,
b: u32,
s: u32,
prefix: u32,
is_context: bool,
) -> Result<f64, AicError> {
msa_attention_sol_ms(
&db.system_spec,
is_context,
b as i128,
s as i128,
prefix as i128,
self.num_heads as i128,
self.num_kv_heads as i128,
self.hidden_size as i128,
self.head_dim as i128,
self.v_head_dim as i128,
self.index_n_heads as i128,
self.index_head_dim as i128,
self.index_topk as i128,
self.block_size as i128,
self.kv_cache_dtype,
self.fmha_quant_mode,
self.gemm_quant_mode,
)
}
fn dsa_context_util(&self, db: &PerfDatabase, b: u32, s: u32, prefix: u32) -> Option<f64> {
let dims = dsa_dims(&self.dsa_architecture);
let flops =
dsa_context_sol_flops(&db.system_spec, self.gemm_quant_mode, self.fmha_quant_mode)
.ok()?;
let sol = dsa_context_sol_ms(
&db.system_spec,
dims,
dims.index_topk,
self.kv_cache_dtype,
self.fmha_quant_mode,
self.gemm_quant_mode,
b as i64,
s as i64,
prefix as i64,
self.num_heads as i64,
false,
flops,
);
let probe = self.dsa_probe(dims.index_topk);
let silicon = probe
.query_context(&db.silicon_view(), b, s, prefix)
.ok()?
.latency_ms;
if sol > 0.0 && silicon > 0.0 {
Some(sol / silicon)
} else {
None
}
}
fn dsa_generation_util(&self, db: &PerfDatabase, b: u32, s: u32) -> Option<f64> {
let dims = dsa_dims(&self.dsa_architecture);
let flops = dsa_generation_sol_flops(&db.system_spec, self.gemm_quant_mode).ok()?;
let sol = dsa_generation_sol_ms(
&db.system_spec,
dims,
self.kv_cache_dtype,
self.gemm_quant_mode,
b as i64,
s as i64,
self.num_heads as i64,
flops,
);
let probe = self.dsa_probe(dims.index_topk);
let silicon = probe
.query_generation(&db.silicon_view(), b, s)
.ok()?
.latency_ms;
if sol > 0.0 && silicon > 0.0 {
Some(sol / silicon)
} else {
None
}
}
fn dsa_probe(&self, index_topk: i64) -> DsaModuleOp {
DsaModuleOp::new(
format!("{}_dsa_probe", self.name),
self.num_heads,
self.kv_cache_dtype,
self.fmha_quant_mode,
self.gemm_quant_mode,
self.dsa_architecture.clone(),
index_topk as u32,
)
}
}
#[derive(Clone, Copy)]
struct MsaSolFlops {
gemm: f64,
indexer_fp8: f64,
attn: f64,
}
fn msa_sol_flops(
spec: &SystemSpec,
gemm_quant: GemmQuantMode,
fmha_quant: FmhaQuantMode,
) -> Result<MsaSolFlops, AicError> {
Ok(MsaSolFlops {
gemm: quant_tc_flops(spec, gemm_quant.mapping())?,
indexer_fp8: quant_tc_flops(spec, FmhaQuantMode::Fp8.mapping())?,
attn: quant_tc_flops(spec, fmha_quant.mapping())?,
})
}
#[allow(clippy::too_many_arguments)]
fn msa_attention_sol_ms(
spec: &SystemSpec,
is_context: bool,
b: i128,
s: i128,
prefix: i128,
num_heads: i128,
num_kv_heads: i128,
hidden_size: i128,
head_dim: i128,
v_head_dim: i128,
index_n_heads: i128,
index_head_dim: i128,
index_topk: i128,
block_size: i128,
kv_quant: KvCacheQuantMode,
fmha_quant: FmhaQuantMode,
gemm_quant: GemmQuantMode,
) -> Result<f64, AicError> {
let flops = msa_sol_flops(spec, gemm_quant, fmha_quant)?;
Ok(msa_attention_sol_ms_with(
spec,
is_context,
b,
s,
prefix,
num_heads,
num_kv_heads,
hidden_size,
head_dim,
v_head_dim,
index_n_heads,
index_head_dim,
index_topk,
block_size,
kv_quant,
fmha_quant,
gemm_quant,
flops,
))
}
#[allow(clippy::too_many_arguments)]
fn msa_attention_sol_ms_with(
spec: &SystemSpec,
is_context: bool,
b: i128,
s: i128,
prefix: i128,
num_heads: i128,
num_kv_heads: i128,
hidden_size: i128,
head_dim: i128,
v_head_dim: i128,
index_n_heads: i128,
index_head_dim: i128,
index_topk: i128,
block_size: i128,
kv_quant: KvCacheQuantMode,
fmha_quant: FmhaQuantMode,
gemm_quant: GemmQuantMode,
flops: MsaSolFlops,
) -> f64 {
let qk_head_dim = head_dim;
let tokens = if is_context { b * s } else { b };
let full_s = if is_context { prefix + s } else { s };
let kv_len = if is_context { full_s } else { (s - 1).max(0) };
let gemm_ops = 2 * tokens * hidden_size * (num_heads * qk_head_dim)
+ 2 * tokens * hidden_size * (2 * num_kv_heads * head_dim)
+ 2 * tokens * (num_heads * v_head_dim) * hidden_size
+ 2 * tokens * hidden_size * (index_n_heads * index_head_dim);
let (pairs, score_len) = if is_context {
let pairs = if full_s <= index_topk {
b * (full_s * (full_s + 1) - prefix * (prefix + 1)) / 2
} else if prefix >= index_topk {
tokens * index_topk
} else {
let ramp = b * (index_topk * (index_topk + 1) - prefix * (prefix + 1)) / 2;
let sat = b * (full_s - index_topk) * index_topk;
ramp + sat
};
(pairs, full_s)
} else {
(tokens * kv_len.min(index_topk), kv_len)
};
let effective_kv = if is_context {
full_s.min(index_topk)
} else {
kv_len.min(index_topk)
};
let attention_ops = 2 * num_heads * (qk_head_dim + v_head_dim) * pairs;
let num_blocks = if score_len > index_topk {
(score_len + block_size - 1) / block_size
} else {
0
};
let indexer_ops = 2 * tokens * index_n_heads * index_head_dim * num_blocks;
let gemm_weight_elems = hidden_size * num_heads * qk_head_dim
+ hidden_size * 2 * num_kv_heads * head_dim
+ num_heads * v_head_dim * hidden_size
+ hidden_size * index_n_heads * index_head_dim;
let gemm_weight_bytes = gemm_weight_elems as f64 * gemm_quant.mapping().memory;
let kv_cache_bytes = (b * num_kv_heads * effective_kv * (qk_head_dim + v_head_dim)) as f64
* kv_quant.mapping().memory;
let indexer_cache_bytes = (b * num_blocks * index_n_heads * index_head_dim) as f64;
let q_io_bytes = (tokens * num_heads * qk_head_dim) as f64 * fmha_quant.mapping().memory * 2.0;
let total_mem = gemm_weight_bytes + kv_cache_bytes + indexer_cache_bytes + q_io_bytes;
let MsaSolFlops {
gemm: gemm_flops,
indexer_fp8: fp8_flops,
attn: attn_flops,
} = flops;
let sol_math = (gemm_ops as f64 / gemm_flops
+ indexer_ops as f64 / fp8_flops
+ attention_ops as f64 / attn_flops)
* 1000.0;
let sol_mem = total_mem / spec.gpu.mem_bw * 1000.0;
sol_math.max(sol_mem)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::common::enums::TransferPolicy;
use crate::perf_database::perf_interp::LeafValue;
use std::path::PathBuf;
const REPO_ROOT_HINT: &str = env!("CARGO_MANIFEST_DIR");
fn db(backend: &str, version: &str) -> PerfDatabase {
let systems_root = PathBuf::from(REPO_ROOT_HINT)
.join("../..")
.join("python/aisimulate/src/aiconfigurator_core/systems");
let mut db =
PerfDatabase::load(&systems_root, "b200_sxm", backend, version).expect("db must load");
db.database_mode = DatabaseMode::Hybrid;
db
}
fn msa_op() -> MsaModuleOp {
MsaModuleOp {
name: "msa".to_string(),
scale_factor: 1.0,
num_heads: 8,
num_kv_heads: 1,
hidden_size: 6144,
head_dim: 128,
v_head_dim: 128,
index_n_heads: 4,
index_head_dim: 128,
index_topk: 2048,
block_size: 128,
kv_cache_dtype: KvCacheQuantMode::Bfloat16,
fmha_quant_mode: FmhaQuantMode::Bfloat16,
gemm_quant_mode: GemmQuantMode::Bfloat16,
dsa_architecture: "GlmMoeDsaForCausalLM".to_string(),
dsa_scale_k: 1.0,
}
}
fn approx(a: f64, b: f64) {
assert!(
(a - b).abs() < 1e-9 * b.abs().max(1.0),
"expected {b}, got {a}"
);
}
#[test]
fn msa_xop_transfer_matches_python_oracles() {
for (backend, version, anchors) in [(
"sglang",
"0.5.14",
[
(1u32, 1024u32, 0u32, true),
(2, 3000, 512, true),
(8, 1025, 0, false),
(4, 7777, 0, false),
],
)] {
let db = db(backend, version);
let op = msa_op();
for (b, s, prefix, is_context) in anchors {
db.reset_provenance();
let result = if is_context {
op.query_context(&db, b, s, prefix)
} else {
op.query_generation(&db, b, s)
}
.unwrap_or_else(|e| panic!("{backend} b={b} s={s}: {e}"));
assert!(result.latency_ms.is_finite() && result.latency_ms > 0.0);
assert_eq!(result.source, Source::Empirical);
assert_eq!(
db.worst_provenance(),
crate::operators::util_empirical::ProvenanceTier::XOp
);
}
}
}
use crate::perf_database::dsa::{DsaGrids, DsaHeadGrid, DsaKey};
use std::collections::BTreeMap;
fn msa_key() -> DsaKey {
DsaKey {
architecture: MSA_ARCHITECTURE.to_string(),
fmha_quant: "bfloat16".to_string(),
kv_quant: "bfloat16".to_string(),
gemm_quant: "bfloat16".to_string(),
}
}
fn inject_msa_grids(db: &PerfDatabase) {
let mut ctx_head = DsaHeadGrid::new();
ctx_head
.entry(8)
.or_default()
.entry(0)
.or_default()
.entry(1024)
.or_default()
.insert(1, LeafValue::latency_only(10.0));
let mut gen_head = DsaHeadGrid::new();
gen_head
.entry(8)
.or_default()
.entry(0)
.or_default()
.entry(4097)
.or_default()
.insert(1, LeafValue::latency_only(0.5));
let context = DsaGrids {
by_keys: BTreeMap::from([(
msa_key(),
BTreeMap::from([("flashmla_kv".to_string(), ctx_head)]),
)]),
};
let generation = DsaGrids {
by_keys: BTreeMap::from([(
msa_key(),
BTreeMap::from([("flashmla_kv".to_string(), gen_head)]),
)]),
};
db.msa.inject_for_test(context, generation);
}
#[test]
fn msa_silicon_table_hit_prefers_silicon_over_xop() {
for mode in [DatabaseMode::Silicon, DatabaseMode::Hybrid] {
let mut db = db("vllm", "0.24.0");
db.database_mode = mode;
inject_msa_grids(&db);
let op = msa_op();
db.reset_provenance();
let ctx = op
.query_context(&db, 1, 1024, 0)
.expect("context silicon hit");
approx(ctx.latency_ms, 10.0);
assert_eq!(ctx.source, Source::Silicon, "{mode:?}");
let generation = op
.query_generation(&db, 1, 4097)
.expect("generation silicon hit");
approx(generation.latency_ms, 0.5);
assert_eq!(generation.source, Source::Silicon, "{mode:?}");
assert_eq!(
db.worst_provenance(),
crate::operators::util_empirical::ProvenanceTier::Silicon,
"{mode:?}"
);
}
}
#[test]
fn msa_missing_quant_slice_falls_back_to_xop_under_hybrid() {
let mut silicon = db("vllm", "0.24.0");
silicon.database_mode = DatabaseMode::Silicon;
inject_msa_grids(&silicon);
let mut op = msa_op();
op.kv_cache_dtype = KvCacheQuantMode::Fp8;
assert!(matches!(
op.query_context(&silicon, 1, 1024, 0),
Err(AicError::PerfDatabase(_))
));
let hybrid = db("vllm", "0.24.0"); inject_msa_grids(&hybrid);
let bf16_op = msa_op(); let mut fp8_kv = msa_op();
fp8_kv.kv_cache_dtype = KvCacheQuantMode::Fp8;
let result = fp8_kv.query_context(&hybrid, 1, 1024, 0);
match result {
Ok(r) => assert_eq!(r.source, Source::Empirical),
Err(AicError::EmpiricalNotImplemented(_)) => {} Err(other) => panic!("unexpected error: {other}"),
}
let hit = bf16_op
.query_context(&hybrid, 1, 1024, 0)
.expect("silicon hit");
assert_eq!(hit.source, Source::Silicon);
}
#[test]
fn msa_policy_and_silicon_contracts() {
let mut hybrid = db("sglang", "0.5.14");
hybrid.transfer_policy = TransferPolicy {
xshape: true,
xquant: true,
xprofile: false,
xop: false,
};
let op = msa_op();
assert!(matches!(
op.query_context(&hybrid, 1, 1024, 0),
Err(AicError::EmpiricalNotImplemented(_))
));
assert!(matches!(
op.query_generation(&hybrid, 8, 1025),
Err(AicError::EmpiricalNotImplemented(_))
));
let mut silicon = db("sglang", "0.5.14");
silicon.database_mode = DatabaseMode::Silicon;
assert!(matches!(
op.query_context(&silicon, 1, 1024, 0),
Err(AicError::PerfDatabase(_))
));
}
#[test]
fn msa_empirical_mode_never_reads_silicon() {
let op = msa_op();
let mut silicon = db("vllm", "0.24.0");
silicon.database_mode = DatabaseMode::Silicon;
inject_msa_grids(&silicon);
assert_eq!(
op.query_context(&silicon, 1, 1024, 0)
.expect("silicon ctx")
.source,
Source::Silicon
);
assert_eq!(
op.query_generation(&silicon, 1, 4097)
.expect("silicon gen")
.source,
Source::Silicon
);
let hybrid = db("vllm", "0.24.0"); inject_msa_grids(&hybrid);
assert_eq!(
op.query_context(&hybrid, 1, 1024, 0)
.expect("hybrid ctx")
.source,
Source::Silicon
);
let mut empirical = db("vllm", "0.24.0");
empirical.database_mode = DatabaseMode::Empirical;
inject_msa_grids(&empirical);
empirical.reset_provenance();
let ctx = op
.query_context(&empirical, 1, 1024, 0)
.expect("empirical ctx");
assert_eq!(ctx.source, Source::Empirical);
assert_eq!(
empirical.worst_provenance(),
crate::operators::util_empirical::ProvenanceTier::XOp
);
empirical.reset_provenance();
let generation = op
.query_generation(&empirical, 8, 1025)
.expect("empirical gen");
assert_eq!(generation.source, Source::Empirical);
assert_eq!(
empirical.worst_provenance(),
crate::operators::util_empirical::ProvenanceTier::XOp
);
assert!((f64::from(ctx.latency_ms) - 10.0).abs() > 1e-6);
}
}