use ruda_model::{module::Module,tensor::{Tensor,Int,FloatDType,backend::Backend}};
use ruda_autodiff::{Autodiff,checkpoint::strategy::CheckpointStrategy};
use crate::{Embedding,Dropout,transformer::{TransformerEmbeddings,DenseTransformerNorm}};
use super::{VocabParallelEmbedding,VocabParallelLossLayout,BroadcastTensorCollective,VocabParallelTransformerHead,VocabParallelAdaptedTransformerHead,head::transformed};
#[derive(Module,Debug)]
pub struct TensorParallelTransformerEmbeddings<B:Backend> {
pub token:VocabParallelEmbedding<B>,
pub position:Option<Embedding<B>>,
pub token_type:Option<Embedding<B>>,
pub normalization:Option<DenseTransformerNorm<B>>,
pub dropout:Dropout,
}
fn independent_token<B:Backend>(base:&TransformerEmbeddings<B>,layout:&VocabParallelLossLayout,rank:usize) {
if layout.interval(rank) == (0..base.token.weight.val().dims()[0]) {return;}
for table in [&base.position,&base.token_type].into_iter().flatten() {
assert_ne!(base.token.weight.id,table.weight.id,"token/position/type sharing across different shard roles needs explicit compatible local loading");
}
}
impl<B:Backend> TensorParallelTransformerEmbeddings<B> {
pub fn from_tables(token:VocabParallelEmbedding<B>,position:Option<Embedding<B>>,token_type:Option<Embedding<B>>,
normalization:Option<DenseTransformerNorm<B>>,dropout:Dropout) -> Self {
let [rows,hidden] = token.local.weight.val().dims();assert!(rows > 0 && hidden > 0,"native parallel input token geometry must be positive");
for table in [&position,&token_type].into_iter().flatten() {
assert!(table.weight.val().dims()[0] > 0 && table.weight.val().dims()[1] == hidden,"parallel input table hidden width differs");
}
if let Some(norm) = &normalization {assert_eq!(norm.width(),hidden,"parallel input normalization width differs");}
assert!(dropout.prob.is_finite() && (0.0..=1.0).contains(&dropout.prob),"invalid parallel input dropout");
Self {token,position,token_type,normalization,dropout}
}
pub fn from_full(base:TransformerEmbeddings<B>,layout:&VocabParallelLossLayout,rank:usize,padding_index:Option<usize>) -> Self {
independent_token(&base,layout,rank);
Self::from_tables(VocabParallelEmbedding::from_full_with_layout(base.token,layout,rank,padding_index),
base.position,base.token_type,base.normalization,base.dropout)
}
pub fn from_full_tied_head(base:TransformerEmbeddings<B>,head:VocabParallelTransformerHead<B>,layout:&VocabParallelLossLayout,rank:usize,padding_index:Option<usize>)
-> (Self,VocabParallelTransformerHead<B>) {
independent_token(&base,layout,rank);
let (token,head) = VocabParallelTransformerHead::from_full_tied(base.token,head,layout,rank,padding_index);
(Self::from_tables(token,base.position,base.token_type,base.normalization,base.dropout),head)
}
pub fn from_full_tied_adapted_head(base:TransformerEmbeddings<B>,head:VocabParallelAdaptedTransformerHead<B>,layout:&VocabParallelLossLayout,rank:usize,padding_index:Option<usize>)
-> (Self,VocabParallelAdaptedTransformerHead<B>) {
independent_token(&base,layout,rank);
let (token,head) = VocabParallelAdaptedTransformerHead::from_full_tied(base.token,head,layout,rank,padding_index);
(Self::from_tables(token,base.position,base.token_type,base.normalization,base.dropout),head)
}
fn validate_inputs(&self,tokens:&Tensor<B,2,Int>,positions:&Option<Tensor<B,2,Int>>,types:&Option<Tensor<B,2,Int>>,compute:Option<FloatDType>) {
let shape = tokens.dims();let device = tokens.device();let storage = self.token.local.weight.val().dtype();
assert_eq!(self.token.local.weight.val().device(),device,"parallel token IDs/table must share the device");
for (table,ids) in [(&self.position,positions),(&self.token_type,types)] {
assert_eq!(table.is_some(),ids.is_some(),"parallel input table/ID presence differs");
if let (Some(table),Some(ids)) = (table,ids) {
assert_eq!(ids.dims(),shape,"parallel input metadata differs from actual token geometry");
assert_eq!(ids.device(),device,"parallel input IDs must share the device");
assert_eq!(table.weight.val().device(),device,"parallel input tables must share the token device");
if compute.is_none() {assert_eq!(table.weight.val().dtype(),storage,"mixed table storage requires an explicit compute dtype");}
}
}
}
fn finish<F>(&self,mut hidden:Tensor<B,3>,positions:Option<Tensor<B,2,Int>>,types:Option<Tensor<B,2,Int>>,compute:Option<FloatDType>,dropout:F) -> Tensor<B,3>
where F:FnOnce(&Dropout,Tensor<B,3>)->Tensor<B,3> {
if let Some(dtype) = compute {hidden = hidden.cast(dtype);}
for (table,ids) in [(&self.position,positions),(&self.token_type,types)] {
if let (Some(table),Some(ids)) = (table,ids) {
let rows = table.forward(ids);let rows = if let Some(dtype) = compute {rows.cast(dtype)} else {rows};hidden = hidden+rows;
}
}
if let Some(norm) = &self.normalization {hidden = norm.forward(hidden);}
transformed(&self.dropout,hidden,dropout)
}
pub fn forward_inference<C:BroadcastTensorCollective<B>>(&self,tokens:Tensor<B,2,Int>,positions:Option<Tensor<B,2,Int>>,types:Option<Tensor<B,2,Int>>,
communicator:C,layout:&VocabParallelLossLayout) -> Result<Tensor<B,3>,C::Error> {
self.validate_inputs(&tokens,&positions,&types,None);
self.token.forward_inference_with_layout(tokens,communicator,layout).map(|hidden|self.finish(hidden,positions,types,None,|dropout,input|dropout.forward(input)))
}
pub fn forward_inference_with_compute_dtype<C:BroadcastTensorCollective<B>>(&self,tokens:Tensor<B,2,Int>,positions:Option<Tensor<B,2,Int>>,types:Option<Tensor<B,2,Int>>,
communicator:C,layout:&VocabParallelLossLayout,compute:FloatDType,output:FloatDType) -> Result<Tensor<B,3>,C::Error> {
self.validate_inputs(&tokens,&positions,&types,Some(compute));
self.token.forward_inference_with_layout(tokens,communicator,layout).map(|hidden|self.finish(hidden,positions,types,Some(compute),|dropout,input|dropout.forward(input)).cast(output))
}
pub fn forward_packed_inference<C:BroadcastTensorCollective<B>>(&self,tokens:Tensor<B,1,Int>,positions:Option<Tensor<B,1,Int>>,types:Option<Tensor<B,1,Int>>,
communicator:C,layout:&VocabParallelLossLayout) -> Result<Tensor<B,2>,C::Error> {
let count = tokens.dims()[0];let convert = |ids:Tensor<B,1,Int>| {assert_eq!(ids.dims(),[count],"packed parallel input metadata length differs");ids.reshape([1,count])};
self.forward_inference(tokens.reshape([1,count]),positions.map(convert),types.map(convert),communicator,layout)
.map(|hidden| {let features = hidden.dims()[2];hidden.reshape([count,features])})
}
}
impl<B:Backend,S:CheckpointStrategy> TensorParallelTransformerEmbeddings<Autodiff<B,S>> {
pub fn forward<C:BroadcastTensorCollective<B>>(&self,tokens:Tensor<Autodiff<B,S>,2,Int>,positions:Option<Tensor<Autodiff<B,S>,2,Int>>,types:Option<Tensor<Autodiff<B,S>,2,Int>>,
communicator:C,layout:&VocabParallelLossLayout) -> Result<Tensor<Autodiff<B,S>,3>,C::Error> {
self.forward_with_dtype_and_dropout(tokens,positions,types,communicator,layout,None,None,|dropout,input|dropout.forward(input))
}
pub fn forward_with_dropout<C,F>(&self,tokens:Tensor<Autodiff<B,S>,2,Int>,positions:Option<Tensor<Autodiff<B,S>,2,Int>>,types:Option<Tensor<Autodiff<B,S>,2,Int>>,
communicator:C,layout:&VocabParallelLossLayout,dropout:F) -> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C:BroadcastTensorCollective<B>,F:FnOnce(&Dropout,Tensor<Autodiff<B,S>,3>)->Tensor<Autodiff<B,S>,3> {
self.forward_with_dtype_and_dropout(tokens,positions,types,communicator,layout,None,None,dropout)
}
pub fn forward_with_compute_dtype<C:BroadcastTensorCollective<B>>(&self,tokens:Tensor<Autodiff<B,S>,2,Int>,positions:Option<Tensor<Autodiff<B,S>,2,Int>>,types:Option<Tensor<Autodiff<B,S>,2,Int>>,
communicator:C,layout:&VocabParallelLossLayout,compute:FloatDType,output:FloatDType) -> Result<Tensor<Autodiff<B,S>,3>,C::Error> {
self.forward_with_dtype_and_dropout(tokens,positions,types,communicator,layout,Some(compute),Some(output),|dropout,input|dropout.forward(input))
}
pub fn forward_with_compute_dtype_and_dropout<C,F>(&self,tokens:Tensor<Autodiff<B,S>,2,Int>,positions:Option<Tensor<Autodiff<B,S>,2,Int>>,types:Option<Tensor<Autodiff<B,S>,2,Int>>,
communicator:C,layout:&VocabParallelLossLayout,compute:FloatDType,output:FloatDType,dropout:F) -> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C:BroadcastTensorCollective<B>,F:FnOnce(&Dropout,Tensor<Autodiff<B,S>,3>)->Tensor<Autodiff<B,S>,3> {
self.forward_with_dtype_and_dropout(tokens,positions,types,communicator,layout,Some(compute),Some(output),dropout)
}
fn forward_with_dtype_and_dropout<C,F>(&self,tokens:Tensor<Autodiff<B,S>,2,Int>,positions:Option<Tensor<Autodiff<B,S>,2,Int>>,types:Option<Tensor<Autodiff<B,S>,2,Int>>,
communicator:C,layout:&VocabParallelLossLayout,compute:Option<FloatDType>,output:Option<FloatDType>,dropout:F) -> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C:BroadcastTensorCollective<B>,F:FnOnce(&Dropout,Tensor<Autodiff<B,S>,3>)->Tensor<Autodiff<B,S>,3> {
self.validate_inputs(&tokens,&positions,&types,compute);
let hidden = self.token.forward_with_layout(tokens,communicator,layout)?;let hidden = self.finish(hidden,positions,types,compute,dropout);
Ok(if let Some(dtype) = output {hidden.cast(dtype)} else {hidden})
}
pub fn forward_packed<C:BroadcastTensorCollective<B>>(&self,tokens:Tensor<Autodiff<B,S>,1,Int>,positions:Option<Tensor<Autodiff<B,S>,1,Int>>,types:Option<Tensor<Autodiff<B,S>,1,Int>>,
communicator:C,layout:&VocabParallelLossLayout) -> Result<Tensor<Autodiff<B,S>,2>,C::Error> {
let count = tokens.dims()[0];let convert = |ids:Tensor<Autodiff<B,S>,1,Int>| {assert_eq!(ids.dims(),[count],"packed parallel input metadata length differs");ids.reshape([1,count])};
self.forward(tokens.reshape([1,count]),positions.map(convert),types.map(convert),communicator,layout)
.map(|hidden| {let features = hidden.dims()[2];hidden.reshape([count,features])})
}
}