use super::*;
use ruda_model::{module::{ModuleMapper,ModuleVisitor,ModuleDisplay},record::{Record,PrecisionSettings}};
use alloc::collections::BTreeSet;
use crate::hybrid_sharded::{FullyShardedColumnParallelLinear,FullyShardedRowParallelLinear,FullyShardedColumnParallelLoRA,
FullyShardedRowParallelLoRA,FullyShardedTensorParallelGatedMlp,FullyShardedVocabParallelEmbedding,FullyShardedVocabParallelProjection};
mod delta;
pub use delta::FullyShardedModuleDeltaRecord;
mod packed;
pub use packed::{FullyShardedPackedParameterRecord,FullyShardedStorageRecord,FullyShardedStorageDeltaRecord};
pub trait FullyShardedModule<B:Backend>:Module<B> {
fn visit_shards<F:FnMut(&ShardedParameter<B>)>(&self,visitor:&mut F);
fn visit_packed_shards<F:FnMut(&ShardedPackedParameter<B>)>(&self,_visitor:&mut F) {}
}
fn require_float_only<B:Backend,M:FullyShardedModule<B>>(module:&M) -> Result<(),FullyShardedParameterError> {
let mut packed=false;module.visit_packed_shards(&mut |_|packed=true);
if packed {return Err(FullyShardedParameterError::Geometry("packed modules require FullyShardedStorageRecord; floating-only records omit packed base words"));}
Ok(())
}
#[derive(Clone)]
pub struct FullyShardedModuleParameterRecord<B:Backend> {
version:u32,
parameters:Vec<FullyShardedParameterRecord<B>>,
}
fn schema<B:Backend,M:FullyShardedModule<B>>(module:&M) -> Result<BTreeMap<ParamId,ShardedParameter<B>>,FullyShardedParameterError> {
let mut shards=BTreeMap::<ParamId,ShardedParameter<B>>::new();let mut error=None;
module.visit_shards(&mut |parameter| {
if error.is_some() {return;}
if let Err(reason)=parameter.parameter_record() {error=Some(reason);return;}
if let Some(previous)=shards.get(¶meter.local.id) {
let first=previous.local.val();let current=parameter.local.val();
if previous.logical_shape!=parameter.logical_shape || previous.rank!=parameter.rank || previous.world_size!=parameter.world_size {
error=Some(FullyShardedParameterError::Record);
} else if first.dtype()!=current.dtype() {error=Some(FullyShardedParameterError::DType);
} else if first.is_require_grad()!=current.is_require_grad() {error=Some(FullyShardedParameterError::Trainability);
} else if first.device()!=current.device() {error=Some(FullyShardedParameterError::Device);}
} else {shards.insert(parameter.local.id,parameter.clone());}
});
if let Some(error)=error {return Err(error);}
struct Check<'a,B:Backend> {shards:&'a BTreeMap<ParamId,ShardedParameter<B>>,seen:BTreeSet<ParamId>,error:Option<FullyShardedParameterError>}
impl<B:Backend> ModuleVisitor<B> for Check<'_,B> {
fn visit_float<const D:usize>(&mut self,parameter:&Param<Tensor<B,D>>) {
let Some(shard)=self.shards.get(¶meter.id) else {self.error=Some(FullyShardedParameterError::Record);return;};
if D!=1 || parameter.val().shape().num_elements()!=shard.local.val().dims()[0] {
self.error=Some(FullyShardedParameterError::Geometry("module contains non-local parameter storage"));
}
if parameter.val().dtype()!=shard.local.val().dtype() {self.error=Some(FullyShardedParameterError::DType);}
if parameter.val().is_require_grad()!=shard.local.val().is_require_grad() {self.error=Some(FullyShardedParameterError::Trainability);}
if parameter.val().device()!=shard.local.val().device() {self.error=Some(FullyShardedParameterError::Device);}
self.seen.insert(parameter.id);
}
}
let mut check=Check {shards:&shards,seen:BTreeSet::new(),error:None};module.visit(&mut check);
if let Some(error)=check.error {return Err(error);}
if check.seen.len()!=shards.len() {return Err(FullyShardedParameterError::Record);}
Ok(shards)
}
fn restore_values<B:Backend,M:FullyShardedModule<B>>(module:M,values:BTreeMap<ParamId,Tensor<B,1>>) -> Result<M,FullyShardedParameterError> {
struct Restore<B:Backend> {values:BTreeMap<ParamId,Tensor<B,1>>,seen:BTreeSet<ParamId>}
impl<B:Backend> ModuleMapper<B> for Restore<B> {
fn map_float<const D:usize>(&mut self,parameter:Param<Tensor<B,D>>) -> Param<Tensor<B,D>> {
let Some(value)=self.values.get(¶meter.id).cloned() else {return parameter;};
assert_eq!(D,1,"validated sharded module leaf rank changed");self.seen.insert(parameter.id);
parameter.map(|_|Tensor::<B,D>::from_primitive(value.into_primitive()))
}
}
let mut restore=Restore {values,seen:BTreeSet::new()};let module=module.map(&mut restore);
if restore.seen.len()!=restore.values.len() {return Err(FullyShardedParameterError::Record);}
Ok(module)
}
impl<B:Backend> FullyShardedModuleParameterRecord<B> {
pub fn capture<M:FullyShardedModule<B>>(module:&M) -> Result<Self,FullyShardedParameterError> {
require_float_only(module)?;
let parameters=schema(module)?.into_values().map(|parameter|parameter.parameter_record()).collect::<Result<Vec<_>,_>>()?;
Ok(Self {version:1,parameters})
}
pub fn parameters(&self) -> &[FullyShardedParameterRecord<B>] {&self.parameters}
pub fn parameter_count(&self) -> usize {self.parameters.len()}
pub fn validate(&self) -> Result<(),FullyShardedParameterError> {
if self.version!=1 {return Err(FullyShardedParameterError::Record);}
let mut ids=BTreeSet::new();
for parameter in &self.parameters {
parameter.validate()?;
if !ids.insert(parameter.id()) {return Err(FullyShardedParameterError::Record);}
}
Ok(())
}
pub fn validate_for<M:FullyShardedModule<B>>(&self,module:&M) -> Result<(),FullyShardedParameterError> {
require_float_only(module)?;
self.validate_float_for(module)
}
fn validate_float_for<M:FullyShardedModule<B>>(&self,module:&M) -> Result<(),FullyShardedParameterError> {
self.validate()?;let shards=schema(module)?;
if shards.len()!=self.parameters.len() {return Err(FullyShardedParameterError::Record);}
for saved in &self.parameters {
let target=shards.get(&saved.id()).ok_or(FullyShardedParameterError::Record)?;
if target.logical_shape!=saved.logical_shape() || target.rank!=saved.rank() || target.world_size!=saved.world_size() {
return Err(FullyShardedParameterError::Record);
}
if target.local.val().dtype()!=saved.local().val().dtype() {return Err(FullyShardedParameterError::DType);}
if B::ad_enabled(&target.local.val().device()) && target.local.val().is_require_grad()!=saved.is_trainable() {
return Err(FullyShardedParameterError::Trainability);
}
}
Ok(())
}
pub fn restore_into<M:FullyShardedModule<B>>(self,module:M) -> Result<M,FullyShardedParameterError> {
self.validate_for(&module)?;let targets=schema(&module)?;
let mut values=BTreeMap::new();
for saved in self.parameters {
let original=targets.get(&saved.id()).ok_or(FullyShardedParameterError::Record)?.local.val();
let value=saved.local().val().to_device(&original.device()).detach().set_require_grad(original.is_require_grad());
values.insert(saved.id(),value);
}
restore_values(module,values)
}
pub fn repartition_from_ranks(sources:&[Self],rank:usize,world:usize) -> Result<Self,FullyShardedParameterError> {
if world==0 || rank>=world {return Err(FullyShardedParameterError::Geometry("valid destination rank/world is required"));}
let first=sources.first().ok_or(FullyShardedParameterError::Geometry("complete module rank set is empty"))?;
let mut ranks=Vec::with_capacity(sources.len());
for source in sources {
source.validate()?;
if source.parameters.len()!=first.parameters.len() {return Err(FullyShardedParameterError::Record);}
ranks.push(source.parameters.iter().map(|record|(record.id(),record)).collect::<BTreeMap<_,_>>());
}
let mut parameters=Vec::with_capacity(first.parameters.len());
for record in &first.parameters {
let records=ranks.iter().map(|source|source.get(&record.id()).ok_or(FullyShardedParameterError::Record)
.map(|saved|(*saved).clone())).collect::<Result<Vec<_>,_>>()?;
parameters.push(FullyShardedParameterRecord::repartition_from_ranks(&records,rank,world)?);
}
Ok(Self {version:1,parameters})
}
}
impl<B:Backend> Record<B> for FullyShardedModuleParameterRecord<B> {
type Item<P:PrecisionSettings>=(u32,Vec<<FullyShardedParameterRecord<B> as Record<B>>::Item<P>>);
fn into_item<P:PrecisionSettings>(self) -> Self::Item<P> {
(self.version,self.parameters.into_iter().map(|parameter|parameter.into_item::<P>()).collect())
}
fn from_item<P:PrecisionSettings>(item:Self::Item<P>,device:&B::Device) -> Self {
Self {version:item.0,parameters:item.1.into_iter().map(|parameter|FullyShardedParameterRecord::<B>::from_item::<P>(parameter,device)).collect()}
}
}
macro_rules! visit_fields {
($module:ident,$($field:ident),+ $(,)?) => {
impl<B:Backend> FullyShardedModule<B> for $module<B> {
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);)+}
}
};
}
impl<B:Backend> FullyShardedModule<B> for ShardedParameter<B> {
fn visit_shards<F:FnMut(&ShardedParameter<B>)>(&self,visitor:&mut F) {visitor(self);}
}
impl<B:Backend,M:FullyShardedModule<B>> FullyShardedModule<B> for Option<M> {
fn visit_shards<F:FnMut(&ShardedParameter<B>)>(&self,visitor:&mut F) {if let Some(module)=self {module.visit_shards(visitor);}}
fn visit_packed_shards<F:FnMut(&ShardedPackedParameter<B>)>(&self,visitor:&mut F) {if let Some(module)=self {module.visit_packed_shards(visitor);}}
}
impl<B:Backend,M:FullyShardedModule<B>> FullyShardedModule<B> for Vec<M> {
fn visit_shards<F:FnMut(&ShardedParameter<B>)>(&self,visitor:&mut F) {for module in self {module.visit_shards(visitor);}}
fn visit_packed_shards<F:FnMut(&ShardedPackedParameter<B>)>(&self,visitor:&mut F) {for module in self {module.visit_packed_shards(visitor);}}
}
visit_fields!(FullyShardedLinear,weight,bias);
visit_fields!(FullyShardedEmbedding,weight);
visit_fields!(FullyShardedProjection,weight,bias);
visit_fields!(FullyShardedLoRALinear,base,adapter_a,adapter_b);
visit_fields!(FullyShardedGatedMLP,gate,up,down);
visit_fields!(FullyShardedLayerNorm,gamma,beta);
visit_fields!(FullyShardedRmsNorm,gamma);
visit_fields!(FullyShardedPRelu,alpha);
visit_fields!(FullyShardedSwiGlu,inner,outer);
visit_fields!(FullyShardedGroupedQueryAttention,query,key,value,output);
visit_fields!(FullyShardedFeedForward,up,gate,down,activation);
visit_fields!(FullyShardedTransformerBlock,attention,feed_forward,attention_norm,feed_forward_norm);
visit_fields!(FullyShardedCrossAttentionBlock,attention,query_norm,memory_norm);
visit_fields!(FullyShardedEncoderDecoderLayer,backbone,cross_attention);
visit_fields!(FullyShardedTransformerStack,blocks);
visit_fields!(FullyShardedEncoderDecoderStack,layers);
visit_fields!(FullyShardedTransformerEmbeddings,token,position,token_type,normalization);
visit_fields!(FullyShardedTransformerHead,projection,normalization);
visit_fields!(FullyShardedTransformerModel,embeddings,backbone,normalization,head);
visit_fields!(FullyShardedEncoderDecoderModel,source_embeddings,encoder,encoder_normalization,target_embeddings,decoder,decoder_normalization,head);
visit_fields!(FullyShardedColumnParallelLinear,local);
visit_fields!(FullyShardedRowParallelLinear,local);
visit_fields!(FullyShardedColumnParallelLoRA,base,adapter_a,adapter_b);
visit_fields!(FullyShardedRowParallelLoRA,base,adapter_a,adapter_b);
visit_fields!(FullyShardedTensorParallelGatedMlp,gate,up,down);
visit_fields!(FullyShardedVocabParallelEmbedding,weight);
visit_fields!(FullyShardedVocabParallelProjection,weight,bias);
visit_fields!(FullyShardedAwqLinear,qweight,qzeros,scales,bias);
visit_fields!(FullyShardedAwqLoRALinear,base,adapter_a,adapter_b);
macro_rules! visit_projected_fields {
($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_projected_fields!(FullyShardedAwqAttention,query,key,value,output);
visit_projected_fields!(FullyShardedAwqFeedForward,up,gate,down,activation);
visit_projected_fields!(FullyShardedAwqTransformerBlock,attention,feed_forward,attention_norm,feed_forward_norm);
visit_projected_fields!(FullyShardedAwqTransformerStack,blocks);
visit_projected_fields!(FullyShardedAwqTransformerHead,projection,normalization);
visit_projected_fields!(FullyShardedAwqTransformerModel,embeddings,backbone,normalization,head);
visit_projected_fields!(FullyShardedProjectedCrossAttention,attention,query_norm,memory_norm);
visit_projected_fields!(FullyShardedProjectedDecoderLayer,backbone,cross_attention);
visit_projected_fields!(FullyShardedProjectedDecoderStack,layers);
visit_projected_fields!(FullyShardedProjectedEncoderDecoderModel,source_embeddings,encoder,encoder_normalization,target_embeddings,decoder,decoder_normalization,head);
visit_fields!(FullyShardedNf4Linear,packed,scales,codebook,bias);
visit_fields!(FullyShardedNf4LoRALinear,base,adapter_a,adapter_b);
visit_fields!(FullyShardedNativeSwiGluExperts,gate,up,down);
visit_projected_fields!(FullyShardedNativeMoeLayer,router,experts,correction_bias);
visit_projected_fields!(FullyShardedNativeMoeFeedForward,routed,shared);
visit_projected_fields!(FullyShardedNativeMoeTransformerBlock,attention,feed_forward,attention_norm,feed_forward_norm);
visit_projected_fields!(FullyShardedNativeMoeTransformerStack,layers);
visit_projected_fields!(FullyShardedNativeMoeTransformerModel,embeddings,backbone,normalization,head);
impl<B:Backend,P:FullyShardedModule<B>+ruda_model::module::ModuleDisplay> FullyShardedModule<B> for FullyShardedNativeMoeTransformerLayer<B,P> {
fn visit_shards<F:FnMut(&ShardedParameter<B>)>(&self,visitor:&mut F) {match self {Self::Dense(block)=>block.visit_shards(visitor),Self::Routed(block)=>block.visit_shards(visitor)}}
fn visit_packed_shards<F:FnMut(&ShardedPackedParameter<B>)>(&self,visitor:&mut F) {match self {Self::Dense(block)=>block.visit_packed_shards(visitor),Self::Routed(block)=>block.visit_packed_shards(visitor)}}
}
impl<B:Backend> FullyShardedModule<B> for ShardedPackedParameter<B> {
fn visit_shards<F:FnMut(&ShardedParameter<B>)>(&self,_visitor:&mut F) {}
fn visit_packed_shards<F:FnMut(&ShardedPackedParameter<B>)>(&self,visitor:&mut F) {visitor(self);}
}
macro_rules! visit_variants {
($module:ident,$($variant:ident),+ $(,)?) => {
impl<B:Backend> FullyShardedModule<B> for $module<B> {
fn visit_shards<F:FnMut(&ShardedParameter<B>)>(&self,visitor:&mut F) {match self {$(Self::$variant(module)=>module.visit_shards(visitor),)+}}
fn visit_packed_shards<F:FnMut(&ShardedPackedParameter<B>)>(&self,visitor:&mut F) {match self {$(Self::$variant(module)=>module.visit_packed_shards(visitor),)+}}
}
};
}
visit_variants!(FullyShardedAdaptedProjection,Dense,LoRA);
visit_variants!(FullyShardedAwqProjection,Dense,LoRA,Awq,AwqLoRA);
visit_variants!(FullyShardedNf4Projection,Dense,LoRA,Nf4,Nf4LoRA);
visit_variants!(FullyShardedMixedProjection,Dense,LoRA,Awq,AwqLoRA,Nf4,Nf4LoRA);
visit_variants!(FullyShardedTransformerNorm,Layer,Rms);
visit_variants!(FullyShardedHeadProjection,Column,RowMajor);
impl<B:Backend> FullyShardedModule<B> for FullyShardedActivation<B> {
fn visit_shards<F:FnMut(&ShardedParameter<B>)>(&self,visitor:&mut F) {
match self {Self::Stateless(_)=>{},Self::PRelu(module)=>module.visit_shards(visitor),Self::SwiGlu(module)=>module.visit_shards(visitor)}
}
}