use crate::common::enums::{BackendKind, GemmQuantMode};
use crate::common::error::AicError;
use crate::common::system_spec::SystemSpec;
use crate::operators::op::Op;
use crate::operators::{
ContextAttentionOp, CustomAllReduceOp, ElementwiseOp, EmbeddingOp, GemmOp,
GenerationAttentionOp, MoEDispatchOp, MoeOp, NcclOp, P2POp,
};
use crate::perf_database::PerfDatabase;
fn floor_div(a: f64, b: f64) -> f64 {
(a / b).floor()
}
fn ceil_div(a: f64, b: f64) -> f64 {
(a / b).ceil()
}
pub(crate) fn op_sol_latency_ms(
op: &Op,
db: &PerfDatabase,
x: f64,
batch: f64,
s: f64,
prefix: f64,
) -> Result<f64, AicError> {
let spec = &db.system_spec;
match op {
Op::Gemm(o) => Ok(gemm_sol(o, spec, x)),
Op::Embedding(o) => Ok(embedding_sol(o, spec, x)),
Op::Elementwise(o) => Ok(elementwise_sol(o, spec, x)),
Op::ContextAttention(o) => Ok(context_attention_sol(o, spec, batch, s, prefix)),
Op::GenerationAttention(o) => Ok(generation_attention_sol(o, spec, batch, s)),
Op::Moe(o) => Ok(moe_sol(o, spec, x)),
Op::MoeDispatch(o) => moe_dispatch_sol(o, spec, x),
Op::CustomAllReduce(o) => Ok(custom_allreduce_op_sol(o, spec, x)),
Op::Nccl(o) => Ok(nccl_op_sol(o, spec, x)),
Op::P2P(o) => Ok(p2p_sol(o, spec, x)),
Op::Overlap(o) => {
let mut total_a = 0.0;
for inner in &o.group_a {
total_a += op_sol_latency_ms(inner, db, x, batch, s, prefix)?;
}
let mut total_b = 0.0;
for inner in &o.group_b {
total_b += op_sol_latency_ms(inner, db, x, batch, s, prefix)?;
}
Ok(total_a.max(total_b))
}
Op::Fallback(o) => match op_sol_latency_ms(&o.primary, db, x, batch, s, prefix) {
Ok(v) => Ok(v),
Err(AicError::PerfDatabase(_)) | Err(AicError::Io { .. }) => {
let mut total = 0.0;
for inner in &o.fallback {
total += op_sol_latency_ms(inner, db, x, batch, s, prefix)?;
}
Ok(total)
}
Err(other) => Err(other),
},
other => Err(AicError::SolNotImplemented(format!(
"forward_model='fpm' SOL roofline has no Rust implementation for op {}",
other.name()
))),
}
}
fn quant_tc_flops(spec: &SystemSpec, quant: GemmQuantMode) -> f64 {
let compute = quant.mapping().compute;
let direct = if compute == 1.0 {
spec.gpu.bfloat16_tc_flops
} else if compute == 2.0 {
spec.gpu.fp8_tc_flops
} else if compute == 4.0 {
spec.gpu.fp4_tc_flops
} else {
None
};
direct.unwrap_or_else(|| spec.gpu.bfloat16_tc_flops.unwrap_or(0.0) * compute)
}
fn gemm_sol(op: &GemmOp, spec: &SystemSpec, x: f64) -> f64 {
let m = ceil_div(
floor_div(x, op.scale_num_tokens.max(1) as f64),
op.seq_split.max(1) as f64,
);
let (n, k) = (op.n as f64, op.k as f64);
let mapping = op.quant_mode.mapping();
let tc_flops = quant_tc_flops(spec, op.quant_mode);
let sol_math = 2.0 * m * n * k / tc_flops * 1000.0;
let sol_mem = mapping.memory * (m * n + m * k + n * k) / spec.gpu.mem_bw * 1000.0;
sol_math.max(sol_mem) * op.scale_factor
}
fn mem_op_sol_ms(spec: &SystemSpec, mem_bytes: f64) -> f64 {
mem_bytes / spec.gpu.mem_bw * 1000.0
}
fn embedding_sol(op: &EmbeddingOp, spec: &SystemSpec, x: f64) -> f64 {
let tokens = ceil_div(x, op.seq_split.max(1) as f64);
mem_op_sol_ms(spec, tokens * op.hidden_size as f64 * 2.0) * op.scale_factor
}
fn elementwise_sol(op: &ElementwiseOp, spec: &SystemSpec, x: f64) -> f64 {
let tokens = ceil_div(
floor_div(x, op.scale_num_tokens.max(1) as f64),
op.seq_split.max(1) as f64,
);
mem_op_sol_ms(spec, op.bytes_per_token * tokens) * op.scale_factor
}
fn context_sol_one(op: &ContextAttentionOp, spec: &SystemSpec, b: f64, s: f64, p: f64) -> f64 {
let (n, n_kv, h, w) = (
op.n as f64,
op.n_kv as f64,
op.head_size as f64,
op.window_size as f64,
);
let full_s = s + p;
let ops = if op.window_size > 0 && full_s > w {
2.0 * b * s * w * n * h * 2.0
} else {
2.0 * b * (full_s * full_s - p * p) * n * h * 2.0 / 2.0
};
let mem = 2.0 * b * (n * s * h + n * s * h)
+ op.kv_cache_dtype.mapping().memory * b * (2.0 * n_kv * full_s * h);
let flops = spec.gpu.bfloat16_tc_flops.unwrap_or(0.0);
let sol_math = ops / flops * 1000.0 / op.fmha_quant_mode.mapping().compute;
let sol_mem = mem / spec.gpu.mem_bw * 1000.0;
sol_math.max(sol_mem)
}
fn context_attention_sol(
op: &ContextAttentionOp,
spec: &SystemSpec,
b: f64,
s: f64,
p: f64,
) -> f64 {
let fmha = if op.cp_size > 1 {
let c = ceil_div(s, 2.0 * op.cp_size as f64).max(1.0);
context_sol_one(op, spec, b, c, p) + context_sol_one(op, spec, b, c, p + s - c)
} else {
context_sol_one(op, spec, b, s, p)
};
let q_num = (op.n * op.head_size) as f64;
let k_num = (op.n_kv * op.head_size) as f64;
let mut extra = 0.0;
if op.use_qk_norm {
let qk_norm =
2.0 * mem_op_sol_ms(spec, q_num * 2.0) + 2.0 * mem_op_sol_ms(spec, k_num * 2.0);
extra += qk_norm * 2.0;
}
extra += 2.0 * mem_op_sol_ms(spec, q_num * 2.0 + k_num * 2.0); let fq_mem = op.fmha_quant_mode.mapping().memory;
extra += mem_op_sol_ms(spec, k_num * fq_mem) + mem_op_sol_ms(spec, k_num * fq_mem); (fmha + extra * 1.1) * op.scale_factor
}
fn generation_attention_sol(op: &GenerationAttentionOp, spec: &SystemSpec, b: f64, s: f64) -> f64 {
let (n, n_kv, h, w) = (
op.n as f64,
op.n_kv as f64,
op.head_size as f64,
op.window_size as f64,
);
let kv_len = if op.window_size > 0 {
(s - 1.0).min(w)
} else {
s - 1.0
};
let compute = if op.kv_cache_dtype == crate::common::enums::KvCacheQuantMode::Fp8 {
2.0
} else {
1.0
};
let kv_mem = op.kv_cache_dtype.mapping().memory;
let ops = 2.0 * b * n * h * 2.0 * kv_len;
let mem = b * (n * h * 2.0 + 2.0 * n_kv * kv_len * h * kv_mem + n * h * 2.0);
let flops = spec.gpu.bfloat16_tc_flops.unwrap_or(0.0);
let sol_math = ops / flops * 1000.0 / compute;
let sol_mem = mem / spec.gpu.mem_bw * 1000.0;
sol_math.max(sol_mem) * op.scale_factor
}
fn moe_sol(op: &MoeOp, spec: &SystemSpec, x: f64) -> f64 {
let dp = op.attention_dp_size.max(1) as f64;
let (h, inter) = (op.hidden_size as f64, op.inter_size as f64);
let num_gemms = if op.is_gated { 3.0 } else { 2.0 };
let (ep, tp) = (op.moe_ep_size.max(1) as f64, op.moe_tp_size.max(1) as f64);
let total_tokens = x * dp * op.topk as f64;
let ops = floor_div(
floor_div(total_tokens * h * inter * num_gemms * 2.0, ep),
tp,
);
let tt_ep = floor_div(total_tokens, ep);
let mem = op.quant_mode.mapping().memory
* (tt_ep * h * 2.0
+ floor_div(tt_ep * inter * num_gemms, tp)
+ floor_div(h * inter * num_gemms, tp)
* floor_div(op.num_experts as f64, ep).min(tt_ep));
let flops = spec.gpu.bfloat16_tc_flops.unwrap_or(0.0);
let sol_math = ops / (flops * op.quant_mode.mapping().compute) * 1000.0;
let sol_mem = mem / spec.gpu.mem_bw * 1000.0;
sol_math.max(sol_mem) * op.scale_factor
}
fn custom_allreduce_sol(spec: &SystemSpec, tp_size: u32, size_elems: f64) -> f64 {
if tp_size <= 1 {
return 0.0;
}
let tp = tp_size as f64;
let bw = spec.get_p2p_bandwidth(tp_size);
2.0 * size_elems * 2.0 / tp * (tp - 1.0) / bw * 1000.0
}
fn nccl_sol(
spec: &SystemSpec,
num_gpus: u32,
operation: &str,
message_size: f64,
bytes_per_elem: f64,
) -> f64 {
let n = num_gpus as f64;
let bw = spec.get_p2p_bandwidth(num_gpus);
match operation {
"all_gather" | "alltoall" | "reduce_scatter" => {
bytes_per_elem * message_size * (n - 1.0) / n / bw * 1000.0
}
"all_reduce" => 2.0 * bytes_per_elem * message_size * (n - 1.0) / n / bw * 1000.0,
_ => 0.0,
}
}
fn custom_allreduce_op_sol(op: &CustomAllReduceOp, spec: &SystemSpec, x: f64) -> f64 {
if op.tp_size == 1 {
return 0.0;
}
let size = ceil_div(x, op.seq_split.max(1) as f64) * op.hidden_size as f64;
custom_allreduce_sol(spec, op.tp_size, size) * op.scale_factor
}
fn nccl_op_sol(op: &NcclOp, spec: &SystemSpec, x: f64) -> f64 {
let msg = ceil_div(x, op.seq_split.max(1) as f64) * op.hidden_size;
nccl_sol(
spec,
op.num_gpus,
&op.operation,
msg,
op.dtype.mapping().memory,
) * op.scale_factor
}
fn p2p_sol(op: &P2POp, spec: &SystemSpec, x: f64) -> f64 {
if op.pp_size == 1 {
return 0.0;
}
let p2p_bytes = ceil_div(x, op.seq_split.max(1) as f64) * op.hidden_size as f64 * 2.0;
p2p_bytes / spec.node.inter_node_bw * 1000.0 * op.scale_factor
}
fn moe_dispatch_sol(op: &MoEDispatchOp, spec: &SystemSpec, x: f64) -> Result<f64, AicError> {
use crate::operators::moe_dispatch::DispatchFlavor;
let volume = x * op.hidden_size as f64; let num_gpus = (op.moe_tp_size * op.moe_ep_size).max(1);
let attn_dp = op.attention_dp_size.max(1);
let attn_tp = (num_gpus / attn_dp).max(1);
let dp = attn_dp as f64;
let pre = op.pre_dispatch;
let half_bytes = 2.0;
let comm = match op.flavor {
DispatchFlavor::RetiredDeepEp => {
return Err(AicError::InvalidEngineConfig(format!(
"MoEDispatch '{}' (moe_backend='deepep_moe') has no native SOL \
(retired with AIC-1601; large-EP comm is modeled by MoeAllToAll)",
op.name
)));
}
DispatchFlavor::CustomAllReduce => match op.backend {
BackendKind::Vllm => {
let mut total = 0.0;
if attn_tp > 1 {
total += custom_allreduce_sol(spec, num_gpus, volume);
}
if attn_dp > 1 {
let op_name = if pre { "all_gather" } else { "reduce_scatter" };
total += nccl_sol(spec, num_gpus, op_name, volume * dp, half_bytes);
}
total
}
BackendKind::Sglang => {
if attn_tp > 1 && attn_dp > 1 {
if pre {
nccl_sol(spec, attn_tp, "reduce_scatter", volume, half_bytes)
+ nccl_sol(spec, num_gpus, "all_gather", volume * dp, half_bytes)
} else {
nccl_sol(spec, num_gpus, "reduce_scatter", volume * dp, half_bytes)
+ nccl_sol(spec, attn_tp, "all_gather", volume, half_bytes)
}
} else if op.attn_cp_size > 1 {
if op.is_context {
let op_name = if pre { "all_gather" } else { "reduce_scatter" };
nccl_sol(spec, num_gpus, op_name, volume, half_bytes)
} else if pre {
0.0
} else {
custom_allreduce_sol(spec, num_gpus, volume)
}
} else if attn_tp > 1 {
custom_allreduce_sol(spec, num_gpus, volume)
} else if attn_dp > 1 {
let op_name = if pre { "all_gather" } else { "reduce_scatter" };
nccl_sol(spec, num_gpus, op_name, volume * dp, half_bytes)
} else {
0.0
}
}
BackendKind::Trtllm => {
if attn_tp > 1 {
custom_allreduce_sol(spec, num_gpus, volume)
} else if attn_dp > 1 {
let op_name = if pre { "all_gather" } else { "reduce_scatter" };
nccl_sol(spec, num_gpus, op_name, volume * dp, half_bytes)
} else {
0.0
}
}
},
DispatchFlavor::TrtllmAlltoall => {
let is_nvl72 = spec.node.num_gpus_per_node >= 72;
let enable_alltoall = op.attention_dp_size > 1 && op.moe_tp_size == 1 && is_nvl72;
if enable_alltoall {
let node_num = if op.moe_ep_size < 4 {
1
} else {
op.moe_ep_size / 4
};
let bw = if node_num > 1 {
spec.node.inter_node_bw
} else {
spec.node.intra_node_bw
};
let remote_ranks =
op.topk
.min(op.num_experts)
.min(op.moe_ep_size.saturating_sub(1)) as f64;
let bytes_per_elem = if pre {
op.moe_quant.mapping().memory
} else {
2.0
};
let data_bytes = x * remote_ranks * op.hidden_size as f64 * bytes_per_elem;
data_bytes / bw * 1000.0
} else if op.attention_dp_size > 1 {
if pre {
let compressed = match op.moe_quant.mapping().name {
"nvfp4" => volume / 4.0 + volume / 4.0 / 8.0,
"fp8" | "fp8_block" => volume / 2.0,
_ => volume,
};
nccl_sol(spec, num_gpus, "all_gather", compressed * dp, half_bytes)
} else {
nccl_sol(spec, num_gpus, "reduce_scatter", volume * dp, half_bytes)
}
} else if attn_tp > 1 {
if spec.node.num_gpus_per_node == 72 && num_gpus > 4 {
nccl_sol(spec, num_gpus, "all_reduce", volume, half_bytes)
} else {
custom_allreduce_sol(spec, num_gpus, volume)
}
} else {
0.0
}
}
};
Ok(comm * op.scale_factor)
}
#[cfg(test)]
mod tests {
use super::*;
use std::path::PathBuf;
use crate::common::enums::{
CommQuantMode, FmhaQuantMode, GemmQuantMode, KvCacheQuantMode, MoeQuantMode,
};
fn spec() -> SystemSpec {
let root = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("../../python/aisimulate/src/aiconfigurator_core/systems/b200_sxm.yaml");
SystemSpec::load(&root).expect("b200 spec")
}
fn db() -> PerfDatabase {
let root = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("../../python/aisimulate/src/aiconfigurator_core/systems");
PerfDatabase::load(&root, "b200_sxm", "vllm", "0.19.0").expect("db")
}
fn approx(got: f64, expected: f64) {
assert!(
(got - expected).abs() <= 1e-9 * expected.abs().max(1e-12),
"got {got}, expected {expected}"
);
}
#[test]
fn gemm_sol_matches_formula() {
let op = GemmOp {
name: "qkv_gemm".into(),
scale_factor: 1.0,
n: 4096,
k: 4096,
quant_mode: GemmQuantMode::Nvfp4,
scale_num_tokens: 1,
low_precision_input: false,
seq_split: 1,
below_grid_sol: false,
};
let s = spec();
let m = 8192.0_f64;
let (n, k) = (4096.0_f64, 4096.0_f64);
let tc = s.gpu.fp4_tc_flops.unwrap();
let expected = (2.0 * m * n * k / tc * 1000.0)
.max(9.0 / 16.0 * (m * n + m * k + n * k) / s.gpu.mem_bw * 1000.0);
approx(gemm_sol(&op, &s, 8192.0), expected);
approx(gemm_sol(&op, &s, 8192.7), expected);
}
#[test]
fn embedding_vs_elementwise_rounding_directions() {
let s = spec();
let emb = EmbeddingOp {
name: "context_embedding".into(),
scale_factor: 1.0,
vocab_size: 128000,
hidden_size: 6144,
quant_mode: GemmQuantMode::Bfloat16,
seq_split: 1,
};
approx(
embedding_sol(&emb, &s, 10.5),
11.0 * 6144.0 * 2.0 / s.gpu.mem_bw * 1000.0,
);
let ew = ElementwiseOp {
name: "add_norm".into(),
scale_factor: 2.0,
bytes_per_token: 8192.0,
scale_num_tokens: 1,
seq_split: 1,
};
approx(
elementwise_sol(&ew, &s, 10.5),
8192.0 * 10.0 / s.gpu.mem_bw * 1000.0 * 2.0,
);
}
#[test]
fn context_attention_sol_prefix_aware() {
let s = spec();
let op = ContextAttentionOp {
name: "context_attention".into(),
scale_factor: 1.0,
n: 48,
n_kv: 8,
head_size: 128,
window_size: 0,
kv_cache_dtype: KvCacheQuantMode::Fp8,
fmha_quant_mode: FmhaQuantMode::Bfloat16,
use_qk_norm: false,
cp_size: 1,
lane_order: crate::operators::attention::b200_vllm_context_lane_order(),
};
let (b, sq, p) = (4.0, 682.6666666666666_f64, 128.5_f64);
let (n, n_kv, h) = (48.0, 8.0, 128.0);
let full = sq + p;
let ops = 2.0 * b * (full * full - p * p) * n * h * 2.0 / 2.0;
let mem = 2.0 * b * (n * sq * h + n * sq * h) + 1.0 * b * (2.0 * n_kv * full * h);
let fmha = (ops / s.gpu.bfloat16_tc_flops.unwrap() * 1000.0 / 1.0)
.max(mem / s.gpu.mem_bw * 1000.0);
let q_num = n * h;
let k_num = n_kv * h;
let extras = 2.0 * (q_num * 2.0 + k_num * 2.0) / s.gpu.mem_bw * 1000.0
+ (k_num * 2.0) / s.gpu.mem_bw * 1000.0
+ (k_num * 2.0) / s.gpu.mem_bw * 1000.0;
approx(
context_attention_sol(&op, &s, b, sq, p),
fmha + extras * 1.1,
);
}
#[test]
fn generation_attention_sol_fp8_kv_uses_fp8_compute() {
let s = spec();
let op = GenerationAttentionOp {
name: "generation_attention".into(),
scale_factor: 1.0,
n: 48,
n_kv: 8,
head_size: 128,
window_size: 0,
kv_cache_dtype: KvCacheQuantMode::Fp8,
lane_order: crate::operators::attention::b200_vllm_generation_lane_order(),
};
let (b, sq) = (256.0, 8441.75_f64);
let kv_len = sq - 1.0;
let ops = 2.0 * b * 48.0 * 128.0 * 2.0 * kv_len;
let mem = b * (48.0 * 128.0 * 2.0 + 2.0 * 8.0 * kv_len * 128.0 * 1.0 + 48.0 * 128.0 * 2.0);
let expected = (ops / s.gpu.bfloat16_tc_flops.unwrap() * 1000.0 / 2.0)
.max(mem / s.gpu.mem_bw * 1000.0);
approx(generation_attention_sol(&op, &s, b, sq), expected);
}
#[test]
fn moe_sol_floor_association() {
let s = spec();
let op = MoeOp {
name: "context_moe".into(),
scale_factor: 1.0,
hidden_size: 6144,
inter_size: 1536,
topk: 8,
num_experts: 256,
moe_tp_size: 1,
moe_ep_size: 4,
attention_dp_size: 1,
quant_mode: MoeQuantMode::Nvfp4,
workload_distribution: "uniform".into(),
is_gated: true,
moe_backend: None,
enable_eplb: false,
is_context: true,
};
let x = 8192.0_f64;
let tt = x * 8.0;
let ops = ((tt * 6144.0 * 1536.0 * 3.0 * 2.0 / 4.0).floor() / 1.0).floor();
let tt_ep = (tt / 4.0).floor();
let mem = 9.0 / 16.0
* (tt_ep * 6144.0 * 2.0
+ (tt_ep * 1536.0 * 3.0 / 1.0).floor()
+ (6144.0_f64 * 1536.0 * 3.0 / 1.0).floor() * (256.0_f64 / 4.0).floor().min(tt_ep));
let expected = (ops / (s.gpu.bfloat16_tc_flops.unwrap() * 4.0) * 1000.0)
.max(mem / s.gpu.mem_bw * 1000.0);
approx(moe_sol(&op, &s, x), expected);
}
#[test]
fn comm_sols_match_formulas() {
let s = spec();
let car = CustomAllReduceOp {
name: "ar".into(),
scale_factor: 1.0,
hidden_size: 6144,
tp_size: 4,
quant: CommQuantMode::Half,
seq_split: 1,
};
let size = 8192.0 * 6144.0;
let bw = s.get_p2p_bandwidth(4);
approx(
custom_allreduce_op_sol(&car, &s, 8192.0),
2.0 * size * 2.0 / 4.0 * 3.0 / bw * 1000.0,
);
let car1 = CustomAllReduceOp {
tp_size: 1,
..car.clone()
};
assert_eq!(custom_allreduce_op_sol(&car1, &s, 8192.0), 0.0);
let p2p = P2POp {
name: "p2p".into(),
scale_factor: 1.0,
pp_size: 2,
hidden_size: 6144,
seq_split: 1,
};
approx(
p2p_sol(&p2p, &s, 8192.0),
8192.0 * 6144.0 * 2.0 / s.node.inter_node_bw * 1000.0,
);
let nccl = NcclOp {
name: "nccl".into(),
scale_factor: 1.0,
hidden_size: 6144.0,
num_gpus: 8,
dtype: CommQuantMode::Half,
operation: "all_reduce".into(),
seq_split: 1,
};
let bw8 = s.get_p2p_bandwidth(8);
approx(
nccl_op_sol(&nccl, &s, 1024.0),
2.0 * 2.0 * (1024.0 * 6144.0) * 7.0 / 8.0 / bw8 * 1000.0,
);
}
#[test]
fn moe_dispatch_vllm_is_additive() {
let s = spec();
let op = MoEDispatchOp {
name: "context_moe_pre_dispatch".into(),
scale_factor: 1.0,
hidden_size: 6144,
topk: 8,
num_experts: 256,
moe_tp_size: 1,
moe_ep_size: 4,
attention_dp_size: 1,
pre_dispatch: true,
backend: BackendKind::Vllm,
flavor: crate::operators::moe_dispatch::DispatchFlavor::CustomAllReduce,
comm_quant: CommQuantMode::Half,
moe_quant: MoeQuantMode::Nvfp4,
attn_cp_size: 1,
is_context: true,
sms: 12,
scale_num_tokens: 1,
attn_ar_modeled: false,
};
let volume = 8192.0 * 6144.0;
let bw = s.get_p2p_bandwidth(4);
approx(
moe_dispatch_sol(&op, &s, 8192.0).unwrap(),
2.0 * volume * 2.0 / 4.0 * 3.0 / bw * 1000.0,
);
}
#[test]
fn overlap_and_fallback_compose() {
let d = db();
let ew = |bpt: f64| {
Op::Elementwise(ElementwiseOp {
name: "e".into(),
scale_factor: 1.0,
bytes_per_token: bpt,
scale_num_tokens: 1,
seq_split: 1,
})
};
let overlap = Op::Overlap(crate::operators::OverlapOp::new(
"ov",
vec![ew(1000.0), ew(2000.0)],
vec![ew(5000.0)],
));
let expected = super::mem_op_sol_ms(&d.system_spec, 5000.0 * 64.0);
approx(
op_sol_latency_ms(&overlap, &d, 64.0, 1.0, 1.0, 0.0).unwrap(),
expected,
);
let fb = Op::Fallback(crate::operators::FallbackOp::new("fb", ew(3000.0), vec![]));
approx(
op_sol_latency_ms(&fb, &d, 64.0, 1.0, 1.0, 0.0).unwrap(),
super::mem_op_sol_ms(&d.system_spec, 3000.0 * 64.0),
);
}
#[test]
fn unsupported_fpm_sol_op_has_typed_error() {
let d = db();
let op = Op::Mamba2(crate::operators::Mamba2Op {
name: "mamba2".into(),
scale_factor: 1.0,
kernel_source: "causal_conv1d_fn".into(),
phase: "context".into(),
d_model: 4096,
d_state: 128,
d_conv: 4,
nheads: 128,
head_dim: 64,
n_groups: 8,
chunk_size: 256,
});
let err = op_sol_latency_ms(&op, &d, 64.0, 1.0, 1.0, 0.0).unwrap_err();
assert!(matches!(&err, AicError::SolNotImplemented(_)));
assert!(
err.to_string()
.contains("no Rust implementation for op mamba2")
);
}
}