Skip to main content

ruda_tensor_device/dispatch/
grouped_nf4.rs

1use 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/// Actual source row mapping, packed projection and original activation dtype.
7#[derive(Clone,Debug)]
8pub struct NativeNf4GroupedState<R:DeviceRuntime> {
9    rows:ReceivedExpertRows<R>,projection:Nf4ExpertProjection<R>,options:Nf4GroupedOptions,dtype:FloatDType,
10}
11/// Actual original selected SwiGLU row mapping and optional first-order cache.
12#[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}