1use super::*;
6
7impl MoeLayer {
8 pub fn set_pre_expert_norm(&mut self, norm: crate::weight_map::DenseWeight) {
18 self.pre_expert_norm = Some(norm);
19 }
20
21 pub fn set_gelu_activation(&mut self, gpu: &dyn GpuBackend) -> Result<()> {
25 self.moe_act_mul = gpu.kernel("gelu", "gelu_mul")?;
26 self.gelu_activation = true;
27 Ok(())
28 }
29
30 pub fn transpose_for_prefill(
31 &mut self,
32 gpu: &dyn GpuBackend,
33 config: &atlas_core::config::ModelConfig,
34 ) -> Result<()> {
35 self.transpose_for_prefill_impl(gpu, config, true)
36 }
37
38 pub fn transpose_gate_up_for_prefill(
46 &mut self,
47 gpu: &dyn GpuBackend,
48 config: &atlas_core::config::ModelConfig,
49 ) -> Result<()> {
50 self.transpose_for_prefill_impl(gpu, config, false)
51 }
52
53 pub(super) fn transpose_for_prefill_impl(
54 &mut self,
55 gpu: &dyn GpuBackend,
56 config: &atlas_core::config::ModelConfig,
57 include_down: bool,
58 ) -> Result<()> {
59 let h = config.hidden_size;
60 let inter = config.moe_intermediate_size;
61 let shared_inter = config.shared_expert_intermediate_size;
62
63 let num_experts = self.weights.experts.len();
65 let mut gate_t = Vec::with_capacity(num_experts);
66 let mut up_t = Vec::with_capacity(num_experts);
67 let mut down_t = Vec::with_capacity(num_experts);
68
69 let routed_gs =
73 if self.experts_scale_kind == crate::weight_map::WeightQuantFormat::Mxfp4E8m0 {
74 32
75 } else {
76 16
77 };
78 for expert in &self.weights.experts {
79 if expert.gate_proj.is_null() {
80 gate_t.push(QuantizedWeight::null());
81 up_t.push(QuantizedWeight::null());
82 if include_down {
83 down_t.push(QuantizedWeight::null());
84 }
85 } else {
86 gate_t.push(
87 expert
88 .gate_proj
89 .transpose_for_gemm_gs(gpu, inter, h, routed_gs)?,
90 );
91 up_t.push(
92 expert
93 .up_proj
94 .transpose_for_gemm_gs(gpu, inter, h, routed_gs)?,
95 );
96 if include_down {
97 down_t.push(
98 expert
99 .down_proj
100 .transpose_for_gemm_gs(gpu, h, inter, routed_gs)?,
101 );
102 }
103 }
104 }
105
106 self.gate_ptrs_t = Some(build_ptr_table_from_qw(&gate_t, gpu)?);
107 self.up_ptrs_t = Some(build_ptr_table_from_qw(&up_t, gpu)?);
108 if include_down {
109 self.down_ptrs_t = Some(build_ptr_table_from_qw(&down_t, gpu)?);
110 }
111
112 if !self.weights.shared_expert.gate_proj.is_null() && shared_inter > 0 {
114 self.shared_gate_t = Some(self.weights.shared_expert.gate_proj.transpose_for_gemm(
115 gpu,
116 shared_inter,
117 h,
118 )?);
119 self.shared_up_t = Some(self.weights.shared_expert.up_proj.transpose_for_gemm(
120 gpu,
121 shared_inter,
122 h,
123 )?);
124 if include_down {
125 self.shared_down_t =
126 Some(self.weights.shared_expert.down_proj.transpose_for_gemm(
127 gpu,
128 h,
129 shared_inter,
130 )?);
131 }
132 }
133
134 Ok(())
135 }
136
137 pub fn transpose_for_prefill_unified(
158 &mut self,
159 gpu: &dyn GpuBackend,
160 config: &atlas_core::config::ModelConfig,
161 ) -> Result<()> {
162 self.transpose_for_prefill_unified_inner(gpu, config, false)
163 }
164
165 pub fn transpose_for_prefill_hybrid(
173 &mut self,
174 gpu: &dyn GpuBackend,
175 config: &atlas_core::config::ModelConfig,
176 ) -> Result<()> {
177 self.transpose_for_prefill_unified_inner(gpu, config, true)
178 }
179
180 pub(super) fn transpose_for_prefill_unified_inner(
186 &mut self,
187 gpu: &dyn GpuBackend,
188 config: &atlas_core::config::ModelConfig,
189 keep_originals: bool,
190 ) -> Result<()> {
191 let h = config.hidden_size;
192 let inter = config.moe_intermediate_size;
193 let shared_inter = config.shared_expert_intermediate_size;
194 let _num_experts = self.weights.experts.len();
195
196 let routed_gs =
199 if self.experts_scale_kind == crate::weight_map::WeightQuantFormat::Mxfp4E8m0 {
200 32
201 } else {
202 16
203 };
204 let gate_src: Vec<QuantizedWeight> = self
205 .weights
206 .experts
207 .iter()
208 .map(|e| {
209 if e.gate_proj.is_null() {
210 QuantizedWeight::null()
211 } else {
212 e.gate_proj
213 }
214 })
215 .collect();
216 let up_src: Vec<QuantizedWeight> = self
217 .weights
218 .experts
219 .iter()
220 .map(|e| {
221 if e.gate_proj.is_null() {
222 QuantizedWeight::null()
223 } else {
224 e.up_proj
225 }
226 })
227 .collect();
228 let gate_t = self.transpose_experts_gpu(gpu, &gate_src, inter, h, routed_gs)?;
229 let up_t = self.transpose_experts_gpu(gpu, &up_src, inter, h, routed_gs)?;
230 self.gate_ptrs_t = Some(build_ptr_table_from_qw(&gate_t, gpu)?);
231 self.up_ptrs_t = Some(build_ptr_table_from_qw(&up_t, gpu)?);
232 if !self.weights.shared_expert.gate_proj.is_null() && shared_inter > 0 {
234 self.shared_gate_t = Some(self.weights.shared_expert.gate_proj.transpose_for_gemm(
235 gpu,
236 shared_inter,
237 h,
238 )?);
239 self.shared_up_t = Some(self.weights.shared_expert.up_proj.transpose_for_gemm(
240 gpu,
241 shared_inter,
242 h,
243 )?);
244 }
245
246 if !keep_originals {
247 for expert in &mut self.weights.experts {
252 if !expert.gate_proj.weight.is_null() {
253 gpu.free(expert.gate_proj.weight)?;
254 gpu.free(expert.gate_proj.weight_scale)?;
255 expert.gate_proj.weight = DevicePtr::NULL;
256 expert.gate_proj.weight_scale = DevicePtr::NULL;
257 }
258 if !expert.up_proj.weight.is_null() {
259 gpu.free(expert.up_proj.weight)?;
260 gpu.free(expert.up_proj.weight_scale)?;
261 expert.up_proj.weight = DevicePtr::NULL;
262 expert.up_proj.weight_scale = DevicePtr::NULL;
263 }
264 }
265 if !self.weights.shared_expert.gate_proj.weight.is_null() && shared_inter > 0 {
266 gpu.free(self.weights.shared_expert.gate_proj.weight)?;
267 gpu.free(self.weights.shared_expert.gate_proj.weight_scale)?;
268 self.weights.shared_expert.gate_proj.weight = DevicePtr::NULL;
269 self.weights.shared_expert.gate_proj.weight_scale = DevicePtr::NULL;
270 gpu.free(self.weights.shared_expert.up_proj.weight)?;
271 gpu.free(self.weights.shared_expert.up_proj.weight_scale)?;
272 self.weights.shared_expert.up_proj.weight = DevicePtr::NULL;
273 self.weights.shared_expert.up_proj.weight_scale = DevicePtr::NULL;
274 }
275 }
276
277 let down_src: Vec<QuantizedWeight> = self
279 .weights
280 .experts
281 .iter()
282 .map(|e| {
283 if e.down_proj.is_null() {
284 QuantizedWeight::null()
285 } else {
286 e.down_proj
287 }
288 })
289 .collect();
290 let down_t = self.transpose_experts_gpu(gpu, &down_src, h, inter, routed_gs)?;
291 self.down_ptrs_t = Some(build_ptr_table_from_qw(&down_t, gpu)?);
292 if !self.weights.shared_expert.down_proj.is_null() && shared_inter > 0 {
293 self.shared_down_t = Some(self.weights.shared_expert.down_proj.transpose_for_gemm(
294 gpu,
295 h,
296 shared_inter,
297 )?);
298 }
299
300 if !keep_originals {
301 for expert in &mut self.weights.experts {
303 if !expert.down_proj.weight.is_null() {
304 gpu.free(expert.down_proj.weight)?;
305 gpu.free(expert.down_proj.weight_scale)?;
306 expert.down_proj.weight = DevicePtr::NULL;
307 expert.down_proj.weight_scale = DevicePtr::NULL;
308 }
309 }
310 if !self.weights.shared_expert.down_proj.weight.is_null() && shared_inter > 0 {
311 gpu.free(self.weights.shared_expert.down_proj.weight)?;
312 gpu.free(self.weights.shared_expert.down_proj.weight_scale)?;
313 self.weights.shared_expert.down_proj.weight = DevicePtr::NULL;
314 self.weights.shared_expert.down_proj.weight_scale = DevicePtr::NULL;
315 }
316 }
317
318 Ok(())
319 }
320
321 #[allow(clippy::too_many_arguments)]
335 fn transpose_experts_gpu(
336 &self,
337 gpu: &dyn GpuBackend,
338 src: &[QuantizedWeight],
339 n: usize,
340 k: usize,
341 group_size: usize,
342 ) -> Result<Vec<QuantizedWeight>> {
343 let num_experts = src.len();
344 let packed_each = n * (k / 2);
345 let scale_each = n * (k / group_size);
346 anyhow::ensure!(
347 packed_each > 0 && scale_each > 0,
348 "transpose_experts_gpu: zero-sized projection (n={n} k={k} gs={group_size})"
349 );
350
351 let packed_slab = gpu.alloc(num_experts * packed_each)?;
353 let scale_slab = gpu.alloc(num_experts * scale_each)?;
354
355 let mut out = Vec::with_capacity(num_experts);
358 for (e, w) in src.iter().enumerate() {
359 if w.is_null() {
360 out.push(QuantizedWeight::null());
361 } else {
362 out.push(QuantizedWeight {
363 weight: packed_slab.offset(e * packed_each),
364 weight_scale: scale_slab.offset(e * scale_each),
365 weight_scale_2: w.weight_scale_2,
366 input_scale: w.input_scale,
367 weight_scale_2_vec: w.weight_scale_2_vec,
368 });
369 }
370 }
371
372 let src_tbl = build_ptr_table_from_qw(src, gpu)?;
373 let dst_tbl = build_ptr_table_from_qw(&out, gpu)?;
374 let stream = gpu.default_stream();
375 crate::layers::ops::moe_transpose_u8_batched(
377 gpu,
378 self.moe_transpose_u8_batched_k,
379 src_tbl.packed_ptrs,
380 dst_tbl.packed_ptrs,
381 n as u32,
382 (k / 2) as u32,
383 num_experts as u32,
384 stream,
385 )?;
386 crate::layers::ops::moe_transpose_u8_batched(
388 gpu,
389 self.moe_transpose_u8_batched_k,
390 src_tbl.scale_ptrs,
391 dst_tbl.scale_ptrs,
392 n as u32,
393 (k / group_size) as u32,
394 num_experts as u32,
395 stream,
396 )?;
397 gpu.synchronize(stream)?;
398 gpu.free(src_tbl.packed_ptrs)?;
400 gpu.free(src_tbl.scale_ptrs)?;
401 gpu.free(src_tbl.scale2_vals)?;
402 gpu.free(dst_tbl.packed_ptrs)?;
403 gpu.free(dst_tbl.scale_ptrs)?;
404 gpu.free(dst_tbl.scale2_vals)?;
405 Ok(out)
406 }
407
408 pub fn build_cutlass_grouped_sfb(
415 &mut self,
416 gpu: &dyn GpuBackend,
417 config: &atlas_core::config::ModelConfig,
418 stream: u64,
419 ) -> Result<()> {
420 let h = config.hidden_size;
421 let inter = config.moe_intermediate_size;
422 let num = self.weights.experts.len();
423 let sfb_len = |n: usize, k: usize| n.div_ceil(128) * 128 * (k / 16).div_ceil(4) * 4;
425 let (gate_scale_dev, up_scale_dev, src_n_major) =
432 match (self.gate_ptrs_t.as_ref(), self.up_ptrs_t.as_ref()) {
433 (Some(g), Some(u)) => (g.scale_ptrs, u.scale_ptrs, false),
434 _ => (self.gate_ptrs.scale_ptrs, self.up_ptrs.scale_ptrs, true),
435 };
436 if gate_scale_dev.is_null() || up_scale_dev.is_null() {
437 return Ok(());
438 }
439 let down_scale_dev = match self.down_ptrs_t.as_ref() {
440 Some(d) => Some(d.scale_ptrs),
441 None if !self.down_ptrs.scale_ptrs.is_null() => Some(self.down_ptrs.scale_ptrs),
442 None => None,
443 };
444 let mut owned: Vec<DevicePtr> = Vec::new();
445 let mut build_one = |scale_ptrs_dev: DevicePtr, n: usize, k: usize| -> Result<Vec<u64>> {
451 let len = sfb_len(n, k);
452 let sp = crate::layers::ops::read_expert_ptrs_u64(gpu, scale_ptrs_dev, num)?;
453 let mut sfb_ptrs = vec![0u64; num];
454 for (e, &sptr) in sp.iter().enumerate() {
455 if sptr == 0 {
456 continue; }
458 let sfb = gpu.alloc(len)?;
459 spark_runtime::cutlass::pack_weight_sfb(
460 sptr,
461 sfb.0,
462 n as u32,
463 k as u32,
464 src_n_major,
465 stream,
466 )?;
467 sfb_ptrs[e] = sfb.0;
468 owned.push(sfb);
469 }
470 gpu.synchronize(stream)?;
471 Ok(sfb_ptrs)
472 };
473 let gate_sfb = build_one(gate_scale_dev, inter, h)?;
474 let up_sfb = build_one(up_scale_dev, inter, h)?;
475 let down = match down_scale_dev {
476 Some(ds) => Some((
477 self.down_ptrs.packed_ptrs,
478 build_one(ds, h, inter)?,
479 self.down_ptrs.scale2_vals,
480 )),
481 None => None,
482 };
483 self.cutlass_grouped_host = Some(crate::layers::ops::MoeCutlassHostTables::snapshot(
484 gpu,
485 num,
486 self.gate_ptrs.packed_ptrs,
487 gate_sfb,
488 self.gate_ptrs.scale2_vals,
489 self.up_ptrs.packed_ptrs,
490 up_sfb,
491 self.up_ptrs.scale2_vals,
492 down,
493 )?);
494 self._cutlass_sfb_owned = owned;
495 tracing::info!(
496 "CUTLASS grouped SFB: built {num} experts gate/up (N={inter} K={h}) + down (N={h} K={inter})"
497 );
498 Ok(())
499 }
500}