use super::*;
#[derive(Clone,Debug)]
pub struct FullyShardedGradientMigration {
pub target:FullyShardedParameterPlacement,
pub source_records:Vec<usize>,
}
impl FullyShardedGradientMigration {
pub fn new(target:FullyShardedParameterPlacement,source_records:Vec<usize>) -> Self {Self {target,source_records}}
}
impl FullyShardedGradientsRecord {
pub fn reshard_explicit<B:Backend>(records:&[Self],migrations:&[FullyShardedGradientMigration],device:&B::Device)
-> Result<Self,RecorderError> {
let sources=records.iter().collect::<Vec<_>>();let migrated=migrate::<B>(&sources,migrations,device)?;
let gradients=migrated.gradients.try_to_record::<B>()?;Ok(migrated.into_record(gradients))
}
pub async fn reshard_explicit_async<B:Backend>(records:&[Self],migrations:&[FullyShardedGradientMigration],device:&B::Device)
-> Result<Self,RecorderError> {
let sources=records.iter().collect::<Vec<_>>();let migrated=migrate::<B>(&sources,migrations,device)?;
let gradients=migrated.gradients.to_record_async::<B>().await?;Ok(migrated.into_record(gradients))
}
}
impl<B:Backend> FullyShardedWeightedGradientsRecord<B> {
pub fn reshard_explicit(records:&[Self],migrations:&[FullyShardedGradientMigration],device:&B::Device)
-> Result<Self,RecorderError> {
let first=Self::validate_migration_weights(records)?;
let windows=records.iter().map(Self::window_record).collect::<Vec<_>>();
let migrated=migrate::<B>(&windows,migrations,device)?;
let gradients=migrated.gradients.try_to_record::<B>()?;let window=migrated.into_record(gradients);
Ok(Self::with_migrated_window(first,window,device))
}
pub async fn reshard_explicit_async(records:&[Self],migrations:&[FullyShardedGradientMigration],device:&B::Device)
-> Result<Self,RecorderError> {
let first=Self::validate_migration_weights_async(records).await?;
let windows=records.iter().map(Self::window_record).collect::<Vec<_>>();
let migrated=migrate::<B>(&windows,migrations,device)?;
let gradients=migrated.gradients.to_record_async::<B>().await?;let window=migrated.into_record(gradients);
Ok(Self::with_migrated_window(first,window,device))
}
}
struct NativeMigration {
gradients:GradientsParams,
state:FullyShardedAccumulationState,
placement:Placement,
}
impl NativeMigration {
fn into_record(self,gradients:GradientsParamsRecord) -> FullyShardedGradientsRecord {
FullyShardedGradientsRecord {version:1,gradients,state:self.state,placement:self.placement}
}
}
fn migrate<B:Backend>(records:&[&FullyShardedGradientsRecord],migrations:&[FullyShardedGradientMigration],device:&B::Device)
-> Result<NativeMigration,RecorderError> {
let invalid=|message:&str|RecorderError::Unknown(message.to_string());
let first=records.first().ok_or_else(||invalid("actual original pending-gradient archives required"))?;
work_dtype(&first.state).map_err(|error|invalid(&error.to_string()))?;
if first.state.microbatches==0 && first.state.global_count!=0 {return Err(invalid("invalid empty global pending-gradient window"));}
let mut indices=Vec::with_capacity(records.len());
for record in records {
if record.version!=1 || record.state!=first.state {return Err(invalid("pending-gradient global windows/options differ"));}
let index: BTreeMap<_,_>=record.placement.iter().map(|entry|(entry.0,entry)).collect();
if index.len()!=record.placement.len() {return Err(invalid("duplicate original pending-gradient parameter identity"));}
let ids=record.gradients.parameter_ids()?;
if ids.iter().any(|id|!index.contains_key(&id.val())) || (record.state.microbatches==0 && !ids.is_empty()) {
return Err(invalid("archived derivatives do not match the original pending-gradient window"));
}
indices.push(index);
}
let mut target=GradientsParams::new();let mut placement=Vec::with_capacity(migrations.len());let mut target_ids=BTreeSet::new();
for migration in migrations {
let binding=&migration.target;let id=binding.parameter.val();
if binding.world_size==0 || binding.rank>=binding.world_size || !target_ids.insert(id) {
return Err(invalid("invalid or duplicate explicit target pending-gradient ownership"));
}
let total=elements(&binding.logical_shape).map_err(|error|invalid(&error.to_string()))?;
let first_index=*migration.source_records.first().ok_or_else(||invalid("explicit complete original source group required"))?;
let reference=indices.get(first_index).and_then(|index|index.get(&id)).copied()
.ok_or_else(||invalid("selected archive does not contain the original parameter placement"))?;
let source_world=reference.3;
if source_world==0 || source_world as usize!=migration.source_records.len() || reference.1!=binding.logical_shape
|| !matches!(reference.4,DType::F32|DType::F16|DType::BF16) {
return Err(invalid("original source group size/logical axes/storage differ"));
}
let old_slots=total.div_ceil(source_world as usize);let new_slots=total.div_ceil(binding.world_size as usize);
if old_slots.checked_mul(source_world as usize).is_none() || new_slots.checked_mul(binding.world_size as usize).is_none() {
return Err(invalid("pending-gradient padded topology overflows"));
}
let start=binding.rank as usize*new_slots;let end=start.saturating_add(new_slots).min(total);
let mut value:Option<Tensor<B,1>>=None;
let mut ranks=BTreeSet::new();let mut present=None;
for source_index in &migration.source_records {
let source=records.get(*source_index).ok_or_else(||invalid("original source archive index out of range"))?;
let spec=indices[*source_index].get(&id).copied().ok_or_else(||invalid("missing explicitly selected original parameter placement"))?;
if spec.1!=reference.1 || spec.3!=source_world || spec.2>=source_world || spec.4!=reference.4 || spec.5!=reference.5
|| !ranks.insert(spec.2) {return Err(invalid("original logical ownership/storage/trainability differs or source rank repeats"));}
let archive=source.gradients.select_parameters(&[binding.parameter])?;
let metadata=archive.parameter_metadata(binding.parameter)?;
if present.is_some_and(|previous|previous!=metadata.is_some()) {return Err(invalid("original gradient presence differs across source data shards"));}
present=Some(metadata.is_some());
let Some((dtype,shape))=metadata else {continue;};
if !reference.5 || shape!=[old_slots] || dtype!=first.state.dtype {
return Err(invalid("actual original pending-gradient storage/precision/trainability differs"));
}
let mut selected=GradientsParams::from_record::<B>(archive,device)?;
let original=selected.remove::<B,1>(binding.parameter).ok_or_else(||invalid("missing restored original pending derivative"))?;
if value.is_none() {value=Some(Tensor::zeros([new_slots],(device,first.state.dtype)));}
let old_start=spec.2 as usize*old_slots;let old_end=old_start.saturating_add(old_slots).min(total);
let overlap_start=start.max(old_start);let overlap_end=end.min(old_end);
if overlap_start<overlap_end {
let target=value.take().ok_or_else(||invalid("missing target pending-gradient storage"))?;
value=Some(target.slice_assign([overlap_start-start..overlap_end-start],original.slice([overlap_start-old_start..overlap_end-old_start])));
}
}
placement.push((id,binding.logical_shape.clone(),binding.rank,binding.world_size,reference.4,reference.5));
if let Some(value)=value {target.register(binding.parameter,value);}
}
placement.sort_by_key(|entry|entry.0);
Ok(NativeMigration {gradients:target,state:first.state.clone(),placement})
}