use core::ops::Range;
use alloc::{collections::BTreeSet,vec::Vec};
use ruda_model::{module::{Module,ModuleVisitor,Param,ParamId},tensor::{DType,Tensor,backend::Backend}};
use crate::{Linear,activation::Activation,attention::GroupedQueryAttention,
transformer::{AdaptedProjection,AdaptedGroupedQueryAttention,DenseFeedForward,AdaptedFeedForward,
DenseTransformerBlock,AdaptedTransformerBlock,DenseTransformerStack,AdaptedTransformerStack,AdaptedStackLayer}};
use super::{ColumnParallelLinear,RowParallelLinear,TensorParallelGroupedQueryAttention,TensorParallelAdaptedGroupedQueryAttention,
TensorParallelFeedForward,TensorParallelAdaptedFeedForward,TensorParallelTransformerBlock,TensorParallelAdaptedTransformerBlock,
TensorParallelTransformerStack,TensorParallelAdaptedTransformerStack,TensorParallelAdaptedStackLayer};
mod decoder;
pub use decoder::*;
#[derive(Clone,Copy,Debug,PartialEq,Eq)]
pub enum TensorParallelProjectionAxis {
Column,
Row,
}
fn check_range(range: &Range<usize>,length: usize) {
assert!(range.start < range.end && range.end <= length,"parallel projection range must be a nonempty actual feature slice");
}
fn check_storage<B: Backend,const D: usize>(value: &Tensor<B,D>) {
assert!(!matches!(value.dtype(),DType::QFloat(_)),"packed quantized projections must be loaded as explicit native local shards");
}
fn partition_linear<B: Backend>(mut layer: Linear<B>,axis: TensorParallelProjectionAxis,range: Range<usize>) -> Linear<B> {
let shape = layer.weight.val().dims();let dimension = if axis == TensorParallelProjectionAxis::Column {1} else {0};
check_range(&range,shape[dimension]);
if let Some(bias) = &layer.bias {assert_eq!(bias.val().dims(),[shape[1]],"full projection output bias width differs");}
if range == (0..shape[dimension]) {return layer;}
check_storage(&layer.weight.val());
if axis == TensorParallelProjectionAxis::Column {if let Some(bias) = &layer.bias {check_storage(&bias.val());}}
layer.weight = layer.weight.map(|value| {
let trainable = value.is_require_grad();
value.slice_dim(dimension,range.clone()).detach().set_require_grad(trainable)
});
if axis == TensorParallelProjectionAxis::Column {
layer.bias = layer.bias.map(|bias|bias.map(|value| {
let trainable = value.is_require_grad();
value.slice_dim(0,range).detach().set_require_grad(trainable)
}));
}
layer
}
impl<B: Backend> ColumnParallelLinear<B> {
pub fn from_full(layer: Linear<B>,columns: Range<usize>) -> Self {
Self::from_shard(partition_linear(layer,TensorParallelProjectionAxis::Column,columns))
}
}
impl<B: Backend> RowParallelLinear<B> {
pub fn from_full(layer: Linear<B>,rows: Range<usize>) -> Self {
Self::from_shard(partition_linear(layer,TensorParallelProjectionAxis::Row,rows))
}
}
pub fn partition_parallel_projection<B: Backend>(layer: AdaptedProjection<B>,axis: TensorParallelProjectionAxis,range: Range<usize>) -> AdaptedProjection<B> {
match layer {
AdaptedProjection::Dense(layer)=>AdaptedProjection::Dense(partition_linear(layer,axis,range)),
AdaptedProjection::LoRA(mut layer)=>{
let [input,output] = layer.base.weight.val().dims();let rank = layer.adapter_a.weight.val().dims()[1];
assert_eq!(layer.adapter_a.weight.val().dims(),[input,rank],"full adapter A geometry differs");
assert_eq!(layer.adapter_b.weight.val().dims(),[rank,output],"full adapter B geometry differs");
assert!(rank > 0 && layer.scale.is_finite(),"invalid full adapter rank/scale");
layer.base = partition_linear(layer.base,axis,range.clone());
if axis == TensorParallelProjectionAxis::Column {layer.adapter_b = partition_linear(layer.adapter_b,axis,range);}
else {layer.adapter_a = partition_linear(layer.adapter_a,axis,range);}
AdaptedProjection::LoRA(layer)
}
}
}
#[derive(Clone,Debug,PartialEq,Eq)]
pub struct TensorParallelHeadPartition {
pub query: Range<usize>,
pub key_value: Range<usize>,
}
impl TensorParallelHeadPartition {
pub fn new(query: Range<usize>,key_value: Range<usize>) -> Self {Self {query,key_value}}
pub fn validate(&self,query_heads: usize,kv_heads: usize,head_dimension: usize) {
assert!(kv_heads > 0 && query_heads > 0 && head_dimension > 0 && query_heads.is_multiple_of(kv_heads),"invalid global GQA geometry");
check_range(&self.query,query_heads);check_range(&self.key_value,kv_heads);
let queries = self.query.end-self.query.start;let keys = self.key_value.end-self.key_value.start;
assert!(queries.is_multiple_of(keys),"local query/KV head counts do not form native GQA groups");
let global_group = query_heads/kv_heads;let local_group = queries/keys;
for key in 0..keys {
let first = self.query.start+key*local_group;let last = first+local_group-1;
assert!(first/global_group == self.key_value.start+key && last/global_group == self.key_value.start+key,
"local GQA partition changes the original query-to-KV head mapping");
}
self.query.end.checked_mul(head_dimension).expect("parallel full query feature range overflow");
self.key_value.end.checked_mul(head_dimension).expect("parallel full KV feature range overflow");
}
fn query_features(&self,width: usize) -> Range<usize> {self.query.start*width..self.query.end*width}
fn kv_features(&self,width: usize) -> Range<usize> {self.key_value.start*width..self.key_value.end*width}
}
fn check_independent<B: Backend>(layers: &[&Linear<B>]) {
let mut ids: Vec<ParamId> = Vec::new();
for layer in layers {
for id in core::iter::once(&layer.weight.id).chain(layer.bias.iter().map(|bias|&bias.id)) {
assert!(!ids.contains(id),"tied full projection parameters require explicit tie-aware local shards");
ids.push(id.clone());
}
}
}
fn adapted_linears<'a,B: Backend>(layers: &[&'a AdaptedProjection<B>]) -> Vec<&'a Linear<B>> {
let mut values = Vec::new();
for layer in layers {match layer {AdaptedProjection::Dense(layer)=>values.push(layer),
AdaptedProjection::LoRA(layer)=>{values.push(&layer.base);values.push(&layer.adapter_a);values.push(&layer.adapter_b);}}}
values
}
impl<B: Backend> TensorParallelGroupedQueryAttention<B> {
pub fn from_full(base: GroupedQueryAttention<B>,partition: &TensorParallelHeadPartition) -> Self {
let mut base = Self::from_shard(base).local;
partition.validate(base.query_heads,base.kv_heads,base.head_dimension);
check_independent(&[&base.query,&base.key,&base.value,&base.output]);
let query = partition.query_features(base.head_dimension);let key_value = partition.kv_features(base.head_dimension);
base.query = ColumnParallelLinear::from_full(base.query,query.clone()).local;
base.key = ColumnParallelLinear::from_full(base.key,key_value.clone()).local;
base.value = ColumnParallelLinear::from_full(base.value,key_value).local;
base.output = RowParallelLinear::from_full(base.output,query).local;
base.query_heads = partition.query.end-partition.query.start;base.kv_heads = partition.key_value.end-partition.key_value.start;
Self::from_shard(base)
}
}
impl<B: Backend> TensorParallelAdaptedGroupedQueryAttention<B> {
pub fn from_full(base: AdaptedGroupedQueryAttention<B>,partition: &TensorParallelHeadPartition) -> Self {
let mut base = Self::from_shard(base).local;
partition.validate(base.query_heads,base.kv_heads,base.head_dimension);
check_independent(&adapted_linears(&[&base.query,&base.key,&base.value,&base.output]));
let query = partition.query_features(base.head_dimension);let key_value = partition.kv_features(base.head_dimension);
base.query = partition_parallel_projection(base.query,TensorParallelProjectionAxis::Column,query.clone());
base.key = partition_parallel_projection(base.key,TensorParallelProjectionAxis::Column,key_value.clone());
base.value = partition_parallel_projection(base.value,TensorParallelProjectionAxis::Column,key_value);
base.output = partition_parallel_projection(base.output,TensorParallelProjectionAxis::Row,query);
base.query_heads = partition.query.end-partition.query.start;base.kv_heads = partition.key_value.end-partition.key_value.start;
Self::from_shard(base)
}
}
impl<B: Backend> TensorParallelFeedForward<B> {
pub fn from_full<F>(base: DenseFeedForward<B>,features: Range<usize>,activation: F) -> Self
where F: FnOnce(Activation<B>,Range<usize>)->Activation<B> {
let base = Self::from_shard(base).local;
let mut layers = alloc::vec![&base.up,&base.down];if let Some(gate) = &base.gate {layers.push(gate);}
check_independent(&layers);check_range(&features,base.up.weight.val().dims()[1]);
Self::from_shard(DenseFeedForward {up:ColumnParallelLinear::from_full(base.up,features.clone()).local,
gate:base.gate.map(|gate|ColumnParallelLinear::from_full(gate,features.clone()).local),
down:RowParallelLinear::from_full(base.down,features.clone()).local,activation:activation(base.activation,features),dropout:base.dropout})
}
}
impl<B: Backend> TensorParallelAdaptedFeedForward<B> {
pub fn from_full<F>(base: AdaptedFeedForward<B>,features: Range<usize>,activation: F) -> Self
where F: FnOnce(Activation<B>,Range<usize>)->Activation<B> {
let base = Self::from_shard(base).local;
let mut layers = alloc::vec![&base.up,&base.down];if let Some(gate) = &base.gate {layers.push(gate);}
check_independent(&adapted_linears(&layers));
let width = match &base.up {AdaptedProjection::Dense(layer)=>layer.weight.val().dims()[1],AdaptedProjection::LoRA(layer)=>layer.base.weight.val().dims()[1]};
check_range(&features,width);
Self::from_shard(AdaptedFeedForward {up:partition_parallel_projection(base.up,TensorParallelProjectionAxis::Column,features.clone()),
gate:base.gate.map(|gate|partition_parallel_projection(gate,TensorParallelProjectionAxis::Column,features.clone())),
down:partition_parallel_projection(base.down,TensorParallelProjectionAxis::Row,features.clone()),
activation:activation(base.activation,features),dropout:base.dropout})
}
}
#[derive(Clone,Debug,PartialEq,Eq)]
pub struct TensorParallelTransformerPartition {
pub attention: TensorParallelHeadPartition,
pub feed_forward: Range<usize>,
}
impl TensorParallelTransformerPartition {
pub fn new(attention: TensorParallelHeadPartition,feed_forward: Range<usize>) -> Self {Self {attention,feed_forward}}
}
struct IndependentParameters {ids: BTreeSet<ParamId>}
impl<B: Backend> ModuleVisitor<B> for IndependentParameters {
fn visit_float<const D: usize>(&mut self,param: &Param<Tensor<B,D>>) {
assert!(self.ids.insert(param.id),"tied full model parameters require explicit tie-aware local shards");
}
}
fn check_independent_module<B: Backend,M: Module<B>>(module: &M) {module.visit(&mut IndependentParameters {ids:BTreeSet::new()});}
impl<B: Backend> TensorParallelTransformerBlock<B> {
pub fn from_full_block<F>(block: DenseTransformerBlock<B>,partition: &TensorParallelTransformerPartition,activation: F) -> Self
where F: FnOnce(Activation<B>,Range<usize>)->Activation<B> {
check_independent_module(&block);
Self::from_sharded_block(DenseTransformerBlock {attention:TensorParallelGroupedQueryAttention::from_full(block.attention,&partition.attention).local,
feed_forward:TensorParallelFeedForward::from_full(block.feed_forward,partition.feed_forward.clone(),activation).local,
attention_norm:block.attention_norm,feed_forward_norm:block.feed_forward_norm,
residual_dropout:block.residual_dropout,norm_first:block.norm_first})
}
}
impl<B: Backend> TensorParallelAdaptedTransformerBlock<B> {
pub fn from_full_block<F>(block: AdaptedTransformerBlock<B>,partition: &TensorParallelTransformerPartition,activation: F) -> Self
where F: FnOnce(Activation<B>,Range<usize>)->Activation<B> {
check_independent_module(&block);
Self::from_sharded_block(AdaptedTransformerBlock {attention:TensorParallelAdaptedGroupedQueryAttention::from_full(block.attention,&partition.attention).local,
feed_forward:TensorParallelAdaptedFeedForward::from_full(block.feed_forward,partition.feed_forward.clone(),activation).local,
attention_norm:block.attention_norm,feed_forward_norm:block.feed_forward_norm,
residual_dropout:block.residual_dropout,norm_first:block.norm_first})
}
}
impl<B: Backend> TensorParallelAdaptedStackLayer<B> {
pub fn from_full_layer<F>(layer: AdaptedStackLayer<B>,partition: &TensorParallelTransformerPartition,activation: F) -> Self
where F: FnOnce(Activation<B>,Range<usize>)->Activation<B> {
match layer {AdaptedStackLayer::Dense(block)=>Self::Dense(TensorParallelTransformerBlock::from_full_block(block,partition,activation)),
AdaptedStackLayer::Adapted(block)=>Self::Adapted(TensorParallelAdaptedTransformerBlock::from_full_block(block,partition,activation))}
}
}
impl<B: Backend> TensorParallelTransformerStack<B> {
pub fn from_full_stack<F>(stack: DenseTransformerStack<B>,partitions: &[TensorParallelTransformerPartition],mut activation: F) -> Self
where F: FnMut(usize,Activation<B>,Range<usize>)->Activation<B> {
assert_eq!(partitions.len(),stack.blocks.len(),"native parallel plans must cover every actual layer exactly once");
check_independent_module(&stack);
Self::new(stack.blocks.into_iter().zip(partitions).enumerate().map(|(index,(block,partition))|
TensorParallelTransformerBlock::from_full_block(block,partition,|module,features|activation(index,module,features))).collect())
}
}
impl<B: Backend> TensorParallelAdaptedTransformerStack<B> {
pub fn from_full_stack<F>(stack: AdaptedTransformerStack<B>,partitions: &[TensorParallelTransformerPartition],mut activation: F) -> Self
where F: FnMut(usize,Activation<B>,Range<usize>)->Activation<B> {
assert_eq!(partitions.len(),stack.layers.len(),"native adapter parallel plans must cover every actual layer exactly once");
check_independent_module(&stack);
Self::new(stack.layers.into_iter().zip(partitions).enumerate().map(|(index,(layer,partition))|
TensorParallelAdaptedStackLayer::from_full_layer(layer,partition,|module,features|activation(index,module,features))).collect())
}
}