use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
use std::sync::{Arc, Mutex, OnceLock};
use super::gemm::quant_tc_flops;
use super::perf_interp::{self, LeafValue, Node, OpInterpConfig};
use super::{SourceResolver, kernel_source_ok};
use crate::common::enums::{FmhaQuantMode, GemmQuantMode, KvCacheQuantMode};
use crate::common::error::AicError;
use crate::common::system_spec::SystemSpec;
use crate::config::{PerfDbSources, PerfSource};
use crate::operators::base::SolComponents;
use crate::perf_database::parquet_loader::PerfReader;
pub struct DsaTable {
data_root: PathBuf,
context_sources: Vec<PerfSource>,
generation_sources: Vec<PerfSource>,
context: OnceLock<Result<DsaGrids, AicError>>,
generation: OnceLock<Result<DsaGrids, AicError>>,
context_skip: OnceLock<Result<DsaGrids, AicError>>,
generation_skip: OnceLock<Result<DsaGrids, AicError>>,
context_skip_nodes: OnceLock<Result<NodeCache, AicError>>,
generation_skip_nodes: OnceLock<Result<NodeCache, AicError>>,
context_nodes: OnceLock<Result<NodeCache, AicError>>,
generation_nodes: OnceLock<Result<NodeCache, AicError>>,
source_resolver: Arc<SourceResolver>,
sparse: Mutex<BTreeMap<(String, u32), Arc<DsaSparseTables>>>,
}
pub type SparseGrid = BTreeMap<u32, BTreeMap<(u32, u32), f64>>;
#[derive(Debug, Default, PartialEq)]
pub struct DsaSparseTables {
pub mqa: SparseGrid,
pub topk_last: SparseGrid,
pub topk_flat: SparseGrid,
pub dsa_attn: SparseGrid,
}
pub fn dsa_sparse_file_prefix(architecture: &str) -> &'static str {
match architecture {
"DeepseekV32ForCausalLM" => "dsv32",
_ => "glm5",
}
}
pub(crate) struct NodeCache {
pub(crate) by_keys: BTreeMap<DsaKey, BTreeMap<String, Node>>,
}
pub type DsaHeadGrid = BTreeMap<u32, BTreeMap<u32, BTreeMap<u32, BTreeMap<u32, LeafValue>>>>;
pub struct DsaGrids {
pub by_keys: BTreeMap<DsaKey, BTreeMap<String, DsaHeadGrid>>,
}
pub(crate) fn select_dsa_backend<'a, T>(
by_backend: &'a BTreeMap<String, T>,
dsa_backend: &str,
) -> Option<&'a T> {
by_backend
.get(dsa_backend)
.or_else(|| by_backend.get("flashmla_kv"))
.or_else(|| by_backend.get("trtllm"))
.or_else(|| by_backend.values().next())
}
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub struct DsaKey {
pub architecture: String,
pub fmha_quant: String,
pub kv_quant: String,
pub gemm_quant: String,
}
pub(crate) struct DsaDims {
pub(crate) hidden_size: i64,
pub(crate) q_lora_rank: i64,
pub(crate) kv_lora_rank: i64,
pub(crate) qk_nope_head_dim: i64,
pub(crate) qk_rope_head_dim: i64,
pub(crate) v_head_dim: i64,
pub(crate) index_topk: i64,
pub(crate) index_head_dim: i64,
pub(crate) index_n_heads: i64,
}
const DSV32_DIMS: DsaDims = DsaDims {
hidden_size: 7168,
q_lora_rank: 1536,
kv_lora_rank: 512,
qk_nope_head_dim: 128,
qk_rope_head_dim: 64,
v_head_dim: 128,
index_topk: 2048,
index_head_dim: 128,
index_n_heads: 64,
};
const GLM_MOE_DSA_DIMS: DsaDims = DsaDims {
hidden_size: 6144,
q_lora_rank: 2048,
qk_nope_head_dim: 192,
kv_lora_rank: 512,
qk_rope_head_dim: 64,
v_head_dim: 256,
index_topk: 2048,
index_head_dim: 128,
index_n_heads: 32,
};
pub(crate) fn dsa_kernel_source_buckets(
kernel_source: &str,
kv_quant: &str,
) -> &'static [&'static str] {
if kv_quant == "bfloat16" {
return &["trtllm", "flashmla_kv"];
}
match kernel_source {
"sglang_dsa_indexer_trtllm" | "sglang_dsa_skip_indexer_trtllm" => &["trtllm"],
"sglang_dsa_indexer_flashmla_sparse" | "sglang_dsa_skip_indexer_flashmla_sparse" => {
&["flashmla_kv"]
}
"sglang_dsa_dense_mha_trtllm_ragged" => &["trtllm", "flashmla_kv"],
ks if ks.contains("trtllm") => &["trtllm"],
_ => &["flashmla_kv"],
}
}
pub(crate) fn dsa_dims(architecture: &str) -> &'static DsaDims {
match architecture {
"GlmMoeDsaForCausalLM" => &GLM_MOE_DSA_DIMS,
_ => &DSV32_DIMS,
}
}
impl DsaTable {
pub fn new(data_root: PathBuf) -> Self {
Self::with_sources(
data_root,
&Arc::new(SourceResolver::fixed(PerfDbSources::default())),
)
.expect("fixed-map resolution is infallible")
}
pub fn with_sources(
data_root: PathBuf,
resolver: &Arc<SourceResolver>,
) -> Result<Self, AicError> {
let context_sources =
resolver.sources_for("dsa_context_module_perf.parquet", &data_root)?;
let generation_sources =
resolver.sources_for("dsa_generation_module_perf.parquet", &data_root)?;
Ok(Self {
data_root,
context_sources,
generation_sources,
context: OnceLock::new(),
generation: OnceLock::new(),
context_skip: OnceLock::new(),
generation_skip: OnceLock::new(),
context_skip_nodes: OnceLock::new(),
generation_skip_nodes: OnceLock::new(),
context_nodes: OnceLock::new(),
generation_nodes: OnceLock::new(),
source_resolver: Arc::clone(resolver),
sparse: Mutex::new(BTreeMap::new()),
})
}
pub fn load_cp_sparse(
&self,
architecture: &str,
num_heads: u32,
) -> Result<Arc<DsaSparseTables>, AicError> {
let file_prefix = dsa_sparse_file_prefix(architecture);
let key = (file_prefix.to_string(), num_heads);
if let Some(tables) = self
.sparse
.lock()
.expect("dsa sparse cache poisoned")
.get(&key)
{
return Ok(Arc::clone(tables));
}
let mut tables = DsaSparseTables::default();
load_sparse_parquet(
&self.source_resolver.sources_for(
&format!("{file_prefix}_mqa_logits_module_perf.parquet"),
&self.data_root,
)?,
num_heads,
|_| SparseKind::Mqa,
&mut tables,
)?;
load_sparse_parquet(
&self.source_resolver.sources_for(
&format!("{file_prefix}_topk_module_perf.parquet"),
&self.data_root,
)?,
num_heads,
|score_mode| {
if score_mode == Some("flat") {
SparseKind::TopkFlat
} else {
SparseKind::TopkLast
}
},
&mut tables,
)?;
load_sparse_parquet(
&self.source_resolver.sources_for(
&format!("{file_prefix}_dsa_attn_module_perf.parquet"),
&self.data_root,
)?,
num_heads,
|_| SparseKind::DsaAttn,
&mut tables,
)?;
let arc = Arc::new(tables);
Ok(Arc::clone(
self.sparse
.lock()
.expect("dsa sparse cache poisoned")
.entry(key)
.or_insert(arc),
))
}
#[allow(clippy::too_many_arguments)]
pub fn query_context(
&self,
spec: &SystemSpec,
b: u32,
isl: u32,
num_heads: u32,
kv_quant: KvCacheQuantMode,
fmha_quant: FmhaQuantMode,
gemm_quant: GemmQuantMode,
architecture: &str,
prefix: u32,
index_topk: u32,
dsa_backend: &str,
skip_indexer: bool,
) -> Result<LeafValue, AicError> {
let flops = dsa_context_sol_flops(spec, gemm_quant, fmha_quant)?;
let nodes = if skip_indexer {
self.load_context_skip_nodes()?
} else {
self.load_context_nodes()?
};
let key = DsaKey {
architecture: architecture.to_string(),
fmha_quant: fmha_quant.name().to_string(),
kv_quant: kv_quant.name().to_string(),
gemm_quant: gemm_quant.name().to_string(),
};
let node = nodes
.by_keys
.get(&key)
.and_then(|by_backend| select_dsa_backend(by_backend, dsa_backend))
.ok_or_else(|| {
AicError::PerfDatabase(format!("context DSA module data missing for {key:?}"))
})?;
let dims = dsa_dims(architecture);
let topk = index_topk as i64;
let sol = move |c: &[f64]| {
dsa_context_sol_ms(
spec,
dims,
topk,
kv_quant,
fmha_quant,
gemm_quant,
c[3] as i64, c[2] as i64, c[1] as i64, c[0] as i64, skip_indexer,
flops,
)
};
let cfg = OpInterpConfig::grid(&["num_heads", "prefix", "seq_len", "batch"], &sol);
perf_interp::query_value(
&cfg,
node,
&[num_heads as f64, prefix as f64, isl as f64, b as f64],
)
}
#[allow(clippy::too_many_arguments)]
pub fn query_generation(
&self,
spec: &SystemSpec,
b: u32,
sequence_tokens: u32,
num_heads: u32,
kv_quant: KvCacheQuantMode,
fmha_quant: FmhaQuantMode,
gemm_quant: GemmQuantMode,
architecture: &str,
dsa_backend: &str,
skip_indexer: bool,
) -> Result<LeafValue, AicError> {
let flops = dsa_generation_sol_flops(spec, gemm_quant)?;
let nodes = if skip_indexer {
self.load_generation_skip_nodes()?
} else {
self.load_generation_nodes()?
};
let key = DsaKey {
architecture: architecture.to_string(),
fmha_quant: fmha_quant.name().to_string(),
kv_quant: kv_quant.name().to_string(),
gemm_quant: gemm_quant.name().to_string(),
};
let node = nodes
.by_keys
.get(&key)
.and_then(|by_backend| select_dsa_backend(by_backend, dsa_backend))
.ok_or_else(|| missing("generation DSA module", &self.data_root, format!("{key:?}")))?;
let dims = dsa_dims(architecture);
let sol = move |c: &[f64]| {
dsa_generation_sol_ms(
spec,
dims,
kv_quant,
gemm_quant,
c[1] as i64, c[2] as i64, c[0] as i64, flops,
)
};
let cfg = OpInterpConfig::grid(&["num_heads", "batch", "seq_len"], &sol);
perf_interp::query_value(
&cfg,
node,
&[num_heads as f64, b as f64, sequence_tokens as f64],
)
}
pub fn context_raw_slice(
&self,
key: &DsaKey,
dsa_backend: &str,
skip_indexer: bool,
) -> Result<&DsaHeadGrid, AicError> {
let grids = if skip_indexer {
self.load_context_skip()?
} else {
self.load_context()?
};
grids
.by_keys
.get(key)
.and_then(|by_backend| select_dsa_backend(by_backend, dsa_backend))
.ok_or_else(|| {
missing(
"raw context DSA module",
&self.data_root,
format!("{key:?} (dsa_backend={dsa_backend})"),
)
})
}
pub fn generation_raw_slice(
&self,
key: &DsaKey,
dsa_backend: &str,
skip_indexer: bool,
) -> Result<&DsaHeadGrid, AicError> {
let grids = if skip_indexer {
self.load_generation_skip()?
} else {
self.load_generation()?
};
grids
.by_keys
.get(key)
.and_then(|by_backend| select_dsa_backend(by_backend, dsa_backend))
.ok_or_else(|| {
missing(
"raw generation DSA module",
&self.data_root,
format!("{key:?} (dsa_backend={dsa_backend})"),
)
})
}
pub fn has_context_skip_rows(&self) -> Result<bool, AicError> {
match self.load_context_skip() {
Ok(grids) => Ok(!grids.by_keys.is_empty()),
Err(AicError::PerfDatabase(msg)) if msg.contains(NO_DSA_ROWS_PREFIX) => Ok(false),
Err(err) => Err(err),
}
}
pub fn has_generation_skip_rows(&self) -> Result<bool, AicError> {
match self.load_generation_skip() {
Ok(grids) => Ok(!grids.by_keys.is_empty()),
Err(AicError::PerfDatabase(msg)) if msg.contains(NO_DSA_ROWS_PREFIX) => Ok(false),
Err(err) => Err(err),
}
}
fn load_context(&self) -> Result<&DsaGrids, AicError> {
let cell = self
.context
.get_or_init(|| load_dsa_parquet(&self.context_sources, false, false));
cell.as_ref().map_err(clone_err)
}
fn load_context_skip(&self) -> Result<&DsaGrids, AicError> {
let cell = self
.context_skip
.get_or_init(|| load_dsa_parquet(&self.context_sources, false, true));
cell.as_ref().map_err(clone_err)
}
fn load_generation(&self) -> Result<&DsaGrids, AicError> {
let cell = self
.generation
.get_or_init(|| load_dsa_parquet(&self.generation_sources, true, false));
cell.as_ref().map_err(clone_err)
}
fn load_generation_skip(&self) -> Result<&DsaGrids, AicError> {
let cell = self
.generation_skip
.get_or_init(|| load_dsa_parquet(&self.generation_sources, true, true));
cell.as_ref().map_err(clone_err)
}
fn load_context_nodes(&self) -> Result<&NodeCache, AicError> {
let cell = self.context_nodes.get_or_init(|| {
let grids = self.load_context()?;
Ok(build_context_nodes(grids))
});
cell.as_ref().map_err(clone_err)
}
fn load_context_skip_nodes(&self) -> Result<&NodeCache, AicError> {
let cell = self.context_skip_nodes.get_or_init(|| {
let grids = self.load_context_skip()?;
Ok(build_context_nodes(grids))
});
cell.as_ref().map_err(clone_err)
}
fn load_generation_nodes(&self) -> Result<&NodeCache, AicError> {
let cell = self.generation_nodes.get_or_init(|| {
let grids = self.load_generation()?;
Ok(build_generation_nodes(grids))
});
cell.as_ref().map_err(clone_err)
}
fn load_generation_skip_nodes(&self) -> Result<&NodeCache, AicError> {
let cell = self.generation_skip_nodes.get_or_init(|| {
let grids = self.load_generation_skip()?;
Ok(build_generation_nodes(grids))
});
cell.as_ref().map_err(clone_err)
}
}
pub(crate) fn build_context_nodes(grids: &DsaGrids) -> NodeCache {
let mut by_keys: BTreeMap<DsaKey, BTreeMap<String, Node>> = BTreeMap::new();
for (key, by_backend) in &grids.by_keys {
let backends = by_keys.entry(key.clone()).or_default();
for (backend, by_heads) in by_backend {
let node = backends.entry(backend.clone()).or_insert_with(Node::branch);
for (&n, by_step) in by_heads {
for (&step, by_isl) in by_step {
for (&isl, by_batch) in by_isl {
for (&bb, &leaf) in by_batch {
node.insert_value(&[n, step, isl, bb], leaf);
}
}
}
}
}
}
NodeCache { by_keys }
}
pub(crate) fn build_generation_nodes(grids: &DsaGrids) -> NodeCache {
let mut by_keys: BTreeMap<DsaKey, BTreeMap<String, Node>> = BTreeMap::new();
for (key, by_backend) in &grids.by_keys {
let backends = by_keys.entry(key.clone()).or_default();
for (backend, by_heads) in by_backend {
let node = backends.entry(backend.clone()).or_insert_with(Node::branch);
for (&n, by_step) in by_heads {
for (&step, by_isl) in by_step {
for (&isl, by_batch) in by_isl {
let seq = isl + step;
for (&bb, &leaf) in by_batch {
node.insert_value(&[n, bb, seq], leaf);
}
}
}
}
}
}
NodeCache { by_keys }
}
#[derive(Clone, Copy)]
pub(crate) struct DsaSolFlops {
pub gemm: f64,
pub indexer_fp8: f64,
pub attn: f64,
}
pub(crate) fn dsa_context_sol_flops(
spec: &SystemSpec,
gemm_quant: GemmQuantMode,
fmha_quant: FmhaQuantMode,
) -> Result<DsaSolFlops, AicError> {
Ok(DsaSolFlops {
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())?,
})
}
pub(crate) fn dsa_generation_sol_flops(
spec: &SystemSpec,
gemm_quant: GemmQuantMode,
) -> Result<DsaSolFlops, AicError> {
Ok(DsaSolFlops {
gemm: quant_tc_flops(spec, gemm_quant.mapping())?,
indexer_fp8: quant_tc_flops(spec, FmhaQuantMode::Fp8.mapping())?,
attn: quant_tc_flops(spec, FmhaQuantMode::Bfloat16.mapping())?,
})
}
fn indexer_cache_entry_bytes(index_head_dim: i64) -> i64 {
index_head_dim + ((index_head_dim + 127) / 128) * 4
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn dsa_context_sol(
spec: &SystemSpec,
dims: &DsaDims,
index_topk: i64,
kv_quant: KvCacheQuantMode,
fmha_quant: FmhaQuantMode,
gemm_quant: GemmQuantMode,
b: i64,
s: i64,
prefix: i64,
num_heads: i64,
skip_indexer: bool,
flops: DsaSolFlops,
) -> SolComponents {
let (hidden, q_lora, kv_lora) = (dims.hidden_size, dims.q_lora_rank, dims.kv_lora_rank);
let (inh, ihd) = (dims.index_n_heads, dims.index_head_dim);
let qk_head_dim = dims.qk_nope_head_dim + dims.qk_rope_head_dim;
let attn_head_dim = kv_lora + dims.qk_rope_head_dim;
let v_dim = dims.v_head_dim;
let (b, s, prefix, num_heads) = (b as i128, s as i128, prefix as i128, num_heads as i128);
let (hidden, q_lora, kv_lora) = (hidden as i128, q_lora as i128, kv_lora as i128);
let (inh, ihd, topk) = (inh as i128, ihd as i128, index_topk as i128);
let (qk_head_dim, attn_head_dim, v_dim) =
(qk_head_dim as i128, attn_head_dim as i128, v_dim as i128);
let (qk_nope, qk_rope) = (dims.qk_nope_head_dim as i128, dims.qk_rope_head_dim as i128);
let full_s = s + prefix;
let tokens = b * s;
let proj_out = q_lora + kv_lora + qk_rope + ihd;
let gemm_group_ops = 2 * tokens * hidden * proj_out
+ 2 * tokens * q_lora * (num_heads * qk_head_dim)
+ 2 * tokens * q_lora * (inh * ihd)
+ 2 * tokens * hidden * inh
+ 2 * tokens * (num_heads * v_dim) * hidden
+ 2 * num_heads * tokens * qk_nope * kv_lora
+ 2 * num_heads * tokens * kv_lora * v_dim;
let indexer_logits_ops = if skip_indexer || full_s <= topk {
0
} else {
2 * tokens * inh * ihd * full_s
};
let effective_kv = full_s.min(topk);
let total_kv_pairs = if full_s <= topk {
b * (full_s * (full_s + 1) - prefix * (prefix + 1)) / 2
} else if prefix >= topk {
tokens * topk
} else {
let ramp = b * (topk * (topk + 1) - prefix * (prefix + 1)) / 2;
let sat = b * (full_s - topk) * topk;
ramp + sat
};
let sparse_attn_ops = 2 * num_heads * (attn_head_dim + kv_lora) * total_kv_pairs;
let gemm_weight_elems = hidden * proj_out
+ q_lora * num_heads * qk_head_dim
+ q_lora * inh * ihd
+ hidden * inh
+ num_heads * v_dim * hidden;
let gemm_weight_bytes = gemm_weight_elems as f64 * gemm_quant.mapping().memory;
let kv_cache_bytes =
(b * num_heads * effective_kv * attn_head_dim) as f64 * kv_quant.mapping().memory;
let indexer_cache_bytes = if skip_indexer || full_s <= topk {
0.0
} else {
(b * full_s * indexer_cache_entry_bytes(dims.index_head_dim) as i128) 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 DsaSolFlops {
gemm: gemm_flops,
indexer_fp8: indexer_fp8_flops,
attn: attn_flops,
} = flops;
let sol_math = (gemm_group_ops as f64 / gemm_flops
+ indexer_logits_ops as f64 / indexer_fp8_flops
+ sparse_attn_ops as f64 / attn_flops)
* 1000.0;
let sol_mem = total_mem / spec.gpu.mem_bw * 1000.0;
SolComponents::new(sol_math, sol_mem)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn dsa_context_sol_ms(
spec: &SystemSpec,
dims: &DsaDims,
index_topk: i64,
kv_quant: KvCacheQuantMode,
fmha_quant: FmhaQuantMode,
gemm_quant: GemmQuantMode,
b: i64,
s: i64,
prefix: i64,
num_heads: i64,
skip_indexer: bool,
flops: DsaSolFlops,
) -> f64 {
dsa_context_sol(
spec,
dims,
index_topk,
kv_quant,
fmha_quant,
gemm_quant,
b,
s,
prefix,
num_heads,
skip_indexer,
flops,
)
.time_ms()
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn dsa_generation_sol(
spec: &SystemSpec,
dims: &DsaDims,
kv_quant: KvCacheQuantMode,
gemm_quant: GemmQuantMode,
b: i64,
s: i64,
num_heads: i64,
flops: DsaSolFlops,
) -> SolComponents {
let (b, s, num_heads) = (b as i128, s as i128, num_heads as i128);
let (hidden, q_lora, kv_lora) = (
dims.hidden_size as i128,
dims.q_lora_rank as i128,
dims.kv_lora_rank as i128,
);
let (inh, ihd, topk) = (
dims.index_n_heads as i128,
dims.index_head_dim as i128,
dims.index_topk as i128,
);
let (qk_nope, qk_rope, v_dim) = (
dims.qk_nope_head_dim as i128,
dims.qk_rope_head_dim as i128,
dims.v_head_dim as i128,
);
let qk_head_dim = qk_nope + qk_rope;
let attn_head_dim = kv_lora + qk_rope;
let tokens = b;
let proj_out = q_lora + kv_lora + qk_rope + ihd;
let effective_kv = s.min(topk);
let gemm_group_ops = 2 * tokens * hidden * proj_out
+ 2 * tokens * q_lora * num_heads * qk_head_dim
+ 2 * tokens * q_lora * inh * ihd
+ 2 * tokens * hidden * inh
+ 2 * tokens * num_heads * v_dim * hidden
+ 2 * num_heads * tokens * qk_nope * kv_lora
+ 2 * num_heads * tokens * kv_lora * v_dim;
let indexer_logits_ops = 2 * tokens * inh * ihd * s;
let sparse_attn_ops = 2 * tokens * num_heads * (attn_head_dim + kv_lora) * effective_kv;
let gemm_weight_elems = hidden * proj_out
+ q_lora * num_heads * qk_head_dim
+ q_lora * inh * ihd
+ hidden * inh
+ num_heads * v_dim * hidden;
let gemm_weight_bytes = gemm_weight_elems as f64 * gemm_quant.mapping().memory;
let indexer_cache_bytes =
(b * s * indexer_cache_entry_bytes(dims.index_head_dim) as i128) as f64;
let kv_cache_bytes = (b * effective_kv * attn_head_dim) as f64 * kv_quant.mapping().memory;
let total_mem = gemm_weight_bytes + indexer_cache_bytes + kv_cache_bytes;
let DsaSolFlops {
gemm: gemm_flops,
indexer_fp8: indexer_fp8_flops,
attn: attn_flops,
} = flops;
let sol_math = (gemm_group_ops as f64 / gemm_flops
+ indexer_logits_ops as f64 / indexer_fp8_flops
+ sparse_attn_ops as f64 / attn_flops)
* 1000.0;
let sol_mem = total_mem / spec.gpu.mem_bw * 1000.0;
SolComponents::new(sol_math, sol_mem)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn dsa_generation_sol_ms(
spec: &SystemSpec,
dims: &DsaDims,
kv_quant: KvCacheQuantMode,
gemm_quant: GemmQuantMode,
b: i64,
s: i64,
num_heads: i64,
flops: DsaSolFlops,
) -> f64 {
dsa_generation_sol(spec, dims, kv_quant, gemm_quant, b, s, num_heads, flops).time_ms()
}
const NO_DSA_ROWS_PREFIX: &str = "no DSA module rows loaded";
pub(crate) fn load_dsa_parquet(
sources: &[PerfSource],
collapse_isl_step_to_seq: bool,
want_skip_rows: bool,
) -> Result<DsaGrids, AicError> {
let mut by_keys: BTreeMap<DsaKey, BTreeMap<String, DsaHeadGrid>> = BTreeMap::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 arch_col = reader.col("architecture")?;
let mla_dtype_col = reader.col("mla_dtype")?;
let kv_cache_dtype_col = reader.col("kv_cache_dtype")?;
let gemm_type_col = reader.col("gemm_type")?;
let num_heads_col = reader.col("num_heads")?;
let batch_size_col = reader.col("batch_size")?;
let isl_col = reader.col("isl")?;
let step_col = reader.col("step")?;
let latency_col = reader.col("latency")?;
let power_col = reader.col_optional("power");
let op_name_col = reader.col_optional("op_name");
let ks_col = reader.col_optional("kernel_source");
let mut source_values: BTreeMap<(DsaKey, String, u32, u32, u32, u32), LeafValue> =
BTreeMap::new();
for row in reader.rows()? {
let row = row?;
if !kernel_source_ok(source.kernel_sources(), ks_col, &row)? {
continue;
}
let is_skip_row = row
.str_optional(op_name_col)?
.unwrap_or("")
.contains("skip_indexer");
if is_skip_row != want_skip_rows {
continue;
}
let key = DsaKey {
architecture: row.str_owned(arch_col)?,
fmha_quant: row.str_owned(mla_dtype_col)?,
kv_quant: row.str_owned(kv_cache_dtype_col)?,
gemm_quant: row.str_owned(gemm_type_col)?,
};
let ks_name = row.str_optional(ks_col)?.unwrap_or("").to_string();
let buckets = dsa_kernel_source_buckets(&ks_name, &key.kv_quant);
let (step, isl) = if collapse_isl_step_to_seq {
(0, row.u32(isl_col)? + row.u32(step_col)?)
} else {
(row.u32(step_col)?, row.u32(isl_col)?)
};
let num_heads = row.u32(num_heads_col)?;
let batch_size = row.u32(batch_size_col)?;
let latency = row.f64(latency_col)?;
let power = row.f64_optional(power_col)?.unwrap_or(0.0);
for dsa_backend in buckets {
source_values.insert(
(
key.clone(),
dsa_backend.to_string(),
num_heads,
step,
isl,
batch_size,
),
LeafValue::with_power(latency, power),
);
}
}
for ((key, dsa_backend, num_heads, step, isl, batch_size), leaf) in source_values {
by_keys
.entry(key)
.or_default()
.entry(dsa_backend)
.or_default()
.entry(num_heads)
.or_default()
.entry(step)
.or_default()
.entry(isl)
.or_default()
.entry(batch_size)
.or_insert(leaf);
}
}
if !any_source || by_keys.is_empty() {
return Err(AicError::PerfDatabase(format!(
"{NO_DSA_ROWS_PREFIX} from {} source(s) (first: {})",
sources.len(),
sources
.first()
.map(|s| s.path().display().to_string())
.unwrap_or_default()
)));
}
Ok(DsaGrids { by_keys })
}
enum SparseKind {
Mqa,
TopkLast,
TopkFlat,
DsaAttn,
}
fn load_sparse_parquet(
sources: &[PerfSource],
num_heads: u32,
kind_of: impl Fn(Option<&str>) -> SparseKind,
tables: &mut DsaSparseTables,
) -> Result<(), AicError> {
for source in sources {
let path = source.path();
if !path.exists() {
continue;
}
let reader = PerfReader::open(path)?;
let batch_size_col = reader.col("batch_size")?;
let isl_col = reader.col("isl")?;
let step_col = reader.col("step")?;
let latency_col = reader.col("latency")?;
let num_heads_col = reader.col_optional("num_heads");
let score_mode_col = reader.col_optional("score_mode");
let ks_col = reader.col_optional("kernel_source");
let mut source_values: BTreeMap<(u8, u32, u32, u32), f64> = BTreeMap::new();
for row in reader.rows()? {
let row = row?;
if !kernel_source_ok(source.kernel_sources(), ks_col, &row)? {
continue;
}
if let Some(col) = num_heads_col {
if row.u32(col)? != num_heads {
continue;
}
}
let kind = kind_of(row.str_optional(score_mode_col)?);
source_values.insert(
(
kind as u8,
row.u32(batch_size_col)?,
row.u32(isl_col)?,
row.u32(step_col)?,
),
row.f64(latency_col)?,
);
}
for ((kind, bs, isl, step), latency) in source_values {
let grid = match kind {
k if k == SparseKind::Mqa as u8 => &mut tables.mqa,
k if k == SparseKind::TopkLast as u8 => &mut tables.topk_last,
k if k == SparseKind::TopkFlat as u8 => &mut tables.topk_flat,
_ => &mut tables.dsa_attn,
};
grid.entry(bs)
.or_default()
.entry((isl, step))
.or_insert(latency);
}
}
Ok(())
}
pub fn bs_slice(by_bs: &SparseGrid, b: u32) -> Option<&BTreeMap<(u32, u32), f64>> {
if let Some(exact) = by_bs.get(&b) {
return Some(exact);
}
let mut best: Option<(u64, &BTreeMap<(u32, u32), f64>)> = None;
for (&bs, grid) in by_bs {
let dist = (i64::from(bs) - i64::from(b)).unsigned_abs();
if best.map_or(true, |(d, _)| dist < d) {
best = Some((dist, grid));
}
}
best.map(|(_, grid)| grid)
}
pub fn lookup_2d(
table: &BTreeMap<(u32, u32), f64>,
isl: u32,
step: u32,
) -> Result<Option<f64>, AicError> {
if table.is_empty() {
return Ok(None);
}
let max_isl = table.keys().map(|&(i, _)| i).max().expect("non-empty");
if isl > max_isl {
return Err(AicError::PerfDatabase(format!(
"DSA CP: isl={isl} exceeds the collected sparse-kernel grid \
(max isl={max_isl}); mqa/topk scale super-linearly with isl, so \
clamping the isl axis would silently under-estimate. Re-collect with \
AIC_CHUNKED_PREFILL_SIZE >= {isl} \
(docs/CONTEXT_PARALLEL_DSA_MODELING.md \u{a7}9.1)."
)));
}
let mut use_isl = None;
for &(i, _) in table.keys() {
let dist = (i64::from(i) - i64::from(isl)).unsigned_abs();
if use_isl.map_or(true, |(d, _)| dist < d) {
use_isl = Some((dist, i));
}
}
let use_isl = use_isl.expect("non-empty").1;
if let Some(&exact) = table.get(&(use_isl, step)) {
return Ok(Some(exact));
}
let steps: Vec<u32> = table
.range((use_isl, u32::MIN)..=(use_isl, u32::MAX))
.map(|(&(_, st), _)| st)
.collect();
let Some((&first, &last)) = steps.first().zip(steps.last()) else {
return Ok(None);
};
let lo = steps
.iter()
.rev()
.find(|&&st| st <= step)
.copied()
.unwrap_or(first);
let hi = steps
.iter()
.find(|&&st| st >= step)
.copied()
.unwrap_or(last);
if lo == hi {
return Ok(Some(table[&(use_isl, lo)]));
}
let a = table[&(use_isl, lo)];
let b = table[&(use_isl, hi)];
Ok(Some(
a + (b - a) * f64::from(step - lo) / f64::from(hi - lo),
))
}
pub(crate) fn missing(table: &str, data_root: &Path, descriptor: String) -> AicError {
AicError::PerfDatabase(format!(
"{table} data missing for {descriptor} at {}",
data_root.display()
))
}
pub(crate) 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")
}
fn b200_sxm_spec() -> SystemSpec {
let systems_yaml = PathBuf::from(REPO_ROOT_HINT)
.join("../..")
.join("python/aisimulate/src/aiconfigurator_core/systems/b200_sxm.yaml");
SystemSpec::load(&systems_yaml).expect("b200_sxm.yaml must parse")
}
const INDEX_TOPK: u32 = 2048;
fn approx_rel(got: f64, want: f64) {
assert!(
((got - want) / want).abs() < 1e-9,
"rust {got} vs python {want}"
);
}
#[test]
fn dsa_context_module_exact_hit() {
let table = DsaTable::new(b200_vllm_data_root());
let spec = b200_sxm_spec();
let latency = table
.query_context(
&spec,
1,
1,
128,
KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
"DeepseekV32ForCausalLM",
0,
INDEX_TOPK,
"trtllm",
false,
)
.expect("DSA context query must succeed")
.latency;
assert!(
(latency - 1.0972).abs() < 1e-6,
"expected recorded latency, got {latency}"
);
}
#[test]
fn dsa_within_file_duplicates_last_row_wins() {
let data_root = PathBuf::from(REPO_ROOT_HINT)
.join("../..")
.join("python/aisimulate/src/aiconfigurator_core/systems/data/b300_sxm/vllm/0.19.0");
let systems_yaml = PathBuf::from(REPO_ROOT_HINT)
.join("../..")
.join("python/aisimulate/src/aiconfigurator_core/systems/b300_sxm.yaml");
let spec = SystemSpec::load(&systems_yaml).expect("b300_sxm.yaml must parse");
let table = DsaTable::new(data_root);
let latency = table
.query_context(
&spec,
1,
8192,
128,
KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
"DeepseekV32ForCausalLM",
0,
INDEX_TOPK,
"trtllm",
false,
)
.expect("DSA context query must succeed")
.latency;
approx_rel(latency, 7.756);
}
#[test]
fn dsa_unknown_architecture_errors() {
let table = DsaTable::new(b200_vllm_data_root());
let spec = b200_sxm_spec();
let err = table
.query_context(
&spec,
1,
1024,
128,
KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
"NotAnArchitecture",
0,
INDEX_TOPK,
"trtllm",
false,
)
.unwrap_err();
assert!(matches!(err, AicError::PerfDatabase(_)));
}
#[test]
fn dsa_context_matches_python_v2_engine() {
let table = DsaTable::new(b200_vllm_data_root());
let spec = b200_sxm_spec();
let q = |b: u32, s: u32, prefix: u32, heads: u32, arch: &str| {
table
.query_context(
&spec,
b,
s,
heads,
KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
arch,
prefix,
INDEX_TOPK,
"trtllm",
false,
)
.unwrap()
.latency
};
let dsv32 = "DeepseekV32ForCausalLM";
let glm = "GlmMoeDsaForCausalLM";
approx_rel(q(4, 2048, 0, 128, dsv32), 7.6471);
approx_rel(q(2, 2560, 0, 128, dsv32), 4.9806);
approx_rel(q(3, 1024, 0, 128, dsv32), 3.0913);
approx_rel(q(1, 128, 64, 16, glm), 1.2492999999999999);
approx_rel(q(1, 65536, 0, 128, dsv32), 93.51797494885695);
approx_rel(q(1, 2048, 4096, 128, dsv32), 3.270467722991338);
}
fn grid(rows: &[(u32, u32, u32, f64)]) -> SparseGrid {
let mut g = SparseGrid::new();
for &(bs, isl, step, lat) in rows {
g.entry(bs).or_default().insert((isl, step), lat);
}
g
}
#[test]
fn cp_lookup_2d_exact_step_interp_clamp_and_empty() {
let t = grid(&[
(1, 4096, 0, 100.0),
(1, 4096, 1024, 200.0),
(1, 8192, 0, 400.0),
]);
let t = &t[&1];
assert_eq!(lookup_2d(t, 4096, 0).unwrap(), Some(100.0)); assert_eq!(lookup_2d(t, 4096, 512).unwrap(), Some(150.0)); assert_eq!(lookup_2d(t, 4096, 4096).unwrap(), Some(200.0)); assert_eq!(lookup_2d(&BTreeMap::new(), 4096, 0).unwrap(), None); }
#[test]
fn cp_lookup_2d_fails_loud_on_out_of_grid_isl() {
let t = grid(&[(1, 4096, 0, 100.0), (1, 8192, 0, 400.0)]);
let err = lookup_2d(&t[&1], 16384, 0).unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("DSA CP: isl=16384 exceeds the collected sparse-kernel grid")
&& msg.contains("max isl=8192")
&& msg.contains("AIC_CHUNKED_PREFILL_SIZE >= 16384"),
"unexpected message: {msg}"
);
}
#[test]
fn cp_bs_slice_exact_nearest_and_empty() {
let g = grid(&[(1, 2048, 0, 25.0), (8, 2048, 0, 90.0)]);
assert_eq!(bs_slice(&g, 8).unwrap()[&(2048, 0)], 90.0); assert_eq!(bs_slice(&g, 6).unwrap()[&(2048, 0)], 90.0); assert_eq!(bs_slice(&g, 3).unwrap()[&(2048, 0)], 25.0); assert!(bs_slice(&SparseGrid::new(), 1).is_none()); }
#[test]
fn cp_sparse_absent_files_load_empty() {
let tmp = tempfile::tempdir().expect("tmpdir");
let table = DsaTable::new(tmp.path().to_path_buf());
let sparse = table
.load_cp_sparse("GlmMoeDsaForCausalLM", 64)
.expect("absent files are not a load error");
assert_eq!(*sparse, DsaSparseTables::default());
}
fn write_sparse_parquet(
path: &Path,
with_score_mode: bool,
rows: &[(i64, i64, i64, i64, f64, &str)],
) {
use parquet::data_type::{ByteArray, ByteArrayType, DoubleType, Int64Type};
use parquet::file::properties::WriterProperties;
use parquet::file::writer::SerializedFileWriter;
use parquet::schema::parser::parse_message_type;
let schema = if with_score_mode {
"message schema {
REQUIRED INT64 num_heads;
REQUIRED INT64 batch_size;
REQUIRED INT64 isl;
REQUIRED INT64 step;
REQUIRED DOUBLE latency;
REQUIRED BINARY score_mode (UTF8);
}"
} else {
"message schema {
REQUIRED INT64 num_heads;
REQUIRED INT64 batch_size;
REQUIRED INT64 isl;
REQUIRED INT64 step;
REQUIRED DOUBLE latency;
}"
};
let schema = Arc::new(parse_message_type(schema).expect("schema must parse"));
let file = std::fs::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");
let int_cols: [Vec<i64>; 4] = [
rows.iter().map(|r| r.0).collect(),
rows.iter().map(|r| r.1).collect(),
rows.iter().map(|r| r.2).collect(),
rows.iter().map(|r| r.3).collect(),
];
for values in &int_cols {
let mut col = rg.next_column().expect("next col").expect("int col");
col.typed::<Int64Type>()
.write_batch(values, None, None)
.expect("write ints");
col.close().expect("close col");
}
let latencies: Vec<f64> = rows.iter().map(|r| r.4).collect();
let mut col = rg.next_column().expect("next col").expect("latency col");
col.typed::<DoubleType>()
.write_batch(&latencies, None, None)
.expect("write latency");
col.close().expect("close col");
if with_score_mode {
let modes: Vec<ByteArray> = rows.iter().map(|r| ByteArray::from(r.5)).collect();
let mut col = rg.next_column().expect("next col").expect("score col");
col.typed::<ByteArrayType>()
.write_batch(&modes, None, None)
.expect("write score");
col.close().expect("close col");
}
rg.close().expect("close row group");
writer.close().expect("close writer");
}
#[test]
fn cp_sparse_loader_matches_python_loader() {
let tmp = tempfile::tempdir().expect("tmpdir");
write_sparse_parquet(
&tmp.path().join("glm5_mqa_logits_module_perf.parquet"),
false,
&[
(64, 1, 2048, 0, 25.0, ""),
(64, 1, 16384, 0, 1600.0, ""),
(64, 1, 16384, 0, 1601.5, ""), (32, 1, 2048, 0, 99.0, ""), (64, 2, 2048, 128, 30.0, ""), ],
);
write_sparse_parquet(
&tmp.path().join("glm5_topk_module_perf.parquet"),
true,
&[
(64, 1, 16384, 0, 800.0, "top_last"),
(64, 1, 2048, 0, 190.0, "top_last"),
(64, 1, 2048, 0, 100.0, "flat"),
],
);
let table = DsaTable::new(tmp.path().to_path_buf());
let sparse = table
.load_cp_sparse("GlmMoeDsaForCausalLM", 64)
.expect("sparse tables must load");
assert_eq!(
sparse.mqa,
grid(&[
(1, 2048, 0, 25.0),
(1, 16384, 0, 1601.5),
(2, 2048, 128, 30.0)
])
);
assert_eq!(
sparse.topk_last,
grid(&[(1, 16384, 0, 800.0), (1, 2048, 0, 190.0)])
);
assert_eq!(sparse.topk_flat, grid(&[(1, 2048, 0, 100.0)]));
assert!(sparse.dsa_attn.is_empty());
}
#[test]
fn dsa_generation_matches_python_v2_engine() {
let table = DsaTable::new(b200_vllm_data_root());
let spec = b200_sxm_spec();
let q = |b: u32, s: u32| {
table
.query_generation(
&spec,
b,
s,
128,
KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
"DeepseekV32ForCausalLM",
"trtllm",
false,
)
.unwrap()
.latency
};
approx_rel(q(16, 4097), 0.2698);
approx_rel(q(16, 3000), 0.261390380859375);
approx_rel(q(24, 4097), 0.27545);
approx_rel(q(16, 300000), 0.5491372293318538);
}
fn write_dsa_module_parquet(path: &Path, rows: &[(&str, &str, f64)]) {
let rows_kv: Vec<(&str, &str, &str, i64, f64)> = rows
.iter()
.map(|r| (r.0, r.1, "bfloat16", 1024, r.2))
.collect();
write_dsa_module_parquet_rows(path, &rows_kv)
}
fn write_dsa_module_parquet_rows(path: &Path, rows: &[(&str, &str, &str, i64, f64)]) {
use parquet::data_type::{ByteArray, ByteArrayType, DoubleType, Int64Type};
use parquet::file::properties::WriterProperties;
use parquet::file::writer::SerializedFileWriter;
use parquet::schema::parser::parse_message_type;
let schema = "message schema {
REQUIRED BINARY op_name (UTF8);
REQUIRED BINARY kernel_source (UTF8);
REQUIRED BINARY architecture (UTF8);
REQUIRED BINARY mla_dtype (UTF8);
REQUIRED BINARY kv_cache_dtype (UTF8);
REQUIRED BINARY gemm_type (UTF8);
REQUIRED INT64 num_heads;
REQUIRED INT64 batch_size;
REQUIRED INT64 isl;
REQUIRED INT64 step;
REQUIRED DOUBLE latency;
}";
let schema = Arc::new(parse_message_type(schema).expect("schema must parse"));
let file = std::fs::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");
let str_cols: [Vec<ByteArray>; 6] = [
rows.iter().map(|r| ByteArray::from(r.0)).collect(),
rows.iter().map(|r| ByteArray::from(r.1)).collect(),
rows.iter()
.map(|_| ByteArray::from("DeepseekV32ForCausalLM"))
.collect(),
rows.iter().map(|_| ByteArray::from("bfloat16")).collect(),
rows.iter().map(|r| ByteArray::from(r.2)).collect(),
rows.iter().map(|_| ByteArray::from("bfloat16")).collect(),
];
for values in &str_cols {
let mut col = rg.next_column().expect("next col").expect("str col");
col.typed::<ByteArrayType>()
.write_batch(values, None, None)
.expect("write str");
col.close().expect("close col");
}
let int_cols: [Vec<i64>; 4] = [
rows.iter().map(|_| 128).collect(), rows.iter().map(|_| 1).collect(), rows.iter().map(|r| r.3).collect(), rows.iter().map(|_| 0).collect(), ];
for values in &int_cols {
let mut col = rg.next_column().expect("next col").expect("int col");
col.typed::<Int64Type>()
.write_batch(values, None, None)
.expect("write ints");
col.close().expect("close col");
}
let latencies: Vec<f64> = rows.iter().map(|r| r.4).collect();
let mut col = rg.next_column().expect("next col").expect("latency col");
col.typed::<DoubleType>()
.write_batch(&latencies, None, None)
.expect("write latency");
col.close().expect("close col");
rg.close().expect("close row group");
writer.close().expect("close writer");
}
#[test]
fn dsa_full_and_skip_indexer_rows_do_not_blend_and_backend_slices_split() {
let tmp = tempfile::tempdir().expect("tmpdir");
write_dsa_module_parquet(
&tmp.path().join("dsa_context_module_perf.parquet"),
&[
("dsa_context_module", "trtllm_gen", 1.0),
("dsa_context_module_skip_indexer", "trtllm_gen", 9.0),
("dsa_context_module", "default", 5.0),
],
);
let table = DsaTable::new(tmp.path().to_path_buf());
let spec = b200_sxm_spec();
let q = |dsa_backend: &str| {
table
.query_context(
&spec,
1,
1024,
128,
KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
"DeepseekV32ForCausalLM",
0,
INDEX_TOPK,
dsa_backend,
false,
)
.expect("query must succeed")
.latency
};
assert_eq!(q("trtllm"), 5.0);
assert_eq!(q("flashmla_kv"), 5.0);
}
#[test]
fn skip_probe_reports_absence_but_propagates_load_errors() {
let absent = tempfile::tempdir().expect("tmpdir");
write_dsa_module_parquet(
&absent.path().join("dsa_context_module_perf.parquet"),
&[("dsa_context_module", "default", 1.0)],
);
let table = DsaTable::new(absent.path().to_path_buf());
assert!(
!table
.has_context_skip_rows()
.expect("full-only is absence, not an error")
);
let present = tempfile::tempdir().expect("tmpdir");
write_dsa_module_parquet(
&present.path().join("dsa_context_module_perf.parquet"),
&[
("dsa_context_module", "default", 1.0),
("dsa_context_module_skip_indexer", "default", 0.5),
],
);
let table = DsaTable::new(present.path().to_path_buf());
assert!(
table
.has_context_skip_rows()
.expect("both variants present")
);
let corrupt = tempfile::tempdir().expect("tmpdir");
write_dsa_module_parquet_rows(
&corrupt.path().join("dsa_context_module_perf.parquet"),
&[
("dsa_context_module", "default", "bfloat16", 1024, 1.0),
(
"dsa_context_module_skip_indexer",
"default",
"bfloat16",
-1,
0.5,
),
],
);
let table = DsaTable::new(corrupt.path().to_path_buf());
let err = table
.has_context_skip_rows()
.expect_err("a malformed skip row must propagate, not read as absence");
assert!(
!err.to_string().contains(NO_DSA_ROWS_PREFIX),
"the propagated error must not be the absence outcome: {err}"
);
let spec = b200_sxm_spec();
table
.query_context(
&spec,
1,
1024,
128,
KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
"DeepseekV32ForCausalLM",
0,
INDEX_TOPK,
"trtllm",
false,
)
.expect("full-variant query must survive a malformed skip row");
}
#[test]
fn dsa_fp8_rows_bucket_by_executed_kernel_name() {
let tmp = tempfile::tempdir().expect("tmpdir");
write_dsa_module_parquet_rows(
&tmp.path().join("dsa_context_module_perf.parquet"),
&[
(
"dsa_context_module",
"sglang_dsa_indexer_trtllm",
"fp8",
4096,
1.0,
),
(
"dsa_context_module",
"sglang_dsa_indexer_flashmla_sparse",
"fp8",
4096,
5.0,
),
(
"dsa_context_module",
"sglang_dsa_dense_mha_trtllm_ragged",
"fp8",
1024,
7.0,
),
],
);
let table = DsaTable::new(tmp.path().to_path_buf());
let spec = b200_sxm_spec();
let q = |dsa_backend: &str, isl: u32| {
table
.query_context(
&spec,
1,
isl,
128,
KvCacheQuantMode::Fp8,
FmhaQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
"DeepseekV32ForCausalLM",
0,
INDEX_TOPK,
dsa_backend,
false,
)
.expect("query must succeed")
.latency
};
assert_eq!(q("trtllm", 4096), 1.0);
assert_eq!(q("flashmla_kv", 4096), 5.0);
assert_eq!(q("trtllm", 1024), 7.0);
assert_eq!(q("flashmla_kv", 1024), 7.0);
}
fn write_dsa_generation_parquet(path: &Path, rows: &[(i64, i64, f64)]) {
use parquet::data_type::{ByteArray, ByteArrayType, DoubleType, Int64Type};
use parquet::file::properties::WriterProperties;
use parquet::file::writer::SerializedFileWriter;
use parquet::schema::parser::parse_message_type;
let schema = "message schema {
REQUIRED BINARY op_name (UTF8);
REQUIRED BINARY kernel_source (UTF8);
REQUIRED BINARY architecture (UTF8);
REQUIRED BINARY mla_dtype (UTF8);
REQUIRED BINARY kv_cache_dtype (UTF8);
REQUIRED BINARY gemm_type (UTF8);
REQUIRED INT64 num_heads;
REQUIRED INT64 batch_size;
REQUIRED INT64 isl;
REQUIRED INT64 step;
REQUIRED DOUBLE latency;
}";
let schema = Arc::new(parse_message_type(schema).expect("schema must parse"));
let file = std::fs::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");
let str_cols: [Vec<ByteArray>; 6] = [
rows.iter()
.map(|_| ByteArray::from("dsa_generation_module"))
.collect(),
rows.iter().map(|_| ByteArray::from("default")).collect(),
rows.iter()
.map(|_| ByteArray::from("DeepseekV32ForCausalLM"))
.collect(),
rows.iter().map(|_| ByteArray::from("bfloat16")).collect(),
rows.iter().map(|_| ByteArray::from("bfloat16")).collect(),
rows.iter().map(|_| ByteArray::from("bfloat16")).collect(),
];
for values in &str_cols {
let mut col = rg.next_column().expect("next col").expect("str col");
col.typed::<ByteArrayType>()
.write_batch(values, None, None)
.expect("write str");
col.close().expect("close col");
}
let int_cols: [Vec<i64>; 4] = [
rows.iter().map(|_| 128).collect(), rows.iter().map(|_| 1).collect(), rows.iter().map(|r| r.0).collect(), rows.iter().map(|r| r.1).collect(), ];
for values in &int_cols {
let mut col = rg.next_column().expect("next col").expect("int col");
col.typed::<Int64Type>()
.write_batch(values, None, None)
.expect("write ints");
col.close().expect("close col");
}
let latencies: Vec<f64> = rows.iter().map(|r| r.2).collect();
let mut col = rg.next_column().expect("next col").expect("latency col");
col.typed::<DoubleType>()
.write_batch(&latencies, None, None)
.expect("write latency");
col.close().expect("close col");
rg.close().expect("close row group");
writer.close().expect("close writer");
}
#[test]
fn dsa_generation_same_seq_ties_resolve_last_file_row() {
let tmp = tempfile::tempdir().expect("tmpdir");
write_dsa_generation_parquet(
&tmp.path().join("dsa_generation_module_perf.parquet"),
&[(80, 20, 111.0), (90, 10, 222.0)],
);
let table = DsaTable::new(tmp.path().to_path_buf());
let spec = b200_sxm_spec();
let got = table
.query_generation(
&spec,
1, 100, 128, KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
"DeepseekV32ForCausalLM",
"flashmla_kv",
false,
)
.expect("query must succeed")
.latency;
assert_eq!(got, 222.0);
}
#[test]
fn dsa_backend_fallback_resolves_single_backend_files() {
let tmp = tempfile::tempdir().expect("tmpdir");
write_dsa_module_parquet(
&tmp.path().join("dsa_context_module_perf.parquet"),
&[("dsa_context_module", "trtllm_gen", 3.5)],
);
let table = DsaTable::new(tmp.path().to_path_buf());
let spec = b200_sxm_spec();
let got = table
.query_context(
&spec,
1,
1024,
128,
KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
"DeepseekV32ForCausalLM",
0,
INDEX_TOPK,
"flashmla_kv", false,
)
.expect("query must succeed")
.latency;
assert_eq!(got, 3.5);
}
#[test]
fn dsa_context_energy_matches_python_oracle() {
use crate::perf_database::energy_test_fixtures::{Col, energy_test_spec, write_parquet};
let tmp = tempfile::tempdir().expect("tmpdir");
write_parquet(
&tmp.path().join("dsa_context_module_perf.parquet"),
&[
Col::Str("architecture", vec!["DeepseekV32ForCausalLM"; 2]),
Col::Str("mla_dtype", vec!["bfloat16"; 2]),
Col::Str("kv_cache_dtype", vec!["bfloat16"; 2]),
Col::Str("gemm_type", vec!["bfloat16"; 2]),
Col::I64("num_heads", vec![128, 128]),
Col::I64("batch_size", vec![1, 1]),
Col::I64("isl", vec![1024, 2048]),
Col::I64("step", vec![0, 0]),
Col::Str("op_name", vec!["dsa_context_module"; 2]),
Col::Str("kernel_source", vec!["sglang_dsa_indexer_trtllm"; 2]),
Col::F64("latency", vec![1.0, 3.0]),
Col::F64("power", vec![100.0, 200.0]),
],
);
let table = DsaTable::new(tmp.path().to_path_buf());
let spec = energy_test_spec();
let v = table
.query_context(
&spec,
1,
1536,
128,
KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Bfloat16,
GemmQuantMode::Bfloat16,
"DeepseekV32ForCausalLM",
0,
INDEX_TOPK,
"trtllm",
false,
)
.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
);
}
}