Skip to main content

ferrox_models/
engine.rs

1//! A trait-based abstraction over the two structurally different but
2//! text-in/text-out-shaped forward passes this crate has: `Decoder`
3//! (GQA+RoPE, used by GLM-5.2/DeepSeek V4 Pro and every real GGUF
4//! checkpoint) and Kimi K3's dedicated hybrid KDA/Gated-MLA stack
5//! (`kimi_decoder`). This lets `ferrox-server` share one generic
6//! generation loop across both engines (see
7//! `ferrox-server::generate::generate_engine`) for the actual
8//! sampling/stop-sequence logic, rather than hand-duplicating it --
9//! while keeping GGUF-only features (the KV block pool, `PrefixCache`)
10//! as `Decoder`-specific code layered on top, not forced into this
11//! trait. The reason: Kimi's KDA state is
12//! a fixed-size recurrent matrix that collapses history irreversibly,
13//! so it cannot support the same restore/truncate operations
14//! `KvCache` can -- unifying those too would mean either a leaky
15//! abstraction or silently pretending Kimi supports something it
16//! doesn't.
17//!
18//! `forward_token`'s `pos` parameter is meaningful for `Decoder` (used
19//! directly for RoPE) but not for `KimiEngine`: Kimi's real forward
20//! pass (`kimi_forward_token`) derives position purely from its own
21//! per-layer state (KDA's recurrent state, MLA's growing K/V buffers)
22//! -- its real signature has no `pos` parameter at all. `KimiEngine`
23//! ignores the argument; this is a real architectural fact about the
24//! model, not an oversight in this trait's design.
25
26use crate::config::{KdaConfig, MlaConfig};
27use crate::decoder::Decoder;
28use crate::deepseek_v4_decoder::{
29    deepseek_v4_forward_token, DeepseekV4DecodeState, DeepseekV4DecoderConfig,
30    DeepseekV4DecoderWeights,
31};
32use crate::glm52_decoder::{
33    glm52_forward_token, Glm52DecodeState, Glm52DecoderConfig, Glm52DecoderWeights,
34};
35use crate::kimi_decoder::{
36    kimi_forward_token, KimiDecodeState, KimiDecoderConfig, KimiDecoderWeights,
37};
38use crate::kimi_tokenizer::KimiTokenizer;
39use ferrox_core::cache::KvCache;
40use ferrox_core::weight_matrix::WeightMatrix;
41
42/// A decoder that can run one incremental forward step given a token id
43/// and position, updating its own per-layer state in place.
44pub trait Engine {
45    type State;
46
47    /// Builds fresh (empty) per-layer state for a new request.
48    fn new_state(&self) -> Self::State;
49
50    fn vocab_size(&self) -> usize;
51
52    fn forward_token(&self, token_id: usize, pos: usize, state: &mut Self::State) -> Vec<f32>;
53}
54
55impl Engine for Decoder {
56    type State = Vec<KvCache>;
57
58    fn new_state(&self) -> Vec<KvCache> {
59        self.layers
60            .iter()
61            .map(|_| KvCache::new(self.config.n_kv_heads, self.config.head_dim))
62            .collect()
63    }
64
65    fn vocab_size(&self) -> usize {
66        self.config.vocab_size
67    }
68
69    fn forward_token(&self, token_id: usize, pos: usize, state: &mut Self::State) -> Vec<f32> {
70        Decoder::forward_token(self, token_id, pos, state)
71    }
72}
73
74/// Bundles Kimi K3's weights with the three real config structs
75/// `kimi_forward_token` needs, so `Engine::forward_token`'s three-
76/// argument shape (`token_id`, `pos`, `state`) can wrap Kimi's real
77/// four-config-argument function.
78pub struct KimiEngine {
79    pub weights: KimiDecoderWeights,
80    pub cfg: KimiDecoderConfig,
81    pub mla_cfg: MlaConfig,
82    pub kda_cfg: KdaConfig,
83}
84
85impl Engine for KimiEngine {
86    type State = KimiDecodeState;
87
88    fn new_state(&self) -> KimiDecodeState {
89        KimiDecodeState::new(&self.weights, &self.kda_cfg)
90    }
91
92    fn vocab_size(&self) -> usize {
93        self.weights.output_head.rows()
94    }
95
96    fn forward_token(&self, token_id: usize, _pos: usize, state: &mut Self::State) -> Vec<f32> {
97        kimi_forward_token(
98            &self.weights,
99            &self.cfg,
100            &self.mla_cfg,
101            &self.kda_cfg,
102            token_id,
103            state,
104        )
105    }
106}
107
108/// GLM-5.2 dedicated DSA stack behind the same [`Engine`] trait as Kimi.
109/// Synthetic / loader-backed weights only — no claim of a full real
110/// ~744B serve path. Lets `generate_engine` exercise GLM without
111/// forcing DSA into the GQA [`Decoder`].
112pub struct Glm52Engine {
113    pub weights: Glm52DecoderWeights,
114    pub cfg: Glm52DecoderConfig,
115}
116
117impl Engine for Glm52Engine {
118    type State = Glm52DecodeState;
119
120    fn new_state(&self) -> Glm52DecodeState {
121        Glm52DecodeState::new(&self.weights)
122    }
123
124    fn vocab_size(&self) -> usize {
125        self.weights.output_head.rows()
126    }
127
128    fn forward_token(&self, token_id: usize, _pos: usize, state: &mut Self::State) -> Vec<f32> {
129        glm52_forward_token(&self.weights, &self.cfg, token_id, state)
130    }
131}
132
133/// Multi-layer MLA stack for DeepSeek-2 / Mistral-4-style GGUF serve.
134///
135/// Uses [`crate::mla::mla_forward_token`] with asymmetric K/V caches
136/// (plain `Vec<f32>`, not [`KvCache`]). Layers
137/// `[0, leading_dense)` use dense SwiGLU; later layers use MoE
138/// (`ferrox_moe`) when the GGUF carries experts — fail-closed at load if
139/// expert tensors are missing.
140pub struct MlaEngine {
141    pub embedding: WeightMatrix,
142    pub layers: Vec<MlaLayerWeights>,
143    pub final_norm: Vec<f32>,
144    pub output_head: WeightMatrix,
145    pub mla_cfg: MlaConfig,
146    pub rms_norm_eps: f32,
147    pub hidden_dim: usize,
148    /// Present when any layer uses [`MlaLayerFfn::Moe`].
149    pub moe: Option<MlaMoeRuntime>,
150}
151
152/// MoE routing knobs shared by every MoE layer (DeepSeek-2 / Mistral-4).
153#[derive(Debug, Clone)]
154pub struct MlaMoeRuntime {
155    pub n_experts_active: usize,
156    pub gating: ferrox_moe::GatingFunction,
157    pub norm_topk_prob: bool,
158    pub expert_weights_scale: f32,
159}
160
161pub struct MlaDenseFfn {
162    pub gate: WeightMatrix,
163    pub up: WeightMatrix,
164    pub down: WeightMatrix,
165}
166
167pub struct MlaMoeFfn {
168    pub router: WeightMatrix,
169    pub experts: Vec<ferrox_moe::ExpertWeights>,
170    pub shared_expert: ferrox_moe::ExpertWeights,
171    /// Optional aux-loss-free bias (`ffn_exp_probs_b.bias`).
172    pub exp_probs_bias: Option<Vec<f32>>,
173}
174
175pub enum MlaLayerFfn {
176    Dense(MlaDenseFfn),
177    Moe(MlaMoeFfn),
178}
179
180pub struct MlaLayerWeights {
181    pub attn_norm: Vec<f32>,
182    pub attn: crate::mla::MlaAttnWeights,
183    pub ffn_norm: Vec<f32>,
184    pub ffn: MlaLayerFfn,
185}
186
187pub struct MlaDecodeState {
188    pub layers: Vec<(Vec<f32>, Vec<f32>)>,
189}
190
191impl MlaEngine {
192    pub fn new_state(&self) -> MlaDecodeState {
193        MlaDecodeState {
194            layers: (0..self.layers.len())
195                .map(|_| (Vec::new(), Vec::new()))
196                .collect(),
197        }
198    }
199
200    fn moe_ffn_forward(&self, ffn: &MlaMoeFfn, x: &[f32]) -> Vec<f32> {
201        use ferrox_moe::{
202            combine_expert_outputs, route_top_k, route_top_k_sigmoid_with_bias, run_expert,
203            GatingFunction,
204        };
205        let moe = self
206            .moe
207            .as_ref()
208            .expect("MlaLayerFfn::Moe requires MlaEngine.moe");
209        let router_logits = ffn.router.apply(x);
210        let decision = match (moe.gating, ffn.exp_probs_bias.as_deref()) {
211            (GatingFunction::Sigmoid, Some(bias)) => route_top_k_sigmoid_with_bias(
212                &router_logits,
213                bias,
214                moe.n_experts_active,
215                moe.norm_topk_prob,
216                moe.expert_weights_scale,
217            ),
218            _ => {
219                let mut d = route_top_k(
220                    &router_logits,
221                    moe.n_experts_active,
222                    moe.gating,
223                    moe.norm_topk_prob,
224                );
225                if (moe.expert_weights_scale - 1.0).abs() > f32::EPSILON {
226                    for w in d.weights.iter_mut() {
227                        *w *= moe.expert_weights_scale;
228                    }
229                }
230                d
231            }
232        };
233        let routed: Vec<(Vec<f32>, f32)> = decision
234            .expert_ids
235            .iter()
236            .zip(decision.weights.iter())
237            .map(|(&e, &w)| (run_expert(x, &ffn.experts[e]), w))
238            .collect();
239        let shared = run_expert(x, &ffn.shared_expert);
240        combine_expert_outputs(&routed, &[shared], x.len())
241    }
242}
243
244impl Engine for MlaEngine {
245    type State = MlaDecodeState;
246
247    fn new_state(&self) -> MlaDecodeState {
248        MlaEngine::new_state(self)
249    }
250
251    fn vocab_size(&self) -> usize {
252        self.output_head.rows()
253    }
254
255    fn forward_token(&self, token_id: usize, _pos: usize, state: &mut Self::State) -> Vec<f32> {
256        use ferrox_core::matmul::{rms_norm, swiglu};
257        let mut hidden = self.embedding.dequant_row(token_id);
258        for (layer, (k_cache, v_cache)) in self.layers.iter().zip(state.layers.iter_mut()) {
259            let normed = rms_norm(&hidden, &layer.attn_norm, self.rms_norm_eps);
260            let attn_out = crate::mla::mla_forward_token(
261                &layer.attn,
262                &self.mla_cfg,
263                &normed,
264                self.rms_norm_eps,
265                k_cache,
266                v_cache,
267            );
268            for (h, a) in hidden.iter_mut().zip(attn_out.iter()) {
269                *h += a;
270            }
271            let ffn_in = rms_norm(&hidden, &layer.ffn_norm, self.rms_norm_eps);
272            let down = match &layer.ffn {
273                MlaLayerFfn::Dense(d) => {
274                    let gate = d.gate.apply(&ffn_in);
275                    let up = d.up.apply(&ffn_in);
276                    d.down.apply(&swiglu(&gate, &up))
277                }
278                MlaLayerFfn::Moe(m) => self.moe_ffn_forward(m, &ffn_in),
279            };
280            for (h, d) in hidden.iter_mut().zip(down.iter()) {
281                *h += d;
282            }
283        }
284        let final_normed = rms_norm(&hidden, &self.final_norm, self.rms_norm_eps);
285        self.output_head.apply(&final_normed)
286    }
287}
288
289/// DeepSeek V4 synthetic stack behind [`Engine`]. Preset `deepseek_v4_pro`
290/// remains a sketch until a real GGUF loader + incremental DSV4 KV land.
291pub struct DeepseekV4Engine {
292    pub weights: DeepseekV4DecoderWeights,
293    pub cfg: DeepseekV4DecoderConfig,
294}
295
296impl Engine for DeepseekV4Engine {
297    type State = DeepseekV4DecodeState;
298
299    fn new_state(&self) -> DeepseekV4DecodeState {
300        DeepseekV4DecodeState::new(self.weights.embedding.shape[1])
301    }
302
303    fn vocab_size(&self) -> usize {
304        self.weights.output_head.rows()
305    }
306
307    fn forward_token(&self, token_id: usize, _pos: usize, state: &mut Self::State) -> Vec<f32> {
308        deepseek_v4_forward_token(&self.weights, &self.cfg, token_id, state)
309    }
310}
311
312/// A minimal text<->token-id interface shared by every real tokenizer
313/// this crate has, regardless of each one's native id width
314/// (`GgufBpeTokenizer`/`GgufSpmTokenizer`/`GgufUnigramTokenizer` use
315/// `u32`, `KimiTokenizer` also uses `u32`) -- lets a generic generation
316/// loop encode/decode without caring which concrete tokenizer it was
317/// given.
318pub trait TextTokenizer {
319    fn encode(&self, text: &str) -> Vec<usize>;
320    fn decode(&self, ids: &[usize]) -> String;
321}
322
323impl TextTokenizer for KimiTokenizer {
324    fn encode(&self, text: &str) -> Vec<usize> {
325        KimiTokenizer::encode(self, text)
326            .into_iter()
327            .map(|id| id as usize)
328            .collect()
329    }
330
331    fn decode(&self, ids: &[usize]) -> String {
332        let ids32: Vec<u32> = ids.iter().map(|&id| id as u32).collect();
333        KimiTokenizer::decode(self, &ids32)
334    }
335}
336
337#[cfg(test)]
338mod tests {
339    use super::*;
340    use crate::config::test_dense_fixture;
341
342    /// Locks in the refactor: calling `Decoder` through the generic
343    /// `Engine` trait must be bit-identical to calling its own
344    /// `forward_token`/`forward_batch` directly -- the whole point of
345    /// the trait is that `ferrox-server`'s generic generation loop can
346    /// use it as a drop-in replacement with zero numeric difference.
347    #[test]
348    fn decoder_via_engine_trait_matches_direct_forward_token_calls() {
349        let decoder = Decoder::new_random_small(test_dense_fixture(), 2, 64);
350        let tokens = [3usize, 7, 1, 9];
351
352        let mut direct_caches: Vec<KvCache> = decoder
353            .layers
354            .iter()
355            .map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
356            .collect();
357        let mut direct_logits = Vec::new();
358        for (pos, &tok) in tokens.iter().enumerate() {
359            direct_logits = decoder.forward_token(tok, pos, &mut direct_caches);
360        }
361
362        let mut engine_state = Engine::new_state(&decoder);
363        let mut engine_logits = Vec::new();
364        for (pos, &tok) in tokens.iter().enumerate() {
365            engine_logits = Engine::forward_token(&decoder, tok, pos, &mut engine_state);
366        }
367
368        assert_eq!(engine_logits, direct_logits);
369        assert_eq!(Engine::vocab_size(&decoder), decoder.config.vocab_size);
370    }
371
372    /// Same equivalence, but against `forward_batch`'s independent
373    /// computation (the ground truth every other test in this
374    /// workspace already uses) rather than a second manual loop.
375    #[test]
376    fn decoder_via_engine_trait_matches_forward_batch_ground_truth() {
377        let decoder = Decoder::new_random_small(test_dense_fixture(), 2, 64);
378        let tokens = vec![2usize, 5, 8];
379
380        let mut batch_caches: Vec<KvCache> = decoder
381            .layers
382            .iter()
383            .map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
384            .collect();
385        let batch_logits = decoder.forward_batch(&tokens, 0, &mut batch_caches);
386        let ground_truth = batch_logits.last().unwrap().clone();
387
388        let mut engine_state = Engine::new_state(&decoder);
389        let mut engine_logits = Vec::new();
390        for (pos, &tok) in tokens.iter().enumerate() {
391            engine_logits = Engine::forward_token(&decoder, tok, pos, &mut engine_state);
392        }
393
394        // Tight tolerance, not bit equality: batched prefill runs the
395        // blocked three-pass softmax while per-token decode keeps the
396        // online accumulator, and the two round differently in the last
397        // ulp (they were bit-identical only while both were online).
398        assert_eq!(engine_logits.len(), ground_truth.len());
399        for (i, (e, g)) in engine_logits.iter().zip(ground_truth.iter()).enumerate() {
400            assert!(
401                (e - g).abs() < 1e-5,
402                "logit {i}: engine {e} vs forward_batch {g}"
403            );
404        }
405    }
406}