use core::{convert::Infallible,fmt};
use ruda_model::{module::{Module,ModuleDisplay},tensor::{FrozenAwqOps,FrozenNf4Ops,Tensor,backend::Backend}};
use crate::{Linear,LoRALinear,FrozenNf4Linear,Nf4LoRALinear,FrozenAwqLinear,AwqLoRALinear,
QuantizedLinear,QuantizedLoRALinear};
use super::{AwqTransformerProjection,AwqGroupedQueryAttention,AwqFeedForward,AwqTransformerBlock,
AwqTransformerStack,AwqTransformerHead,AwqTransformerModel};
use super::AdaptedProjection;
pub trait TransformerProjectionStorage<B:Backend>:Module<B> {
type Stored:Module<B>;
}
impl<B:Backend,P:Module<B>> TransformerProjectionStorage<B> for P {type Stored=P;}
pub type BackendProjection<B,P> = <P as TransformerProjectionStorage<B>>::Stored;
pub trait TransformerProjectionShape<B:Backend>:Module<B>+ModuleDisplay {
fn dimensions(&self) -> [usize;2];
}
pub trait TransformerProjection<B:Backend>:TransformerProjectionShape<B> {
type Error:fmt::Debug;
fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,Self::Error>;
}
impl<B:Backend> TransformerProjectionShape<B> for AwqTransformerProjection<B> {
fn dimensions(&self) -> [usize;2] {self.dimensions()}
}
impl<B:FrozenAwqOps> TransformerProjection<B> for AwqTransformerProjection<B> {
type Error=B::AwqError;
fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,Self::Error> {self.forward(input)}
}
impl<B:Backend> TransformerProjectionShape<B> for FrozenAwqLinear<B> {
fn dimensions(&self) -> [usize;2] {self.dimensions()}
}
impl<B:FrozenAwqOps> TransformerProjection<B> for FrozenAwqLinear<B> {
type Error=B::AwqError;
fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,Self::Error> {self.forward(input)}
}
impl<B:Backend> TransformerProjectionShape<B> for AwqLoRALinear<B> {
fn dimensions(&self) -> [usize;2] {self.base.dimensions()}
}
impl<B:FrozenAwqOps> TransformerProjection<B> for AwqLoRALinear<B> {
type Error=B::AwqError;
fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,Self::Error> {self.forward(input)}
}
impl<B:Backend> TransformerProjectionShape<B> for Linear<B> {
fn dimensions(&self) -> [usize;2] {self.weight.val().dims()}
}
impl<B:Backend> TransformerProjection<B> for Linear<B> {
type Error=Infallible;
fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,Self::Error> {Ok(self.forward(input))}
}
impl<B:Backend> TransformerProjectionShape<B> for LoRALinear<B> {
fn dimensions(&self) -> [usize;2] {self.base.weight.val().dims()}
}
impl<B:Backend> TransformerProjection<B> for LoRALinear<B> {
type Error=Infallible;
fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,Self::Error> {Ok(self.forward(input))}
}
impl<B:Backend> TransformerProjectionShape<B> for QuantizedLinear<B> {
fn dimensions(&self) -> [usize;2] {let [output,input]=self.weight.val().dims();[input,output]}
}
impl<B:Backend> TransformerProjection<B> for QuantizedLinear<B> {
type Error=Infallible;
fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,Self::Error> {Ok(self.forward(input))}
}
impl<B:Backend> TransformerProjectionShape<B> for QuantizedLoRALinear<B> {
fn dimensions(&self) -> [usize;2] {self.base.dimensions()}
}
impl<B:Backend> TransformerProjection<B> for QuantizedLoRALinear<B> {
type Error=Infallible;
fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,Self::Error> {Ok(self.forward(input))}
}
#[derive(Module,Debug)]
pub enum QuantizedTransformerProjection<B:Backend> {
Dense(Linear<B>),
LoRA(LoRALinear<B>),
Quantized(QuantizedLinear<B>),
QuantizedLoRA(QuantizedLoRALinear<B>),
}
impl<B:Backend> TransformerProjectionShape<B> for QuantizedTransformerProjection<B> {
fn dimensions(&self) -> [usize;2] {
match self {Self::Dense(layer)=>layer.dimensions(),Self::LoRA(layer)=>layer.dimensions(),
Self::Quantized(layer)=>layer.dimensions(),Self::QuantizedLoRA(layer)=>layer.dimensions()}
}
}
impl<B:Backend> TransformerProjection<B> for QuantizedTransformerProjection<B> {
type Error=Infallible;
fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,Self::Error> {
Ok(match self {Self::Dense(layer)=>layer.forward(input),Self::LoRA(layer)=>layer.forward(input),
Self::Quantized(layer)=>layer.forward(input),Self::QuantizedLoRA(layer)=>layer.forward(input)})
}
}
impl<B:Backend> QuantizedTransformerProjection<B> {
pub fn pack_float(self, scheme:&ruda_model::tensor::quantization::QuantScheme,
calibration_dtype:ruda_model::tensor::FloatDType) -> Self {
match self {
Self::Dense(layer)=>Self::Quantized(QuantizedLinear::from_linear(layer,scheme,calibration_dtype)),
Self::LoRA(layer)=>Self::QuantizedLoRA(QuantizedLoRALinear {
base:QuantizedLinear::from_linear(layer.base,scheme,calibration_dtype),
adapter_a:layer.adapter_a,adapter_b:layer.adapter_b,dropout:layer.dropout,scale:layer.scale}),
_=>panic!("pack_float requires an actual floating base, not an already packed checkpoint"),
}
}
}
impl<B:Backend> TransformerProjectionShape<B> for AdaptedProjection<B> {
fn dimensions(&self) -> [usize;2] {
match self {Self::Dense(layer)=>layer.weight.val().dims(),Self::LoRA(layer)=>layer.base.weight.val().dims()}
}
}
impl<B:Backend> TransformerProjection<B> for AdaptedProjection<B> {
type Error=Infallible;
fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,Self::Error> {Ok(self.forward(input))}
}
impl<B:Backend> TransformerProjectionShape<B> for FrozenNf4Linear<B> {
fn dimensions(&self) -> [usize;2] {[self.input_features,self.output_features]}
}
impl<B:FrozenNf4Ops> TransformerProjection<B> for FrozenNf4Linear<B> {
type Error=B::Nf4Error;
fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,Self::Error> {self.forward(input)}
}
impl<B:Backend> TransformerProjectionShape<B> for Nf4LoRALinear<B> {
fn dimensions(&self) -> [usize;2] {[self.base.input_features,self.base.output_features]}
}
impl<B:FrozenNf4Ops> TransformerProjection<B> for Nf4LoRALinear<B> {
type Error=B::Nf4Error;
fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,Self::Error> {self.forward(input)}
}
#[derive(Module,Debug)]
pub enum Nf4TransformerProjection<B:Backend> {
Dense(Linear<B>),
LoRA(LoRALinear<B>),
Nf4(FrozenNf4Linear<B>),
Nf4LoRA(Nf4LoRALinear<B>),
}
impl<B:Backend> TransformerProjectionShape<B> for Nf4TransformerProjection<B> {
fn dimensions(&self) -> [usize;2] {
match self {Self::Dense(layer)=>layer.dimensions(),Self::LoRA(layer)=>layer.dimensions(),
Self::Nf4(layer)=>layer.dimensions(),Self::Nf4LoRA(layer)=>layer.dimensions()}
}
}
impl<B:FrozenNf4Ops> TransformerProjection<B> for Nf4TransformerProjection<B> {
type Error=B::Nf4Error;
fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,Self::Error> {
match self {Self::Dense(layer)=>Ok(layer.forward(input)),Self::LoRA(layer)=>Ok(layer.forward(input)),
Self::Nf4(layer)=>layer.forward(input),Self::Nf4LoRA(layer)=>layer.forward(input)}
}
}
#[derive(Debug)]
pub enum MixedProjectionError<A:fmt::Debug,N:fmt::Debug> {
Awq(A),
Nf4(N),
}
impl<A:fmt::Debug,N:fmt::Debug> fmt::Display for MixedProjectionError<A,N> {
fn fmt(&self,f:&mut fmt::Formatter<'_>) -> fmt::Result {
match self {Self::Awq(error)=>write!(f,"AWQ projection: {error:?}"),Self::Nf4(error)=>write!(f,"NF4 projection: {error:?}")}
}
}
impl<A:fmt::Debug,N:fmt::Debug> core::error::Error for MixedProjectionError<A,N> {}
#[derive(Module,Debug)]
pub enum MixedTransformerProjection<B:Backend> {
Dense(Linear<B>),
LoRA(LoRALinear<B>),
Awq(FrozenAwqLinear<B>),
AwqLoRA(AwqLoRALinear<B>),
Nf4(FrozenNf4Linear<B>),
Nf4LoRA(Nf4LoRALinear<B>),
}
impl<B:Backend> TransformerProjectionShape<B> for MixedTransformerProjection<B> {
fn dimensions(&self) -> [usize;2] {
match self {Self::Dense(layer)=>layer.dimensions(),Self::LoRA(layer)=>layer.dimensions(),
Self::Awq(layer)=>layer.dimensions(),Self::AwqLoRA(layer)=>layer.base.dimensions(),
Self::Nf4(layer)=>layer.dimensions(),Self::Nf4LoRA(layer)=>layer.dimensions()}
}
}
impl<B:FrozenAwqOps+FrozenNf4Ops> TransformerProjection<B> for MixedTransformerProjection<B> {
type Error=MixedProjectionError<B::AwqError,B::Nf4Error>;
fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,Self::Error> {
match self {Self::Dense(layer)=>Ok(layer.forward(input)),Self::LoRA(layer)=>Ok(layer.forward(input)),
Self::Awq(layer)=>layer.forward(input).map_err(MixedProjectionError::Awq),
Self::AwqLoRA(layer)=>layer.forward(input).map_err(MixedProjectionError::Awq),
Self::Nf4(layer)=>layer.forward(input).map_err(MixedProjectionError::Nf4),
Self::Nf4LoRA(layer)=>layer.forward(input).map_err(MixedProjectionError::Nf4)}
}
}
#[derive(Module,Debug)]
pub enum UniversalTransformerProjection<B:Backend> {
Mixed(MixedTransformerProjection<B>),
Generic(QuantizedTransformerProjection<B>),
}
impl<B:Backend> From<MixedTransformerProjection<B>> for UniversalTransformerProjection<B> {
fn from(value:MixedTransformerProjection<B>) -> Self {Self::Mixed(value)}
}
impl<B:Backend> From<QuantizedTransformerProjection<B>> for UniversalTransformerProjection<B> {
fn from(value:QuantizedTransformerProjection<B>) -> Self {Self::Generic(value)}
}
impl<B:Backend> TransformerProjectionShape<B> for UniversalTransformerProjection<B> {
fn dimensions(&self) -> [usize;2] {
match self {Self::Mixed(layer)=>layer.dimensions(),Self::Generic(layer)=>layer.dimensions()}
}
}
impl<B:FrozenAwqOps+FrozenNf4Ops> TransformerProjection<B> for UniversalTransformerProjection<B> {
type Error=MixedProjectionError<B::AwqError,B::Nf4Error>;
fn forward<const D:usize>(&self,input:Tensor<B,D>) -> Result<Tensor<B,D>,Self::Error> {
match self {Self::Mixed(layer)=>layer.forward(input),
Self::Generic(layer)=>layer.forward(input).map_err(|never| match never {})}
}
}
pub type ProjectedGroupedQueryAttention<B,P> = AwqGroupedQueryAttention<B,P>;
pub type ProjectedFeedForward<B,P> = AwqFeedForward<B,P>;
pub type ProjectedTransformerBlock<B,P> = AwqTransformerBlock<B,P>;
pub type ProjectedTransformerStack<B,P> = AwqTransformerStack<B,P>;
pub type ProjectedTransformerHead<B,P> = AwqTransformerHead<B,P>;
pub type ProjectedTransformerModel<B,P> = AwqTransformerModel<B,P>;
pub type Nf4TransformerModel<B> = ProjectedTransformerModel<B,Nf4TransformerProjection<B>>;
pub type MixedTransformerModel<B> = ProjectedTransformerModel<B,MixedTransformerProjection<B>>;