1use 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
42pub trait Engine {
45 type State;
46
47 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
74pub 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
108pub 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
133pub 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 pub moe: Option<MlaMoeRuntime>,
150}
151
152#[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 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
289pub 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
312pub 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 #[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 #[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 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}