ruda_tensor_device/dispatch/
frozen_nf4.rs1use 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}