use super::{AttnWeights, Decoder, LayerWeights};
use crate::config::ModelConfig;
#[derive(Clone, Copy)]
pub(crate) struct FusedAttnExtras<'a> {
pub q_bias: Option<&'a [f32]>,
pub k_bias: Option<&'a [f32]>,
pub v_bias: Option<&'a [f32]>,
pub q_norm: Option<&'a [f32]>,
pub k_norm: Option<&'a [f32]>,
}
impl Decoder {
pub(crate) fn fused_attn_extras<'a>(layer: &'a LayerWeights) -> Option<FusedAttnExtras<'a>> {
let AttnWeights {
q_proj: _,
k_proj: _,
v_proj: _,
o_proj: _,
norm_weight: _,
q_norm,
k_norm,
q_bias,
k_bias,
v_bias,
post_attn_norm: _,
post_ffn_norm: _,
output_gate,
sinks,
attn_sub_norm,
o_scale,
o_bias,
shortconv,
ssm,
q_gate_interleaved,
} = &layer.attn;
if output_gate.is_some()
|| sinks.is_some()
|| attn_sub_norm.is_some()
|| o_scale.is_some()
|| o_bias.is_some()
|| shortconv.is_some()
|| ssm.is_some()
|| *q_gate_interleaved
{
return None;
}
Some(FusedAttnExtras {
q_bias: q_bias.as_deref(),
k_bias: k_bias.as_deref(),
v_bias: v_bias.as_deref(),
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
})
}
pub(crate) fn layer_supports_fused_attn(&self, layer: &LayerWeights) -> bool {
use crate::config::RopeLayout;
if !Self::metal_can_serve_model(&self.config, self.lora_attached()) {
return false;
}
if self.gpt_oss.is_some() {
return false;
}
if Self::fused_attn_extras(layer).is_none() {
return false;
}
if !matches!(self.config.rope_layout, RopeLayout::Norm | RopeLayout::Neox) {
return false;
}
if self.config.qk_norm_style == crate::capability::QkNormStyle::PerHeadDistinct {
return false;
}
let q_len = self.config.n_heads * self.config.head_dim;
let k_len = self.config.n_kv_heads * self.config.head_dim;
let qk_norm_ok = |w: Option<&Vec<f32>>, vec_len: usize| -> bool {
match w {
None => true,
Some(w) if w.len() == self.config.head_dim => true,
Some(w) if w.len() == vec_len => true,
_ => false,
}
};
if !qk_norm_ok(layer.attn.q_norm.as_ref(), q_len)
|| !qk_norm_ok(layer.attn.k_norm.as_ref(), k_len)
{
return false;
}
if self.qk_norm_after_rope {
return false;
}
if self.config.attention_scale.is_some() {
return false;
}
if self.config.head_dim > 256 {
return false;
}
!self
.config
.rope_dim
.is_some_and(|rot| rot == 0 || rot % 2 != 0 || rot > self.config.head_dim)
}
pub(crate) fn fused_prefill_dense_layer_eligible(
layer: &LayerWeights,
config: &ModelConfig,
lora_attached: bool,
) -> bool {
Self::is_dense_layer(layer)
&& layer.moe.down_scale.is_none()
&& layer.moe.dense_bias.is_none()
&& Self::metal_can_serve_model(config, lora_attached)
}
}