use super::*;
use crate::OptimizerTensorShard;
#[derive(Clone,Debug,PartialEq,Eq)]
pub struct LBFGSMasterTensorShard {
pub source_parameter:u64,
pub local_parameter:u64,
pub shard:Option<OptimizerTensorShard>,
}
fn axis_view(shape:&[usize],shard:&OptimizerTensorShard) -> Result<([usize;3],usize),LBFGSShardError> {
let product = |axes:&[usize]|axes.iter().try_fold(1usize,|total,axis|total.checked_mul(*axis))
.ok_or(LBFGSShardError::Shape("master axis view overflows"));
let before = product(&shape[..shard.axis])?;
let after = product(&shape[shard.axis+1..])?;
let length = before.checked_mul(shard.interval.len()).and_then(|length|length.checked_mul(after))
.ok_or(LBFGSShardError::Shape("master local axis view overflows"))?;
Ok(([before,shape[shard.axis],after],length))
}
impl<B:Backend> LBFGSFp32MasterState<B> {
pub fn partition_from_full(
&self,
placements:&[Vec<LBFGSMasterTensorShard>],
rank:u32,
) -> Result<Self,LBFGSShardError> {
self.validate()?;
if self.placement.is_some() {return Err(LBFGSShardError::Record);}
let full_master = self.master.as_ref().ok_or(LBFGSShardError::Record)?;
let world = u32::try_from(placements.len()).map_err(|_|LBFGSShardError::Layout("master destination rank count overflows"))?;
if world == 0 || rank >= world {return Err(LBFGSShardError::Layout("master destination rank is outside placements"));}
let mut originals = HashMap::new();
let mut offset = 0usize;
for parameter in &self.parameters {
let length = parameter_length(core::slice::from_ref(parameter))?;
let end = offset.checked_add(length).ok_or(LBFGSShardError::Shape("complete parameter interval overflows"))?;
originals.insert(parameter.id,(parameter,offset..end));offset = end;
}
let mut coverage:HashMap<u64,Vec<Option<(usize,core::ops::Range<usize>)>>> = HashMap::new();
let mut rank_parameters = Vec::new();
let mut lengths = Vec::new();
for placement in placements {
let mut parameters = Vec::new();
for entry in placement {
let (original,_) = originals.get(&entry.source_parameter).ok_or(LBFGSShardError::Record)?;
let shape = if let Some(shard) = &entry.shard {
shard.validate().map_err(|_|LBFGSShardError::Layout("invalid master tensor coordinates"))?;
if shard.global_shape != original.shape {return Err(LBFGSShardError::Shape("master tensor source axes"));}
axis_view(&original.shape,shard)?;
coverage.entry(entry.source_parameter).or_default().push(Some((shard.axis,shard.interval.clone())));
shard.local_shape().map_err(|_|LBFGSShardError::Shape("master local parameter axes"))?
} else {
coverage.entry(entry.source_parameter).or_default().push(None);
original.shape.clone()
};
parameters.push(LBFGSMasterParameter {id:entry.local_parameter,shape,storage:original.storage});
}
lengths.push(parameter_length(¶meters)?);
rank_parameters.push(parameters);
}
let layout = LBFGSShardLayout::new(lengths);
layout.validate(rank,world)?;
if layout.global_len()? != full_master.dims()[0] {return Err(LBFGSShardError::Shape("destination changes complete master length"));}
for original in &self.parameters {
let pieces = coverage.get_mut(&original.id).ok_or(LBFGSShardError::Layout("complete master parameter has no owner"))?;
if pieces.iter().any(Option::is_none) {
if pieces.len() != 1 {return Err(LBFGSShardError::Layout("complete parameter ownership overlaps another slice"));}
continue;
}
let axis = pieces[0].as_ref().expect("actual tensor interval").0;
if pieces.iter().any(|piece|piece.as_ref().expect("actual tensor interval").0 != axis) {
return Err(LBFGSShardError::Layout("one original tensor must use the same partition axis"));
}
pieces.sort_by_key(|piece|piece.as_ref().expect("actual tensor interval").1.start);
let mut cursor = 0usize;
for piece in pieces {
let interval = &piece.as_ref().expect("actual tensor interval").1;
if interval.start != cursor {return Err(LBFGSShardError::Layout("master tensor partition has overlap or missing coordinates"));}
cursor = interval.end;
}
if cursor != original.shape[axis] {return Err(LBFGSShardError::Layout("master tensor partition is incomplete"));}
}
let destination = &placements[rank as usize];
let parameters = rank_parameters.swap_remove(rank as usize);
let convert = |value:Tensor<B,1>| {
let mut pieces = Vec::new();
for entry in destination {
let (original,interval) = originals.get(&entry.source_parameter).expect("validated original parameter");
let parameter = if interval.start == 0 && interval.end == value.dims()[0] {
value.clone()
} else {value.clone().slice(interval.clone())};
let local = if let Some(shard) = &entry.shard {
if shard.interval == (0..original.shape[shard.axis]) {parameter} else {
let (shape,length) = axis_view(&original.shape,shard).expect("validated exact tensor axis view");
parameter.reshape(shape)
.slice_dim(1,shard.interval.clone()).reshape([length])
}
} else {parameter};
pieces.push(local);
}
if pieces.len() == 1 {pieces.pop().expect("one actual destination parameter")} else {Tensor::cat(pieces,0)}
};
let state = LBFGSFp32MasterState {
version:1,
parameters,
master:Some(convert(full_master.clone())),
optimizer:LBFGSState {
history_s:self.optimizer.history_s.iter().cloned().map(&convert).collect(),
history_y:self.optimizer.history_y.iter().cloned().map(&convert).collect(),
d:self.optimizer.d.clone().map(&convert),
t:self.optimizer.t,
prev_flat_grad:self.optimizer.prev_flat_grad.clone().map(&convert),
prev_loss:self.optimizer.prev_loss,
g_iter:self.optimizer.g_iter,
},
placement:Some((rank,layout)),
};
state.validate()?;Ok(state)
}
}