ruda_tensor_device/dispatch/
grouped_nf4.rs1use crate::{DeviceBackend,DeviceRuntime,FloatElement,IntElement,element::BoolElement};
2use ruda_tensor::{FloatDType,grouped_nf4::*,ops::FloatTensorOps,tensor::{FloatTensor,IntTensor}};
3use rudnn::moe::{MoeError,ReceivedExpertRows,Nf4ExpertProjection,Nf4ExpertError,Nf4SwiGluExperts,Nf4SwiGluCache,Nf4ExpertExecution};
4use rublas::tensor_nf4::Nf4Layout;
5
6#[derive(Clone,Debug)]
8pub struct NativeNf4GroupedState<R:DeviceRuntime> {
9 rows:ReceivedExpertRows<R>,projection:Nf4ExpertProjection<R>,options:Nf4GroupedOptions,dtype:FloatDType,
10}
11#[derive(Clone,Debug)]
13pub struct NativeNf4SwiGluState<R:DeviceRuntime> {rows:ReceivedExpertRows<R>,cache:Option<Nf4SwiGluCache<R>>,dtype:FloatDType}
14fn execution(options:Nf4GroupedOptions) -> Nf4ExpertExecution {
15 Nf4ExpertExecution {tile_rows:options.projection.tile_rows,use_tensor_core:options.projection.use_tensor_core}
16}
17fn projection<R,F,I,BT>(payload:Nf4ExpertPayload<DeviceBackend<R,F,I,BT>>) -> Result<Nf4ExpertProjection<R>,Nf4ExpertError>
18 where R:DeviceRuntime,F:FloatElement,I:IntElement,BT:BoolElement {
19 let o=payload.options;
20 if o.projection.tile_rows==0 {return Err(MoeError("NF4 expert decoded tile rows must be positive").into());}
21 Ok(Nf4ExpertProjection::new(payload.packed,payload.scales,payload.codebook,o.experts,
22 Nf4Layout::new(o.projection.input_features,o.projection.output_features,o.projection.block_size)?)?)
23}
24impl<R,F,I,BT> FrozenNf4GroupedOps for DeviceBackend<R,F,I,BT> where R:DeviceRuntime,F:FloatElement,I:IntElement,BT:BoolElement {
25 type Nf4GroupedError=Nf4ExpertError;
26 type Nf4GroupedState=NativeNf4GroupedState<R>;
27 fn frozen_nf4_grouped_forward(input:FloatTensor<Self>,global_ids:IntTensor<Self>,payload:Nf4ExpertPayload<Self>)
28 -> Result<(FloatTensor<Self>,Self::Nf4GroupedState),Self::Nf4GroupedError> {
29 let options=payload.options;let dtype=input.dtype.into();let projection=projection(payload)?;
30 let rows=ReceivedExpertRows::new(input,global_ids,options.expert_start,options.experts)?;
31 let output=projection.forward(rows.grouped(),options.projection.tile_rows,options.projection.use_tensor_core)?;
32 Ok((rows.restore(output)?,NativeNf4GroupedState {rows,projection,options,dtype}))
33 }
34 fn frozen_nf4_grouped_input_backward(state:Self::Nf4GroupedState,gradient:FloatTensor<Self>) -> Result<FloatTensor<Self>,Self::Nf4GroupedError> {
35 let gradient=state.rows.sort_gradient(Self::float_cast(gradient,state.dtype))?;
36 let gradient=state.projection.input_backward_f32(state.rows.grouped(),gradient,state.options.projection.tile_rows,state.options.projection.use_tensor_core)?;
37 Ok(state.rows.restore(Self::float_cast(gradient,state.dtype))?)
38 }
39}
40impl<R,F,I,BT> FrozenNf4SwiGluOps for DeviceBackend<R,F,I,BT> where R:DeviceRuntime,F:FloatElement,I:IntElement,BT:BoolElement {
41 type Nf4SwiGluState=NativeNf4SwiGluState<R>;
42 fn frozen_nf4_swiglu_forward(input:FloatTensor<Self>,global_ids:IntTensor<Self>,gate:Nf4ExpertPayload<Self>,up:Nf4ExpertPayload<Self>,down:Nf4ExpertPayload<Self>,retain_input:bool)
43 -> Result<(FloatTensor<Self>,Self::Nf4SwiGluState),Self::Nf4GroupedError> {
44 let options=[gate.options,up.options,down.options];
45 if options.iter().any(|o|o.experts!=options[0].experts || o.expert_start!=options[0].expert_start) {
46 return Err(MoeError("NF4 gate/up/down payloads must describe the same explicit global expert range").into());
47 }
48 let experts=Nf4SwiGluExperts::new(projection(gate)?,projection(up)?,projection(down)?)?;
49 let dtype=input.dtype.into();let rows=ReceivedExpertRows::new(input,global_ids,options[0].expert_start,options[0].experts)?;
50 let execution=options.map(execution);
51 let (output,cache)=if retain_input {let (output,cache)=experts.forward_training(rows.grouped(),execution)?;(output,Some(cache))}
52 else {(experts.forward(rows.grouped(),execution)?,None)};
53 Ok((rows.restore(output)?,NativeNf4SwiGluState {rows,cache,dtype}))
54 }
55 fn frozen_nf4_swiglu_input_backward(state:Self::Nf4SwiGluState,gradient:FloatTensor<Self>) -> Result<FloatTensor<Self>,Self::Nf4GroupedError> {
56 let cache=state.cache.ok_or(MoeError("NF4 SwiGLU forward did not retain the input VJP cache"))?;
57 let gradient=state.rows.sort_gradient(Self::float_cast(gradient,state.dtype))?;
58 Ok(state.rows.restore(cache.input_backward(gradient)?)?)
59 }
60}