use ruda_model::{config::Config,module::Module,tensor::{DType,Tensor,backend::Backend}};
use crate::{Linear,LoRALinear,LoRALinearConfig,Dropout,activation::Activation,
attention::{GroupedQueryAttention,DenseAttentionMask,DenseAttentionOptions,dense_scaled_dot_product_attention}};
use super::{DenseFeedForward,DenseTransformerBlock,DenseTransformerNorm};
use super::dense::residual_branch;
#[derive(Config,Debug)]
pub struct TransformerAdapterConfig {
pub lora: LoRALinearConfig,
pub adapter_dtype: Option<DType>,
#[config(default = false)]
pub use_rslora: bool,
}
#[derive(Module,Debug)]
pub enum AdaptedProjection<B: Backend> {
Dense(Linear<B>),
LoRA(LoRALinear<B>),
}
impl<B: Backend> AdaptedProjection<B> {
pub fn forward<const D: usize>(&self,input: Tensor<B,D>) -> Tensor<B,D> {
match self {Self::Dense(layer)=>layer.forward(input),Self::LoRA(layer)=>layer.forward(input)}
}
pub fn merge(self) -> Linear<B> {
match self {Self::Dense(layer)=>layer,Self::LoRA(layer)=>layer.merge()}
}
}
impl TransformerAdapterConfig {
fn wrap<B: Backend>(&self,base: Linear<B>,selected: bool) -> AdaptedProjection<B> {
if selected {
let dtype = self.adapter_dtype.unwrap_or_else(||base.weight.val().dtype());
AdaptedProjection::LoRA(self.lora.init_with_options(base,dtype,self.use_rslora))
} else { AdaptedProjection::Dense(base) }
}
}
#[derive(Config,Debug,Copy,PartialEq,Eq)]
pub enum AttentionAdapterTarget {
Query,
Key,
Value,
Output,
}
#[derive(Module,Debug)]
pub struct AdaptedGroupedQueryAttention<B: Backend> {
pub query: AdaptedProjection<B>,
pub key: AdaptedProjection<B>,
pub value: AdaptedProjection<B>,
pub output: AdaptedProjection<B>,
pub dropout: Dropout,
pub query_heads: usize,
pub kv_heads: usize,
pub head_dimension: usize,
}
impl<B: Backend> AdaptedGroupedQueryAttention<B> {
pub fn from_dense(base: GroupedQueryAttention<B>,config: &TransformerAdapterConfig,
targets: &[AttentionAdapterTarget]) -> Self {
for (i,target) in targets.iter().enumerate() {
assert!(!targets[..i].contains(target),"duplicate attention adapter target");
}
Self {query:config.wrap(base.query,targets.contains(&AttentionAdapterTarget::Query)),
key:config.wrap(base.key,targets.contains(&AttentionAdapterTarget::Key)),
value:config.wrap(base.value,targets.contains(&AttentionAdapterTarget::Value)),
output:config.wrap(base.output,targets.contains(&AttentionAdapterTarget::Output)),
dropout:base.dropout,query_heads:base.query_heads,kv_heads:base.kv_heads,head_dimension:base.head_dimension}
}
pub fn project(&self,query: Tensor<B,3>,key: Tensor<B,3>,value: Tensor<B,3>)
-> (Tensor<B,4>,Tensor<B,4>,Tensor<B,4>) {
let [batch,queries,_] = query.dims();
let [key_batch,keys,_] = key.dims();
let [value_batch,values,_] = value.dims();
assert_eq!((batch,keys),(key_batch,values),"adapted attention batches/key lengths differ");
assert_eq!(batch,value_batch,"adapted value batch differs");
(self.query.forward(query).reshape([batch,queries,self.query_heads,self.head_dimension]).swap_dims(1,2),
self.key.forward(key).reshape([batch,keys,self.kv_heads,self.head_dimension]).swap_dims(1,2),
self.value.forward(value).reshape([batch,keys,self.kv_heads,self.head_dimension]).swap_dims(1,2))
}
pub fn forward_projected(&self,query: Tensor<B,4>,key: Tensor<B,4>,value: Tensor<B,4>,
masks: DenseAttentionMask<B>,options: DenseAttentionOptions) -> Tensor<B,3> {
let [batch,heads,queries,width] = query.dims();
assert_eq!((heads,width),(self.query_heads,self.head_dimension),"adapted query head geometry differs");
assert_eq!((key.dims()[1],key.dims()[3]),(self.kv_heads,self.head_dimension),"adapted key head geometry differs");
assert_eq!((value.dims()[1],value.dims()[3]),(self.kv_heads,self.head_dimension),"adapted value head geometry differs");
let context = dense_scaled_dot_product_attention(query,key,value,masks,options,Some(&self.dropout));
self.output.forward(context.swap_dims(1,2).reshape([batch,queries,heads*width]))
}
pub fn forward(&self,query: Tensor<B,3>,key: Tensor<B,3>,value: Tensor<B,3>,
masks: DenseAttentionMask<B>,options: DenseAttentionOptions) -> Tensor<B,3> {
let (query,key,value) = self.project(query,key,value);
self.forward_projected(query,key,value,masks,options)
}
pub fn merge(self) -> GroupedQueryAttention<B> {
GroupedQueryAttention {query:self.query.merge(),key:self.key.merge(),value:self.value.merge(),output:self.output.merge(),
dropout:self.dropout,query_heads:self.query_heads,kv_heads:self.kv_heads,head_dimension:self.head_dimension}
}
}
#[derive(Config,Debug,Copy,PartialEq,Eq)]
pub enum FeedForwardAdapterTarget {
Up,
Gate,
Down,
}
#[derive(Module,Debug)]
pub struct AdaptedFeedForward<B: Backend> {
pub up: AdaptedProjection<B>,
pub gate: Option<AdaptedProjection<B>>,
pub down: AdaptedProjection<B>,
pub activation: Activation<B>,
pub dropout: Dropout,
}
impl<B: Backend> AdaptedFeedForward<B> {
pub fn from_dense(base: DenseFeedForward<B>,config: &TransformerAdapterConfig,
targets: &[FeedForwardAdapterTarget]) -> Self {
for (i,target) in targets.iter().enumerate() {
assert!(!targets[..i].contains(target),"duplicate feed-forward adapter target");
}
assert!(base.gate.is_some() || !targets.contains(&FeedForwardAdapterTarget::Gate),"cannot adapt an absent gate projection");
Self {up:config.wrap(base.up,targets.contains(&FeedForwardAdapterTarget::Up)),
gate:base.gate.map(|gate|config.wrap(gate,targets.contains(&FeedForwardAdapterTarget::Gate))),
down:config.wrap(base.down,targets.contains(&FeedForwardAdapterTarget::Down)),activation:base.activation,dropout:base.dropout}
}
pub fn forward<const D: usize>(&self,input: Tensor<B,D>) -> Tensor<B,D> {
let up = self.up.forward(input.clone());
let value = if let Some(gate) = &self.gate {
let activated = self.activation.forward(gate.forward(input));
assert_eq!(activated.dims(),up.dims(),"adapted gate activation must preserve intermediate geometry");
activated*up
}
else { self.activation.forward(up) };
self.down.forward(self.dropout.forward(value))
}
pub fn merge(self) -> DenseFeedForward<B> {
DenseFeedForward {up:self.up.merge(),gate:self.gate.map(|gate|gate.merge()),down:self.down.merge(),
activation:self.activation,dropout:self.dropout}
}
}
#[derive(Module,Debug)]
pub struct AdaptedTransformerBlock<B: Backend> {
pub attention: AdaptedGroupedQueryAttention<B>,
pub feed_forward: AdaptedFeedForward<B>,
pub attention_norm: DenseTransformerNorm<B>,
pub feed_forward_norm: DenseTransformerNorm<B>,
pub residual_dropout: Dropout,
pub norm_first: bool,
}
impl<B: Backend> AdaptedTransformerBlock<B> {
pub fn from_dense(base: DenseTransformerBlock<B>,config: &TransformerAdapterConfig,
attention_targets: &[AttentionAdapterTarget],feed_forward_targets: &[FeedForwardAdapterTarget]) -> Self {
assert!(!attention_targets.is_empty() || !feed_forward_targets.is_empty(),"at least one adapter target is required");
Self {attention:AdaptedGroupedQueryAttention::from_dense(base.attention,config,attention_targets),
feed_forward:AdaptedFeedForward::from_dense(base.feed_forward,config,feed_forward_targets),
attention_norm:base.attention_norm,feed_forward_norm:base.feed_forward_norm,
residual_dropout:base.residual_dropout,norm_first:base.norm_first}
}
pub fn forward(&self,input: Tensor<B,3>,masks: DenseAttentionMask<B>,options: DenseAttentionOptions) -> Tensor<B,3> {
self.forward_with_positions(input,masks,options,|query,key|(query,key))
}
pub fn forward_with_positions<F>(&self,input: Tensor<B,3>,masks: DenseAttentionMask<B>,
options: DenseAttentionOptions,positions: F) -> Tensor<B,3>
where F: FnOnce(Tensor<B,4>,Tensor<B,4>)->(Tensor<B,4>,Tensor<B,4>) {
self.forward_feed_forward(self.forward_attention_with_positions(input,masks,options,positions))
}
pub fn forward_attention_with_positions<F>(&self,input: Tensor<B,3>,masks: DenseAttentionMask<B>,
options: DenseAttentionOptions,positions: F) -> Tensor<B,3>
where F: FnOnce(Tensor<B,4>,Tensor<B,4>)->(Tensor<B,4>,Tensor<B,4>) {
residual_branch(input,&self.attention_norm,&self.residual_dropout,self.norm_first,|source| {
let (query,key,value) = self.attention.project(source.clone(),source.clone(),source);
let query_shape = query.dims();
let key_shape = key.dims();
let (query,key) = positions(query,key);
assert_eq!(query.dims(),query_shape,"adapted query position transform changed geometry");
assert_eq!(key.dims(),key_shape,"adapted key position transform changed geometry");
self.attention.forward_projected(query,key,value,masks,options)
})
}
pub fn forward_feed_forward(&self,hidden: Tensor<B,3>) -> Tensor<B,3> {
residual_branch(hidden,&self.feed_forward_norm,&self.residual_dropout,self.norm_first,
|source|self.feed_forward.forward(source))
}
pub fn merge(self) -> DenseTransformerBlock<B> {
DenseTransformerBlock {attention:self.attention.merge(),feed_forward:self.feed_forward.merge(),
attention_norm:self.attention_norm,feed_forward_norm:self.feed_forward_norm,
residual_dropout:self.residual_dropout,norm_first:self.norm_first}
}
}