use crate::format::layout_contract::block_sizes;
use crate::format::model_family::{
AttentionType, DeltaNetShape, MlpType, ModelConstraints, ModelSizeConfig,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct LayerParams {
pub d_attn: u64,
pub d_ffn: u64,
pub d_norm: u64,
}
impl LayerParams {
#[must_use]
pub const fn total(&self) -> u64 {
self.d_attn
.saturating_add(self.d_ffn)
.saturating_add(self.d_norm)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ParameterBreakdown {
pub embedding: u64,
pub layers: u64,
pub final_norm: u64,
pub unembedding: u64,
pub total: u64,
}
#[must_use]
pub fn attention_layer_params(
size: &ModelSizeConfig,
constraints: &ModelConstraints,
) -> LayerParams {
let d = size.hidden_dim as u64;
let d_k = size.head_dim as u64;
let q_dim = (size.num_heads as u64).saturating_mul(d_k);
let kv_dim = (size.num_kv_heads as u64).saturating_mul(d_k);
let q_out = if matches!(
constraints.attention_type,
AttentionType::HybridGatedDeltaNet
) {
q_dim.saturating_mul(2)
} else {
q_dim
};
let projections = d
.saturating_mul(q_out)
.saturating_add(d.saturating_mul(kv_dim).saturating_mul(2))
.saturating_add(q_dim.saturating_mul(d));
let biases = if constraints.has_bias {
q_out
.saturating_add(kv_dim.saturating_mul(2))
.saturating_add(d)
} else {
0
};
let qk_norms = if constraints.qk_norm {
d_k.saturating_mul(2)
} else {
0
};
LayerParams {
d_attn: projections.saturating_add(biases).saturating_add(qk_norms),
d_ffn: ffn_params(size, constraints),
d_norm: d.saturating_mul(2),
}
}
fn ffn_params(size: &ModelSizeConfig, constraints: &ModelConstraints) -> u64 {
let matrices = if matches!(constraints.mlp_type, MlpType::SwiGlu | MlpType::GatedMlp) {
3
} else {
2
};
(size.hidden_dim as u64)
.saturating_mul(size.intermediate_dim as u64)
.saturating_mul(matrices)
}
#[must_use]
pub fn gated_deltanet_layer_params(
size: &ModelSizeConfig,
constraints: &ModelConstraints,
shape: &DeltaNetShape,
) -> LayerParams {
let d = size.hidden_dim as u64;
let inner = shape.inner_size as u64;
let heads = shape.group_count as u64;
let qkv = d.saturating_mul(inner).saturating_mul(3);
let gate = d.saturating_mul(inner);
let conv = (shape.conv_kernel as u64)
.saturating_mul(inner)
.saturating_mul(3);
let alpha_beta = d.saturating_mul(heads).saturating_mul(2);
let per_head = heads.saturating_mul(2);
let out = inner.saturating_mul(d);
let d_attn = qkv
.saturating_add(gate)
.saturating_add(conv)
.saturating_add(alpha_beta)
.saturating_add(per_head)
.saturating_add(shape.state_size as u64)
.saturating_add(out);
LayerParams {
d_attn,
d_ffn: ffn_params(size, constraints),
d_norm: d.saturating_mul(2),
}
}
#[must_use]
pub fn hybrid_layers(size: &ModelSizeConfig, constraints: &ModelConstraints) -> Vec<LayerParams> {
let Some(shape) = constraints.deltanet else {
return uniform_layers(size, constraints);
};
if shape.full_attention_interval == 0 {
return uniform_layers(size, constraints);
}
let attention = attention_layer_params(size, constraints);
let deltanet = gated_deltanet_layer_params(size, constraints, &shape);
(0..size.num_layers)
.map(|i| {
if (i + 1) % shape.full_attention_interval == 0 {
attention
} else {
deltanet
}
})
.collect()
}
#[must_use]
pub fn uniform_layers(size: &ModelSizeConfig, constraints: &ModelConstraints) -> Vec<LayerParams> {
vec![attention_layer_params(size, constraints); size.num_layers]
}
#[must_use]
pub fn model_parameter_count(
size: &ModelSizeConfig,
constraints: &ModelConstraints,
layers: &[LayerParams],
) -> ParameterBreakdown {
let v = size.vocab_size as u64;
let d = size.hidden_dim as u64;
let embedding = v.saturating_mul(d);
let layer_total = layers
.iter()
.fold(0u64, |acc, l| acc.saturating_add(l.total()));
let unembedding = if constraints.tied_embeddings {
0
} else {
embedding
};
let total = embedding
.saturating_add(layer_total)
.saturating_add(d)
.saturating_add(unembedding);
ParameterBreakdown {
embedding,
layers: layer_total,
final_norm: d,
unembedding,
total,
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct FlopsPerToken {
pub dense: u64,
pub attention: u64,
pub total: u64,
}
#[must_use]
pub fn flops_per_token(
parameter_count: u64,
seq_len: u64,
hidden_dim: u64,
attn_layers: u64,
) -> FlopsPerToken {
let dense = parameter_count.saturating_mul(2);
let attention = attn_layers
.saturating_mul(seq_len)
.saturating_mul(hidden_dim)
.saturating_mul(2);
FlopsPerToken {
dense,
attention,
total: dense.saturating_add(attention),
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
pub enum Precision {
Q4K,
Q5K,
Q6K,
F16,
F32,
}
impl Precision {
#[must_use]
pub const fn bytes_for(self, elements: u64) -> u64 {
let qk = block_sizes::QK_K as u64;
match self {
Self::Q4K => elements
.div_ceil(qk)
.saturating_mul(block_sizes::Q4_K as u64),
Self::Q5K => elements
.div_ceil(qk)
.saturating_mul(block_sizes::Q5_K as u64),
Self::Q6K => elements
.div_ceil(qk)
.saturating_mul(block_sizes::Q6_K as u64),
Self::F16 => elements.saturating_mul(2),
Self::F32 => elements.saturating_mul(4),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct InferencePlan {
pub seq_len: u64,
pub batch_size: u64,
pub kv_layers: u64,
pub weights: Precision,
pub kv_cache: Precision,
pub activations: Precision,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct MemoryBreakdown {
pub weights: u64,
pub kv: u64,
pub activations: u64,
pub total: u64,
}
#[must_use]
pub fn memory_breakdown(
size: &ModelSizeConfig,
parameter_count: u64,
plan: &InferencePlan,
) -> MemoryBreakdown {
let d = size.hidden_dim as u64;
let n_kv = size.num_kv_heads as u64;
let d_k = size.head_dim as u64;
let kv_elements = plan
.kv_layers
.saturating_mul(n_kv)
.saturating_mul(d_k)
.saturating_mul(plan.seq_len)
.saturating_mul(2);
let act_elements = plan
.batch_size
.saturating_mul(plan.seq_len)
.saturating_mul(d);
let weights = plan.weights.bytes_for(parameter_count);
let kv = plan.kv_cache.bytes_for(kv_elements);
let activations = plan.activations.bytes_for(act_elements);
MemoryBreakdown {
weights,
kv,
activations,
total: weights.saturating_add(kv).saturating_add(activations),
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RooflineLimit {
MemoryBound,
ComputeBound,
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Throughput {
pub memory_bound: f64,
pub compute_bound: f64,
pub tokens_per_second: f64,
pub limit: RooflineLimit,
}
#[must_use]
pub fn throughput_model(
bandwidth_bytes_per_s: f64,
compute_flops_per_s: f64,
bytes_per_token: f64,
flops_per_token: f64,
) -> Throughput {
let ratio = |numerator: f64, denominator: f64| {
if denominator > 0.0 && numerator > 0.0 {
numerator / denominator
} else {
0.0
}
};
let memory_bound = ratio(bandwidth_bytes_per_s, bytes_per_token);
let compute_bound = ratio(compute_flops_per_s, flops_per_token);
let (tokens_per_second, limit) = if memory_bound <= compute_bound {
(memory_bound, RooflineLimit::MemoryBound)
} else {
(compute_bound, RooflineLimit::ComputeBound)
};
Throughput {
memory_bound,
compute_bound,
tokens_per_second,
limit,
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CompositionStage {
pub component: String,
pub input: Vec<usize>,
pub output: Vec<usize>,
}
impl CompositionStage {
#[must_use]
pub fn preserves_shape(&self) -> bool {
self.input == self.output
}
}
#[must_use]
pub fn contract_composition(size: &ModelSizeConfig, seq_len: usize) -> Vec<CompositionStage> {
let d = size.hidden_dim;
let v = size.vocab_size;
let residual = vec![seq_len, d];
let mut stages = Vec::with_capacity(size.num_layers.saturating_add(3));
stages.push(CompositionStage {
component: "embedding".to_string(),
input: vec![seq_len],
output: residual.clone(),
});
for l in 0..size.num_layers {
stages.push(CompositionStage {
component: format!("block_{l}"),
input: residual.clone(),
output: residual.clone(),
});
}
stages.push(CompositionStage {
component: "final_norm".to_string(),
input: residual.clone(),
output: residual.clone(),
});
stages.push(CompositionStage {
component: "unembed".to_string(),
input: residual,
output: vec![seq_len, v],
});
stages
}
#[must_use]
pub fn composition_is_well_formed(stages: &[CompositionStage]) -> bool {
stages.windows(2).all(|w| w[0].output == w[1].input)
}
#[cfg(test)]
#[path = "model_arithmetic_tests.rs"]
mod model_arithmetic_tests;