use std::fmt;
use serde::{Deserialize, Serialize};
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum BackendKind {
Trtllm,
Sglang,
Vllm,
}
impl BackendKind {
pub fn as_str(self) -> &'static str {
match self {
Self::Trtllm => "trtllm",
Self::Sglang => "sglang",
Self::Vllm => "vllm",
}
}
}
impl fmt::Display for BackendKind {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
pub enum DatabaseMode {
Silicon,
Hybrid,
Empirical,
Sol,
SolFull,
}
impl Default for DatabaseMode {
fn default() -> Self {
Self::Silicon
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum TransferKind {
XShape,
XQuant,
XProfile,
XOp,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct TransferPolicy {
pub xshape: bool,
pub xquant: bool,
pub xprofile: bool,
pub xop: bool,
}
impl Default for TransferPolicy {
fn default() -> Self {
Self::ALL
}
}
impl TransferPolicy {
pub const ALL: TransferPolicy = TransferPolicy {
xshape: true,
xquant: true,
xprofile: true,
xop: true,
};
pub const OFF: TransferPolicy = TransferPolicy {
xshape: false,
xquant: false,
xprofile: false,
xop: false,
};
pub fn contains(&self, kind: TransferKind) -> bool {
match kind {
TransferKind::XShape => self.xshape,
TransferKind::XQuant => self.xquant,
TransferKind::XProfile => self.xprofile,
TransferKind::XOp => self.xop,
}
}
pub fn from_wire(kinds: Option<&[String]>) -> Result<TransferPolicy, String> {
let Some(kinds) = kinds else {
return Ok(TransferPolicy::ALL);
};
let mut policy = TransferPolicy::OFF;
for token in kinds {
match token.as_str() {
"xshape" => policy.xshape = true,
"xquant" => policy.xquant = true,
"xprofile" => policy.xprofile = true,
"xop" => policy.xop = true,
other => return Err(format!("unknown transfer kind {other:?}")),
}
}
Ok(policy)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ComputeDtype {
Bfloat16,
Int8,
Fp8,
Fp4,
}
impl ComputeDtype {
pub fn flops_key(self) -> &'static str {
match self {
Self::Bfloat16 => "bfloat16_tc_flops",
Self::Int8 => "int8_tc_flops",
Self::Fp8 => "fp8_tc_flops",
Self::Fp4 => "fp4_tc_flops",
}
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct QuantMapping {
pub memory: f64,
pub compute: f64,
pub name: &'static str,
pub compute_dtype: Option<ComputeDtype>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum GemmQuantMode {
Bfloat16,
Int8Wo,
Int4Wo,
Fp8,
Fp8Static,
Sq,
Fp8Block,
Fp8Ootb,
Nvfp4,
Nvfp4Wo,
W4a16Nvfp4,
}
impl GemmQuantMode {
pub fn mapping(self) -> QuantMapping {
match self {
Self::Bfloat16 => QuantMapping {
memory: 2.0,
compute: 1.0,
name: "bfloat16",
compute_dtype: Some(ComputeDtype::Bfloat16),
},
Self::Int8Wo => QuantMapping {
memory: 1.0,
compute: 1.0,
name: "int8_wo",
compute_dtype: Some(ComputeDtype::Bfloat16),
},
Self::Int4Wo => QuantMapping {
memory: 0.5,
compute: 1.0,
name: "int4_wo",
compute_dtype: Some(ComputeDtype::Bfloat16),
},
Self::Fp8 => QuantMapping {
memory: 1.0,
compute: 2.0,
name: "fp8",
compute_dtype: Some(ComputeDtype::Fp8),
},
Self::Fp8Static => QuantMapping {
memory: 1.0,
compute: 2.0,
name: "fp8_static",
compute_dtype: Some(ComputeDtype::Fp8),
},
Self::Sq => QuantMapping {
memory: 1.0,
compute: 2.0,
name: "sq",
compute_dtype: Some(ComputeDtype::Int8),
},
Self::Fp8Block => QuantMapping {
memory: 1.0,
compute: 2.0,
name: "fp8_block",
compute_dtype: Some(ComputeDtype::Fp8),
},
Self::Fp8Ootb => QuantMapping {
memory: 1.0,
compute: 2.0,
name: "fp8_ootb",
compute_dtype: Some(ComputeDtype::Fp8),
},
Self::Nvfp4 => QuantMapping {
memory: 9.0 / 16.0,
compute: 4.0,
name: "nvfp4",
compute_dtype: Some(ComputeDtype::Fp4),
},
Self::Nvfp4Wo => QuantMapping {
memory: 9.0 / 16.0,
compute: 1.0,
name: "nvfp4_wo",
compute_dtype: Some(ComputeDtype::Bfloat16),
},
Self::W4a16Nvfp4 => QuantMapping {
memory: 9.0 / 16.0,
compute: 1.0,
name: "w4a16_nvfp4",
compute_dtype: Some(ComputeDtype::Bfloat16),
},
}
}
pub fn name(self) -> &'static str {
self.mapping().name
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum MoeQuantMode {
Bfloat16,
Fp8,
Int4Wo,
Fp8Block,
W4afp8,
Nvfp4,
Nvfp4Wo,
W4a16Mxfp4,
W4a8Mxfp4Mxfp8,
W4a8Mxfp4Mxfp8Trtllm,
W4a16Mxfp4Cutlass,
W4a16Nvfp4,
}
impl MoeQuantMode {
pub fn mapping(self) -> QuantMapping {
match self {
Self::Bfloat16 => QuantMapping {
memory: 2.0,
compute: 1.0,
name: "bfloat16",
compute_dtype: Some(ComputeDtype::Bfloat16),
},
Self::Fp8 => QuantMapping {
memory: 1.0,
compute: 2.0,
name: "fp8",
compute_dtype: Some(ComputeDtype::Fp8),
},
Self::Int4Wo => QuantMapping {
memory: 0.5,
compute: 1.0,
name: "int4_wo",
compute_dtype: Some(ComputeDtype::Bfloat16),
},
Self::Fp8Block => QuantMapping {
memory: 1.0,
compute: 2.0,
name: "fp8_block",
compute_dtype: Some(ComputeDtype::Fp8),
},
Self::W4afp8 => QuantMapping {
memory: 0.5,
compute: 2.0,
name: "w4afp8",
compute_dtype: Some(ComputeDtype::Fp8),
},
Self::Nvfp4 => QuantMapping {
memory: 9.0 / 16.0,
compute: 4.0,
name: "nvfp4",
compute_dtype: Some(ComputeDtype::Fp4),
},
Self::Nvfp4Wo => QuantMapping {
memory: 9.0 / 16.0,
compute: 1.0,
name: "nvfp4_wo",
compute_dtype: Some(ComputeDtype::Bfloat16),
},
Self::W4a16Mxfp4 => QuantMapping {
memory: 0.5,
compute: 1.0,
name: "w4a16_mxfp4",
compute_dtype: Some(ComputeDtype::Bfloat16),
},
Self::W4a8Mxfp4Mxfp8 => QuantMapping {
memory: 0.5,
compute: 2.0,
name: "w4a8_mxfp4_mxfp8",
compute_dtype: Some(ComputeDtype::Fp8),
},
Self::W4a8Mxfp4Mxfp8Trtllm => QuantMapping {
memory: 0.5,
compute: 2.0,
name: "w4a8_mxfp4_mxfp8_trtllm",
compute_dtype: Some(ComputeDtype::Fp8),
},
Self::W4a16Mxfp4Cutlass => QuantMapping {
memory: 0.5,
compute: 1.0,
name: "w4a16_mxfp4_cutlass",
compute_dtype: Some(ComputeDtype::Bfloat16),
},
Self::W4a16Nvfp4 => QuantMapping {
memory: 9.0 / 16.0,
compute: 1.0,
name: "w4a16_nvfp4",
compute_dtype: Some(ComputeDtype::Bfloat16),
},
}
}
pub fn name(self) -> &'static str {
self.mapping().name
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum FmhaQuantMode {
Bfloat16,
Fp8,
Fp8Block,
}
impl FmhaQuantMode {
pub fn mapping(self) -> QuantMapping {
match self {
Self::Bfloat16 => QuantMapping {
memory: 2.0,
compute: 1.0,
name: "bfloat16",
compute_dtype: Some(ComputeDtype::Bfloat16),
},
Self::Fp8 => QuantMapping {
memory: 1.0,
compute: 2.0,
name: "fp8",
compute_dtype: Some(ComputeDtype::Fp8),
},
Self::Fp8Block => QuantMapping {
memory: 1.0,
compute: 2.0,
name: "fp8_block",
compute_dtype: Some(ComputeDtype::Fp8),
},
}
}
pub fn name(self) -> &'static str {
self.mapping().name
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum KvCacheQuantMode {
Bfloat16,
Int8,
Fp8,
}
impl KvCacheQuantMode {
pub fn mapping(self) -> QuantMapping {
match self {
Self::Bfloat16 => QuantMapping {
memory: 2.0,
compute: 0.0,
name: "bfloat16",
compute_dtype: None,
},
Self::Int8 => QuantMapping {
memory: 1.0,
compute: 0.0,
name: "int8",
compute_dtype: None,
},
Self::Fp8 => QuantMapping {
memory: 1.0,
compute: 0.0,
name: "fp8",
compute_dtype: None,
},
}
}
pub fn name(self) -> &'static str {
self.mapping().name
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum CommQuantMode {
Half,
Int8,
Fp8,
}
impl CommQuantMode {
pub fn mapping(self) -> QuantMapping {
match self {
Self::Half => QuantMapping {
memory: 2.0,
compute: 0.0,
name: "half",
compute_dtype: None,
},
Self::Int8 => QuantMapping {
memory: 1.0,
compute: 0.0,
name: "int8",
compute_dtype: None,
},
Self::Fp8 => QuantMapping {
memory: 1.0,
compute: 0.0,
name: "fp8",
compute_dtype: None,
},
}
}
pub fn name(self) -> &'static str {
self.mapping().name
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
pub enum ModelFamily {
Gpt,
Llama,
Moe,
DeepSeek,
DeepSeekV32,
DeepSeekV4,
KimiK25,
NemotronNas,
NemotronH,
HybridMoe,
Qwen35,
Gemma4Mix,
MinimaxM3,
Qwen3Vl,
Qwen3VlMoe,
}
impl ModelFamily {
pub fn as_str(self) -> &'static str {
match self {
Self::Gpt => "GPT",
Self::Llama => "LLAMA",
Self::Moe => "MOE",
Self::DeepSeek => "DEEPSEEK",
Self::DeepSeekV32 => "DEEPSEEKV32",
Self::DeepSeekV4 => "DEEPSEEKV4",
Self::KimiK25 => "KIMIK25",
Self::NemotronNas => "NEMOTRONNAS",
Self::NemotronH => "NEMOTRONH",
Self::HybridMoe => "HYBRIDMOE",
Self::Qwen35 => "QWEN35",
Self::Gemma4Mix => "GEMMA4MIX",
Self::MinimaxM3 => "MINIMAXM3",
Self::Qwen3Vl => "QWEN3VL",
Self::Qwen3VlMoe => "QWEN3VL_MOE",
}
}
}
impl fmt::Display for ModelFamily {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum PerfDataFilename {
Gemm,
Nccl,
Oneccl,
GenerationAttention,
ContextAttention,
EncoderAttention,
ContextMla,
GenerationMla,
MlaBmm,
Moe,
CustomAllreduce,
WideepContextMla,
WideepGenerationMla,
TrtllmAlltoall,
ComputeScale,
ScaleMatrix,
Mamba2,
Gdn,
Kda,
MlaContextModule,
MlaGenerationModule,
DsaContextModule,
DsaGenerationModule,
MhcModule,
Dsv4CsaContextModule,
Dsv4HcaContextModule,
Dsv4CsaGenerationModule,
Dsv4HcaGenerationModule,
Dsv4PagedMqaLogitsModule,
Dsv4HcaAttnModule,
Dsv4CsaAttnModule,
Dsv4CsaTopkCalib,
Dsv4MegamoeModule,
MoeA2a,
MoeExpertCompute,
}
impl PerfDataFilename {
pub fn as_str(self) -> &'static str {
match self {
Self::Gemm => "gemm_perf.parquet",
Self::Nccl => "nccl_perf.parquet",
Self::Oneccl => "oneccl_perf.parquet",
Self::GenerationAttention => "generation_attention_perf.parquet",
Self::ContextAttention => "context_attention_perf.parquet",
Self::EncoderAttention => "encoder_attention_perf.parquet",
Self::ContextMla => "context_mla_perf.parquet",
Self::GenerationMla => "generation_mla_perf.parquet",
Self::MlaBmm => "mla_bmm_perf.parquet",
Self::Moe => "moe_perf.parquet",
Self::CustomAllreduce => "custom_allreduce_perf.parquet",
Self::WideepContextMla => "wideep_context_mla_perf.parquet",
Self::WideepGenerationMla => "wideep_generation_mla_perf.parquet",
Self::TrtllmAlltoall => "trtllm_alltoall_perf.parquet",
Self::ComputeScale => "computescale_perf.parquet",
Self::ScaleMatrix => "scale_matrix_perf.parquet",
Self::Mamba2 => "mamba2_perf.parquet",
Self::Gdn => "gdn_perf.parquet",
Self::Kda => "kda_perf.parquet",
Self::MlaContextModule => "mla_context_module_perf.parquet",
Self::MlaGenerationModule => "mla_generation_module_perf.parquet",
Self::DsaContextModule => "dsa_context_module_perf.parquet",
Self::DsaGenerationModule => "dsa_generation_module_perf.parquet",
Self::MhcModule => "mhc_module_perf.parquet",
Self::Dsv4CsaContextModule => "dsv4_csa_context_module_perf.parquet",
Self::Dsv4HcaContextModule => "dsv4_hca_context_module_perf.parquet",
Self::Dsv4CsaGenerationModule => "dsv4_csa_generation_module_perf.parquet",
Self::Dsv4HcaGenerationModule => "dsv4_hca_generation_module_perf.parquet",
Self::Dsv4PagedMqaLogitsModule => "dsv4_paged_mqa_logits_module_perf.parquet",
Self::Dsv4HcaAttnModule => "dsv4_hca_attn_module_perf.parquet",
Self::Dsv4CsaAttnModule => "dsv4_csa_attn_module_perf.parquet",
Self::Dsv4CsaTopkCalib => "dsv4_csa_topk_calib_perf.parquet",
Self::Dsv4MegamoeModule => "dsv4_megamoe_module_perf.parquet",
Self::MoeA2a => "moe_a2a_perf.parquet",
Self::MoeExpertCompute => "moe_expert_compute_perf.parquet",
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn backend_kind_string_keys_match_python() {
assert_eq!(BackendKind::Trtllm.as_str(), "trtllm");
assert_eq!(BackendKind::Sglang.as_str(), "sglang");
assert_eq!(BackendKind::Vllm.as_str(), "vllm");
}
#[test]
fn database_mode_default_is_silicon() {
assert_eq!(DatabaseMode::default(), DatabaseMode::Silicon);
}
#[test]
fn gemm_quant_payloads_match_python_quant_mapping() {
assert_eq!(
GemmQuantMode::Bfloat16.mapping(),
QuantMapping {
memory: 2.0,
compute: 1.0,
name: "bfloat16",
compute_dtype: Some(ComputeDtype::Bfloat16)
}
);
assert_eq!(
GemmQuantMode::Fp8.mapping(),
QuantMapping {
memory: 1.0,
compute: 2.0,
name: "fp8",
compute_dtype: Some(ComputeDtype::Fp8)
}
);
assert_eq!(
GemmQuantMode::Nvfp4.mapping(),
QuantMapping {
memory: 9.0 / 16.0,
compute: 4.0,
name: "nvfp4",
compute_dtype: Some(ComputeDtype::Fp4)
}
);
}
#[test]
fn moe_quant_payloads_match_python_quant_mapping() {
assert_eq!(
MoeQuantMode::W4afp8.mapping(),
QuantMapping {
memory: 0.5,
compute: 2.0,
name: "w4afp8",
compute_dtype: Some(ComputeDtype::Fp8)
}
);
assert_eq!(
MoeQuantMode::W4a8Mxfp4Mxfp8.mapping(),
QuantMapping {
memory: 0.5,
compute: 2.0,
name: "w4a8_mxfp4_mxfp8",
compute_dtype: Some(ComputeDtype::Fp8)
}
);
}
#[test]
fn kvcache_compute_is_zero() {
for mode in [
KvCacheQuantMode::Bfloat16,
KvCacheQuantMode::Int8,
KvCacheQuantMode::Fp8,
] {
assert_eq!(mode.mapping().compute, 0.0);
}
}
#[test]
fn model_family_string_round_trip() {
assert_eq!(ModelFamily::Qwen3VlMoe.as_str(), "QWEN3VL_MOE");
assert_eq!(ModelFamily::DeepSeekV32.as_str(), "DEEPSEEKV32");
assert_eq!(ModelFamily::KimiK25.as_str(), "KIMIK25");
}
#[test]
fn perf_data_filenames_match_python_enum() {
assert_eq!(PerfDataFilename::Gemm.as_str(), "gemm_perf.parquet");
assert_eq!(
PerfDataFilename::ContextAttention.as_str(),
"context_attention_perf.parquet"
);
assert_eq!(
PerfDataFilename::CustomAllreduce.as_str(),
"custom_allreduce_perf.parquet"
);
assert_eq!(
PerfDataFilename::TrtllmAlltoall.as_str(),
"trtllm_alltoall_perf.parquet"
);
assert_eq!(
PerfDataFilename::Dsv4HcaGenerationModule.as_str(),
"dsv4_hca_generation_module_perf.parquet"
);
}
}