use ruda_autodiff::{Autodiff, checkpoint::strategy::CheckpointStrategy, tensor_parallel as region};
use ruda_model::{module::{Module, Param}, tensor::{Tensor, Int, backend::Backend, module::linear, activation::silu}};
use region::BroadcastTensorCollective;
use super::{fully_sharded::{ShardedParameter, FullyShardedLinear, ShardingContext}, tensor_parallel};
impl<B:Backend> ShardingContext<B> {
pub fn tensor_column(&mut self,layer:tensor_parallel::ColumnParallelLinear<B>)->FullyShardedColumnParallelLinear<B> {
FullyShardedColumnParallelLinear{local:self.linear(layer.local)}
}
pub fn tensor_row(&mut self,layer:tensor_parallel::RowParallelLinear<B>)->FullyShardedRowParallelLinear<B> {
FullyShardedRowParallelLinear{local:self.linear(layer.local)}
}
pub fn tensor_column_lora(&mut self,layer:tensor_parallel::ColumnParallelLoRA<B>)->FullyShardedColumnParallelLoRA<B> {
FullyShardedColumnParallelLoRA{base:self.tensor_column(layer.base),
adapter_a:self.linear(layer.adapter_a),adapter_b:self.linear(layer.adapter_b),scale:layer.scale}
}
pub fn tensor_row_lora(&mut self,layer:tensor_parallel::RowParallelLoRA<B>)->FullyShardedRowParallelLoRA<B> {
FullyShardedRowParallelLoRA{base:self.tensor_row(layer.base),
adapter_a:self.linear(layer.adapter_a),adapter_b:self.linear(layer.adapter_b),scale:layer.scale}
}
pub fn tensor_gated_mlp(&mut self,layer:tensor_parallel::TensorParallelGatedMlp<B>)->FullyShardedTensorParallelGatedMlp<B> {
FullyShardedTensorParallelGatedMlp{gate:self.tensor_column(layer.gate),
up:self.tensor_column(layer.up),down:self.tensor_row(layer.down)}
}
pub fn tensor_embedding(&mut self,layer:tensor_parallel::VocabParallelEmbedding<B>)->FullyShardedVocabParallelEmbedding<B> {
FullyShardedVocabParallelEmbedding{weight:self.parameter(layer.local.weight),
vocabulary_start:layer.vocabulary_start,vocabulary_size:layer.vocabulary_size,padding_index:layer.padding_index}
}
}
#[derive(Debug)]
pub enum HybridParallelError<D, T> {
Data(D),
Tensor(T),
}
#[derive(Module, Debug)]
pub struct FullyShardedColumnParallelLinear<B: Backend> {
pub local: FullyShardedLinear<B>,
}
#[derive(Module, Debug)]
pub struct FullyShardedRowParallelLinear<B: Backend> {
pub local: FullyShardedLinear<B>,
}
impl<B: Backend> FullyShardedColumnParallelLinear<B> {
pub fn from_tensor_shard(layer: tensor_parallel::ColumnParallelLinear<B>, data_rank: usize, data_world: usize) -> Self {
ShardingContext::new(data_rank,data_world).tensor_column(layer)
}
}
impl<B: Backend> FullyShardedRowParallelLinear<B> {
pub fn from_tensor_shard(layer: tensor_parallel::RowParallelLinear<B>, data_rank: usize, data_world: usize) -> Self {
ShardingContext::new(data_rank,data_world).tensor_row(layer)
}
}
impl<B: Backend, S: CheckpointStrategy> FullyShardedColumnParallelLinear<Autodiff<B, S>> {
pub fn forward<C: BroadcastTensorCollective<B>, T: BroadcastTensorCollective<B>, const D: usize>(
&self, input: Tensor<Autodiff<B, S>, D>, data: C, tensor: T, gather_output: bool,
) -> Result<Tensor<Autodiff<B, S>, D>, HybridParallelError<C::Error, T::Error>> {
assert!(D > 0, "projection needs a feature axis");
let input = region::copy_to_region(input, tensor.clone()).map_err(HybridParallelError::Tensor)?;
let output = self.local.forward(input, data).map_err(HybridParallelError::Data)?;
if gather_output { region::gather_from_region(output, tensor, D - 1).map_err(HybridParallelError::Tensor) }
else { Ok(output) }
}
}
impl<B: Backend, S: CheckpointStrategy> FullyShardedRowParallelLinear<Autodiff<B, S>> {
pub fn forward<C: BroadcastTensorCollective<B>, T: BroadcastTensorCollective<B>, const D: usize>(
&self, input: Tensor<Autodiff<B, S>, D>, data: C, tensor: T, input_is_parallel: bool,
) -> Result<Tensor<Autodiff<B, S>, D>, HybridParallelError<C::Error, T::Error>> {
assert!(D > 0, "projection needs a feature axis");
let input = if input_is_parallel { input } else {
region::scatter_to_region(input, tensor.clone(), D - 1).map_err(HybridParallelError::Tensor)?
};
let weight = self.local.weight.gather::<C, 2>(data.clone()).map_err(HybridParallelError::Data)?;
let partial = linear(input, weight, None);
let output = region::reduce_from_region(partial, tensor).map_err(HybridParallelError::Tensor)?;
Ok(match &self.local.bias {
Some(bias) => {
let bias = bias.gather::<C, 1>(data).map_err(HybridParallelError::Data)?;
let mut shape = [1; D];
shape[D - 1] = bias.dims()[0];
output + bias.reshape(shape)
}
None => output,
})
}
}
#[derive(Module, Debug)]
pub struct FullyShardedColumnParallelLoRA<B: Backend> {
pub base: FullyShardedColumnParallelLinear<B>,
pub adapter_a: FullyShardedLinear<B>,
pub adapter_b: FullyShardedLinear<B>,
pub scale: f64,
}
#[derive(Module, Debug)]
pub struct FullyShardedRowParallelLoRA<B: Backend> {
pub base: FullyShardedRowParallelLinear<B>,
pub adapter_a: FullyShardedLinear<B>,
pub adapter_b: FullyShardedLinear<B>,
pub scale: f64,
}
impl<B: Backend> FullyShardedColumnParallelLoRA<B> {
pub fn from_tensor_shard(layer: tensor_parallel::ColumnParallelLoRA<B>, rank: usize, world: usize) -> Self {
ShardingContext::new(rank,world).tensor_column_lora(layer)
}
}
impl<B: Backend> FullyShardedRowParallelLoRA<B> {
pub fn from_tensor_shard(layer: tensor_parallel::RowParallelLoRA<B>, rank: usize, world: usize) -> Self {
ShardingContext::new(rank,world).tensor_row_lora(layer)
}
}
impl<B: Backend, S: CheckpointStrategy> FullyShardedColumnParallelLoRA<Autodiff<B, S>> {
pub fn forward<C: BroadcastTensorCollective<B>, T: BroadcastTensorCollective<B>, const D: usize>(
&self, input: Tensor<Autodiff<B, S>, D>, adapter_input: Option<Tensor<Autodiff<B, S>, D>>,
data: C, tensor: T, gather_output: bool,
) -> Result<Tensor<Autodiff<B, S>, D>, HybridParallelError<C::Error, T::Error>> {
assert!(D > 0, "adapter needs a feature axis");
let adapted = adapter_input.unwrap_or_else(|| input.clone());
assert_eq!(adapted.dims(), input.dims(), "adapter input geometry differs");
let base = self.base.forward(input, data.clone(), tensor.clone(), false)?;
let dtype = base.dtype();
let a = self.adapter_a.weight.gather::<C, 2>(data.clone()).map_err(HybridParallelError::Data)?;
let adapted = region::copy_to_region(adapted.cast(a.dtype()), tensor.clone()).map_err(HybridParallelError::Tensor)?;
let a = region::copy_to_region(a, tensor.clone()).map_err(HybridParallelError::Tensor)?;
let hidden = linear(adapted, a, None).cast(self.adapter_b.weight.local.val().dtype());
let update = self.adapter_b.forward(hidden, data).map_err(HybridParallelError::Data)?;
let output = base + update.mul_scalar(self.scale).cast(dtype);
if gather_output { region::gather_from_region(output, tensor, D - 1).map_err(HybridParallelError::Tensor) }
else { Ok(output) }
}
}
impl<B: Backend, S: CheckpointStrategy> FullyShardedRowParallelLoRA<Autodiff<B, S>> {
pub fn forward<C: BroadcastTensorCollective<B>, T: BroadcastTensorCollective<B>, const D: usize>(
&self, input: Tensor<Autodiff<B, S>, D>, adapter_input: Option<Tensor<Autodiff<B, S>, D>>,
data: C, tensor: T, input_is_parallel: bool,
) -> Result<Tensor<Autodiff<B, S>, D>, HybridParallelError<C::Error, T::Error>> {
assert!(D > 0, "adapter needs a feature axis");
let adapted = adapter_input.unwrap_or_else(|| input.clone());
assert_eq!(adapted.dims(), input.dims(), "adapter input geometry differs");
let (input, adapted) = if input_is_parallel { (input, adapted) } else {
(region::scatter_to_region(input, tensor.clone(), D - 1).map_err(HybridParallelError::Tensor)?,
region::scatter_to_region(adapted, tensor.clone(), D - 1).map_err(HybridParallelError::Tensor)?)
};
let base = self.base.forward(input, data.clone(), tensor.clone(), true)?;
let dtype = base.dtype();
let adapted = adapted.cast(self.adapter_a.weight.local.val().dtype());
let hidden = self.adapter_a.forward(adapted, data.clone()).map_err(HybridParallelError::Data)?;
let hidden = region::reduce_from_region(hidden, tensor).map_err(HybridParallelError::Tensor)?;
let hidden = hidden.cast(self.adapter_b.weight.local.val().dtype());
let update = self.adapter_b.forward(hidden, data).map_err(HybridParallelError::Data)?;
Ok(base + update.mul_scalar(self.scale).cast(dtype))
}
}
#[derive(Module, Debug)]
pub struct FullyShardedTensorParallelGatedMlp<B: Backend> {
pub gate: FullyShardedColumnParallelLinear<B>,
pub up: FullyShardedColumnParallelLinear<B>,
pub down: FullyShardedRowParallelLinear<B>,
}
impl<B: Backend> FullyShardedTensorParallelGatedMlp<B> {
pub fn from_tensor_shard(layer: tensor_parallel::TensorParallelGatedMlp<B>, rank: usize, world: usize) -> Self {
ShardingContext::new(rank,world).tensor_gated_mlp(layer)
}
}
impl<B: Backend, S: CheckpointStrategy> FullyShardedTensorParallelGatedMlp<Autodiff<B, S>> {
pub fn forward<C: BroadcastTensorCollective<B>, T: BroadcastTensorCollective<B>, const D: usize>(
&self, input: Tensor<Autodiff<B, S>, D>, data: C, tensor: T,
) -> Result<Tensor<Autodiff<B, S>, D>, HybridParallelError<C::Error, T::Error>> {
self.forward_with(input, data, tensor, silu)
}
pub fn forward_with<C: BroadcastTensorCollective<B>, T: BroadcastTensorCollective<B>, F, const D: usize>(
&self, input: Tensor<Autodiff<B, S>, D>, data: C, tensor: T, activation: F,
) -> Result<Tensor<Autodiff<B, S>, D>, HybridParallelError<C::Error, T::Error>>
where F: FnOnce(Tensor<Autodiff<B, S>, D>) -> Tensor<Autodiff<B, S>, D> {
let gate = self.gate.forward(input.clone(), data.clone(), tensor.clone(), false)?;
let up = self.up.forward(input, data.clone(), tensor.clone(), false)?;
self.down.forward(activation(gate) * up, data, tensor, true)
}
}
#[derive(Module, Debug)]
pub struct FullyShardedVocabParallelEmbedding<B: Backend> {
pub weight: ShardedParameter<B>,
pub vocabulary_start: usize,
pub vocabulary_size: usize,
pub padding_index: Option<usize>,
}
impl<B: Backend> FullyShardedVocabParallelEmbedding<B> {
pub fn from_tensor_shard(layer: tensor_parallel::VocabParallelEmbedding<B>, rank: usize, world: usize) -> Self {
ShardingContext::new(rank,world).tensor_embedding(layer)
}
}
impl<B: Backend, S: CheckpointStrategy> FullyShardedVocabParallelEmbedding<Autodiff<B, S>> {
pub fn forward<C: BroadcastTensorCollective<B>, T: BroadcastTensorCollective<B>>(
&self, tokens: Tensor<Autodiff<B, S>, 2, Int>, data: C, tensor: T,
) -> Result<Tensor<Autodiff<B, S>, 3>, HybridParallelError<C::Error, T::Error>> {
let weight = self.weight.gather::<C, 2>(data).map_err(HybridParallelError::Data)?;
let layer = tensor_parallel::VocabParallelEmbedding {
local: crate::Embedding { weight: Param::initialized(self.weight.local.id, weight) },
vocabulary_start: self.vocabulary_start, vocabulary_size: self.vocabulary_size, padding_index: self.padding_index,
};
layer.forward(tokens, tensor).map_err(HybridParallelError::Tensor)
}
}
#[derive(Module, Debug)]
pub struct FullyShardedVocabParallelProjection<B: Backend> {
pub weight: ShardedParameter<B>,
pub bias: Option<ShardedParameter<B>>,
}
impl<B: Backend> FullyShardedVocabParallelProjection<B> {
pub fn from_embedding(embedding: &FullyShardedVocabParallelEmbedding<B>, bias: Option<ShardedParameter<B>>) -> Self {
if let Some(bias) = &bias {
assert_eq!(bias.logical_shape.as_slice(), &[embedding.weight.logical_shape[0]], "vocabulary bias width differs");
assert_eq!((bias.rank, bias.world_size), (embedding.weight.rank, embedding.weight.world_size), "bias data topology differs");
}
Self { weight: embedding.weight.clone(), bias }
}
}
impl<B: Backend, S: CheckpointStrategy> FullyShardedVocabParallelProjection<Autodiff<B, S>> {
pub fn forward<C: BroadcastTensorCollective<B>, T: BroadcastTensorCollective<B>, const D: usize>(
&self, input: Tensor<Autodiff<B, S>, D>, data: C, tensor: T, gather_output: bool,
) -> Result<Tensor<Autodiff<B, S>, D>, HybridParallelError<C::Error, T::Error>> {
assert!(D > 0, "projection needs a feature axis");
let weight = self.weight.gather::<C, 2>(data.clone()).map_err(HybridParallelError::Data)?;
let bias = match &self.bias {
Some(bias) => Some(Param::initialized(bias.local.id, bias.gather::<C, 1>(data).map_err(HybridParallelError::Data)?)),
None => None,
};
let layer = tensor_parallel::VocabParallelProjection { weight: Param::initialized(self.weight.local.id, weight), bias };
layer.forward(input, tensor, gather_output).map_err(HybridParallelError::Tensor)
}
}