1use 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#[derive(Debug, Clone)]
20pub struct Ds4FlashConfig {
21 pub vocab_size: u32,
23 pub hidden_dim: u32,
25 pub num_layers: u32,
27 pub num_heads: u32,
29 pub head_dim: u32,
31 pub kv_lora_rank: u32,
33 pub qk_rope_head_dim: u32,
35 pub num_experts: u32,
37 pub moe_top_k: u32,
39 pub shared_expert_hidden_dim: u32,
41 pub rms_norm_eps: f32,
43 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
66pub 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 let embed_program = embedding(
99 "embed_table",
100 "tokens",
101 "embed_out",
102 max_seq_len,
103 hidden_dim,
104 );
105
106 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 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 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 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 let rms_norm_program = rms_norm("rms_in", "rms_out", hidden_dim, rms_norm_eps);
195
196 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 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 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 assert!(lm_head.buffers().len() >= 4);
326 }
327}