use ruda_model::{module::Module,tensor::{FloatDType,Int,Tensor,backend::Backend}};
use crate::{attention::{PackedSequenceLayout,DenseAttentionMask,DenseAttentionOptions,PackedDocumentAttentionMask,PackedAttentionOptions},
cache::{EncoderDecoderKvCache,ProjectedKvCache},loss::{CausalCrossEntropyConfig,CausalLoss}};
use super::{TransformerProjectionShape,TransformerProjection,TransformerEmbeddings,DenseTransformerNorm,
ProjectedTransformerBlock,ProjectedTransformerStack,ProjectedTransformerHead,ProjectedEncoderDecoderLayer,ProjectedEncoderDecoderStack};
#[derive(Clone,Debug)]
pub struct ProjectedTransformerInput<B:Backend,const D:usize=2> {
pub tokens:Tensor<B,D,Int>,
pub positions:Option<Tensor<B,D,Int>>,
pub token_types:Option<Tensor<B,D,Int>>,
pub embedding_dtypes:Option<(FloatDType,FloatDType)>,
}
impl<B:Backend,const D:usize> ProjectedTransformerInput<B,D> {
pub fn new(tokens:Tensor<B,D,Int>) -> Self {Self {tokens,positions:None,token_types:None,embedding_dtypes:None}}
}
pub(super) fn embed_projected<B:Backend>(embeddings:&TransformerEmbeddings<B>,input:ProjectedTransformerInput<B>) -> Tensor<B,3> {
match input.embedding_dtypes {
Some((compute,output))=>embeddings.forward_with_compute_dtype(input.tokens,input.positions,input.token_types,compute,output),
None=>embeddings.forward(input.tokens,input.positions,input.token_types),
}
}
pub(super) fn embed_packed_projected<B:Backend>(embeddings:&TransformerEmbeddings<B>,input:ProjectedTransformerInput<B,1>,layout:&PackedSequenceLayout) -> Tensor<B,2> {
let tokens=layout.tokens();assert_eq!(input.tokens.dims(),[tokens],"packed IDs/document boundaries differ");
for ids in input.positions.iter().chain(input.token_types.iter()) {assert_eq!(ids.dims(),[tokens],"packed optional ID rows differ");}
let input=ProjectedTransformerInput {tokens:input.tokens.reshape([1,tokens]),positions:input.positions.map(|ids|ids.reshape([1,tokens])),
token_types:input.token_types.map(|ids|ids.reshape([1,tokens])),embedding_dtypes:input.embedding_dtypes};
embed_projected(embeddings,input).reshape([tokens,embeddings.token.weight.val().dims()[1]])
}
fn norm<B:Backend,const D:usize>(hidden:Tensor<B,D>,normalization:&Option<DenseTransformerNorm<B>>) -> Tensor<B,D> {
if let Some(normalization)=normalization {normalization.forward(hidden)} else {hidden}
}
pub(super) fn check_block<B:Backend,P:TransformerProjectionShape<B>>(block:&ProjectedTransformerBlock<B,P>,width:usize) {
for projection in [&block.attention.query,&block.attention.key,&block.attention.value] {assert_eq!(projection.dimensions()[0],width,"original self-attention input width differs");}
assert_eq!(block.attention.output.dimensions()[1],width,"original attention residual width differs");
assert_eq!((block.feed_forward.up.dimensions()[0],block.feed_forward.down.dimensions()[1]),(width,width),"original FFN residual width differs");
assert_eq!((block.attention_norm.width(),block.feed_forward_norm.width()),(width,width),"original block norm width differs");
}
#[derive(Module,Debug)]
pub struct ProjectedEncoderDecoderModel<B:Backend,P:Module<B>> {
pub source_embeddings:TransformerEmbeddings<B>,
pub encoder:ProjectedTransformerStack<B,P>,
pub encoder_normalization:Option<DenseTransformerNorm<B>>,
pub target_embeddings:TransformerEmbeddings<B>,
pub decoder:ProjectedEncoderDecoderStack<B,P>,
pub decoder_normalization:Option<DenseTransformerNorm<B>>,
pub head:ProjectedTransformerHead<B,P>,
}
impl<B:Backend,P:TransformerProjectionShape<B>> ProjectedEncoderDecoderModel<B,P> {
pub fn from_parts(source_embeddings:TransformerEmbeddings<B>,encoder:ProjectedTransformerStack<B,P>,encoder_normalization:Option<DenseTransformerNorm<B>>,
target_embeddings:TransformerEmbeddings<B>,decoder:ProjectedEncoderDecoderStack<B,P>,decoder_normalization:Option<DenseTransformerNorm<B>>,
head:ProjectedTransformerHead<B,P>) -> Self {
let source=source_embeddings.token.weight.val().dims()[1];let target=target_embeddings.token.weight.val().dims()[1];
assert_eq!(head.projection.dimensions()[0],target,"original paired head width differs");
if let Some(norm)=&encoder_normalization {assert_eq!(norm.width(),source,"original encoder final norm width differs");}
if let Some(norm)=&decoder_normalization {assert_eq!(norm.width(),target,"original decoder final norm width differs");}
for block in &encoder.blocks {check_block(block,source);}
for layer in &decoder.layers {
check_block(&layer.backbone,target);let cross=&layer.cross_attention;
assert_eq!(cross.query_norm.width(),target,"original cross query norm width differs");
if let Some(norm)=&cross.memory_norm {assert_eq!(norm.width(),source,"original cross memory norm width differs");}
assert_eq!(cross.attention.query.dimensions()[0],target,"original cross query width differs");
assert_eq!(cross.attention.key.dimensions()[0],source,"original cross key width differs");
assert_eq!(cross.attention.value.dimensions()[0],source,"original cross value width differs");
assert_eq!(cross.attention.output.dimensions()[1],target,"original cross output width differs");
}
Self {source_embeddings,encoder,encoder_normalization,target_embeddings,decoder,decoder_normalization,head}
}
}
impl<B:Backend,P:TransformerProjection<B>> ProjectedEncoderDecoderModel<B,P> {
pub fn forward_with_positions<E,F,G>(&self,source:ProjectedTransformerInput<B>,target:ProjectedTransformerInput<B>,source_masks:DenseAttentionMask<B>,
source_options:DenseAttentionOptions,self_masks:DenseAttentionMask<B>,self_options:DenseAttentionOptions,cross_masks:DenseAttentionMask<B>,
cross_options:DenseAttentionOptions,mut source_positions:E,mut self_positions:F,mut cross_positions:G) -> Result<Tensor<B,3>,P::Error>
where E:FnMut(usize,Tensor<B,4>,Tensor<B,4>)->(Tensor<B,4>,Tensor<B,4>),F:FnMut(usize,Tensor<B,4>,Tensor<B,4>)->(Tensor<B,4>,Tensor<B,4>),
G:FnMut(usize,Tensor<B,4>,Tensor<B,4>)->(Tensor<B,4>,Tensor<B,4>) {
self.forward_with(source,target,|index,block,hidden|block.forward_with_positions(hidden,source_masks.clone(),source_options,|query,key|source_positions(index,query,key)),
|index,layer,hidden,memory|layer.forward_with_positions(hidden,memory,self_masks.clone(),self_options,cross_masks.clone(),cross_options,
|query,key|self_positions(index,query,key),|query,key|cross_positions(index,query,key)))
}
pub fn forward_causal_with_positions<E,F,G>(&self,source:ProjectedTransformerInput<B>,target:ProjectedTransformerInput<B>,labels:Tensor<B,2,Int>,
source_masks:DenseAttentionMask<B>,source_options:DenseAttentionOptions,self_masks:DenseAttentionMask<B>,self_options:DenseAttentionOptions,
cross_masks:DenseAttentionMask<B>,cross_options:DenseAttentionOptions,criterion:&CausalCrossEntropyConfig,label_smoothing:f64,
mut source_positions:E,mut self_positions:F,mut cross_positions:G) -> Result<CausalLoss<B>,P::Error>
where E:FnMut(usize,Tensor<B,4>,Tensor<B,4>)->(Tensor<B,4>,Tensor<B,4>),F:FnMut(usize,Tensor<B,4>,Tensor<B,4>)->(Tensor<B,4>,Tensor<B,4>),
G:FnMut(usize,Tensor<B,4>,Tensor<B,4>)->(Tensor<B,4>,Tensor<B,4>) {
self.forward_causal_with(source,target,labels,criterion,label_smoothing,
|index,block,hidden|block.forward_with_positions(hidden,source_masks.clone(),source_options,|query,key|source_positions(index,query,key)),
|index,layer,hidden,memory|layer.forward_with_positions(hidden,memory,self_masks.clone(),self_options,cross_masks.clone(),cross_options,
|query,key|self_positions(index,query,key),|query,key|cross_positions(index,query,key)))
}
pub fn forward_packed_with_positions<E,F,G>(&self,source:ProjectedTransformerInput<B,1>,target:ProjectedTransformerInput<B,1>,source_layout:&PackedSequenceLayout,
target_layout:&PackedSequenceLayout,source_masks:Option<&[PackedDocumentAttentionMask<B>]>,source_options:PackedAttentionOptions,
self_masks:Option<&[PackedDocumentAttentionMask<B>]>,self_options:PackedAttentionOptions,cross_masks:Option<&[PackedDocumentAttentionMask<B>]>,
cross_options:PackedAttentionOptions,mut source_positions:E,mut self_positions:F,mut cross_positions:G) -> Result<Tensor<B,2>,P::Error>
where E:FnMut(usize,Tensor<B,3>,Tensor<B,3>)->(Tensor<B,3>,Tensor<B,3>),F:FnMut(usize,Tensor<B,3>,Tensor<B,3>)->(Tensor<B,3>,Tensor<B,3>),
G:FnMut(usize,Tensor<B,3>,Tensor<B,3>)->(Tensor<B,3>,Tensor<B,3>) {
self.forward_packed_with(source,target,source_layout,target_layout,
|index,block,hidden|block.forward_packed_with_positions(hidden,source_layout,source_masks,source_options,|query,key|source_positions(index,query,key)),
|index,layer,hidden,memory|layer.forward_packed_with_positions(hidden,memory,target_layout,source_layout,self_masks,self_options,cross_masks,cross_options,
|query,key|self_positions(index,query,key),|query,key|cross_positions(index,query,key)))
}
pub fn forward_packed_causal_with_positions<E,F,G>(&self,source:ProjectedTransformerInput<B,1>,target:ProjectedTransformerInput<B,1>,labels:Tensor<B,1,Int>,
source_layout:&PackedSequenceLayout,target_layout:&PackedSequenceLayout,source_masks:Option<&[PackedDocumentAttentionMask<B>]>,source_options:PackedAttentionOptions,
self_masks:Option<&[PackedDocumentAttentionMask<B>]>,self_options:PackedAttentionOptions,cross_masks:Option<&[PackedDocumentAttentionMask<B>]>,cross_options:PackedAttentionOptions,
criterion:&CausalCrossEntropyConfig,label_smoothing:f64,mut source_positions:E,mut self_positions:F,mut cross_positions:G) -> Result<CausalLoss<B>,P::Error>
where E:FnMut(usize,Tensor<B,3>,Tensor<B,3>)->(Tensor<B,3>,Tensor<B,3>),F:FnMut(usize,Tensor<B,3>,Tensor<B,3>)->(Tensor<B,3>,Tensor<B,3>),
G:FnMut(usize,Tensor<B,3>,Tensor<B,3>)->(Tensor<B,3>,Tensor<B,3>) {
self.forward_packed_causal_with(source,target,labels,source_layout,target_layout,criterion,label_smoothing,
|index,block,hidden|block.forward_packed_with_positions(hidden,source_layout,source_masks,source_options,|query,key|source_positions(index,query,key)),
|index,layer,hidden,memory|layer.forward_packed_with_positions(hidden,memory,target_layout,source_layout,self_masks,self_options,cross_masks,cross_options,
|query,key|self_positions(index,query,key),|query,key|cross_positions(index,query,key)))
}
pub fn encode_with<E>(&self,input:ProjectedTransformerInput<B>,encoder:E) -> Result<Tensor<B,3>,P::Error>
where E:FnMut(usize,&ProjectedTransformerBlock<B,P>,Tensor<B,3>)->Result<Tensor<B,3>,P::Error> {
let rows=input.tokens.dims();let hidden=self.encoder.forward_with(embed_projected(&self.source_embeddings,input),encoder)?;
assert_eq!(hidden.dims(),[rows[0],rows[1],self.source_embeddings.token.weight.val().dims()[1]],"encoder changed original source rows");
Ok(norm(hidden,&self.encoder_normalization))
}
pub fn decode_hidden_with<F>(&self,input:ProjectedTransformerInput<B>,memory:Tensor<B,3>,decoder:F) -> Result<Tensor<B,3>,P::Error>
where F:FnMut(usize,&ProjectedEncoderDecoderLayer<B,P>,Tensor<B,3>,Tensor<B,3>)->Result<Tensor<B,3>,P::Error> {
let rows=input.tokens.dims();let width=self.target_embeddings.token.weight.val().dims()[1];
assert_eq!(memory.dims()[0],rows[0],"paired source/target batch rows differ");
assert_eq!(memory.dims()[2],self.source_embeddings.token.weight.val().dims()[1],"paired memory width differs");
assert_eq!(memory.device(),input.tokens.device(),"paired source/target devices differ");
let hidden=self.decoder.forward_with(embed_projected(&self.target_embeddings,input),memory,decoder)?;
assert_eq!(hidden.dims(),[rows[0],rows[1],width],"decoder changed original target rows");Ok(norm(hidden,&self.decoder_normalization))
}
pub fn forward_hidden_with<E,F>(&self,source:ProjectedTransformerInput<B>,target:ProjectedTransformerInput<B>,encoder:E,decoder:F) -> Result<Tensor<B,3>,P::Error>
where E:FnMut(usize,&ProjectedTransformerBlock<B,P>,Tensor<B,3>)->Result<Tensor<B,3>,P::Error>,
F:FnMut(usize,&ProjectedEncoderDecoderLayer<B,P>,Tensor<B,3>,Tensor<B,3>)->Result<Tensor<B,3>,P::Error> {
self.decode_hidden_with(target,self.encode_with(source,encoder)?,decoder)
}
pub fn forward_with<E,F>(&self,source:ProjectedTransformerInput<B>,target:ProjectedTransformerInput<B>,encoder:E,decoder:F) -> Result<Tensor<B,3>,P::Error>
where E:FnMut(usize,&ProjectedTransformerBlock<B,P>,Tensor<B,3>)->Result<Tensor<B,3>,P::Error>,
F:FnMut(usize,&ProjectedEncoderDecoderLayer<B,P>,Tensor<B,3>,Tensor<B,3>)->Result<Tensor<B,3>,P::Error> {
self.head.forward(self.forward_hidden_with(source,target,encoder,decoder)?)
}
pub fn forward_causal_with<E,F>(&self,source:ProjectedTransformerInput<B>,target:ProjectedTransformerInput<B>,labels:Tensor<B,2,Int>,
criterion:&CausalCrossEntropyConfig,label_smoothing:f64,encoder:E,decoder:F) -> Result<CausalLoss<B>,P::Error>
where E:FnMut(usize,&ProjectedTransformerBlock<B,P>,Tensor<B,3>)->Result<Tensor<B,3>,P::Error>,
F:FnMut(usize,&ProjectedEncoderDecoderLayer<B,P>,Tensor<B,3>,Tensor<B,3>)->Result<Tensor<B,3>,P::Error> {
let hidden=self.forward_hidden_with(source,target,encoder,decoder)?;
criterion.try_forward_hidden_with_smoothing(hidden,labels,|rows|self.head.forward(rows),label_smoothing)
}
pub fn forward_packed_hidden_with<E,F>(&self,source:ProjectedTransformerInput<B,1>,target:ProjectedTransformerInput<B,1>,source_layout:&PackedSequenceLayout,
target_layout:&PackedSequenceLayout,encoder:E,decoder:F) -> Result<Tensor<B,2>,P::Error>
where E:FnMut(usize,&ProjectedTransformerBlock<B,P>,Tensor<B,2>)->Result<Tensor<B,2>,P::Error>,
F:FnMut(usize,&ProjectedEncoderDecoderLayer<B,P>,Tensor<B,2>,Tensor<B,2>)->Result<Tensor<B,2>,P::Error> {
assert_eq!(source_layout.documents(),target_layout.documents(),"paired packed document counts differ");
assert_eq!(source.tokens.device(),target.tokens.device(),"paired packed devices differ");
let memory=self.encoder.forward_with(embed_packed_projected(&self.source_embeddings,source,source_layout),encoder)?;
assert_eq!(memory.dims(),[source_layout.tokens(),self.source_embeddings.token.weight.val().dims()[1]],"packed encoder changed source rows");
let memory=norm(memory,&self.encoder_normalization);let hidden=embed_packed_projected(&self.target_embeddings,target,target_layout);
let hidden=self.decoder.forward_packed_with(hidden,memory,target_layout,source_layout,decoder)?;
assert_eq!(hidden.dims(),[target_layout.tokens(),self.target_embeddings.token.weight.val().dims()[1]],"packed decoder changed target rows");
Ok(norm(hidden,&self.decoder_normalization))
}
pub fn forward_packed_causal_with<E,F>(&self,source:ProjectedTransformerInput<B,1>,target:ProjectedTransformerInput<B,1>,labels:Tensor<B,1,Int>,
source_layout:&PackedSequenceLayout,target_layout:&PackedSequenceLayout,criterion:&CausalCrossEntropyConfig,label_smoothing:f64,encoder:E,decoder:F)
-> Result<CausalLoss<B>,P::Error>
where E:FnMut(usize,&ProjectedTransformerBlock<B,P>,Tensor<B,2>)->Result<Tensor<B,2>,P::Error>,
F:FnMut(usize,&ProjectedEncoderDecoderLayer<B,P>,Tensor<B,2>,Tensor<B,2>)->Result<Tensor<B,2>,P::Error> {
let hidden=self.forward_packed_hidden_with(source,target,source_layout,target_layout,encoder,decoder)?;
criterion.try_forward_packed_hidden_with_smoothing(hidden,labels,target_layout,|rows|self.head.forward(rows),label_smoothing)
}
pub fn forward_packed_with<E,F>(&self,source:ProjectedTransformerInput<B,1>,target:ProjectedTransformerInput<B,1>,source_layout:&PackedSequenceLayout,
target_layout:&PackedSequenceLayout,encoder:E,decoder:F) -> Result<Tensor<B,2>,P::Error>
where E:FnMut(usize,&ProjectedTransformerBlock<B,P>,Tensor<B,2>)->Result<Tensor<B,2>,P::Error>,
F:FnMut(usize,&ProjectedEncoderDecoderLayer<B,P>,Tensor<B,2>,Tensor<B,2>)->Result<Tensor<B,2>,P::Error> {
self.head.forward(self.forward_packed_hidden_with(source,target,source_layout,target_layout,encoder,decoder)?)
}
pub fn decode_cached_with<F>(&self,input:ProjectedTransformerInput<B>,cache:&mut EncoderDecoderKvCache<B>,decoder:F) -> Result<Tensor<B,3>,P::Error>
where F:FnMut(usize,&ProjectedEncoderDecoderLayer<B,P>,Tensor<B,3>,&mut ProjectedKvCache<B>,&ProjectedKvCache<B>)->Result<Tensor<B,3>,P::Error> {
let hidden=self.decoder.forward_cached_with(embed_projected(&self.target_embeddings,input),cache,decoder)?;
self.head.forward(norm(hidden,&self.decoder_normalization))
}
}