use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
use std::sync::OnceLock;
use super::gemm::quant_tc_flops;
use super::interpolation::Grid3;
use super::perf_interp::{self, LeafValue, Node, OpInterpConfig};
use super::{SourceResolver, kernel_source_ok};
use crate::common::enums::{FmhaQuantMode, 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 AttentionTable {
data_root: PathBuf,
system_spec: SystemSpec,
context_sources: Vec<PerfSource>,
generation_sources: Vec<PerfSource>,
encoder_sources: Vec<PerfSource>,
context: OnceLock<Result<ContextGrids, AicError>>,
generation: OnceLock<Result<GenerationGrids, AicError>>,
encoder: OnceLock<Result<EncoderGrids, AicError>>,
}
struct ContextGrids {
by_keys: BTreeMap<ContextKey, Node>,
}
struct GenerationGrids {
by_keys: BTreeMap<GenerationKey, Node>,
}
struct EncoderGrids {
by_keys: BTreeMap<EncoderKey, Node>,
}
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord)]
struct ContextKey {
fmha_quant: String,
kv_quant: String,
n_kv_lookup: u32,
head_size: u32,
window_size: u32,
}
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord)]
struct GenerationKey {
kv_quant: String,
n_kv_lookup: u32,
head_size: u32,
window_size: u32,
}
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord)]
struct EncoderKey {
fmha_quant: String,
head_size: u32,
}
impl AttentionTable {
pub fn new(data_root: PathBuf, system_spec: SystemSpec) -> Self {
Self::with_sources(
data_root,
system_spec,
&SourceResolver::fixed(PerfDbSources::default()),
)
.expect("fixed-map resolution is infallible")
}
pub fn with_sources(
data_root: PathBuf,
system_spec: SystemSpec,
resolver: &SourceResolver,
) -> Result<Self, AicError> {
let context_sources = resolver.sources_for("context_attention_perf.parquet", &data_root)?;
let generation_sources =
resolver.sources_for("generation_attention_perf.parquet", &data_root)?;
let encoder_sources = resolver.sources_for("encoder_attention_perf.parquet", &data_root)?;
Ok(Self {
data_root,
system_spec,
context_sources,
generation_sources,
encoder_sources,
context: OnceLock::new(),
generation: OnceLock::new(),
encoder: OnceLock::new(),
})
}
pub fn query_context(
&self,
b: u32,
full_seq_tokens: u32,
n: u32,
n_kv: u32,
head_size: u32,
window_size: u32,
kv_quant: KvCacheQuantMode,
fmha_quant: FmhaQuantMode,
) -> Result<LeafValue, AicError> {
let attn_flops = quant_tc_flops(&self.system_spec, fmha_quant.mapping())?;
let grids = self.load_context()?;
let key = ContextKey {
fmha_quant: fmha_quant.name().to_string(),
kv_quant: kv_quant.name().to_string(),
n_kv_lookup: normalize_kv(n, n_kv),
head_size,
window_size,
};
let node = grids
.by_keys
.get(&key)
.ok_or_else(|| missing_key(&self.data_root, &key))?;
let spec = &self.system_spec;
let n_kv_lookup = key.n_kv_lookup;
let sol = move |c: &[f64]| {
context_attention_sol_ms(
spec,
n_kv_lookup,
head_size,
window_size,
kv_quant,
c[0],
c[1],
c[2],
attn_flops,
)
};
let cfg = OpInterpConfig::grid_sqrt_axis(&["num_heads", "seq_len", "batch"], 1, &sol);
perf_interp::query_value(&cfg, node, &[n as f64, full_seq_tokens as f64, b as f64])
}
pub fn query_generation(
&self,
b: u32,
kv_seq_tokens: u32,
n: u32,
n_kv: u32,
head_size: u32,
window_size: u32,
kv_quant: KvCacheQuantMode,
) -> Result<LeafValue, AicError> {
let attn_flops = generation_attn_flops(&self.system_spec, kv_quant)?;
let grids = self.load_generation()?;
let key = GenerationKey {
kv_quant: kv_quant.name().to_string(),
n_kv_lookup: normalize_kv(n, n_kv),
head_size,
window_size,
};
let node = grids
.by_keys
.get(&key)
.ok_or_else(|| missing_gen_key(&self.data_root, &key))?;
let spec = &self.system_spec;
let n_kv_lookup = key.n_kv_lookup;
let sol = move |c: &[f64]| {
generation_attention_sol_ms(
spec,
n_kv_lookup,
head_size,
window_size,
kv_quant,
c[0],
c[1],
c[2],
attn_flops,
)
};
let cfg = OpInterpConfig::grid(&["num_heads", "batch", "seq_len"], &sol);
let s = kv_seq_tokens;
let s_min = ((s as f64 * 0.9) as u32).max(1);
let s_max = ((s as f64 * 1.1) as u32).max(s_min);
const SAMPLE_CNT: u32 = 5;
let mut latency_sum = 0.0_f64;
let mut energy_sum = 0.0_f64;
for i in 0..SAMPLE_CNT {
let s_i = s_min
+ ((u64::from(s_max - s_min) * u64::from(i)) / u64::from(SAMPLE_CNT - 1)) as u32;
let sample = perf_interp::query_value(&cfg, node, &[n as f64, b as f64, s_i as f64])?;
latency_sum += sample.latency;
energy_sum += sample.energy;
}
Ok(LeafValue {
latency: latency_sum / SAMPLE_CNT as f64,
power: 0.0, energy: energy_sum / SAMPLE_CNT as f64,
})
}
pub fn query_encoder(
&self,
b: u32,
s: u32,
n: u32,
head_size: u32,
fmha_quant: FmhaQuantMode,
) -> Result<LeafValue, AicError> {
let attn_flops = quant_tc_flops(&self.system_spec, fmha_quant.mapping())?;
let grids = self.load_encoder()?;
let key = EncoderKey {
fmha_quant: fmha_quant.name().to_string(),
head_size,
};
let node = grids
.by_keys
.get(&key)
.ok_or_else(|| missing_encoder_key(&self.data_root, &key))?;
let spec = &self.system_spec;
let sol = move |c: &[f64]| {
encoder_attention_sol_ms(spec, head_size, c[0], c[1], c[2], attn_flops)
};
let cfg = OpInterpConfig::grid_sqrt_axis(&["num_heads", "seq_len", "batch"], 1, &sol);
perf_interp::query_value(&cfg, node, &[n as f64, s as f64, b as f64])
}
pub fn context_points(
&self,
fmha_quant: FmhaQuantMode,
kv_quant: KvCacheQuantMode,
n_kv_lookup: u32,
head_size: u32,
window_size: u32,
) -> Result<Vec<(Vec<f64>, f64)>, AicError> {
let grids = self.load_context()?;
let key = ContextKey {
fmha_quant: fmha_quant.name().to_string(),
kv_quant: kv_quant.name().to_string(),
n_kv_lookup,
head_size,
window_size,
};
let node = grids
.by_keys
.get(&key)
.ok_or_else(|| missing_key(&self.data_root, &key))?;
let points = perf_interp::node_points(node);
if points.is_empty() {
return Err(missing_key(&self.data_root, &key));
}
Ok(points)
}
pub fn context_head_sizes(
&self,
fmha_quant: FmhaQuantMode,
kv_quant: KvCacheQuantMode,
n_kv_lookup: u32,
) -> Result<Vec<u32>, AicError> {
let grids = self.load_context()?;
let fmha = fmha_quant.name();
let kv = kv_quant.name();
let mut sizes: Vec<u32> = Vec::new();
for key in grids.by_keys.keys() {
if key.fmha_quant == fmha
&& key.kv_quant == kv
&& key.n_kv_lookup == n_kv_lookup
&& !sizes.contains(&key.head_size)
{
sizes.push(key.head_size);
}
}
if sizes.is_empty() {
return Err(AicError::PerfDatabase(format!(
"context attention data missing for fmha={fmha}, kv={kv}, \
n_kv={n_kv_lookup} at {}",
self.data_root.display()
)));
}
Ok(sizes)
}
pub fn generation_points(
&self,
kv_quant: KvCacheQuantMode,
n_kv_lookup: u32,
head_size: u32,
window_size: u32,
) -> Result<Vec<(Vec<f64>, f64)>, AicError> {
let grids = self.load_generation()?;
let key = GenerationKey {
kv_quant: kv_quant.name().to_string(),
n_kv_lookup,
head_size,
window_size,
};
let node = grids
.by_keys
.get(&key)
.ok_or_else(|| missing_gen_key(&self.data_root, &key))?;
let points = perf_interp::node_points(node);
if points.is_empty() {
return Err(missing_gen_key(&self.data_root, &key));
}
Ok(points)
}
pub fn generation_head_sizes(
&self,
kv_quant: KvCacheQuantMode,
n_kv_lookup: u32,
) -> Result<Vec<u32>, AicError> {
let grids = self.load_generation()?;
let kv = kv_quant.name();
let mut sizes: Vec<u32> = Vec::new();
for key in grids.by_keys.keys() {
if key.kv_quant == kv
&& key.n_kv_lookup == n_kv_lookup
&& !sizes.contains(&key.head_size)
{
sizes.push(key.head_size);
}
}
if sizes.is_empty() {
return Err(AicError::PerfDatabase(format!(
"generation attention data missing for kv={kv}, n_kv={n_kv_lookup} at {}",
self.data_root.display()
)));
}
Ok(sizes)
}
pub fn encoder_points(
&self,
fmha_quant: FmhaQuantMode,
head_size: u32,
) -> Result<Vec<(Vec<f64>, f64)>, AicError> {
let grids = self.load_encoder()?;
let key = EncoderKey {
fmha_quant: fmha_quant.name().to_string(),
head_size,
};
let node = grids
.by_keys
.get(&key)
.ok_or_else(|| missing_encoder_key(&self.data_root, &key))?;
let points = perf_interp::node_points(node);
if points.is_empty() {
return Err(missing_encoder_key(&self.data_root, &key));
}
Ok(points)
}
fn load_context(&self) -> Result<&ContextGrids, AicError> {
let cell = self.context.get_or_init(|| {
let raw = load_context_parquet(&self.context_sources)?;
Ok(ContextGrids {
by_keys: raw
.into_iter()
.map(|(k, g)| (k, grid3_to_node(&g)))
.collect(),
})
});
cell.as_ref().map_err(clone_err)
}
fn load_generation(&self) -> Result<&GenerationGrids, AicError> {
let cell = self.generation.get_or_init(|| {
let mut raw = load_generation_parquet(&self.generation_sources)?;
clamp_generation_attention_grids_to_sol(&self.system_spec, &mut raw);
Ok(GenerationGrids {
by_keys: raw
.into_iter()
.map(|(k, g)| (k, grid3_to_node(&g)))
.collect(),
})
});
cell.as_ref().map_err(clone_err)
}
fn load_encoder(&self) -> Result<&EncoderGrids, AicError> {
let cell = self.encoder.get_or_init(|| {
let raw = load_encoder_parquet(&self.encoder_sources)?;
Ok(EncoderGrids {
by_keys: raw
.into_iter()
.map(|(k, g)| (k, grid3_to_node(&g)))
.collect(),
})
});
cell.as_ref().map_err(clone_err)
}
}
fn normalize_kv(n: u32, n_kv: u32) -> u32 {
if n_kv == n { 0 } else { n_kv }
}
fn grid3_to_node(grid: &Grid3<LeafValue>) -> Node {
let mut node = Node::branch();
for (&x, by_y) in grid {
for (&y, by_z) in by_y {
for (&z, &leaf) in by_z {
node.insert_value(&[x, y, z], leaf);
}
}
}
node
}
fn load_context_parquet(
sources: &[PerfSource],
) -> Result<BTreeMap<ContextKey, Grid3<LeafValue>>, AicError> {
let mut by_keys: BTreeMap<ContextKey, Grid3<LeafValue>> = 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 batch_size_col = reader.col("batch_size")?;
let isl_col = reader.col("isl")?;
let num_heads_col = reader.col("num_heads")?;
let num_kv_col = reader.col("num_key_value_heads")?;
let head_dim_col = reader.col("head_dim")?;
let attn_dtype_col = reader.col("attn_dtype")?;
let kv_cache_dtype_col = reader.col("kv_cache_dtype")?;
let latency_col = reader.col("latency")?;
let power_col = reader.col_optional("power");
let window_size_col = reader.col_optional("window_size");
let ks_col = reader.col_optional("kernel_source");
for row in reader.rows()? {
let row = row?;
if !kernel_source_ok(source.kernel_sources(), ks_col, &row)? {
continue;
}
let num_heads = row.u32(num_heads_col)?;
let num_kv = row.u32(num_kv_col)?;
let key = ContextKey {
fmha_quant: row.str_owned(attn_dtype_col)?,
kv_quant: row.str_owned(kv_cache_dtype_col)?,
n_kv_lookup: normalize_kv(num_heads, num_kv),
head_size: row.u32(head_dim_col)?,
window_size: row.u32_optional(window_size_col)?.unwrap_or(0),
};
let latency = row.f64(latency_col)?;
let power = row.f64_optional(power_col)?.unwrap_or(0.0);
by_keys
.entry(key)
.or_default()
.entry(num_heads)
.or_default()
.entry(row.u32(isl_col)?)
.or_default()
.entry(row.u32(batch_size_col)?)
.or_insert(LeafValue::with_power(latency, power));
}
}
if !any_source || by_keys.is_empty() {
return Err(AicError::PerfDatabase(format!(
"no context-attention rows loaded from {} source(s) (first: {})",
sources.len(),
sources
.first()
.map(|s| s.path().display().to_string())
.unwrap_or_default()
)));
}
Ok(by_keys)
}
fn load_generation_parquet(
sources: &[PerfSource],
) -> Result<BTreeMap<GenerationKey, Grid3<LeafValue>>, AicError> {
let mut by_keys: BTreeMap<GenerationKey, Grid3<LeafValue>> = 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 batch_size_col = reader.col("batch_size")?;
let isl_col = reader.col("isl")?;
let num_heads_col = reader.col("num_heads")?;
let num_kv_col = reader.col("num_key_value_heads")?;
let head_dim_col = reader.col("head_dim")?;
let kv_cache_dtype_col = reader.col("kv_cache_dtype")?;
let step_col = reader.col("step")?;
let latency_col = reader.col("latency")?;
let power_col = reader.col_optional("power");
let window_size_col = reader.col_optional("window_size");
let ks_col = reader.col_optional("kernel_source");
for row in reader.rows()? {
let row = row?;
if !kernel_source_ok(source.kernel_sources(), ks_col, &row)? {
continue;
}
let num_heads = row.u32(num_heads_col)?;
let num_kv = row.u32(num_kv_col)?;
let key = GenerationKey {
kv_quant: row.str_owned(kv_cache_dtype_col)?,
n_kv_lookup: normalize_kv(num_heads, num_kv),
head_size: row.u32(head_dim_col)?,
window_size: row.u32_optional(window_size_col)?.unwrap_or(0),
};
let sequence_tokens = row.u32(isl_col)? + row.u32(step_col)?;
let latency = row.f64(latency_col)?;
let power = row.f64_optional(power_col)?.unwrap_or(0.0);
by_keys
.entry(key)
.or_default()
.entry(num_heads)
.or_default()
.entry(row.u32(batch_size_col)?)
.or_default()
.entry(sequence_tokens)
.or_insert(LeafValue::with_power(latency, power));
}
}
if !any_source || by_keys.is_empty() {
return Err(AicError::PerfDatabase(format!(
"no generation-attention rows loaded from {} source(s) (first: {})",
sources.len(),
sources
.first()
.map(|s| s.path().display().to_string())
.unwrap_or_default()
)));
}
Ok(by_keys)
}
#[allow(clippy::too_many_arguments)]
#[allow(clippy::too_many_arguments)]
pub(crate) fn context_attention_sol(
spec: &SystemSpec,
n_kv_lookup: u32,
head_size: u32,
window_size: u32,
kv_quant: KvCacheQuantMode,
n: f64,
s: f64,
b: f64,
attn_flops: f64,
) -> SolComponents {
let h = head_size as f64;
let w = window_size as f64;
let n_kv = if n_kv_lookup == 0 {
n
} else {
n_kv_lookup as f64
};
let ops = if window_size > 0 && s > w {
2.0 * b * s * w * n * h * 2.0
} else {
2.0 * b * (s * s) * n * h * 2.0 / 2.0
};
let mem_bytes =
2.0 * b * (n * s * h + n * s * h) + kv_quant.mapping().memory * b * (2.0 * n_kv * s * h);
let sol_math = ops / attn_flops * 1000.0;
let sol_mem = mem_bytes / spec.gpu.mem_bw * 1000.0;
SolComponents::new(sol_math, sol_mem)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn context_attention_sol_ms(
spec: &SystemSpec,
n_kv_lookup: u32,
head_size: u32,
window_size: u32,
kv_quant: KvCacheQuantMode,
n: f64,
s: f64,
b: f64,
attn_flops: f64,
) -> f64 {
context_attention_sol(
spec,
n_kv_lookup,
head_size,
window_size,
kv_quant,
n,
s,
b,
attn_flops,
)
.time_ms()
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn context_attention_sol_with_prefix(
spec: &SystemSpec,
b: f64,
s: f64,
prefix: f64,
n: f64,
n_kv: f64,
head_size: u32,
window_size: u32,
kv_quant: KvCacheQuantMode,
attn_flops: f64,
) -> SolComponents {
let h = head_size as f64;
let w = window_size as f64;
let full_s = s + prefix;
let ops = if window_size > 0 && full_s > w {
2.0 * b * (full_s - prefix) * w * n * h * 2.0
} else {
2.0 * b * (full_s * full_s - prefix * prefix) * n * h * 2.0 / 2.0
};
let mem_bytes = 2.0 * b * (n * (full_s - prefix) * h + n * (full_s - prefix) * h)
+ kv_quant.mapping().memory * b * (2.0 * n_kv * full_s * h);
let sol_math = ops / attn_flops * 1000.0;
let sol_mem = mem_bytes / spec.gpu.mem_bw * 1000.0;
SolComponents::new(sol_math, sol_mem)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn context_attention_sol_with_prefix_ms(
spec: &SystemSpec,
b: f64,
s: f64,
prefix: f64,
n: f64,
n_kv: f64,
head_size: u32,
window_size: u32,
kv_quant: KvCacheQuantMode,
attn_flops: f64,
) -> f64 {
context_attention_sol_with_prefix(
spec,
b,
s,
prefix,
n,
n_kv,
head_size,
window_size,
kv_quant,
attn_flops,
)
.time_ms()
}
pub(crate) fn encoder_attention_sol(
spec: &SystemSpec,
head_size: u32,
n: f64,
s: f64,
b: f64,
attn_flops: f64,
) -> SolComponents {
let h = head_size as f64;
let ops = 2.0 * b * s * s * n * h * 2.0; let mem_bytes = 2.0 * b * (3.0 * n * s * h + n * s * h); let sol_math = ops / attn_flops * 1000.0;
let sol_mem = mem_bytes / spec.gpu.mem_bw * 1000.0;
SolComponents::new(sol_math, sol_mem)
}
pub(crate) fn encoder_attention_sol_ms(
spec: &SystemSpec,
head_size: u32,
n: f64,
s: f64,
b: f64,
attn_flops: f64,
) -> f64 {
encoder_attention_sol(spec, head_size, n, s, b, attn_flops).time_ms()
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn generation_attention_sol(
spec: &SystemSpec,
n_kv_lookup: u32,
head_size: u32,
window_size: u32,
kv_quant: KvCacheQuantMode,
n: f64,
b: f64,
s: f64,
attn_flops: f64,
) -> SolComponents {
let n_kv = if n_kv_lookup == 0 {
n
} else {
n_kv_lookup as f64
};
let kv_len = if window_size > 0 {
(s - 1.0).min(window_size as f64)
} else {
s - 1.0
};
let h = head_size as f64;
let kv_mem = kv_quant.mapping().memory;
let ops = 2.0 * b * n * h * 2.0 * kv_len;
let mem_bytes = b * (n * h * 2.0 + 2.0 * n_kv * kv_len * h * kv_mem + n * h * 2.0);
let sol_math = ops / attn_flops * 1000.0;
let sol_mem = mem_bytes / spec.gpu.mem_bw * 1000.0;
SolComponents::new(sol_math, sol_mem)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn generation_attention_sol_ms(
spec: &SystemSpec,
n_kv_lookup: u32,
head_size: u32,
window_size: u32,
kv_quant: KvCacheQuantMode,
n: f64,
b: f64,
s: f64,
attn_flops: f64,
) -> f64 {
generation_attention_sol(
spec,
n_kv_lookup,
head_size,
window_size,
kv_quant,
n,
b,
s,
attn_flops,
)
.time_ms()
}
pub(crate) fn generation_attn_mode(spec: &SystemSpec, kv_quant: KvCacheQuantMode) -> FmhaQuantMode {
let has_fp8_mma = spec.gpu.sm_version.is_some_and(|sm| sm >= 89);
if kv_quant == KvCacheQuantMode::Fp8 && has_fp8_mma {
FmhaQuantMode::Fp8
} else {
FmhaQuantMode::Bfloat16
}
}
pub(crate) fn generation_attn_flops(
spec: &SystemSpec,
kv_quant: KvCacheQuantMode,
) -> Result<f64, AicError> {
quant_tc_flops(spec, generation_attn_mode(spec, kv_quant).mapping())
}
fn clamp_generation_attention_grids_to_sol(
spec: &SystemSpec,
grids: &mut BTreeMap<GenerationKey, Grid3<LeafValue>>,
) {
for (key, grid) in grids.iter_mut() {
let Some(kv_quant) = kv_cache_quant_by_name(&key.kv_quant) else {
continue;
};
let Ok(attn_flops) = generation_attn_flops(spec, kv_quant) else {
continue;
};
for (&n, by_b) in grid.iter_mut() {
for (&b, by_s) in by_b.iter_mut() {
for (&s, leaf) in by_s.iter_mut() {
let sol = generation_attention_sol_ms(
spec,
key.n_kv_lookup,
key.head_size,
key.window_size,
kv_quant,
n as f64,
b as f64,
s as f64,
attn_flops,
);
if sol > leaf.latency {
leaf.latency = sol;
}
}
}
}
}
}
fn kv_cache_quant_by_name(name: &str) -> Option<KvCacheQuantMode> {
use KvCacheQuantMode::*;
Some(match name {
"bfloat16" => Bfloat16,
"int8" => Int8,
"fp8" => Fp8,
_ => return None,
})
}
fn load_encoder_parquet(
sources: &[PerfSource],
) -> Result<BTreeMap<EncoderKey, Grid3<LeafValue>>, AicError> {
let mut by_keys: BTreeMap<EncoderKey, Grid3<LeafValue>> = 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 batch_size_col = reader.col("batch_size")?;
let isl_col = reader.col("isl")?;
let num_heads_col = reader.col("num_heads")?;
let head_dim_col = reader.col("head_dim")?;
let attn_dtype_col = reader.col("attn_dtype")?;
let latency_col = reader.col("latency")?;
let power_col = reader.col_optional("power");
let ks_col = reader.col_optional("kernel_source");
for row in reader.rows()? {
let row = row?;
if !kernel_source_ok(source.kernel_sources(), ks_col, &row)? {
continue;
}
let key = EncoderKey {
fmha_quant: row.str_owned(attn_dtype_col)?,
head_size: row.u32(head_dim_col)?,
};
let latency = row.f64(latency_col)?;
let power = row.f64_optional(power_col)?.unwrap_or(0.0);
by_keys
.entry(key)
.or_default()
.entry(row.u32(num_heads_col)?)
.or_default()
.entry(row.u32(isl_col)?)
.or_default()
.entry(row.u32(batch_size_col)?)
.or_insert(LeafValue::with_power(latency, power));
}
}
if !any_source || by_keys.is_empty() {
return Err(AicError::PerfDatabase(format!(
"no encoder-attention rows loaded from {} source(s) (first: {})",
sources.len(),
sources
.first()
.map(|s| s.path().display().to_string())
.unwrap_or_default()
)));
}
Ok(by_keys)
}
fn missing_key(data_root: &Path, key: &ContextKey) -> AicError {
AicError::PerfDatabase(format!(
"context attention data missing for {key:?} at {}",
data_root.display()
))
}
fn missing_gen_key(data_root: &Path, key: &GenerationKey) -> AicError {
AicError::PerfDatabase(format!(
"generation attention data missing for {key:?} at {}",
data_root.display()
))
}
fn missing_encoder_key(data_root: &Path, key: &EncoderKey) -> AicError {
AicError::PerfDatabase(format!(
"encoder attention data missing for {key:?} at {}",
data_root.display()
))
}
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")
}
fn gb200_vllm_data_root() -> PathBuf {
PathBuf::from(REPO_ROOT_HINT)
.join("../..")
.join("python/aisimulate/src/aiconfigurator_core/systems/data/gb200/vllm/0.19.0")
}
fn gb200_spec() -> SystemSpec {
let systems_yaml = PathBuf::from(REPO_ROOT_HINT)
.join("../..")
.join("python/aisimulate/src/aiconfigurator_core/systems/gb200.yaml");
SystemSpec::load(&systems_yaml).expect("gb200.yaml must parse")
}
#[test]
fn generation_query_ragged_corner_matches_python_v2_engine() {
let table = AttentionTable::new(gb200_vllm_data_root(), gb200_spec());
let latency = table
.query_generation(256, 2561, 32, 8, 128, 0, KvCacheQuantMode::Bfloat16)
.expect("ragged-corner query must succeed")
.latency;
let expected = 0.37153384771269;
assert!(
((latency - expected) / expected).abs() < 1e-9,
"rust {latency} vs python {expected}"
);
}
#[test]
fn context_attention_exact_hit() {
let table = AttentionTable::new(b200_vllm_data_root(), b200_sxm_spec());
let latency = table
.query_context(
8,
16384,
64,
1,
128,
0,
KvCacheQuantMode::Fp8,
FmhaQuantMode::Bfloat16,
)
.expect("query must succeed")
.latency;
assert!(
(latency - 19.820667266845703).abs() < 1e-9,
"expected recorded latency, got {latency}"
);
}
#[test]
fn generation_attention_query_matches_python_v2_engine() {
let table = AttentionTable::new(b200_vllm_data_root(), b200_sxm_spec());
let latency = table
.query_generation(32, 2, 64, 4, 128, 0, KvCacheQuantMode::Fp8)
.expect("query must succeed")
.latency;
let expected = 0.009131092737966444;
assert!(
((latency - expected) / expected).abs() < 1e-9,
"rust {latency} vs python {expected}"
);
}
#[test]
fn context_attention_query_matches_python_v2_engine() {
let table = AttentionTable::new(b200_vllm_data_root(), b200_sxm_spec());
let cases: &[(u32, u32, f64)] = &[
(8, 16384, 19.820667266845703), (8, 12000, 11.515825737734879), (64, 16384, 184.03017609528183), ];
for &(b, s, expected) in cases {
let got = table
.query_context(
b,
s,
64,
1,
128,
0,
KvCacheQuantMode::Fp8,
FmhaQuantMode::Bfloat16,
)
.unwrap()
.latency;
assert!(
((got - expected) / expected).abs() < 1e-9,
"(b={b},s={s}): rust {got} vs python {expected}"
);
}
}
#[test]
fn context_attention_mha_normalizes_n_kv_to_zero() {
let table = AttentionTable::new(b200_vllm_data_root(), b200_sxm_spec());
let latency = table
.query_context(
4,
16384,
64,
64,
128,
0,
KvCacheQuantMode::Fp8,
FmhaQuantMode::Bfloat16,
)
.expect("MHA lookup must normalize and find the row")
.latency;
assert!(
(latency - 9.983466466267904).abs() < 1e-9,
"expected recorded MHA latency, got {latency}"
);
}
#[test]
fn context_attention_missing_quant_combo_errors() {
let table = AttentionTable::new(b200_vllm_data_root(), b200_sxm_spec());
match table.query_context(
1,
1024,
64,
1,
128,
0,
KvCacheQuantMode::Fp8,
FmhaQuantMode::Fp8,
) {
Err(AicError::PerfDatabase(_)) => {}
other => panic!("expected PerfDatabase error, got {other:?}"),
}
}
#[test]
fn encoder_attention_query_matches_python_v2_engine() {
let table = AttentionTable::new(b200_vllm_data_root(), b200_sxm_spec());
let cases: &[(u32, u32, f64)] = &[
(1, 1024, 0.03258133431275686), (2, 1400, 0.0779337721462867), (64, 65536, 10944.346873534367), ];
for &(b, s, expected) in cases {
let got = table
.query_encoder(b, s, 16, 64, FmhaQuantMode::Bfloat16)
.unwrap()
.latency;
assert!(
((got - expected) / expected).abs() < 1e-9,
"(b={b},s={s}): rust {got} vs python {expected}"
);
}
}
#[test]
fn encoder_attention_missing_head_size_errors() {
let table = AttentionTable::new(b200_vllm_data_root(), b200_sxm_spec());
match table.query_encoder(1, 1024, 16, 128, FmhaQuantMode::Bfloat16) {
Err(AicError::PerfDatabase(_)) => {}
other => panic!("expected PerfDatabase error, got {other:?}"),
}
}
#[test]
fn context_attention_lazy_loads_once() {
let table = AttentionTable::new(b200_vllm_data_root(), b200_sxm_spec());
let first = table
.query_context(
8,
16384,
64,
1,
128,
0,
KvCacheQuantMode::Fp8,
FmhaQuantMode::Bfloat16,
)
.unwrap();
let second = table
.query_context(
8,
16384,
64,
1,
128,
0,
KvCacheQuantMode::Fp8,
FmhaQuantMode::Bfloat16,
)
.unwrap();
assert_eq!(first, second);
}
#[test]
fn context_attention_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("context_attention_perf.parquet"),
&[
Col::Str("attn_dtype", vec!["bfloat16", "bfloat16"]),
Col::Str("kv_cache_dtype", vec!["bfloat16", "bfloat16"]),
Col::I64("batch_size", vec![2, 2]),
Col::I64("isl", vec![1024, 2048]),
Col::I64("num_heads", vec![16, 16]),
Col::I64("num_key_value_heads", vec![16, 16]),
Col::I64("head_dim", vec![128, 128]),
Col::I64("step", vec![0, 0]),
Col::F64("latency", vec![1.0, 3.0]),
Col::F64("power", vec![100.0, 200.0]),
],
);
let table = AttentionTable::new(tmp.path().to_path_buf(), energy_test_spec());
let v = table
.query_context(
2,
1536,
16,
16,
128,
0,
KvCacheQuantMode::Bfloat16,
FmhaQuantMode::Bfloat16,
)
.unwrap();
assert!(
((v.latency - 1.8660254037844386) / 1.8660254037844386).abs() < 1e-9,
"latency {}",
v.latency
);
assert!(
((v.energy - 279.9038105676658) / 279.9038105676658).abs() < 1e-9,
"energy {}",
v.energy
);
}
}