use super::*;
use alloc::{collections::BTreeMap,vec::Vec};
use core::ops::Range;
use crate::{activation::Activation,transformer::{TransformerEmbeddings,AdaptedTransformerStack,DenseTransformerStack}};
use ruda_model::module::{ModuleVisitor,Param,ParamId};
use super::super::{TensorParallelTransformerPartition,VocabParallelTransformerHead,VocabParallelAdaptedTransformerHead,
TensorParallelTransformerHead,TensorParallelAdaptedTransformerHead};
#[derive(Clone,Debug,PartialEq,Eq)]
pub struct TensorParallelTransformerModelPartition {
pub layers:Vec<TensorParallelTransformerPartition>,
pub input_vocabulary:VocabParallelLossLayout,
pub input_rank:usize,
pub output_vocabulary:VocabParallelLossLayout,
pub output_rank:usize,
pub padding_index:Option<usize>,
}
pub(super) fn head_weight<B:Backend>(head:&TensorParallelOutputHead<B>) -> ParamId {
match head {
TensorParallelOutputHead::Linear(head)=>head.local.projection.weight.id,
TensorParallelOutputHead::AdaptedLinear(head)=>head.local.projection.base.weight.id,
TensorParallelOutputHead::Vocabulary(head)=>head.projection.weight.id,
TensorParallelOutputHead::AdaptedVocabulary(head)=>head.projection.base.weight.id,
}
}
struct Sharing {counts:BTreeMap<ParamId,usize>,token:ParamId,tied:bool}
impl<B:Backend> ModuleVisitor<B> for Sharing {
fn visit_float<const D:usize>(&mut self,param:&Param<Tensor<B,D>>) {
let count = self.counts.entry(param.id).or_insert(0);*count += 1;
let limit = if self.tied && param.id == self.token {2} else {1};
assert!(*count <= limit,"other shared full model roles require explicit tie-aware local loading");
}
}
impl<B:Backend> TensorParallelOutputHead<B> {
pub fn from_full(self,layout:&VocabParallelLossLayout,rank:usize) -> Self {
match self {
Self::Linear(head)=>Self::Linear(TensorParallelTransformerHead::from_full(head.local,layout,rank)),
Self::AdaptedLinear(head)=>Self::AdaptedLinear(TensorParallelAdaptedTransformerHead::from_full(head.local,layout,rank)),
Self::Vocabulary(head)=>Self::Vocabulary(VocabParallelTransformerHead::from_full(head,layout,rank)),
Self::AdaptedVocabulary(head)=>Self::AdaptedVocabulary(VocabParallelAdaptedTransformerHead::from_full(head,layout,rank)),
}
}
}
impl<B:Backend> TensorParallelTransformerModel<B> {
pub fn from_full<F>(embeddings:TransformerEmbeddings<B>,backbone:AdaptedTransformerStack<B>,
final_normalization:Option<DenseTransformerNorm<B>>,head:TensorParallelOutputHead<B>,
partition:&TensorParallelTransformerModelPartition,activation:F) -> Self
where F:FnMut(usize,Activation<B>,Range<usize>)->Activation<B> {
assert_eq!(partition.layers.len(),backbone.layers.len(),"full parallel model plans must cover every actual layer exactly once");
partition.input_vocabulary.interval(partition.input_rank);partition.output_vocabulary.interval(partition.output_rank);
if let Some(padding) = partition.padding_index {assert!(padding < partition.input_vocabulary.vocabulary_size(),"full model padding row is outside real input classes");}
let token = embeddings.token.weight.id;let tied = token == head_weight(&head);
if tied {
assert!(matches!(&head,TensorParallelOutputHead::Vocabulary(_)|TensorParallelOutputHead::AdaptedVocabulary(_)),
"shared full column-major output roles require explicitly prepared compatible local weights");
assert!(partition.input_vocabulary == partition.output_vocabulary && partition.input_rank == partition.output_rank,
"tied native input/output storage must use the same actual rank/layout");
}
let mut sharing = Sharing {counts:BTreeMap::new(),token,tied};
embeddings.visit(&mut sharing);backbone.visit(&mut sharing);final_normalization.visit(&mut sharing);head.visit(&mut sharing);
let (embeddings,head) = match head {
TensorParallelOutputHead::Vocabulary(head) if tied=>{
let (embeddings,head) = TensorParallelTransformerEmbeddings::from_full_tied_head(embeddings,head,
&partition.input_vocabulary,partition.input_rank,partition.padding_index);
(embeddings,TensorParallelOutputHead::Vocabulary(head))
},
TensorParallelOutputHead::AdaptedVocabulary(head) if tied=>{
let (embeddings,head) = TensorParallelTransformerEmbeddings::from_full_tied_adapted_head(embeddings,head,
&partition.input_vocabulary,partition.input_rank,partition.padding_index);
(embeddings,TensorParallelOutputHead::AdaptedVocabulary(head))
},
head=>(TensorParallelTransformerEmbeddings::from_full(embeddings,&partition.input_vocabulary,partition.input_rank,partition.padding_index),
head.from_full(&partition.output_vocabulary,partition.output_rank)),
};
Self::from_parts(embeddings,TensorParallelAdaptedTransformerStack::from_full_stack(backbone,&partition.layers,activation),final_normalization,head)
}
pub fn from_full_dense<F>(embeddings:TransformerEmbeddings<B>,backbone:DenseTransformerStack<B>,
final_normalization:Option<DenseTransformerNorm<B>>,head:TensorParallelOutputHead<B>,
partition:&TensorParallelTransformerModelPartition,activation:F) -> Self
where F:FnMut(usize,Activation<B>,Range<usize>)->Activation<B> {
let layers = backbone.blocks.into_iter().map(crate::transformer::AdaptedStackLayer::Dense).collect();
Self::from_full(embeddings,AdaptedTransformerStack::new(layers),final_normalization,head,partition,activation)
}
}