use super::*;
#[derive(Clone,Debug,PartialEq,Eq)]
pub struct FullyShardedParameterPlacement {
pub parameter:ParamId,
pub logical_shape:Vec<usize>,
pub rank:u32,
pub world_size:u32,
}
impl FullyShardedParameterPlacement {
pub fn new(parameter:ParamId,logical_shape:Vec<usize>,rank:u32,world_size:u32) -> Self {
Self {parameter,logical_shape,rank,world_size}
}
pub fn from_shard<B,C>(binding:&FullyShardedOptimizerParameter<C>) -> Self
where B:AutodiffBackend,C:BroadcastTensorCollective<B::InnerBackend> {
Self::new(binding.parameter,binding.logical_shape.clone(),binding.communicator.rank(),binding.communicator.world_size())
}
}
impl<B:Backend> Record<B> for FullyShardedParameterPlacement {
type Item<P:PrecisionSettings>=(u64,Vec<usize>,u32,u32);
fn into_item<P:PrecisionSettings>(self) -> Self::Item<P> {
(self.parameter.val(),self.logical_shape,self.rank,self.world_size)
}
fn from_item<P:PrecisionSettings>(item:Self::Item<P>,_device:&B::Device) -> Self {
Self::new(ParamId::from(item.0),item.1,item.2,item.3)
}
}
fn metadata(placement:&Placement) -> Vec<FullyShardedParameterPlacement> {
placement.iter().map(|entry|FullyShardedParameterPlacement::new(ParamId::from(entry.0),entry.1.clone(),entry.2,entry.3)).collect()
}
impl FullyShardedAccumulationContract {
pub fn placements(&self) -> Vec<FullyShardedParameterPlacement> {metadata(&self.placement)}
}
impl FullyShardedGradientsRecord {
pub fn placements(&self) -> Vec<FullyShardedParameterPlacement> {metadata(&self.placement)}
}
impl<M> FullyShardedGradientsAccumulator<M> {
pub fn from_placements<B:AutodiffBackend>(module:&M,parameters:&[FullyShardedParameterPlacement],
dtype:FloatDType,loss_scale:f64) -> Result<Self,FullyShardedAccumulationError>
where M:AutodiffModule<B> {
validate_work_dtype(dtype)?;
let state=FullyShardedAccumulationState {dtype:dtype.into(),loss_scale,global_count:0,microbatches:0};
work_dtype(&state)?;
let mut placement=Vec::with_capacity(parameters.len());let mut ids=BTreeSet::new();
for parameter in parameters {
if !ids.insert(parameter.parameter) {
return Err(FullyShardedAccumulationError::Placement("duplicate canonical parameter binding"));
}
placement.push((parameter.parameter.val(),parameter.logical_shape.clone(),parameter.rank,parameter.world_size,DType::F32,false));
}
Self::bind_placement::<B>(module,placement,state)
}
pub fn placements(&self) -> Vec<FullyShardedParameterPlacement> {metadata(&self.placement)}
}
impl<M:AutodiffModule<B>,B:AutodiffBackend> FullyShardedWeightedGradientsAccumulator<M,B> {
pub fn from_placements(module:&M,parameters:&[FullyShardedParameterPlacement],dtype:FloatDType,
loss_scale:f64,device:&B::Device) -> Result<Self,FullyShardedAccumulationError> {
let window=FullyShardedGradientsAccumulator::from_placements::<B>(module,parameters,dtype,loss_scale)?;
Ok(Self::from_window(window,device))
}
}
impl HybridShardedGradientNormParameter {
pub fn from_placement(placement:&FullyShardedParameterPlacement,replicas:u32) -> Self {
Self::new(placement.parameter,placement.logical_shape.clone(),placement.rank,placement.world_size,replicas)
}
}