use super::{Decoder, LayerWeights};
use crate::fused_layer::LayerFfnParts;
use crate::layer_shapes::AttnShape;
use crate::ssm_block::SsmBlock;
use frink_core::recurrent_state::RecurrentState;
pub(crate) type PendingQkv = Option<(usize, (Vec<f32>, Vec<f32>, Vec<f32>))>;
impl Decoder {
#[cfg(feature = "metal")]
pub(crate) fn fused_recurrent_layer(
&self,
l: usize,
layer: &LayerWeights,
normed: &[f32],
hidden: &[f32],
slot: &mut Option<RecurrentState>,
) -> Option<Vec<f32>> {
let shape = self.config.layer_shape(l);
if !matches!(shape.attention, AttnShape::Gdn)
|| shape.ffn_dim == 0
|| self.config.residual_scale.is_some()
|| self.config.normed_residual_scale.is_some()
|| self.config.skip_stream
|| self.gpt_oss.is_some()
|| !self.config.layer_ffn_acts(l).all_swiglu()
{
return None;
}
let SsmBlock::Gdn(gdn) = layer.attn.ssm.as_ref()? else {
return None;
};
let ffn = LayerFfnParts::for_layer(layer, self.config.rms_norm_eps, true)?;
let state = slot.get_or_insert_with(|| gdn.zero_state());
let attn_norm = layer.attn.norm_weight.rms_weights()?;
let out = gdn.fused_layer(
attn_norm,
normed,
state,
self.config.rms_norm_eps,
&ffn,
hidden,
)?;
layer.moe.record_activations_dense();
Some(out)
}
}
impl Decoder {
#[cfg(feature = "metal")]
pub(crate) fn fused_recurrent_run(
&self,
start: usize,
hidden: &mut Vec<f32>,
kv_caches: &mut [frink_core::KvCache],
pending: &mut PendingQkv,
) -> Option<usize> {
let mut end = start;
while end < kv_caches.len() && self.fused_layer_parts(end).is_some() {
end += 1;
}
if end - start < 2 {
return None;
}
let mut run = frink_metal::gdn_branch::GdnRun::start(hidden).ok()?;
self.run_layers(&mut run, start, end, kv_caches)?;
let head = self.encode_next_attn_head(&mut run, end, kv_caches.len());
let (out, qkv) = run.finish_with_head().ok()?;
*hidden = out;
*pending = head.then_some(qkv).flatten().map(|q| (end, q));
Some(end)
}
#[cfg(feature = "metal")]
pub(crate) fn fused_attention_tail_then_run(
&self,
l: usize,
layer: &LayerWeights,
branch: &[f32],
hidden: &mut Vec<f32>,
kv_caches: &mut [frink_core::KvCache],
pending: &mut PendingQkv,
) -> Option<usize> {
let (out_proj, fold_branch, ffn, launches) = self.attn_tail_launch(l, layer, branch)?;
let mut end = l + 1;
while end < kv_caches.len() && self.fused_layer_parts(end).is_some() {
end += 1;
}
if end == l + 1 {
return None;
}
let mut run = frink_metal::gdn_branch::GdnRun::start(hidden).ok()?;
run.attn_tail(
&out_proj,
fold_branch.as_ref(),
&launches.as_metal(),
branch,
)
.ok()?;
let _ = ffn;
layer.moe.record_activations_dense();
self.run_layers(&mut run, l + 1, end, kv_caches)?;
let head = self.encode_next_attn_head(&mut run, end, kv_caches.len());
let (out, qkv) = run.finish_with_head().ok()?;
*hidden = out;
*pending = head.then_some(qkv).flatten().map(|q| (end, q));
Some(end)
}
#[cfg(feature = "metal")]
pub(crate) fn encode_next_attn_head(
&self,
run: &mut frink_metal::gdn_branch::GdnRun,
next: usize,
n_layers: usize,
) -> bool {
if next >= n_layers {
return false;
}
let layer = self.layer_for(next);
let Some((norm, q, k, v, fold_x)) = self.attn_head_launch(next, layer) else {
return false;
};
run.attn_head(norm, self.config.rms_norm_eps, &q, &k, &v, fold_x.as_ref())
.is_ok()
}
#[cfg(feature = "metal")]
pub(crate) fn run_layers(
&self,
run: &mut frink_metal::gdn_branch::GdnRun,
start: usize,
end: usize,
kv_caches: &mut [frink_core::KvCache],
) -> Option<()> {
for (l, cache) in (start..end).zip(kv_caches[start..end].iter_mut()) {
let layer = self.layer_for(l);
let (gdn, attn_norm, ffn) = self.fused_layer_parts(l)?;
let state = cache.recurrent.get_or_insert_with(|| gdn.zero_state());
unsafe { gdn.run_layer(run, attn_norm, self.config.rms_norm_eps, &ffn, state) }?;
cache
.advance_len(1)
.expect("unbounded/planned KvCache growth is infallible");
layer.moe.record_activations_dense();
}
Some(())
}
#[cfg(feature = "metal")]
pub(crate) fn fused_layer_parts(
&self,
l: usize,
) -> Option<(&crate::gdn::Gdn, &[f32], LayerFfnParts<'_>)> {
let layer = self.layer_for(l);
let shape = self.config.layer_shape(l);
if !matches!(shape.attention, AttnShape::Gdn)
|| shape.ffn_dim == 0
|| self.config.residual_scale.is_some()
|| self.config.normed_residual_scale.is_some()
|| self.config.skip_stream
|| self.gpt_oss.is_some()
|| !self.config.layer_ffn_acts(l).all_swiglu()
{
return None;
}
let SsmBlock::Gdn(gdn) = layer.attn.ssm.as_ref()? else {
return None;
};
let attn_norm = layer.attn.norm_weight.rms_weights()?;
let ffn = LayerFfnParts::for_layer(layer, self.config.rms_norm_eps, true)?;
Some((gdn, attn_norm, ffn))
}
}