Skip to main content

vyre_libs/nn/
inference_graph.rs

1//! DeepSeek V4 Flash inference graph construction.
2//!
3//! Builds the forward pass as a sequence of [`Program`]s, one per layer type.
4//! Each Program is a Category-A composition over existing `vyre-libs` primitives.
5
6use vyre_foundation::ir::Program;
7
8use super::{
9    activation::embedding,
10    attention::mla::mla_decode,
11    linear::linear,
12    moe::{expert_mlp, moe_layer::moe_layer_route_and_accumulate},
13    norm::rms_norm,
14};
15
16/// Hyperparameters for DeepSeek V4 Flash.
17///
18/// Defaults match the VyreOffload reference configuration.
19#[derive(Debug, Clone)]
20pub struct Ds4FlashConfig {
21    /// Vocabulary size (including special tokens).
22    pub vocab_size: u32,
23    /// Hidden dimension / model dimension (`d_model`).
24    pub hidden_dim: u32,
25    /// Number of transformer layers.
26    pub num_layers: u32,
27    /// Number of attention heads.
28    pub num_heads: u32,
29    /// Dimension per attention head.
30    pub head_dim: u32,
31    /// Compressed KV latent rank (MLA).
32    pub kv_lora_rank: u32,
33    /// RoPE dimension for decoupled Q/K.
34    pub qk_rope_head_dim: u32,
35    /// Number of routed experts in the MoE layer.
36    pub num_experts: u32,
37    /// Top-k experts selected by the router.
38    pub moe_top_k: u32,
39    /// Hidden dimension of the shared expert MLP.
40    pub shared_expert_hidden_dim: u32,
41    /// Epsilon for RMSNorm.
42    pub rms_norm_eps: f32,
43    /// Maximum sequence length for prefill.
44    pub max_seq_len: u32,
45}
46
47impl Default for Ds4FlashConfig {
48    fn default() -> Self {
49        Self {
50            vocab_size: 129_280,
51            hidden_dim: 7_168,
52            num_layers: 61,
53            num_heads: 128,
54            head_dim: 128,
55            kv_lora_rank: 512,
56            qk_rope_head_dim: 64,
57            num_experts: 256,
58            moe_top_k: 8,
59            shared_expert_hidden_dim: 18_432,
60            rms_norm_eps: 1e-6,
61            max_seq_len: 4_096,
62        }
63    }
64}
65
66/// Build the full forward-pass graph for DeepSeek V4 Flash.
67///
68/// Returns one [`Program`] per layer type, in conceptual execution order:
69///
70/// 1. Token embedding lookup (`embed_program`)
71/// 2. MLA prefill attention (`mla_prefill_program`)
72/// 3. MLA single-token decode attention (`mla_decode_program`)
73/// 4. MoE layer dispatch (`moe_layer_program`)
74/// 5. Shared dense expert MLP (`shared_expert_program`)
75/// 6. RMSNorm pre/post normalization (`rms_norm_program`)
76/// 7. LM head logits projection (`lm_head_program`)
77///
78/// Each Program references canonical buffer names (e.g. `"tokens"`, `"q"`,
79/// `"moe_x"`) so the runtime can wire them together in a sequential
80/// dispatch table.
81pub fn build_forward_graph(config: &Ds4FlashConfig) -> Vec<Program> {
82    let Ds4FlashConfig {
83        vocab_size,
84        hidden_dim,
85        num_layers: _,
86        num_heads,
87        head_dim,
88        kv_lora_rank,
89        qk_rope_head_dim,
90        num_experts,
91        moe_top_k,
92        shared_expert_hidden_dim,
93        rms_norm_eps,
94        max_seq_len,
95    } = *config;
96
97    // 1. Embedding: lookup tokens -> hidden_dim vectors.
98    let embed_program = embedding(
99        "embed_table",
100        "tokens",
101        "embed_out",
102        max_seq_len,
103        hidden_dim,
104    );
105
106    // 2. MLA prefill: full-context attention during prompt ingestion.
107    let mla_prefill_program = mla_decode(
108        "q",
109        "kv_cache",
110        "kr_cache",
111        "w_uk",
112        "w_uv",
113        "mla_prefill_out",
114        max_seq_len,
115        num_heads,
116        head_dim,
117        kv_lora_rank,
118        qk_rope_head_dim,
119    )
120    .unwrap_or_else(|e| {
121        crate::invalid_program(
122            "vyre-libs::nn::mla_prefill",
123            format!("Fix: mla_prefill build failed: {e}"),
124        )
125    });
126
127    // 3. MLA decode: single-token autoregressive step.
128    let mla_decode_program = mla_decode(
129        "q",
130        "kv_cache",
131        "kr_cache",
132        "w_uk",
133        "w_uv",
134        "mla_decode_out",
135        1,
136        num_heads,
137        head_dim,
138        kv_lora_rank,
139        qk_rope_head_dim,
140    )
141    .unwrap_or_else(|e| {
142        crate::invalid_program(
143            "vyre-libs::nn::mla_decode",
144            format!("Fix: mla_decode build failed: {e}"),
145        )
146    });
147
148    // 4. MoE layer: weighted accumulation over top-k expert outputs.
149    // The router softmax + top-k selection (via `softmax_top_k`) is
150    // expected to run immediately before this kernel to populate
151    // `expert_indices` and `expert_weights`.
152    let moe_layer_program = moe_layer_route_and_accumulate(
153        "moe_x",
154        "w_router",
155        "b_router",
156        "expert_indices",
157        "expert_weights",
158        "expert_outputs",
159        "moe_out",
160        hidden_dim,
161        num_experts,
162        hidden_dim,
163        moe_top_k,
164    )
165    .unwrap_or_else(|e| {
166        crate::invalid_program(
167            "vyre-libs::nn::moe_layer",
168            format!("Fix: moe_layer build failed: {e}"),
169        )
170    });
171
172    // 5. Shared expert: dense SwiGLU MLP used alongside the routed experts.
173    let shared_expert_program = expert_mlp(
174        "shared_x",
175        "shared_w_gate",
176        "shared_b_gate",
177        "shared_w_up",
178        "shared_b_up",
179        "shared_w_down",
180        "shared_b_down",
181        "shared_out",
182        hidden_dim,
183        shared_expert_hidden_dim,
184        hidden_dim,
185    )
186    .unwrap_or_else(|e| {
187        crate::invalid_program(
188            "vyre-libs::nn::shared_expert",
189            format!("Fix: shared_expert build failed: {e}"),
190        )
191    });
192
193    // 6. RMSNorm: pre/post layer normalization.
194    let rms_norm_program = rms_norm("rms_in", "rms_out", hidden_dim, rms_norm_eps);
195
196    // 7. LM head: project final hidden state to vocabulary logits.
197    let lm_head_program = linear(
198        "lm_head_x",
199        "lm_head_w",
200        "lm_head_b",
201        "lm_head_out",
202        hidden_dim,
203        vocab_size,
204    )
205    .unwrap_or_else(|e| {
206        crate::invalid_program(
207            "vyre-libs::nn::lm_head",
208            format!("Fix: lm_head build failed: {e}"),
209        )
210    });
211
212    vec![
213        embed_program,
214        mla_prefill_program,
215        mla_decode_program,
216        moe_layer_program,
217        shared_expert_program,
218        rms_norm_program,
219        lm_head_program,
220    ]
221}
222
223#[cfg(test)]
224mod tests {
225    use super::*;
226
227    #[test]
228    fn forward_graph_default_config_builds() {
229        let config = Ds4FlashConfig::default();
230        let programs = build_forward_graph(&config);
231        assert_eq!(
232            programs.len(),
233            7,
234            "expected 7 programs (one per layer type)"
235        );
236
237        let expected_names = [
238            "embed",
239            "mla_prefill",
240            "mla_decode",
241            "moe_layer",
242            "shared_expert",
243            "rms_norm",
244            "lm_head",
245        ];
246        for (i, program) in programs.iter().enumerate() {
247            assert!(
248                !program.buffers().is_empty(),
249                "{} program should declare at least one buffer",
250                expected_names[i]
251            );
252        }
253    }
254
255    #[test]
256    fn forward_graph_small_config_builds() {
257        let config = Ds4FlashConfig {
258            vocab_size: 1_024,
259            hidden_dim: 256,
260            num_layers: 2,
261            num_heads: 4,
262            head_dim: 64,
263            kv_lora_rank: 32,
264            qk_rope_head_dim: 16,
265            num_experts: 8,
266            moe_top_k: 2,
267            shared_expert_hidden_dim: 512,
268            rms_norm_eps: 1e-5,
269            max_seq_len: 128,
270        };
271        let programs = build_forward_graph(&config);
272        assert_eq!(programs.len(), 7);
273
274        // Verify each program has a non-empty buffer table.
275        for program in &programs {
276            assert!(!program.buffers().is_empty());
277        }
278    }
279
280    #[test]
281    fn embed_program_has_correct_buffer_count() {
282        let config = Ds4FlashConfig::default();
283        let programs = build_forward_graph(&config);
284        let embed = &programs[0];
285        assert_eq!(embed.buffers().len(), 3);
286    }
287
288    #[test]
289    fn mla_prefill_and_decode_are_distinct() {
290        let config = Ds4FlashConfig::default();
291        let programs = build_forward_graph(&config);
292        let prefill = &programs[1];
293        let decode = &programs[2];
294        // Prefill operates over max_seq_len; decode over seq_len=1.
295        // The buffer counts should be identical (same inputs/outputs)
296        // but the workgroup logic differs internally.
297        assert_eq!(prefill.buffers().len(), decode.buffers().len());
298        assert!(
299            prefill.workgroup_size() == decode.workgroup_size(),
300            "prefill and decode use the same workgroup dispatch shape"
301        );
302    }
303
304    #[test]
305    fn rms_norm_program_is_f32() {
306        let config = Ds4FlashConfig::default();
307        let programs = build_forward_graph(&config);
308        let rms = &programs[5];
309        for buf in rms.buffers() {
310            assert_eq!(
311                buf.element,
312                vyre_foundation::ir::DataType::F32,
313                "rms_norm uses F32 buffers"
314            );
315        }
316    }
317
318    #[test]
319    fn lm_head_program_has_expected_buffers() {
320        let config = Ds4FlashConfig::default();
321        let programs = build_forward_graph(&config);
322        let lm_head = &programs[6];
323        // The tiled linear path adds workgroup scratch buffers, so we
324        // expect at least the 4 core buffers (x, w, b, out).
325        assert!(lm_head.buffers().len() >= 4);
326    }
327}