ruda_tensor_device/dispatch/
packed_experts.rs1use crate::{DeviceBackend,DeviceRuntime,FloatElement,IntElement,element::BoolElement};
2use ruda_tensor::{FloatDType,packed_experts::*,ops::FloatTensorOps,tensor::{FloatTensor,IntTensor}};
3use rudnn::moe::{MoeError,ReceivedExpertRows,PackedExpertError,PackedExpertProjection,PackedSwiGluExperts,PackedSwiGluCache,Nf4ExpertProjection,Nf4ExpertExecution};
4use rublas::{tensor_int4::AwqGroupedGemm,tensor_nf4::Nf4Layout};
5
6#[derive(Clone,Debug)]
8pub struct NativePackedProjectionState<R:DeviceRuntime> {rows:ReceivedExpertRows<R>,projection:PackedExpertProjection<R>,dtype:FloatDType}
9#[derive(Clone,Debug)]
11pub struct NativePackedSwiGluState<R:DeviceRuntime> {rows:ReceivedExpertRows<R>,cache:Option<PackedSwiGluCache<R>>,dtype:FloatDType}
12fn projection<R,F,I,BT>(payload:PackedExpertPayload<DeviceBackend<R,F,I,BT>>) -> Result<PackedExpertProjection<R>,PackedExpertError>
13 where R:DeviceRuntime,F:FloatElement,I:IntElement,BT:BoolElement {
14 Ok(match payload {
15 PackedExpertPayload::Nf4(value)=>{
16 let o=value.options.projection;if o.tile_rows==0 {return Err(MoeError("packed NF4 expert tile rows must be positive").into());}
17 let layout=Nf4Layout::new(o.input_features,o.output_features,o.block_size).map_err(rudnn::moe::Nf4ExpertError::from)?;
18 PackedExpertProjection::Nf4 {projection:Nf4ExpertProjection::new(value.packed,value.scales,value.codebook,value.options.experts,layout)?,
19 execution:Nf4ExpertExecution {tile_rows:o.tile_rows,use_tensor_core:o.use_tensor_core}}
20 },
21 PackedExpertPayload::Awq(value)=>{
22 let projection=AwqGroupedGemm::new(value.qweight,value.qzeros,value.scales,value.bias,value.options.group_size)?;
23 let (e,l)=projection.layout();
24 if [e,l.input_features,l.output_features]!=[value.options.experts,value.options.input_features,value.options.output_features] {
25 return Err(MoeError("AWQ expert options differ from original source packed cube geometry").into());}
26 PackedExpertProjection::Awq(projection)
27 },
28 PackedExpertPayload::Nf4Window {payload:value,element_offset}=>{
29 let o=value.options.projection;if o.tile_rows==0 {return Err(MoeError("packed NF4 window tile rows must be positive").into());}
30 let layout=Nf4Layout::new(o.input_features,o.output_features,o.block_size).map_err(rudnn::moe::Nf4ExpertError::from)?;
31 PackedExpertProjection::Nf4 {projection:Nf4ExpertProjection::from_window(value.packed,value.scales,value.codebook,value.options.experts,layout,element_offset)?,
32 execution:Nf4ExpertExecution {tile_rows:o.tile_rows,use_tensor_core:o.use_tensor_core}}
33 },
34 })
35}
36impl<R,F,I,BT> FrozenPackedExpertOps for DeviceBackend<R,F,I,BT> where R:DeviceRuntime,F:FloatElement,I:IntElement,BT:BoolElement {
37 type PackedExpertError=PackedExpertError;
38 type PackedProjectionState=NativePackedProjectionState<R>;
39 type PackedSwiGluState=NativePackedSwiGluState<R>;
40 fn packed_expert_forward(input:FloatTensor<Self>,global_ids:IntTensor<Self>,payload:PackedExpertPayload<Self>)
41 -> Result<(FloatTensor<Self>,Self::PackedProjectionState),Self::PackedExpertError> {
42 let (e,start)=payload.expert_range();let dtype=input.dtype.into();let projection=projection(payload)?;
43 let rows=ReceivedExpertRows::new(input,global_ids,start,e)?;let output=projection.forward(rows.grouped())?;
44 Ok((rows.restore(output)?,NativePackedProjectionState {rows,projection,dtype}))
45 }
46 fn packed_expert_input_backward(state:Self::PackedProjectionState,gradient:FloatTensor<Self>) -> Result<FloatTensor<Self>,Self::PackedExpertError> {
47 let gradient=state.rows.sort_gradient(Self::float_cast(gradient,state.dtype))?;
48 Ok(state.rows.restore(state.projection.input_backward(state.rows.grouped(),gradient)?)?)
49 }
50 fn packed_swiglu_forward(input:FloatTensor<Self>,global_ids:IntTensor<Self>,gate:PackedExpertPayload<Self>,up:PackedExpertPayload<Self>,down:PackedExpertPayload<Self>,retain_input:bool)
51 -> Result<(FloatTensor<Self>,Self::PackedSwiGluState),Self::PackedExpertError> {
52 let range=gate.expert_range();if up.expert_range()!=range || down.expert_range()!=range {
53 return Err(MoeError("packed gate/up/down must describe the same explicit resident global expert range").into());}
54 let experts=PackedSwiGluExperts::new(projection(gate)?,projection(up)?,projection(down)?)?;let dtype=input.dtype.into();
55 let rows=ReceivedExpertRows::new(input,global_ids,range.1,range.0)?;
56 let (output,cache)=if retain_input {let (output,cache)=experts.forward_training(rows.grouped())?;(output,Some(cache))}
57 else {(experts.forward(rows.grouped())?,None)};
58 Ok((rows.restore(output)?,NativePackedSwiGluState {rows,cache,dtype}))
59 }
60 fn packed_swiglu_input_backward(state:Self::PackedSwiGluState,gradient:FloatTensor<Self>) -> Result<FloatTensor<Self>,Self::PackedExpertError> {
61 let cache=state.cache.ok_or(MoeError("packed expert forward did not retain actual input VJP intermediates"))?;
62 let gradient=state.rows.sort_gradient(Self::float_cast(gradient,state.dtype))?;
63 Ok(state.rows.restore(cache.input_backward(gradient)?)?)
64 }
65}