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
42mod entry;
43
44pub use entry::Engine;
45
46#[cfg(test)]
50pub(crate) use entry::{assert_one_pool_entry_per_step, on_workers_promotes_here};
51
52impl Engine for Decoder {
53 type State = Vec<KvCache>;
54
55 fn new_state(&self) -> Vec<KvCache> {
56 self.layers
57 .iter()
58 .map(|_| KvCache::new(self.config.n_kv_heads, self.config.head_dim))
59 .collect()
60 }
61
62 fn vocab_size(&self) -> usize {
63 self.config.vocab_size
64 }
65
66 fn forward_token_on_worker(
72 &self,
73 token_id: usize,
74 pos: usize,
75 state: &mut Self::State,
76 ) -> Vec<f32> {
77 Decoder::forward_token(self, token_id, pos, state)
78 }
79}
80
81pub struct KimiEngine {
86 pub weights: KimiDecoderWeights,
87 pub cfg: KimiDecoderConfig,
88 pub mla_cfg: MlaConfig,
89 pub kda_cfg: KdaConfig,
90}
91
92impl Engine for KimiEngine {
93 type State = KimiDecodeState;
94
95 fn new_state(&self) -> KimiDecodeState {
96 KimiDecodeState::new(&self.weights, &self.kda_cfg)
97 }
98
99 fn vocab_size(&self) -> usize {
100 self.weights.output_head.rows()
101 }
102
103 fn forward_token_on_worker(
104 &self,
105 token_id: usize,
106 _pos: usize,
107 state: &mut Self::State,
108 ) -> Vec<f32> {
109 kimi_forward_token(
110 &self.weights,
111 &self.cfg,
112 &self.mla_cfg,
113 &self.kda_cfg,
114 token_id,
115 state,
116 )
117 }
118}
119
120pub struct Glm52Engine {
125 pub weights: Glm52DecoderWeights,
126 pub cfg: Glm52DecoderConfig,
127}
128
129impl Engine for Glm52Engine {
130 type State = Glm52DecodeState;
131
132 fn new_state(&self) -> Glm52DecodeState {
133 Glm52DecodeState::new(&self.weights)
134 }
135
136 fn vocab_size(&self) -> usize {
137 self.weights.output_head.rows()
138 }
139
140 fn forward_token_on_worker(
141 &self,
142 token_id: usize,
143 _pos: usize,
144 state: &mut Self::State,
145 ) -> Vec<f32> {
146 glm52_forward_token(&self.weights, &self.cfg, token_id, state)
147 }
148}
149
150pub struct MlaEngine {
158 pub embedding: WeightMatrix,
159 pub layers: Vec<MlaLayerWeights>,
160 pub final_norm: Vec<f32>,
161 pub output_head: WeightMatrix,
162 pub mla_cfg: MlaConfig,
163 pub rms_norm_eps: f32,
164 pub hidden_dim: usize,
165 pub moe: Option<MlaMoeRuntime>,
167}
168
169#[derive(Debug, Clone)]
171pub struct MlaMoeRuntime {
172 pub n_experts_active: usize,
173 pub gating: ferrox_moe::GatingFunction,
174 pub norm_topk_prob: bool,
175 pub expert_weights_scale: f32,
176}
177
178pub struct MlaDenseFfn {
179 pub gate: WeightMatrix,
180 pub up: WeightMatrix,
181 pub down: WeightMatrix,
182}
183
184pub struct MlaMoeFfn {
185 pub router: WeightMatrix,
186 pub experts: Vec<ferrox_moe::ExpertWeights>,
187 pub shared_expert: ferrox_moe::ExpertWeights,
188 pub exp_probs_bias: Option<Vec<f32>>,
190}
191
192pub enum MlaLayerFfn {
193 Dense(MlaDenseFfn),
194 Moe(MlaMoeFfn),
195}
196
197pub struct MlaLayerWeights {
198 pub attn_norm: Vec<f32>,
199 pub attn: crate::mla::MlaAttnWeights,
200 pub ffn_norm: Vec<f32>,
201 pub ffn: MlaLayerFfn,
202}
203
204pub struct MlaDecodeState {
205 pub layers: Vec<(Vec<f32>, Vec<f32>)>,
206}
207
208impl MlaEngine {
209 pub fn new_state(&self) -> MlaDecodeState {
210 MlaDecodeState {
211 layers: (0..self.layers.len())
212 .map(|_| (Vec::new(), Vec::new()))
213 .collect(),
214 }
215 }
216
217 fn moe_ffn_forward(&self, ffn: &MlaMoeFfn, x: &[f32]) -> Vec<f32> {
218 use ferrox_moe::{
219 combine_expert_outputs, route_top_k, route_top_k_sigmoid_with_bias, run_expert,
220 GatingFunction, GluAct,
221 };
222 let moe = self
223 .moe
224 .as_ref()
225 .expect("MlaLayerFfn::Moe requires MlaEngine.moe");
226 let router_logits = ffn.router.apply(x);
227 let decision = match (moe.gating, ffn.exp_probs_bias.as_deref()) {
228 (GatingFunction::Sigmoid, Some(bias)) => route_top_k_sigmoid_with_bias(
229 &router_logits,
230 bias,
231 moe.n_experts_active,
232 moe.norm_topk_prob,
233 moe.expert_weights_scale,
234 ),
235 _ => {
236 let mut d = route_top_k(
237 &router_logits,
238 moe.n_experts_active,
239 moe.gating,
240 moe.norm_topk_prob,
241 );
242 if (moe.expert_weights_scale - 1.0).abs() > f32::EPSILON {
243 for w in d.weights.iter_mut() {
244 *w *= moe.expert_weights_scale;
245 }
246 }
247 d
248 }
249 };
250 let routed: Vec<(Vec<f32>, f32)> = decision
251 .expert_ids
252 .iter()
253 .zip(decision.weights.iter())
254 .map(|(&e, &w)| (run_expert(x, &ffn.experts[e], GluAct::Swiglu), w))
255 .collect();
256 let shared = run_expert(x, &ffn.shared_expert, GluAct::Swiglu);
257 combine_expert_outputs(&routed, &[shared], x.len())
258 }
259}
260
261impl Engine for MlaEngine {
262 type State = MlaDecodeState;
263
264 fn new_state(&self) -> MlaDecodeState {
265 MlaEngine::new_state(self)
266 }
267
268 fn vocab_size(&self) -> usize {
269 self.output_head.rows()
270 }
271
272 fn forward_token_on_worker(
273 &self,
274 token_id: usize,
275 _pos: usize,
276 state: &mut Self::State,
277 ) -> Vec<f32> {
278 use ferrox_core::matmul::{rms_norm, swiglu};
279 let mut hidden = self.embedding.dequant_row(token_id);
280 for (layer, (k_cache, v_cache)) in self.layers.iter().zip(state.layers.iter_mut()) {
281 let normed = rms_norm(&hidden, &layer.attn_norm, self.rms_norm_eps);
282 let attn_out = crate::mla::mla_forward_token(
283 &layer.attn,
284 &self.mla_cfg,
285 &normed,
286 self.rms_norm_eps,
287 k_cache,
288 v_cache,
289 );
290 for (h, a) in hidden.iter_mut().zip(attn_out.iter()) {
291 *h += a;
292 }
293 let ffn_in = rms_norm(&hidden, &layer.ffn_norm, self.rms_norm_eps);
294 let down = match &layer.ffn {
295 MlaLayerFfn::Dense(d) => {
296 let gate = d.gate.apply(&ffn_in);
297 let up = d.up.apply(&ffn_in);
298 d.down.apply(&swiglu(&gate, &up))
299 }
300 MlaLayerFfn::Moe(m) => self.moe_ffn_forward(m, &ffn_in),
301 };
302 for (h, d) in hidden.iter_mut().zip(down.iter()) {
303 *h += d;
304 }
305 }
306 let final_normed = rms_norm(&hidden, &self.final_norm, self.rms_norm_eps);
307 self.output_head.apply(&final_normed)
308 }
309}
310
311pub struct DeepseekV4Engine {
314 pub weights: DeepseekV4DecoderWeights,
315 pub cfg: DeepseekV4DecoderConfig,
316}
317
318impl Engine for DeepseekV4Engine {
319 type State = DeepseekV4DecodeState;
320
321 fn new_state(&self) -> DeepseekV4DecodeState {
322 DeepseekV4DecodeState::new(self.weights.embedding.shape[1])
323 }
324
325 fn vocab_size(&self) -> usize {
326 self.weights.output_head.rows()
327 }
328
329 fn forward_token_on_worker(
330 &self,
331 token_id: usize,
332 _pos: usize,
333 state: &mut Self::State,
334 ) -> Vec<f32> {
335 deepseek_v4_forward_token(&self.weights, &self.cfg, token_id, state)
336 }
337}
338
339pub trait TextTokenizer {
346 fn encode(&self, text: &str) -> Vec<usize>;
347 fn decode(&self, ids: &[usize]) -> String;
348
349 fn decode_bytes(&self, ids: &[usize]) -> Vec<u8> {
359 self.decode(ids).into_bytes()
360 }
361}
362
363impl TextTokenizer for KimiTokenizer {
364 fn encode(&self, text: &str) -> Vec<usize> {
365 KimiTokenizer::encode(self, text)
366 .into_iter()
367 .map(|id| id as usize)
368 .collect()
369 }
370
371 fn decode(&self, ids: &[usize]) -> String {
372 let ids32: Vec<u32> = ids.iter().map(|&id| id as u32).collect();
373 KimiTokenizer::decode(self, &ids32)
374 }
375}
376
377#[cfg(test)]
378mod tests {
379 use super::*;
380 use crate::config::test_dense_fixture;
381
382 #[test]
388 fn decoder_via_engine_trait_matches_direct_forward_token_calls() {
389 let decoder = Decoder::new_random_small(test_dense_fixture(), 2, 64);
390 let tokens = [3usize, 7, 1, 9];
391
392 let mut direct_caches: Vec<KvCache> = decoder
393 .layers
394 .iter()
395 .map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
396 .collect();
397 let mut direct_logits = Vec::new();
398 for (pos, &tok) in tokens.iter().enumerate() {
399 direct_logits = decoder.forward_token(tok, pos, &mut direct_caches);
400 }
401
402 let mut engine_state = Engine::new_state(&decoder);
403 let mut engine_logits = Vec::new();
404 for (pos, &tok) in tokens.iter().enumerate() {
405 engine_logits = Engine::forward_token(&decoder, tok, pos, &mut engine_state);
406 }
407
408 assert_eq!(engine_logits, direct_logits);
409 assert_eq!(Engine::vocab_size(&decoder), decoder.config.vocab_size);
410 }
411
412 #[test]
416 fn decoder_via_engine_trait_matches_forward_batch_ground_truth() {
417 let decoder = Decoder::new_random_small(test_dense_fixture(), 2, 64);
418 let tokens = vec![2usize, 5, 8];
419
420 let mut batch_caches: Vec<KvCache> = decoder
421 .layers
422 .iter()
423 .map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
424 .collect();
425 let batch_logits = decoder.forward_batch(&tokens, 0, &mut batch_caches);
426 let ground_truth = batch_logits.last().unwrap().clone();
427
428 let mut engine_state = Engine::new_state(&decoder);
429 let mut engine_logits = Vec::new();
430 for (pos, &tok) in tokens.iter().enumerate() {
431 engine_logits = Engine::forward_token(&decoder, tok, pos, &mut engine_state);
432 }
433
434 assert_eq!(engine_logits.len(), ground_truth.len());
439 for (i, (e, g)) in engine_logits.iter().zip(ground_truth.iter()).enumerate() {
440 assert!(
441 (e - g).abs() < 1e-5,
442 "logit {i}: engine {e} vs forward_batch {g}"
443 );
444 }
445 }
446}