use super::*;
use ruda_model::tensor::TensorData;
pub struct FullyShardedWeightedAccumulationContract<B:AutodiffBackend> {
continuation:FullyShardedAccumulationContract,
global_weight:Tensor<B::InnerBackend,1>,
}
impl<B:AutodiffBackend> Record<B> for FullyShardedWeightedAccumulationContract<B> {
type Item<S:PrecisionSettings>=(FullyShardedAccumulationContract,TensorData);
fn into_item<S:PrecisionSettings>(self) -> Self::Item<S> {(self.continuation,self.global_weight.into_data())}
fn from_item<S:PrecisionSettings>(item:Self::Item<S>,device:&B::Device) -> Self {
let dtype=item.0.state.dtype;
Self {continuation:item.0,global_weight:Tensor::from_data(item.1.convert_dtype(dtype),(device,dtype))}
}
}
#[derive(Clone)]
pub struct FullyShardedWeightedGradientsRecord<B:Backend> {
window:FullyShardedGradientsRecord,
global_weight:Tensor<B,1>,
}
impl<B:Backend> Record<B> for FullyShardedWeightedGradientsRecord<B> {
type Item<S:PrecisionSettings>=(<FullyShardedGradientsRecord as Record<B>>::Item<S>,TensorData);
fn into_item<S:PrecisionSettings>(self) -> Self::Item<S> {
(<FullyShardedGradientsRecord as Record<B>>::into_item::<S>(self.window),self.global_weight.into_data())
}
fn from_item<S:PrecisionSettings>(item:Self::Item<S>,device:&B::Device) -> Self {
let window=<FullyShardedGradientsRecord as Record<B>>::from_item::<S>(item.0,device);
let global_weight=Tensor::from_data(item.1.convert_dtype(window.state.dtype),(device,window.state.dtype));
Self {window,global_weight}
}
}
impl<B:Backend> FullyShardedWeightedGradientsRecord<B> {
pub fn state(&self) -> &FullyShardedAccumulationState {&self.window.state}
pub fn global_weight(&self) -> Tensor<B,1> {self.global_weight.clone()}
pub fn reshard(records:&[Self],target_rank:u32,target_world:u32,device:&B::Device) -> Result<Self,RecorderError> {
let invalid=|message:&str|RecorderError::Unknown(message.to_string());
let first=records.first().ok_or_else(||invalid("complete original weighted pending-gradient rank set required"))?;
let expected=first.global_weight.clone().into_data().iter::<f64>().next().ok_or_else(||invalid("actual global weight scalar required"))?;
if !expected.is_finite() || expected<0.0 {return Err(invalid("saved effective weight must be finite and nonnegative"));}
for record in records {
if record.global_weight.dims()!=[1] || record.global_weight.dtype()!=record.window.state.dtype {
return Err(invalid("saved global weight geometry/work precision differs"));
}
let weight=record.global_weight.clone().into_data().iter::<f64>().next().ok_or_else(||invalid("actual global weight scalar required"))?;
if weight!=expected || (record.window.state.microbatches==0 && weight!=0.0) {return Err(invalid("actual global weights differ across saved rank windows"));}
}
let windows=records.iter().map(|record|record.window.clone()).collect::<Vec<_>>();
let window=FullyShardedGradientsRecord::reshard::<B>(&windows,target_rank,target_world,device)?;
Ok(Self {window,global_weight:first.global_weight.clone().to_device(device)})
}
}
pub struct FullyShardedWeightedAccumulatedGradients<B:Backend> {
pub window:FullyShardedAccumulatedGradients,
pub global_weight:Tensor<B,1>,
}
pub struct FullyShardedWeightedGradientsAccumulator<M,B:AutodiffBackend> {
window:FullyShardedGradientsAccumulator<M>,
global_weight:Tensor<B::InnerBackend,1>,
}
impl<M:AutodiffModule<B>,B:AutodiffBackend> FullyShardedWeightedGradientsAccumulator<M,B> {
pub fn continuation(&self) -> FullyShardedWeightedAccumulationContract<B> {
FullyShardedWeightedAccumulationContract {continuation:self.window.continuation(),global_weight:self.global_weight.clone()}
}
pub(crate) fn inner(&self) -> &GradientsAccumulator<M> {self.window.inner()}
pub fn from_accumulator(module:&M,accumulator:GradientsAccumulator<M>,continuation:FullyShardedWeightedAccumulationContract<B>)
-> Result<Self,FullyShardedAccumulationError> {
if continuation.global_weight.dims()!=[1] || continuation.global_weight.dtype()!=continuation.continuation.state.dtype {
return Err(FullyShardedAccumulationError::State);
}
let window=FullyShardedGradientsAccumulator::from_accumulator::<B>(module,accumulator,continuation.continuation)?;
Ok(Self {window,global_weight:continuation.global_weight})
}
pub fn new<C:BroadcastTensorCollective<B::InnerBackend>>(module:&M,parameters:&[FullyShardedOptimizerParameter<C>],
dtype:FloatDType,loss_scale:f64,device:&B::Device) -> Result<Self,FullyShardedAccumulationError> {
let window=FullyShardedGradientsAccumulator::new::<B,C>(module,parameters,dtype,loss_scale)?;
Ok(Self {window,global_weight:Tensor::zeros([1],(device,DType::from(dtype)))})
}
pub fn state(&self) -> &FullyShardedAccumulationState {self.window.state()}
pub fn global_weight(&self) -> Tensor<B::InnerBackend,1> {self.global_weight.clone()}
pub fn pending(&self) -> &GradientsParams {self.window.pending()}
fn next_weight(&self,weight:Tensor<B::InnerBackend,1>) -> Result<Tensor<B::InnerBackend,1>,FullyShardedAccumulationError> {
if weight.dims()!=[1] || weight.device()!=self.global_weight.device() || !matches!(weight.dtype(),DType::F32|DType::F64) {
return Err(FullyShardedAccumulationError::Placement("actual global loss weight scalar/device/work storage differs"));
}
Ok(self.global_weight.clone()+weight.cast(self.window.state.dtype))
}
pub fn accumulate_sum(&mut self,module:&M,gradients:&GradientsParams,global_weight:Tensor<B::InnerBackend,1>,global_count:u64)
-> Result<(),FullyShardedAccumulationError> {
let weight=self.next_weight(global_weight)?;
self.window.accumulate_sum::<B>(module,gradients,global_count)?;self.global_weight=weight;Ok(())
}
pub fn backward_sum(&mut self,module:&M,loss_sum:Tensor<B,1>,global_weight:Tensor<B::InnerBackend,1>,global_count:u64)
-> Result<(),FullyShardedAccumulationError> {
let weight=self.next_weight(global_weight)?;
self.window.backward_sum::<B>(module,loss_sum,global_count)?;self.global_weight=weight;Ok(())
}
pub fn finish_sums(&mut self) -> FullyShardedWeightedAccumulatedGradients<B::InnerBackend> {
let global_weight=self.global_weight.clone();
let result=FullyShardedWeightedAccumulatedGradients {window:self.window.finish_sums(),global_weight};
self.global_weight=Tensor::zeros([1],(&self.global_weight.device(),self.global_weight.dtype()));result
}
pub fn finish_mean(&mut self,module:&M) -> Result<FullyShardedWeightedAccumulatedGradients<B::InnerBackend>,FullyShardedAccumulationError> {
if self.window.state.microbatches==0 {return Err(FullyShardedAccumulationError::EmptyWindow);}
inspect::<B,M>(module,&self.window.placement,true)?;let dtype=work_dtype(&self.window.state)?;
let mut gradients=self.window.pending().unscaled_for::<B,M>(module,self.window.state.loss_scale,dtype)?;
for id in gradients.container.ids().into_iter().copied().collect::<Vec<_>>() {
let gradient=gradients.remove::<B::InnerBackend,1>(id).ok_or(FullyShardedAccumulationError::State)?;
let weight=self.global_weight.clone().to_device(&gradient.device());let empty=weight.clone().equal_elem(0);
let shape=gradient.dims();
let gradient=gradient.mask_fill(empty.clone().expand(shape),0)/weight.mask_fill(empty,1);
gradients.register(id,gradient);
}
let mut result=self.finish_sums();result.window.gradients=gradients;Ok(result)
}
pub fn try_to_record(&self) -> Result<FullyShardedWeightedGradientsRecord<B::InnerBackend>,RecorderError> {
Ok(FullyShardedWeightedGradientsRecord {window:self.window.try_to_record::<B>()?,global_weight:self.global_weight.clone()})
}
pub async fn to_record_async(&self) -> Result<FullyShardedWeightedGradientsRecord<B::InnerBackend>,RecorderError> {
Ok(FullyShardedWeightedGradientsRecord {window:self.window.to_record_async::<B>().await?,global_weight:self.global_weight.clone()})
}
pub fn load_record(&mut self,module:&M,record:FullyShardedWeightedGradientsRecord<B::InnerBackend>,device:&B::Device) -> Result<(),RecorderError> {
if record.global_weight.dims()!=[1] || record.global_weight.dtype()!=self.window.state.dtype {
return Err(RecorderError::Unknown(FullyShardedAccumulationError::State.to_string()));
}
let weight=record.global_weight.to_device(&self.global_weight.device());
self.window.load_record::<B>(module,record.window,device)?;self.global_weight=weight;Ok(())
}
}