use super::*;
use crate::{attention::{DenseAttentionMask,DenseAttentionOptions,PackedSequenceLayout,PackedAttentionOptions},
cache::{TransformerKvCache,ProjectedKvCache},transformer::{DenseTransformerStack,AdaptedTransformerStack,DenseEncoderDecoderStack,AdaptedEncoderDecoderStack}};
use ruda_model::tensor::Bool;
#[derive(Module,Debug)]
pub struct FullyShardedTransformerStack<B:Backend> {
pub blocks:Vec<FullyShardedTransformerBlock<B>>,
}
#[derive(Module,Debug)]
pub struct FullyShardedEncoderDecoderStack<B:Backend> {
pub layers:Vec<FullyShardedEncoderDecoderLayer<B>>,
}
impl<B:Backend> ShardingContext<B> {
pub fn transformer_stack(&mut self,stack:DenseTransformerStack<B>) -> FullyShardedTransformerStack<B> {
FullyShardedTransformerStack {blocks:stack.blocks.into_iter().map(|block|self.transformer(block)).collect()}
}
pub fn adapted_transformer_stack(&mut self,stack:AdaptedTransformerStack<B>) -> FullyShardedTransformerStack<B> {
FullyShardedTransformerStack {blocks:stack.layers.into_iter().map(|block|self.transformer_choice(block)).collect()}
}
pub fn encoder_decoder_stack(&mut self,stack:DenseEncoderDecoderStack<B>) -> FullyShardedEncoderDecoderStack<B> {
FullyShardedEncoderDecoderStack {layers:stack.layers.into_iter().map(|layer|self.encoder_decoder_layer(layer)).collect()}
}
pub fn adapted_encoder_decoder_stack(&mut self,stack:AdaptedEncoderDecoderStack<B>) -> FullyShardedEncoderDecoderStack<B> {
FullyShardedEncoderDecoderStack {layers:stack.layers.into_iter().map(|layer|self.adapted_encoder_decoder_layer(layer)).collect()}
}
}
impl<B:Backend> FullyShardedTransformerStack<B> {
pub fn from_full(stack:DenseTransformerStack<B>,rank:usize,world:usize) -> Self {
ShardingContext::new(rank,world).transformer_stack(stack)
}
pub fn from_full_adapted(stack:AdaptedTransformerStack<B>,rank:usize,world:usize) -> Self {
ShardingContext::new(rank,world).adapted_transformer_stack(stack)
}
pub fn forward_with<E,F>(&self,mut input:Tensor<B,3>,mut layer:F) -> Result<Tensor<B,3>,E>
where F:FnMut(usize,&FullyShardedTransformerBlock<B>,Tensor<B,3>)->Result<Tensor<B,3>,E> {
for (index,block) in self.blocks.iter().enumerate() {input = layer(index,block,input)?;}
Ok(input)
}
pub fn forward_packed_with<E,F>(&self,mut input:Tensor<B,2>,mut layer:F) -> Result<Tensor<B,2>,E>
where F:FnMut(usize,&FullyShardedTransformerBlock<B>,Tensor<B,2>)->Result<Tensor<B,2>,E> {
for (index,block) in self.blocks.iter().enumerate() {input = layer(index,block,input)?;}
Ok(input)
}
pub fn new_kv_cache(&self,initial_capacity:usize) -> TransformerKvCache<B> {
TransformerKvCache::new(self.blocks.len(),initial_capacity)
}
pub fn forward_cached_inference_with<E,F>(&self,mut input:Tensor<B,3>,cache:&mut TransformerKvCache<B>,mut layer:F)
-> Result<Tensor<B,3>,E>
where F:FnMut(usize,&FullyShardedTransformerBlock<B>,Tensor<B,3>,&mut ProjectedKvCache<B>)->Result<Tensor<B,3>,E> {
cache.validate_layers(self.blocks.len());
let rows = (input.dims()[0],input.dims()[1]);
let next = cache.position().checked_add(rows.1).expect("fully sharded cached stack position overflow");
for (index,block) in self.blocks.iter().enumerate() {
input = layer(index,block,input,&mut cache.layers_mut()[index])?;
assert_eq!((input.dims()[0],input.dims()[1]),rows,"fully sharded cached layer changed actual chunk rows");
}
cache.finish_chunk(next);Ok(input)
}
}
impl<B:Backend> FullyShardedEncoderDecoderStack<B> {
pub fn from_full(stack:DenseEncoderDecoderStack<B>,rank:usize,world:usize) -> Self {
ShardingContext::new(rank,world).encoder_decoder_stack(stack)
}
pub fn from_full_adapted(stack:AdaptedEncoderDecoderStack<B>,rank:usize,world:usize) -> Self {
ShardingContext::new(rank,world).adapted_encoder_decoder_stack(stack)
}
pub fn forward_with<E,F>(&self,mut input:Tensor<B,3>,memory:Tensor<B,3>,mut layer:F) -> Result<Tensor<B,3>,E>
where F:FnMut(usize,&FullyShardedEncoderDecoderLayer<B>,Tensor<B,3>,Tensor<B,3>)->Result<Tensor<B,3>,E> {
for (index,block) in self.layers.iter().enumerate() {input = layer(index,block,input,memory.clone())?;}
Ok(input)
}
pub fn forward_packed_with<E,F>(&self,mut input:Tensor<B,2>,memory:Tensor<B,2>,mut layer:F) -> Result<Tensor<B,2>,E>
where F:FnMut(usize,&FullyShardedEncoderDecoderLayer<B>,Tensor<B,2>,Tensor<B,2>)->Result<Tensor<B,2>,E> {
for (index,block) in self.layers.iter().enumerate() {input = layer(index,block,input,memory.clone())?;}
Ok(input)
}
}
macro_rules! stack_execution {
($backend:ty,[$($generics:tt)*],$block_forward:ident,$block_packed:ident,$forward:ident,$packed:ident) => {
impl<$($generics)*> FullyShardedTransformerStack<$backend> {
pub fn $forward<C,F>(&self,input:Tensor<$backend,3>,masks:DenseAttentionMask<$backend>,options:DenseAttentionOptions,
communicator:C,mut positions:F) -> Result<Tensor<$backend,3>,C::Error>
where C:BroadcastTensorCollective<B>,F:FnMut(usize,Tensor<$backend,4>,Tensor<$backend,4>)->(Tensor<$backend,4>,Tensor<$backend,4>) {
self.forward_with(input,|index,block,input|block.$block_forward(input,masks.clone(),options,communicator.clone(),
|query,key|positions(index,query,key)))
}
pub fn $packed<C,F>(&self,input:Tensor<$backend,2>,layout:&PackedSequenceLayout,options:PackedAttentionOptions,
communicator:C,mut positions:F) -> Result<Tensor<$backend,2>,C::Error>
where C:BroadcastTensorCollective<B>,F:FnMut(usize,Tensor<$backend,3>,Tensor<$backend,3>)->(Tensor<$backend,3>,Tensor<$backend,3>) {
self.forward_packed_with(input,|index,block,input|block.$block_packed(input,layout,options,communicator.clone(),
|query,key|positions(index,query,key)))
}
}
impl<$($generics)*> FullyShardedEncoderDecoderStack<$backend> {
pub fn $forward<C,F,G>(&self,input:Tensor<$backend,3>,memory:Tensor<$backend,3>,self_masks:DenseAttentionMask<$backend>,
self_options:DenseAttentionOptions,cross_masks:DenseAttentionMask<$backend>,cross_options:DenseAttentionOptions,communicator:C,
mut self_positions:F,mut cross_positions:G) -> Result<Tensor<$backend,3>,C::Error>
where C:BroadcastTensorCollective<B>,F:FnMut(usize,Tensor<$backend,4>,Tensor<$backend,4>)->(Tensor<$backend,4>,Tensor<$backend,4>),
G:FnMut(usize,Tensor<$backend,4>,Tensor<$backend,4>)->(Tensor<$backend,4>,Tensor<$backend,4>) {
self.forward_with(input,memory,|index,block,input,memory|block.$block_forward(input,memory,self_masks.clone(),self_options,
cross_masks.clone(),cross_options,communicator.clone(),|query,key|self_positions(index,query,key),|query,key|cross_positions(index,query,key)))
}
pub fn $packed<C,F,G>(&self,input:Tensor<$backend,2>,memory:Tensor<$backend,2>,query_layout:&PackedSequenceLayout,memory_layout:&PackedSequenceLayout,
self_options:PackedAttentionOptions,cross_options:PackedAttentionOptions,communicator:C,mut self_positions:F,mut cross_positions:G)
-> Result<Tensor<$backend,2>,C::Error>
where C:BroadcastTensorCollective<B>,F:FnMut(usize,Tensor<$backend,3>,Tensor<$backend,3>)->(Tensor<$backend,3>,Tensor<$backend,3>),
G:FnMut(usize,Tensor<$backend,3>,Tensor<$backend,3>)->(Tensor<$backend,3>,Tensor<$backend,3>) {
self.forward_packed_with(input,memory,|index,block,input,memory|block.$block_packed(input,memory,query_layout,memory_layout,
self_options,cross_options,communicator.clone(),|query,key|self_positions(index,query,key),|query,key|cross_positions(index,query,key)))
}
}
};
}
stack_execution!(Autodiff<B,S>,[B:Backend,S:CheckpointStrategy],forward,forward_packed,forward,forward_packed);
stack_execution!(B,[B:Backend],forward_inference,forward_packed_inference,forward_inference,forward_packed_inference);
impl<B:Backend> FullyShardedTransformerStack<B> {
pub fn forward_cached_inference<C,F>(&self,input:Tensor<B,3>,visible:Option<Tensor<B,2,Bool>>,cache:&mut TransformerKvCache<B>,
masks:DenseAttentionMask<B>,options:DenseAttentionOptions,communicator:C,mut positions:F) -> Result<Tensor<B,3>,C::Error>
where C:BroadcastTensorCollective<B>,F:FnMut(usize,Tensor<B,4>,Tensor<B,4>,usize)->(Tensor<B,4>,Tensor<B,4>) {
self.forward_cached_inference_with(input,cache,|index,block,input,cache|
block.forward_cached_inference(input,visible.clone(),cache,masks.clone(),options,communicator.clone(),|query,key,offset|positions(index,query,key,offset)))
}
}