use crate::{Backend,FloatDType,tensor::{FloatTensor,IntTensor}};
use core::fmt;
#[derive(Clone,Copy,Debug,PartialEq,Eq)]
pub struct Nf4ProjectionOptions {
pub input_features:usize,
pub output_features:usize,
pub block_size:usize,
pub tile_rows:usize,
pub use_tensor_core:bool,
}
#[derive(Debug)]
pub enum FrozenNf4Error<E:fmt::Debug> {
Native(E),
TrainableBase,
InputGradientNotDifferentiable,
}
impl<E:fmt::Debug> fmt::Display for FrozenNf4Error<E> {
fn fmt(&self,f:&mut fmt::Formatter<'_>) -> fmt::Result {
match self {Self::Native(error)=>write!(f,"native NF4: {error:?}"),Self::TrainableBase=>f.write_str("NF4 scales/codebook/bias must be frozen"),
Self::InputGradientNotDifferentiable=>f.write_str("NF4 provides first-order input gradients only")}
}
}
impl<E:fmt::Debug> core::error::Error for FrozenNf4Error<E> {}
pub trait FrozenNf4Ops:Backend {
type Nf4Error:fmt::Debug;
fn frozen_nf4_forward(input:FloatTensor<Self>,packed:IntTensor<Self>,scales:FloatTensor<Self>,codebook:FloatTensor<Self>,
bias:Option<FloatTensor<Self>>,options:Nf4ProjectionOptions) -> Result<FloatTensor<Self>,Self::Nf4Error>;
fn frozen_nf4_input_backward(gradient:FloatTensor<Self>,packed:IntTensor<Self>,scales:FloatTensor<Self>,codebook:FloatTensor<Self>,
options:Nf4ProjectionOptions,activation_dtype:FloatDType) -> Result<FloatTensor<Self>,Self::Nf4Error>;
}