use super::*;
use ruda_model::tensor::TensorPrimitive;
#[derive(Clone,Debug,PartialEq,Eq)]
pub struct HybridShardedGradientNormParameter {
pub parameter:ParamId,
pub logical_shape:Vec<usize>,
pub shard_rank:u32,
pub shard_world:u32,
pub replicas:u32,
}
impl HybridShardedGradientNormParameter {
pub fn new(parameter:ParamId,logical_shape:Vec<usize>,shard_rank:u32,shard_world:u32,replicas:u32) -> Self {
Self {parameter,logical_shape,shard_rank,shard_world,replicas}
}
pub fn from_shard<B,C>(binding:&FullyShardedOptimizerParameter<C>,replicas:u32) -> Self
where B:AutodiffBackend,C:BroadcastTensorCollective<B::InnerBackend> {
Self::new(binding.parameter,binding.logical_shape.clone(),binding.communicator.rank(),
binding.communicator.world_size(),replicas)
}
}
impl<B:Backend> Record<B> for HybridShardedGradientNormParameter {
type Item<P:PrecisionSettings>=(u64,Vec<usize>,u32,u32,u32);
fn into_item<P:PrecisionSettings>(self) -> Self::Item<P> {
(self.parameter.val(),self.logical_shape,self.shard_rank,self.shard_world,self.replicas)
}
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,item.4)
}
}
#[derive(Debug)]
pub struct HybridShardedGradientNorm<B:AutodiffBackend> {
norm:Tensor<B::InnerBackend,1>,
safe_maximum:Tensor<B::InnerBackend,1>,
scaled_root:Tensor<B::InnerBackend,1>,
dtype:FloatDType,
}
impl<B:AutodiffBackend> HybridShardedGradientNorm<B> {
pub fn norm(&self) -> Tensor<B::InnerBackend,1> {self.norm.clone()}
pub fn clipping_coefficient(&self,max_norm:f64,epsilon:f64)
-> Result<Tensor<B::InnerBackend,1>,FullyShardedAccumulationError> {
validate_limit(max_norm,epsilon,self.dtype)?;
let device=self.norm.device();let dtype=DType::from(self.dtype);
let limit=Tensor::<B::InnerBackend,1>::ones([1],(&device,dtype)).mul_scalar(max_norm)/self.safe_maximum.clone();
let offset=Tensor::<B::InnerBackend,1>::ones([1],(&device,dtype)).mul_scalar(epsilon)/self.safe_maximum.clone();
Ok((limit/(self.scaled_root.clone()+offset)).clamp_max(1))
}
}
pub fn measure_hybrid_sharded_gradient_norm<B,M,C>(module:&M,gradients:&GradientsParams,
parameters:&[HybridShardedGradientNormParameter],communicator:C,dtype:FloatDType,device:&B::Device)
-> Result<HybridShardedGradientNorm<B>,FullyShardedGradientNormError<C::Error>>
where B:AutodiffBackend,M:AutodiffModule<B>,C:BroadcastTensorCollective<B::InnerBackend> {
measure::<B,M,C>(module,gradients,parameters,&communicator,dtype,device).map(|(_,statistics)|statistics)
}
pub fn clip_hybrid_sharded_gradient_norm<B,M,C>(module:&M,gradients:&mut GradientsParams,
parameters:&[HybridShardedGradientNormParameter],communicator:C,max_norm:f64,epsilon:f64,
dtype:FloatDType,device:&B::Device)
-> Result<Tensor<B::InnerBackend,1>,FullyShardedGradientNormError<C::Error>>
where B:AutodiffBackend,M:AutodiffModule<B>,C:BroadcastTensorCollective<B::InnerBackend> {
validate_work_dtype(dtype).map_err(|error|FullyShardedGradientNormError::Arguments(error.into()))?;
validate_limit(max_norm,epsilon,dtype).map_err(FullyShardedGradientNormError::Arguments)?;
let (mut proposed,statistics)=measure::<B,M,C>(module,gradients,parameters,&communicator,dtype,device)?;
let coefficient=statistics.clipping_coefficient(max_norm,epsilon).map_err(FullyShardedGradientNormError::Arguments)?;
for id in proposed.container.ids().into_iter().copied().collect::<Vec<_>>() {
let value=proposed.remove::<B::InnerBackend,1>(id)
.ok_or(FullyShardedGradientNormError::Arguments(FullyShardedAccumulationError::State))?;
let device=value.device();proposed.register(id,value*coefficient.clone().to_device(&device));
}
*gradients=proposed;Ok(statistics.norm)
}
fn validate_limit(max_norm:f64,epsilon:f64,dtype:FloatDType) -> Result<(),FullyShardedAccumulationError> {
if !representable(max_norm,dtype) || max_norm<=0.0 || !representable(epsilon,dtype) || epsilon<0.0
|| (dtype==FloatDType::F32 && max_norm as f32==0.0) {
return Err(GradientTransformError::InvalidScalar.into());
}
Ok(())
}
fn measure<B,M,C>(module:&M,gradients:&GradientsParams,parameters:&[HybridShardedGradientNormParameter],
communicator:&C,dtype:FloatDType,device:&B::Device)
-> Result<(GradientsParams,HybridShardedGradientNorm<B>),FullyShardedGradientNormError<C::Error>>
where B:AutodiffBackend,M:AutodiffModule<B>,C:BroadcastTensorCollective<B::InnerBackend> {
let invalid=FullyShardedGradientNormError::Arguments;
validate_work_dtype(dtype).map_err(|error|invalid(error.into()))?;
let world=communicator.world_size();
if world==0 || communicator.rank()>=world {
return Err(FullyShardedGradientNormError::Protocol("invalid explicit hybrid norm coordinator"));
}
let mut declared=Vec::with_capacity(parameters.len());let mut copies=BTreeMap::new();
for parameter in parameters {
if parameter.replicas==0 || copies.insert(parameter.parameter,parameter.replicas).is_some() {
return Err(FullyShardedGradientNormError::Protocol("invalid replica multiplicity or duplicate local binding"));
}
declared.push((parameter.parameter.val(),parameter.logical_shape.clone(),parameter.shard_rank,
parameter.shard_world,DType::F32,false));
}
declared.sort_by_key(|entry|entry.0);
let placement=inspect::<B,M>(module,&declared,false).map_err(invalid)?;
let placement:BTreeMap<_,_>=placement.into_iter().map(|entry|(entry.0,entry)).collect();
let mut proposed=gradients.cast_for::<B,M>(module,dtype).map_err(|error|invalid(error.into()))?;
let ids=proposed.container.ids().into_iter().copied().collect::<Vec<_>>();
let mut local_maximum=Tensor::<B::InnerBackend,1>::zeros([1],(device,DType::from(dtype)));
for id in &ids {
let spec=placement.get(&id.val()).ok_or_else(||invalid(FullyShardedAccumulationError::State))?;
if !spec.5 {return Err(invalid(FullyShardedAccumulationError::Placement("frozen parameter has a supplied norm derivative")));}
let value=proposed.get::<B::InnerBackend,1>(*id).ok_or_else(||invalid(FullyShardedAccumulationError::State))?;
let total=elements(&spec.1).map_err(invalid)?;let slots=total.div_ceil(spec.3 as usize);
let real=total.saturating_sub(spec.2 as usize*slots).min(slots);
if real>0 {local_maximum=local_maximum.max_pair(value.clone().slice([0..real]).abs().max().to_device(device));}
if real<slots {
let zeros=Tensor::zeros([slots-real],(&value.device(),value.dtype()));
proposed.register(*id,value.slice_assign([real..slots],zeros));
}
}
let maximum=if world==1 {local_maximum} else {Tensor::<B::InnerBackend,1>::from_primitive(TensorPrimitive::Float(
communicator.all_gather_float(local_maximum.into_primitive().tensor()).map_err(FullyShardedGradientNormError::Collective)?))};
if maximum.dims()!=[world as usize] || maximum.dtype()!=DType::from(dtype) || maximum.device()!=*device {
return Err(FullyShardedGradientNormError::Protocol("hybrid norm maximum transport changed scalar storage/device"));
}
let maximum=maximum.max();let safe_maximum=maximum.clone().mask_fill(maximum.clone().equal_elem(0),1);
let mut square_sum=Tensor::<B::InnerBackend,1>::zeros([1],(device,DType::from(dtype)));
for id in &ids {
let value=proposed.get::<B::InnerBackend,1>(*id).ok_or_else(||invalid(FullyShardedAccumulationError::State))?;
let gradient_device=value.device();let value=value/safe_maximum.clone().to_device(&gradient_device);
let replicas=copies.get(id).ok_or_else(||invalid(FullyShardedAccumulationError::State))?;
square_sum=square_sum+value.square().sum().div_scalar(*replicas as f64).to_device(device);
}
let square_sum=if world==1 {square_sum} else {Tensor::<B::InnerBackend,1>::from_primitive(TensorPrimitive::Float(
communicator.all_reduce_sum(square_sum.into_primitive().tensor()).map_err(FullyShardedGradientNormError::Collective)?))};
if square_sum.dims()!=[1] || square_sum.dtype()!=DType::from(dtype) || square_sum.device()!=*device {
return Err(FullyShardedGradientNormError::Protocol("hybrid norm sum transport changed scalar storage/device"));
}
let scaled_root=square_sum.sqrt();let norm=maximum*scaled_root.clone();
Ok((proposed,HybridShardedGradientNorm {norm,safe_maximum,scaled_root,dtype}))
}