use crate::config::{KdaConfig, MlaConfig};
use crate::decoder::Decoder;
use crate::deepseek_v4_decoder::{
deepseek_v4_forward_token, DeepseekV4DecodeState, DeepseekV4DecoderConfig,
DeepseekV4DecoderWeights,
};
use crate::glm52_decoder::{
glm52_forward_token, Glm52DecodeState, Glm52DecoderConfig, Glm52DecoderWeights,
};
use crate::kimi_decoder::{
kimi_forward_token, KimiDecodeState, KimiDecoderConfig, KimiDecoderWeights,
};
use crate::kimi_tokenizer::KimiTokenizer;
use ferrox_core::cache::KvCache;
use ferrox_core::weight_matrix::WeightMatrix;
pub trait Engine {
type State;
fn new_state(&self) -> Self::State;
fn vocab_size(&self) -> usize;
fn forward_token(&self, token_id: usize, pos: usize, state: &mut Self::State) -> Vec<f32>;
}
impl Engine for Decoder {
type State = Vec<KvCache>;
fn new_state(&self) -> Vec<KvCache> {
self.layers
.iter()
.map(|_| KvCache::new(self.config.n_kv_heads, self.config.head_dim))
.collect()
}
fn vocab_size(&self) -> usize {
self.config.vocab_size
}
fn forward_token(&self, token_id: usize, pos: usize, state: &mut Self::State) -> Vec<f32> {
Decoder::forward_token(self, token_id, pos, state)
}
}
pub struct KimiEngine {
pub weights: KimiDecoderWeights,
pub cfg: KimiDecoderConfig,
pub mla_cfg: MlaConfig,
pub kda_cfg: KdaConfig,
}
impl Engine for KimiEngine {
type State = KimiDecodeState;
fn new_state(&self) -> KimiDecodeState {
KimiDecodeState::new(&self.weights, &self.kda_cfg)
}
fn vocab_size(&self) -> usize {
self.weights.output_head.rows()
}
fn forward_token(&self, token_id: usize, _pos: usize, state: &mut Self::State) -> Vec<f32> {
kimi_forward_token(
&self.weights,
&self.cfg,
&self.mla_cfg,
&self.kda_cfg,
token_id,
state,
)
}
}
pub struct Glm52Engine {
pub weights: Glm52DecoderWeights,
pub cfg: Glm52DecoderConfig,
}
impl Engine for Glm52Engine {
type State = Glm52DecodeState;
fn new_state(&self) -> Glm52DecodeState {
Glm52DecodeState::new(&self.weights)
}
fn vocab_size(&self) -> usize {
self.weights.output_head.rows()
}
fn forward_token(&self, token_id: usize, _pos: usize, state: &mut Self::State) -> Vec<f32> {
glm52_forward_token(&self.weights, &self.cfg, token_id, state)
}
}
pub struct MlaEngine {
pub embedding: WeightMatrix,
pub layers: Vec<MlaLayerWeights>,
pub final_norm: Vec<f32>,
pub output_head: WeightMatrix,
pub mla_cfg: MlaConfig,
pub rms_norm_eps: f32,
pub hidden_dim: usize,
pub moe: Option<MlaMoeRuntime>,
}
#[derive(Debug, Clone)]
pub struct MlaMoeRuntime {
pub n_experts_active: usize,
pub gating: ferrox_moe::GatingFunction,
pub norm_topk_prob: bool,
pub expert_weights_scale: f32,
}
pub struct MlaDenseFfn {
pub gate: WeightMatrix,
pub up: WeightMatrix,
pub down: WeightMatrix,
}
pub struct MlaMoeFfn {
pub router: WeightMatrix,
pub experts: Vec<ferrox_moe::ExpertWeights>,
pub shared_expert: ferrox_moe::ExpertWeights,
pub exp_probs_bias: Option<Vec<f32>>,
}
pub enum MlaLayerFfn {
Dense(MlaDenseFfn),
Moe(MlaMoeFfn),
}
pub struct MlaLayerWeights {
pub attn_norm: Vec<f32>,
pub attn: crate::mla::MlaAttnWeights,
pub ffn_norm: Vec<f32>,
pub ffn: MlaLayerFfn,
}
pub struct MlaDecodeState {
pub layers: Vec<(Vec<f32>, Vec<f32>)>,
}
impl MlaEngine {
pub fn new_state(&self) -> MlaDecodeState {
MlaDecodeState {
layers: (0..self.layers.len())
.map(|_| (Vec::new(), Vec::new()))
.collect(),
}
}
fn moe_ffn_forward(&self, ffn: &MlaMoeFfn, x: &[f32]) -> Vec<f32> {
use ferrox_moe::{
combine_expert_outputs, route_top_k, route_top_k_sigmoid_with_bias, run_expert,
GatingFunction,
};
let moe = self
.moe
.as_ref()
.expect("MlaLayerFfn::Moe requires MlaEngine.moe");
let router_logits = ffn.router.apply(x);
let decision = match (moe.gating, ffn.exp_probs_bias.as_deref()) {
(GatingFunction::Sigmoid, Some(bias)) => route_top_k_sigmoid_with_bias(
&router_logits,
bias,
moe.n_experts_active,
moe.norm_topk_prob,
moe.expert_weights_scale,
),
_ => {
let mut d = route_top_k(
&router_logits,
moe.n_experts_active,
moe.gating,
moe.norm_topk_prob,
);
if (moe.expert_weights_scale - 1.0).abs() > f32::EPSILON {
for w in d.weights.iter_mut() {
*w *= moe.expert_weights_scale;
}
}
d
}
};
let routed: Vec<(Vec<f32>, f32)> = decision
.expert_ids
.iter()
.zip(decision.weights.iter())
.map(|(&e, &w)| (run_expert(x, &ffn.experts[e]), w))
.collect();
let shared = run_expert(x, &ffn.shared_expert);
combine_expert_outputs(&routed, &[shared], x.len())
}
}
impl Engine for MlaEngine {
type State = MlaDecodeState;
fn new_state(&self) -> MlaDecodeState {
MlaEngine::new_state(self)
}
fn vocab_size(&self) -> usize {
self.output_head.rows()
}
fn forward_token(&self, token_id: usize, _pos: usize, state: &mut Self::State) -> Vec<f32> {
use ferrox_core::matmul::{rms_norm, swiglu};
let mut hidden = self.embedding.dequant_row(token_id);
for (layer, (k_cache, v_cache)) in self.layers.iter().zip(state.layers.iter_mut()) {
let normed = rms_norm(&hidden, &layer.attn_norm, self.rms_norm_eps);
let attn_out = crate::mla::mla_forward_token(
&layer.attn,
&self.mla_cfg,
&normed,
self.rms_norm_eps,
k_cache,
v_cache,
);
for (h, a) in hidden.iter_mut().zip(attn_out.iter()) {
*h += a;
}
let ffn_in = rms_norm(&hidden, &layer.ffn_norm, self.rms_norm_eps);
let down = match &layer.ffn {
MlaLayerFfn::Dense(d) => {
let gate = d.gate.apply(&ffn_in);
let up = d.up.apply(&ffn_in);
d.down.apply(&swiglu(&gate, &up))
}
MlaLayerFfn::Moe(m) => self.moe_ffn_forward(m, &ffn_in),
};
for (h, d) in hidden.iter_mut().zip(down.iter()) {
*h += d;
}
}
let final_normed = rms_norm(&hidden, &self.final_norm, self.rms_norm_eps);
self.output_head.apply(&final_normed)
}
}
pub struct DeepseekV4Engine {
pub weights: DeepseekV4DecoderWeights,
pub cfg: DeepseekV4DecoderConfig,
}
impl Engine for DeepseekV4Engine {
type State = DeepseekV4DecodeState;
fn new_state(&self) -> DeepseekV4DecodeState {
DeepseekV4DecodeState::new(self.weights.embedding.shape[1])
}
fn vocab_size(&self) -> usize {
self.weights.output_head.rows()
}
fn forward_token(&self, token_id: usize, _pos: usize, state: &mut Self::State) -> Vec<f32> {
deepseek_v4_forward_token(&self.weights, &self.cfg, token_id, state)
}
}
pub trait TextTokenizer {
fn encode(&self, text: &str) -> Vec<usize>;
fn decode(&self, ids: &[usize]) -> String;
}
impl TextTokenizer for KimiTokenizer {
fn encode(&self, text: &str) -> Vec<usize> {
KimiTokenizer::encode(self, text)
.into_iter()
.map(|id| id as usize)
.collect()
}
fn decode(&self, ids: &[usize]) -> String {
let ids32: Vec<u32> = ids.iter().map(|&id| id as u32).collect();
KimiTokenizer::decode(self, &ids32)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::test_dense_fixture;
#[test]
fn decoder_via_engine_trait_matches_direct_forward_token_calls() {
let decoder = Decoder::new_random_small(test_dense_fixture(), 2, 64);
let tokens = [3usize, 7, 1, 9];
let mut direct_caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let mut direct_logits = Vec::new();
for (pos, &tok) in tokens.iter().enumerate() {
direct_logits = decoder.forward_token(tok, pos, &mut direct_caches);
}
let mut engine_state = Engine::new_state(&decoder);
let mut engine_logits = Vec::new();
for (pos, &tok) in tokens.iter().enumerate() {
engine_logits = Engine::forward_token(&decoder, tok, pos, &mut engine_state);
}
assert_eq!(engine_logits, direct_logits);
assert_eq!(Engine::vocab_size(&decoder), decoder.config.vocab_size);
}
#[test]
fn decoder_via_engine_trait_matches_forward_batch_ground_truth() {
let decoder = Decoder::new_random_small(test_dense_fixture(), 2, 64);
let tokens = vec![2usize, 5, 8];
let mut batch_caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let batch_logits = decoder.forward_batch(&tokens, 0, &mut batch_caches);
let ground_truth = batch_logits.last().unwrap().clone();
let mut engine_state = Engine::new_state(&decoder);
let mut engine_logits = Vec::new();
for (pos, &tok) in tokens.iter().enumerate() {
engine_logits = Engine::forward_token(&decoder, tok, pos, &mut engine_state);
}
assert_eq!(engine_logits.len(), ground_truth.len());
for (i, (e, g)) in engine_logits.iter().zip(ground_truth.iter()).enumerate() {
assert!(
(e - g).abs() < 1e-5,
"logit {i}: engine {e} vs forward_batch {g}"
);
}
}
}