use crate::{Backend, tensor::{FloatTensor, IntTensor}};
use core::fmt::{Debug, Display, Formatter};
#[derive(Debug)]
pub enum FrozenAwqError<E: Debug> {
Native(E),
TrainableBase,
}
impl<E: Debug> Display for FrozenAwqError<E> {
fn fmt(&self, f: &mut Formatter<'_>) -> core::fmt::Result {
match self {
Self::Native(error) => write!(f, "native AWQ: {error:?}"),
Self::TrainableBase => f.write_str("frozen AWQ scales and bias must not require gradients"),
}
}
}
impl<E: Debug> core::error::Error for FrozenAwqError<E> {}
pub trait FrozenAwqOps: Backend {
type AwqError: Debug;
fn frozen_awq_forward(
input: FloatTensor<Self>, qweight: IntTensor<Self>, qzeros: IntTensor<Self>,
scales: FloatTensor<Self>, bias: Option<FloatTensor<Self>>, group_size: usize,
) -> Result<FloatTensor<Self>, Self::AwqError>;
fn frozen_awq_input_backward(
gradient: FloatTensor<Self>, qweight: IntTensor<Self>, qzeros: IntTensor<Self>,
scales: FloatTensor<Self>, group_size: usize,
) -> Result<FloatTensor<Self>, Self::AwqError>;
}