use frink_core::matmul::rms_norm;
use super::{Decoder, GptOssLayer, LayerWeights};
use crate::norm::NormOp;
use crate::router_input::RouterInput;
use crate::scalar_multipliers::residual_add;
use crate::skip_stream::SkipStream;
#[derive(Debug)]
pub(crate) enum RouterOperand {
FfnInput,
Precomputed(Vec<f32>),
BranchInput(Vec<f32>),
}
#[derive(Debug)]
pub(crate) enum FfnInput {
PostAttnResidual,
LayerInput(Vec<f32>),
}
#[derive(Debug)]
pub(crate) struct BranchInputs {
pub(crate) router: RouterOperand,
pub(crate) ffn: FfnInput,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum BatchedFfnKernels {
Prefill,
PerRow,
}
impl Decoder {
pub fn layer_for(&self, l: usize) -> &LayerWeights {
&self.layers[self.physical_index(l)]
}
pub(crate) fn physical_index(&self, l: usize) -> usize {
match self.config.layer_loops {
Some(loops) => loops.physical(l),
None => l,
}
}
pub(crate) fn run_dense_expert(
layer: &LayerWeights,
normed2: &[f32],
act: frink_moe::GluAct,
eps: f32,
) -> Vec<f32> {
layer.moe.with_expert(0, |ex| {
match (&layer.moe.ffn_sub_norm, &layer.moe.dense_bias) {
(None, None) => frink_moe::run_expert(normed2, ex, act),
(Some(w), None) => frink_moe::run_expert_sub_normed(normed2, ex, act, w, eps),
(None, Some(bias)) => frink_moe::run_expert_biased(normed2, ex, act, bias),
(Some(_), Some(_)) => {
unreachable!(
"a dense layer with both an inner norm and biases is refused at load"
)
}
}
})
}
pub(crate) fn branch_inputs(
&self,
layer: &LayerWeights,
hidden_before_attn: &[f32],
batch_size: usize,
) -> BranchInputs {
BranchInputs {
router: self.router_operand(layer, hidden_before_attn, batch_size),
ffn: self.ffn_input(layer, hidden_before_attn, batch_size),
}
}
fn ffn_input(
&self,
layer: &LayerWeights,
hidden_before_attn: &[f32],
batch_size: usize,
) -> FfnInput {
use crate::parallel_residual::ParallelNorm;
let norm = match layer.moe.parallel {
None => return FfnInput::PostAttnResidual,
Some(ParallelNorm::SharedNorm) => {
debug_assert!(
matches!(layer.moe.norm_weight, NormOp::None),
"a shared-norm parallel layer has no pre-FFN tensor"
);
&layer.attn.norm_weight
}
Some(ParallelNorm::TwoNorms) => &layer.moe.norm_weight,
};
debug_assert_eq!(
hidden_before_attn.len(),
batch_size * self.config.hidden_dim
);
let eps = self.config.rms_norm_eps;
FfnInput::LayerInput(
hidden_before_attn
.chunks(self.config.hidden_dim)
.flat_map(|row| norm.apply(row, eps))
.collect(),
)
}
fn router_operand(
&self,
layer: &LayerWeights,
hidden_before_attn: &[f32],
batch_size: usize,
) -> RouterOperand {
match self.config.router_input {
RouterInput::NormedFfnInput => RouterOperand::FfnInput,
RouterInput::NormedLayerInput => {
if Self::is_dense_layer(layer) {
return RouterOperand::FfnInput;
}
debug_assert_eq!(
hidden_before_attn.len(),
batch_size * self.config.hidden_dim
);
let w = layer
.moe
.exps_norm
.as_deref()
.expect("NormedLayerInput layer loaded without ffn_norm_exps");
let eps = self.config.rms_norm_eps;
RouterOperand::BranchInput(
hidden_before_attn
.chunks(self.config.hidden_dim)
.flat_map(|row| rms_norm(row, w, eps))
.collect(),
)
}
RouterInput::RawLayerInput => {
debug_assert!(
self.gpt_oss.is_none(),
"gpt-oss routes on the normed input; a RawLayerInput gpt-oss model is not a \
shape llama.cpp has"
);
if Self::is_dense_layer(layer) {
return RouterOperand::FfnInput;
}
debug_assert_eq!(
hidden_before_attn.len(),
batch_size * self.config.hidden_dim
);
RouterOperand::Precomputed(if batch_size == 1 {
layer.moe.router.apply(hidden_before_attn)
} else {
layer.moe.router.apply_batch(hidden_before_attn, batch_size)
})
}
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn ffn_block_row(
&self,
layer_idx: usize,
layer: &LayerWeights,
hidden: &mut [f32],
oai: Option<&GptOssLayer>,
plan: Option<&frink_moe::PlacementPlan>,
inputs: BranchInputs,
skip: Option<SkipStream<'_>>,
) {
if self.config.layer_shape(layer_idx).ffn_dim == 0 {
return;
}
let hidden_dim = self.config.hidden_dim;
let BranchInputs {
router: operand,
ffn,
} = inputs;
let normed2 = match ffn {
FfnInput::PostAttnResidual => self.pre_norm_residual(&layer.moe.norm_weight, hidden, 1),
FfnInput::LayerInput(x) => x,
};
let mut ffn_out = match oai {
Some(oai) => Self::gpt_oss_ffn(layer, oai, &normed2, &self.config, hidden_dim),
None => Self::run_ffn_block(
layer_idx,
layer,
&normed2,
&self.config,
hidden_dim,
plan,
operand,
),
};
Self::apply_parallel_sum_scale(layer, &mut ffn_out);
Self::apply_down_scale(layer, &mut ffn_out);
if let Some(post) = &layer.attn.post_ffn_norm {
ffn_out = rms_norm(&ffn_out, post, self.config.post_norm_eps());
}
residual_add(hidden, &ffn_out, self.config.residual_scale);
Self::apply_skip_stream(layer, hidden, skip, 1, hidden_dim);
self.apply_loop_norm(layer_idx, hidden, 1);
}
fn apply_parallel_sum_scale(layer: &LayerWeights, ffn_out: &mut [f32]) {
if let Some(scale) = layer.moe.parallel_sum_scale {
for x in ffn_out.iter_mut() {
*x *= scale;
}
}
}
fn apply_down_scale(layer: &LayerWeights, ffn_out: &mut [f32]) {
if let Some(scale) = layer.moe.down_scale {
for x in ffn_out.iter_mut() {
*x *= scale;
}
}
}
fn apply_skip_stream(
layer: &LayerWeights,
hidden: &mut [f32],
skip: Option<SkipStream<'_>>,
rows: usize,
hidden_dim: usize,
) {
match (skip, layer.out_scale) {
(Some(skip), Some(scale)) => {
debug_assert_eq!(skip.rows.len(), rows * hidden_dim);
debug_assert_eq!(hidden.len(), rows * hidden_dim);
for (h, s) in hidden.iter_mut().zip(skip.rows.iter()) {
*h += s * scale;
}
}
(None, None) => {}
(skip, scale) => unreachable!(
"the skip stream and the layer's out_scale come from one config fact; \
got skip={} out_scale={scale:?}",
skip.is_some()
),
}
}
fn apply_loop_norm(&self, layer_idx: usize, hidden: &mut [f32], rows: usize) {
let Some(loops) = self.config.layer_loops else {
return;
};
let Some(kind) = loops.loop_norm_after(layer_idx) else {
return;
};
let width = self.config.hidden_dim;
debug_assert_eq!(hidden.len(), rows * width);
for row in hidden.chunks_mut(width) {
let normed = match kind {
crate::layer_loops::LoopNorm::Output => {
self.final_norm.apply(row, self.config.rms_norm_eps)
}
crate::layer_loops::LoopNorm::Weightless => {
crate::norm::NormOp::RmsNoParams.apply(row, self.config.rms_norm_eps)
}
};
row.copy_from_slice(&normed);
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn ffn_block_batch(
&self,
layer_idx: usize,
layer: &LayerWeights,
hidden_batch: &mut [f32],
batch_size: usize,
oai: Option<&GptOssLayer>,
plan: Option<&frink_moe::PlacementPlan>,
inputs: BranchInputs,
kernels: BatchedFfnKernels,
skip: Option<SkipStream<'_>>,
) {
if self.config.layer_shape(layer_idx).ffn_dim == 0 {
return;
}
let hidden_dim = self.config.hidden_dim;
let config = &self.config;
let BranchInputs {
router: operand,
ffn,
} = inputs;
let normed2_batch: Vec<f32> = match ffn {
FfnInput::PostAttnResidual => {
self.pre_norm_residual(&layer.moe.norm_weight, hidden_batch, batch_size)
}
FfnInput::LayerInput(x) => {
debug_assert_eq!(x.len(), batch_size * hidden_dim);
x
}
};
if let Some(oai) = oai {
for b in 0..batch_size {
let normed2 = &normed2_batch[b * hidden_dim..(b + 1) * hidden_dim];
let ffn_out = Self::gpt_oss_ffn(layer, oai, normed2, config, hidden_dim);
let hidden_row = &mut hidden_batch[b * hidden_dim..(b + 1) * hidden_dim];
residual_add(hidden_row, &ffn_out, config.residual_scale);
}
Self::apply_skip_stream(layer, hidden_batch, skip, batch_size, hidden_dim);
self.apply_loop_norm(layer_idx, hidden_batch, batch_size);
return;
}
let dense = Self::is_dense_layer(layer);
let (router_logits_batch, routed_batch): (Vec<f32>, &[f32]) = match &operand {
_ if dense => (Vec::new(), normed2_batch.as_slice()),
RouterOperand::FfnInput => (
layer.moe.router.apply_batch(&normed2_batch, batch_size),
normed2_batch.as_slice(),
),
RouterOperand::Precomputed(logits) => (logits.clone(), normed2_batch.as_slice()),
RouterOperand::BranchInput(x) => {
(layer.moe.router.apply_batch(x, batch_size), x.as_slice())
}
};
let batched: Option<Vec<f32>> = match kernels {
BatchedFfnKernels::PerRow => None,
BatchedFfnKernels::Prefill => {
#[cfg(feature = "metal")]
let metal_ffn = if !dense {
Self::try_metal_moe_prefill_batch(
layer_idx,
layer,
&normed2_batch,
&router_logits_batch,
batch_size,
hidden_dim,
config,
)
} else {
None
};
#[cfg(not(feature = "metal"))]
let metal_ffn: Option<Vec<f32>> = None;
metal_ffn
.or_else(|| {
Self::dense_ffn_batch(layer_idx, layer, &normed2_batch, batch_size, config)
})
.or_else(|| {
Self::moe_ffn_batch(
layer_idx,
layer,
&normed2_batch,
routed_batch,
&router_logits_batch,
batch_size,
config,
plan,
)
})
}
};
if let Some(mut ffn_batch) = batched {
Self::apply_parallel_sum_scale(layer, &mut ffn_batch);
Self::apply_down_scale(layer, &mut ffn_batch);
if let Some(post) = &layer.attn.post_ffn_norm {
ffn_batch = ffn_batch
.chunks(hidden_dim)
.flat_map(|row| rms_norm(row, post, config.post_norm_eps()))
.collect();
}
residual_add(hidden_batch, &ffn_batch, config.residual_scale);
Self::apply_skip_stream(layer, hidden_batch, skip, batch_size, hidden_dim);
self.apply_loop_norm(layer_idx, hidden_batch, batch_size);
return;
}
let n_experts = layer.moe.n_experts().max(1);
for b in 0..batch_size {
let normed2 = &normed2_batch[b * hidden_dim..(b + 1) * hidden_dim];
let mut ffn_out = if dense {
Self::run_ffn_block(
layer_idx,
layer,
normed2,
config,
hidden_dim,
plan,
RouterOperand::FfnInput,
)
} else {
let router_logits = &router_logits_batch[b * n_experts..(b + 1) * n_experts];
Self::combine_ffn_outputs_for_position(
layer_idx,
layer,
normed2,
&routed_batch[b * hidden_dim..(b + 1) * hidden_dim],
router_logits,
config,
hidden_dim,
plan,
)
};
Self::apply_parallel_sum_scale(layer, &mut ffn_out);
Self::apply_down_scale(layer, &mut ffn_out);
if let Some(post) = &layer.attn.post_ffn_norm {
ffn_out = rms_norm(&ffn_out, post, config.post_norm_eps());
}
let hidden_row = &mut hidden_batch[b * hidden_dim..(b + 1) * hidden_dim];
residual_add(hidden_row, &ffn_out, config.residual_scale);
}
Self::apply_skip_stream(layer, hidden_batch, skip, batch_size, hidden_dim);
self.apply_loop_norm(layer_idx, hidden_batch, batch_size);
}
}