use alloc::collections::BTreeMap;
use ruda_model::tensor::{DType,backend::Backend};
use crate::{Linear,attention::GroupedQueryAttention};
use super::{TransformerAdapterConfig,LayerAdapterConfig,AttentionAdapterTarget,FeedForwardAdapterTarget,
TransformerProjectionShape,AwqTransformerProjection,Nf4TransformerProjection,MixedTransformerProjection,
AwqGroupedQueryAttention,AwqFeedForward,AwqTransformerBlock,AwqTransformerStack,AwqTransformerModel,
DenseFeedForward,DenseTransformerBlock,DenseTransformerStack,TransformerEmbeddings,TransformerHead,DenseTransformerNorm,
ProjectedCrossAttentionBlock,ProjectedEncoderDecoderLayer,ProjectedEncoderDecoderStack,ProjectedEncoderDecoderModel,
DenseCrossAttentionBlock,DenseEncoderDecoderLayer,DenseEncoderDecoderStack,DecoderLayerAdapterConfig};
pub trait AdaptTransformerProjection<B:Backend>:TransformerProjectionShape<B> {
fn validate_adapter(&self,config:&TransformerAdapterConfig);
fn with_adapter(self,config:&TransformerAdapterConfig) -> Self;
}
fn validate_config(config:&TransformerAdapterConfig,packed:bool) {
assert!(config.lora.rank>0 && config.lora.alpha.is_finite(),"invalid native projection adapter rank/alpha");
assert!(config.lora.dropout.is_finite() && (0.0..1.0).contains(&config.lora.dropout),"invalid native projection adapter dropout");
if packed {assert!(config.adapter_dtype.is_some(),"packed projection adapters require an explicit floating adapter_dtype");}
if let Some(dtype)=config.adapter_dtype {
assert!(if packed {matches!(dtype,DType::F16|DType::BF16|DType::F32)} else {dtype.is_float()},"unsupported native projection adapter storage");
}
}
fn unique<T:PartialEq>(targets:&[T]) {
for (index,target) in targets.iter().enumerate() {assert!(!targets[..index].contains(target),"duplicate native projection adapter target");}
}
macro_rules! adapt_original_projection {
($projection:ident,[$($base:ident=>$adapted:ident,$init:ident),*]) => {
impl<B:Backend> From<Linear<B>> for $projection<B> {
fn from(layer:Linear<B>) -> Self {Self::Dense(layer)}
}
impl<B:Backend> AdaptTransformerProjection<B> for $projection<B> {
fn validate_adapter(&self,config:&TransformerAdapterConfig) {
match self {
Self::Dense(_)=>validate_config(config,false),
$(Self::$base(layer)=>{validate_config(config,true);layer.validate();},)*
_=>panic!("selected native projection already has an adapter"),
}
}
fn with_adapter(self,config:&TransformerAdapterConfig) -> Self {
self.validate_adapter(config);
match self {
Self::Dense(layer)=>{
let dtype=config.adapter_dtype.unwrap_or_else(||layer.weight.val().dtype());
Self::LoRA(config.lora.init_with_options(layer,dtype,config.use_rslora))
},
$(Self::$base(layer)=>Self::$adapted(config.lora.$init(layer,config.adapter_dtype.expect("validated packed adapter dtype"),config.use_rslora)),)*
_=>unreachable!("validated projection is not already adapted"),
}
}
}
};
}
adapt_original_projection!(AwqTransformerProjection,[Awq=>AwqLoRA,init_awq]);
adapt_original_projection!(Nf4TransformerProjection,[Nf4=>Nf4LoRA,init_nf4]);
adapt_original_projection!(MixedTransformerProjection,[Awq=>AwqLoRA,init_awq,Nf4=>Nf4LoRA,init_nf4]);
impl<B:Backend,P:TransformerProjectionShape<B>+From<Linear<B>>> AwqGroupedQueryAttention<B,P> {
pub fn from_dense(attention:GroupedQueryAttention<B>) -> Self {
Self::from_projections(attention.query.into(),attention.key.into(),attention.value.into(),attention.output.into(),
attention.query_heads,attention.kv_heads,attention.head_dimension,attention.dropout)
}
}
impl<B:Backend,P:AdaptTransformerProjection<B>> AwqGroupedQueryAttention<B,P> {
pub(super) fn validate_adapters(&self,config:&TransformerAdapterConfig,targets:&[AttentionAdapterTarget]) {
unique(targets);
for target in targets {
match target {AttentionAdapterTarget::Query=>self.query.validate_adapter(config),AttentionAdapterTarget::Key=>self.key.validate_adapter(config),
AttentionAdapterTarget::Value=>self.value.validate_adapter(config),AttentionAdapterTarget::Output=>self.output.validate_adapter(config)}
}
}
pub fn with_adapters(mut self,config:&TransformerAdapterConfig,targets:&[AttentionAdapterTarget]) -> Self {
self.validate_adapters(config,targets);
if targets.contains(&AttentionAdapterTarget::Query) {self.query=self.query.with_adapter(config);}
if targets.contains(&AttentionAdapterTarget::Key) {self.key=self.key.with_adapter(config);}
if targets.contains(&AttentionAdapterTarget::Value) {self.value=self.value.with_adapter(config);}
if targets.contains(&AttentionAdapterTarget::Output) {self.output=self.output.with_adapter(config);}
self
}
}
impl<B:Backend,P:TransformerProjectionShape<B>+From<Linear<B>>> AwqFeedForward<B,P> {
pub fn from_dense(feed:DenseFeedForward<B>) -> Self {
Self::from_projections(feed.up.into(),feed.gate.map(Into::into),feed.down.into(),feed.activation,feed.dropout)
}
}
impl<B:Backend,P:AdaptTransformerProjection<B>> AwqFeedForward<B,P> {
pub(super) fn validate_adapters(&self,config:&TransformerAdapterConfig,targets:&[FeedForwardAdapterTarget]) {
unique(targets);
for target in targets {
match target {FeedForwardAdapterTarget::Up=>self.up.validate_adapter(config),FeedForwardAdapterTarget::Down=>self.down.validate_adapter(config),
FeedForwardAdapterTarget::Gate=>self.gate.as_ref().expect("selected original FFN has no gate").validate_adapter(config)}
}
}
pub fn with_adapters(mut self,config:&TransformerAdapterConfig,targets:&[FeedForwardAdapterTarget]) -> Self {
self.validate_adapters(config,targets);
if targets.contains(&FeedForwardAdapterTarget::Up) {self.up=self.up.with_adapter(config);}
if targets.contains(&FeedForwardAdapterTarget::Gate) {self.gate=self.gate.map(|gate|gate.with_adapter(config));}
if targets.contains(&FeedForwardAdapterTarget::Down) {self.down=self.down.with_adapter(config);}
self
}
}
impl<B:Backend,P:TransformerProjectionShape<B>+From<Linear<B>>> AwqTransformerBlock<B,P> {
pub fn from_dense(block:DenseTransformerBlock<B>) -> Self {
Self {attention:AwqGroupedQueryAttention::from_dense(block.attention),feed_forward:AwqFeedForward::from_dense(block.feed_forward),
attention_norm:block.attention_norm,feed_forward_norm:block.feed_forward_norm,residual_dropout:block.residual_dropout,norm_first:block.norm_first}
}
}
impl<B:Backend,P:AdaptTransformerProjection<B>> AwqTransformerBlock<B,P> {
pub(super) fn validate_adapters(&self,config:&TransformerAdapterConfig,attention:&[AttentionAdapterTarget],feed_forward:&[FeedForwardAdapterTarget]) {
assert!(!attention.is_empty() || !feed_forward.is_empty(),"selected native block requires an actual adapter target");
self.attention.validate_adapters(config,attention);self.feed_forward.validate_adapters(config,feed_forward);
}
pub fn with_adapters(mut self,config:&TransformerAdapterConfig,attention:&[AttentionAdapterTarget],feed_forward:&[FeedForwardAdapterTarget]) -> Self {
self.validate_adapters(config,attention,feed_forward);
self.attention=self.attention.with_adapters(config,attention);self.feed_forward=self.feed_forward.with_adapters(config,feed_forward);self
}
}
impl<B:Backend,P:TransformerProjectionShape<B>+From<Linear<B>>> AwqTransformerStack<B,P> {
pub fn from_dense(stack:DenseTransformerStack<B>) -> Self {Self {blocks:stack.blocks.into_iter().map(AwqTransformerBlock::from_dense).collect()}}
}
impl<B:Backend,P:AdaptTransformerProjection<B>> AwqTransformerStack<B,P> {
pub(super) fn selected_adapters<'a>(&self,targets:&'a [LayerAdapterConfig]) -> BTreeMap<usize,&'a LayerAdapterConfig> {
let mut selected=BTreeMap::new();
for target in targets {
assert!(target.layer<self.blocks.len(),"native adapter layer index is outside the original stack");
assert!(selected.insert(target.layer,target).is_none(),"duplicate native adapter layer index");
self.blocks[target.layer].validate_adapters(&target.adapter,&target.attention,&target.feed_forward);
}
selected
}
pub fn with_adapters(self,targets:&[LayerAdapterConfig]) -> Self {
let selected=self.selected_adapters(targets);
let blocks=self.blocks.into_iter().enumerate().map(|(index,block)| {
if let Some(target)=selected.get(&index) {block.with_adapters(&target.adapter,&target.attention,&target.feed_forward)} else {block}
}).collect();Self {blocks}
}
}
impl<B:Backend,P:TransformerProjectionShape<B>+From<Linear<B>>> AwqTransformerModel<B,P> {
pub fn from_dense_parts(embeddings:TransformerEmbeddings<B>,backbone:DenseTransformerStack<B>,normalization:Option<DenseTransformerNorm<B>>,
head:TransformerHead<B>) -> Self {
let head=super::AwqTransformerHead::from_projection(head.projection.into(),head.normalization,head.dropout);
Self::from_parts(embeddings,AwqTransformerStack::from_dense(backbone),normalization,head)
}
}
impl<B:Backend,P:AdaptTransformerProjection<B>> AwqTransformerModel<B,P> {
pub fn with_adapters(mut self,targets:&[LayerAdapterConfig],head:Option<&TransformerAdapterConfig>) -> Self {
let _=self.backbone.selected_adapters(targets);
if let Some(config)=head {self.head.projection.validate_adapter(config);}
self.backbone=self.backbone.with_adapters(targets);
if let Some(config)=head {self.head.projection=self.head.projection.with_adapter(config);}
self
}
}
impl<B:Backend,P:TransformerProjectionShape<B>+From<Linear<B>>> ProjectedCrossAttentionBlock<B,P> {
pub fn from_dense(cross:DenseCrossAttentionBlock<B>) -> Self {
Self::from_parts(AwqGroupedQueryAttention::from_dense(cross.attention),cross.query_norm,cross.memory_norm,cross.residual_dropout,cross.norm_first)
}
}
impl<B:Backend,P:AdaptTransformerProjection<B>> ProjectedCrossAttentionBlock<B,P> {
pub fn with_adapters(mut self,config:&TransformerAdapterConfig,targets:&[AttentionAdapterTarget]) -> Self {
self.attention=self.attention.with_adapters(config,targets);self
}
}
impl<B:Backend,P:TransformerProjectionShape<B>+From<Linear<B>>> ProjectedEncoderDecoderLayer<B,P> {
pub fn from_dense(layer:DenseEncoderDecoderLayer<B>) -> Self {
Self::from_parts(AwqTransformerBlock::from_dense(layer.backbone),ProjectedCrossAttentionBlock::from_dense(layer.cross_attention))
}
}
impl<B:Backend,P:AdaptTransformerProjection<B>> ProjectedEncoderDecoderLayer<B,P> {
fn validate_adapters(&self,config:&DecoderLayerAdapterConfig) {
assert!(config.backbone.is_some() || config.cross_attention.is_some(),"selected paired layer requires an actual adapter stage");
if let Some(config)=&config.backbone {self.backbone.validate_adapters(&config.adapter,&config.attention,&config.feed_forward);}
if let Some(config)=&config.cross_attention {
assert!(!config.attention.is_empty(),"selected cross stage requires an actual projection target");
self.cross_attention.attention.validate_adapters(&config.adapter,&config.attention);
}
}
pub fn with_adapters(mut self,config:&DecoderLayerAdapterConfig) -> Self {
self.validate_adapters(config);
if let Some(config)=&config.backbone {self.backbone=self.backbone.with_adapters(&config.adapter,&config.attention,&config.feed_forward);}
if let Some(config)=&config.cross_attention {self.cross_attention=self.cross_attention.with_adapters(&config.adapter,&config.attention);}
self
}
}
impl<B:Backend,P:TransformerProjectionShape<B>+From<Linear<B>>> ProjectedEncoderDecoderStack<B,P> {
pub fn from_dense(stack:DenseEncoderDecoderStack<B>) -> Self {Self {layers:stack.layers.into_iter().map(ProjectedEncoderDecoderLayer::from_dense).collect()}}
}
impl<B:Backend,P:AdaptTransformerProjection<B>> ProjectedEncoderDecoderStack<B,P> {
fn selected_adapters<'a>(&self,targets:&'a [DecoderLayerAdapterConfig]) -> BTreeMap<usize,&'a DecoderLayerAdapterConfig> {
let mut selected=BTreeMap::new();
for target in targets {
assert!(target.layer<self.layers.len(),"paired adapter layer is outside the original decoder stack");
assert!(selected.insert(target.layer,target).is_none(),"duplicate paired adapter layer index");self.layers[target.layer].validate_adapters(target);
}
selected
}
pub fn with_adapters(self,targets:&[DecoderLayerAdapterConfig]) -> Self {
let selected=self.selected_adapters(targets);
let layers=self.layers.into_iter().enumerate().map(|(index,layer)| {
if let Some(config)=selected.get(&index) {layer.with_adapters(config)} else {layer}
}).collect();Self {layers}
}
}
impl<B:Backend,P:TransformerProjectionShape<B>+From<Linear<B>>> ProjectedEncoderDecoderModel<B,P> {
pub fn from_dense_parts(source_embeddings:TransformerEmbeddings<B>,encoder:DenseTransformerStack<B>,encoder_normalization:Option<DenseTransformerNorm<B>>,
target_embeddings:TransformerEmbeddings<B>,decoder:DenseEncoderDecoderStack<B>,decoder_normalization:Option<DenseTransformerNorm<B>>,head:TransformerHead<B>) -> Self {
let head=super::AwqTransformerHead::from_projection(head.projection.into(),head.normalization,head.dropout);
Self::from_parts(source_embeddings,AwqTransformerStack::from_dense(encoder),encoder_normalization,target_embeddings,
ProjectedEncoderDecoderStack::from_dense(decoder),decoder_normalization,head)
}
}
impl<B:Backend,P:AdaptTransformerProjection<B>> ProjectedEncoderDecoderModel<B,P> {
pub fn with_adapters(mut self,encoder:&[LayerAdapterConfig],decoder:&[DecoderLayerAdapterConfig],head:Option<&TransformerAdapterConfig>) -> Self {
let _=self.encoder.selected_adapters(encoder);let _=self.decoder.selected_adapters(decoder);
if let Some(config)=head {self.head.projection.validate_adapter(config);}
self.encoder=self.encoder.with_adapters(encoder);self.decoder=self.decoder.with_adapters(decoder);
if let Some(config)=head {self.head.projection=self.head.projection.with_adapter(config);}
self
}
}