use super::*;
use ruda_model::{module::ModuleDisplay, tensor::IntegerTensorCollective};
use crate::{Mhc, attention::{LearnedKVCompressor, LightningIndexer, CompressedAttention,
CompressedAttentionParts, CompressedAttentionProjection, SparseRotaryEmbedding}};
#[derive(Module, Debug)]
pub struct FullyShardedKVCompressor<B: Backend, P: Module<B>> {
pub value: P,
pub gate: P,
pub position_bias: ShardedParameter<B>,
pub norm_weight: ShardedParameter<B>,
pub overlap: bool,
pub epsilon: f64,
}
#[derive(Module, Debug)]
pub struct FullyShardedLightningIndexer<B: Backend, P: Module<B>> {
pub query: P,
pub head_weight: P,
pub key: Option<P>,
pub key_norm: Option<FullyShardedLayerNorm<B>>,
pub rotary: SparseRotaryEmbedding,
pub topk: usize,
pub query_chunk_size: usize,
pub key_chunk_size: usize,
pub detach_inputs: bool,
}
#[derive(Module, Debug)]
pub struct FullyShardedCompressedAttention<B: Backend, P: Module<B>> {
pub query_down: P,
pub query_up: P,
pub query_norm: ShardedParameter<B>,
pub local_kv: P,
pub local_norm: ShardedParameter<B>,
pub compressor: FullyShardedKVCompressor<B, P>,
pub output_down: Vec<P>,
pub output_up: P,
pub sink: Option<ShardedParameter<B>>,
pub indexer: Option<FullyShardedLightningIndexer<B, P>>,
pub index_compressor: Option<FullyShardedKVCompressor<B, P>>,
pub rotary: SparseRotaryEmbedding,
pub width: usize,
pub window_size: usize,
pub query_chunk_size: usize,
pub epsilon: f64,
}
#[derive(Module, Debug)]
pub struct FullyShardedMhc<B: Backend> {
pub mapping: ShardedParameter<B>,
pub alpha: ShardedParameter<B>,
pub bias: ShardedParameter<B>,
pub width: usize,
pub streams: usize,
pub sinkhorn_iterations: usize,
pub epsilon: f64,
}
impl<B: Backend> ShardingContext<B> {
pub fn kv_compressor<P: CompressedAttentionProjection<B> + ShardTransformerProjection<B>>(&mut self,
source: LearnedKVCompressor<B, P>) -> FullyShardedKVCompressor<B, P::Sharded> {
FullyShardedKVCompressor { value: source.value.shard(self), gate: source.gate.shard(self),
position_bias: self.parameter(source.position_bias), norm_weight: self.parameter(source.norm_weight),
overlap: source.overlap, epsilon: source.epsilon }
}
pub fn lightning_indexer<P: CompressedAttentionProjection<B> + ShardTransformerProjection<B>>(&mut self,
source: LightningIndexer<B, P>) -> FullyShardedLightningIndexer<B, P::Sharded> {
FullyShardedLightningIndexer { query: source.query.shard(self), head_weight: source.head_weight.shard(self),
key: source.key.map(|key| key.shard(self)), key_norm: source.key_norm.map(|norm| self.layer_norm(norm)),
rotary: source.rotary, topk: source.topk, query_chunk_size: source.query_chunk_size,
key_chunk_size: source.key_chunk_size, detach_inputs: source.detach_inputs }
}
pub fn compressed_attention<P: CompressedAttentionProjection<B> + ShardTransformerProjection<B>>(&mut self,
source: CompressedAttention<B, P>) -> FullyShardedCompressedAttention<B, P::Sharded> {
let parts = source.parts;
FullyShardedCompressedAttention { query_down: parts.query_down.shard(self), query_up: parts.query_up.shard(self),
query_norm: self.parameter(parts.query_norm), local_kv: parts.local_kv.shard(self), local_norm: self.parameter(parts.local_norm),
compressor: self.kv_compressor(parts.compressor), output_down: parts.output_down.into_iter().map(|value| value.shard(self)).collect(),
output_up: parts.output_up.shard(self), sink: parts.sink.map(|value| self.parameter(value)),
indexer: parts.indexer.map(|value| self.lightning_indexer(value)),
index_compressor: parts.index_compressor.map(|value| self.kv_compressor(value)), rotary: source.rotary, width: source.width,
window_size: source.window_size, query_chunk_size: source.query_chunk_size, epsilon: source.epsilon }
}
pub fn mhc(&mut self, source: Mhc<B>) -> FullyShardedMhc<B> {
FullyShardedMhc { mapping: self.parameter(source.mapping), alpha: self.parameter(source.alpha), bias: self.parameter(source.bias),
width: source.width, streams: source.streams, sinkhorn_iterations: source.sinkhorn_iterations, epsilon: source.epsilon }
}
}
macro_rules! gather_compressed {
($backend:ty, [$($generics:tt)*], $gather:ident) => {
impl<$($generics)*, P: GatherTransformerProjection<$backend, B>> FullyShardedKVCompressor<$backend, P>
where P::Gathered: CompressedAttentionProjection<$backend> {
pub fn $gather<C: IntegerTensorCollective<B>>(&self, communicator: C) -> Result<LearnedKVCompressor<$backend, P::Gathered>, C::Error> {
Ok(LearnedKVCompressor::from_parts(self.value.gather_projection(communicator.clone())?, self.gate.gather_projection(communicator.clone())?,
Param::initialized(self.position_bias.local.id, self.position_bias.$gather::<C, 2>(communicator.clone())?),
Param::initialized(self.norm_weight.local.id, self.norm_weight.$gather::<C, 1>(communicator)?), self.overlap, self.epsilon))
}
}
impl<$($generics)*, P: GatherTransformerProjection<$backend, B>> FullyShardedLightningIndexer<$backend, P>
where P::Gathered: CompressedAttentionProjection<$backend> {
pub fn $gather<C: IntegerTensorCollective<B>>(&self, communicator: C) -> Result<LightningIndexer<$backend, P::Gathered>, C::Error> {
Ok(LightningIndexer::from_parts(self.query.gather_projection(communicator.clone())?, self.head_weight.gather_projection(communicator.clone())?,
self.key.as_ref().map(|key| key.gather_projection(communicator.clone())).transpose()?,
self.key_norm.as_ref().map(|norm| norm.$gather(communicator)).transpose()?, self.rotary.clone(), self.topk,
self.query_chunk_size, self.key_chunk_size, self.detach_inputs))
}
}
impl<$($generics)*, P: GatherTransformerProjection<$backend, B>> FullyShardedCompressedAttention<$backend, P>
where P::Gathered: CompressedAttentionProjection<$backend> {
pub fn $gather<C: IntegerTensorCollective<B>>(&self, communicator: C) -> Result<CompressedAttention<$backend, P::Gathered>, C::Error> {
let parts = CompressedAttentionParts { query_down: self.query_down.gather_projection(communicator.clone())?,
query_up: self.query_up.gather_projection(communicator.clone())?,
query_norm: Param::initialized(self.query_norm.local.id, self.query_norm.$gather::<C, 1>(communicator.clone())?),
local_kv: self.local_kv.gather_projection(communicator.clone())?,
local_norm: Param::initialized(self.local_norm.local.id, self.local_norm.$gather::<C, 1>(communicator.clone())?),
compressor: self.compressor.$gather(communicator.clone())?,
output_down: self.output_down.iter().map(|value| value.gather_projection(communicator.clone())).collect::<Result<_, _>>()?,
output_up: self.output_up.gather_projection(communicator.clone())?,
sink: self.sink.as_ref().map(|value| value.$gather::<C, 1>(communicator.clone()).map(|full| Param::initialized(value.local.id, full))).transpose()?,
indexer: self.indexer.as_ref().map(|value| value.$gather(communicator.clone())).transpose()?,
index_compressor: self.index_compressor.as_ref().map(|value| value.$gather(communicator)).transpose()? };
Ok(CompressedAttention::from_parts(parts, self.rotary.clone(), self.window_size, self.query_chunk_size, self.epsilon))
}
}
impl<$($generics)*> FullyShardedMhc<$backend> {
pub fn $gather<C: BroadcastTensorCollective<B>>(&self, communicator: C) -> Result<Mhc<$backend>, C::Error> {
Ok(Mhc { mapping: Param::initialized(self.mapping.local.id, self.mapping.$gather::<C, 2>(communicator.clone())?),
alpha: Param::initialized(self.alpha.local.id, self.alpha.$gather::<C, 1>(communicator.clone())?),
bias: Param::initialized(self.bias.local.id, self.bias.$gather::<C, 1>(communicator)?), width: self.width, streams: self.streams,
sinkhorn_iterations: self.sinkhorn_iterations, epsilon: self.epsilon })
}
}
};
}
gather_compressed!(B, [B: Backend], gather_inference);
gather_compressed!(Autodiff<B, S>, [B: Backend, S: CheckpointStrategy], gather);
macro_rules! visit_compressed {
($module:ident, [$($field:ident),+]) => {
impl<B: Backend, P: FullyShardedModule<B> + ModuleDisplay> FullyShardedModule<B> for $module<B, P> {
fn visit_shards<F: FnMut(&ShardedParameter<B>)>(&self, visitor: &mut F) { $(self.$field.visit_shards(visitor);)+ }
fn visit_packed_shards<F: FnMut(&ShardedPackedParameter<B>)>(&self, visitor: &mut F) { $(self.$field.visit_packed_shards(visitor);)+ }
}
};
}
visit_compressed!(FullyShardedKVCompressor, [value, gate, position_bias, norm_weight]);
visit_compressed!(FullyShardedLightningIndexer, [query, head_weight, key, key_norm]);
visit_compressed!(FullyShardedCompressedAttention, [query_down, query_up, query_norm, local_kv, local_norm, compressor,
output_down, output_up, sink, indexer, index_compressor]);
impl<B: Backend> FullyShardedModule<B> for FullyShardedMhc<B> {
fn visit_shards<F: FnMut(&ShardedParameter<B>)>(&self, visitor: &mut F) {
self.mapping.visit_shards(visitor); self.alpha.visit_shards(visitor); self.bias.visit_shards(visitor);
}
}
macro_rules! compressed_adapters {
($module:ident, [$($field:ident),+]) => {
impl<B: Backend, P: FullyShardedAdapterModule<B> + ModuleDisplay> FullyShardedAdapterModule<B> for $module<B, P> {
fn visit_adapter_shards<F: FnMut(&ShardedParameter<B>)>(&self, visitor: &mut F) { $(self.$field.visit_adapter_shards(visitor);)+ }
}
};
}
compressed_adapters!(FullyShardedKVCompressor, [value, gate]);
compressed_adapters!(FullyShardedLightningIndexer, [query, head_weight, key]);
compressed_adapters!(FullyShardedCompressedAttention, [query_down, query_up, local_kv, compressor, output_down, output_up, indexer, index_compressor]);