Skip to main content

ruda_tensor_device/dispatch/
frozen_nf4.rs

1use crate::{DeviceBackend,DeviceRuntime,FloatElement,IntElement,element::BoolElement};
2use ruda_tensor::{FloatDType,frozen_nf4::{FrozenNf4Ops,Nf4ProjectionOptions},ops::FloatTensorOps,tensor::{FloatTensor,IntTensor}};
3use rublas::tensor_nf4::{Nf4Gemm,Nf4Layout,Nf4Error};
4
5impl<R,F,I,BT> FrozenNf4Ops for DeviceBackend<R,F,I,BT>
6    where R:DeviceRuntime,F:FloatElement,I:IntElement,BT:BoolElement {
7    type Nf4Error=Nf4Error;
8    fn frozen_nf4_forward(input:FloatTensor<Self>,packed:IntTensor<Self>,scales:FloatTensor<Self>,codebook:FloatTensor<Self>,
9        bias:Option<FloatTensor<Self>>,options:Nf4ProjectionOptions) -> Result<FloatTensor<Self>,Self::Nf4Error> {
10        let layout=Nf4Layout::new(options.input_features,options.output_features,options.block_size)?;
11        Nf4Gemm::new(packed,scales,codebook,bias,layout)?.forward(input,options.tile_rows,options.use_tensor_core)
12    }
13    fn frozen_nf4_input_backward(gradient:FloatTensor<Self>,packed:IntTensor<Self>,scales:FloatTensor<Self>,codebook:FloatTensor<Self>,
14        options:Nf4ProjectionOptions,activation_dtype:FloatDType) -> Result<FloatTensor<Self>,Self::Nf4Error> {
15        let layout=Nf4Layout::new(options.input_features,options.output_features,options.block_size)?;
16        let gradient=Self::float_cast(gradient,activation_dtype);
17        let gradient=Nf4Gemm::new(packed,scales,codebook,None,layout)?.input_backward_f32(gradient,options.tile_rows,options.use_tensor_core)?;
18        Ok(Self::float_cast(gradient,activation_dtype))
19    }
20}