use super::{
AutodiffBackend, AutodiffModule, Backend, GradientsParams, LBFGS, LBFGSState,
LearningRate, Tensor, ToElement, flatten_params_inner,
reductions::VectorReductions,
};
use alloc::vec::Vec;
use core::{fmt, ops::Range};
use ruda_model::{
record::{PrecisionSettings, Record},
tensor::{BroadcastTensorCollective, TensorPrimitive},
};
use serde::{Deserialize, Serialize};
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct LBFGSShardLayout {
pub lengths: Vec<usize>,
}
impl LBFGSShardLayout {
pub fn new(lengths: Vec<usize>) -> Self {
Self { lengths }
}
pub fn validate(&self, rank: u32, world: u32) -> Result<(), LBFGSShardError> {
if world == 0 || rank >= world || self.lengths.len() != world as usize {
return Err(LBFGSShardError::Layout("rank/world does not match vector layout"));
}
if self.lengths.contains(&0) {
return Err(LBFGSShardError::Layout("each vector shard must be nonempty"));
}
self.global_len()?;
Ok(())
}
pub fn global_len(&self) -> Result<usize, LBFGSShardError> {
self.lengths.iter().try_fold(0usize, |total, length| {
total.checked_add(*length)
.ok_or(LBFGSShardError::Layout("complete vector length overflows"))
})
}
pub fn range(&self, rank: u32) -> Result<Range<usize>, LBFGSShardError> {
let rank = rank as usize;
if rank >= self.lengths.len() {
return Err(LBFGSShardError::Layout("rank is outside vector layout"));
}
let start = self.lengths[..rank].iter().try_fold(0usize, |total, length| {
total.checked_add(*length)
.ok_or(LBFGSShardError::Layout("vector interval overflows"))
})?;
let end = start.checked_add(self.lengths[rank])
.ok_or(LBFGSShardError::Layout("vector interval end overflows"))?;
Ok(start..end)
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum LBFGSShardError {
Layout(&'static str),
Shape(&'static str),
DType,
Device,
Record,
}
impl fmt::Display for LBFGSShardError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Layout(message) => write!(formatter, "L-BFGS shard layout: {message}"),
Self::Shape(message) => write!(formatter, "L-BFGS vector shape: {message}"),
Self::DType => write!(formatter, "L-BFGS vector precision does not match"),
Self::Device => write!(formatter, "L-BFGS vector device does not match"),
Self::Record => write!(formatter, "L-BFGS shard checkpoint placement does not match"),
}
}
}
impl core::error::Error for LBFGSShardError {}
#[derive(Debug)]
pub enum LBFGSShardedError<E: fmt::Debug> {
State(LBFGSShardError),
Collective(E),
}
impl<E: fmt::Debug> From<LBFGSShardError> for LBFGSShardedError<E> {
fn from(error: LBFGSShardError) -> Self {
Self::State(error)
}
}
impl<E: fmt::Debug> fmt::Display for LBFGSShardedError<E> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::State(error) => write!(formatter, "{error}"),
Self::Collective(error) => write!(formatter, "L-BFGS collective failed: {error:?}"),
}
}
}
impl<E: fmt::Debug> core::error::Error for LBFGSShardedError<E> {}
impl<B: Backend> LBFGSState<B> {
fn vectors(&self) -> impl Iterator<Item = &Tensor<B, 1>> {
self.history_s.iter().chain(self.history_y.iter())
.chain(self.d.iter()).chain(self.prev_flat_grad.iter())
}
pub(super) fn validate_vectors(
&self,
length: usize,
reference: Option<&Tensor<B, 1>>,
) -> Result<(), LBFGSShardError> {
if self.history_s.len() != self.history_y.len() {
return Err(LBFGSShardError::Shape("history displacement/gradient counts differ"));
}
let reference = reference.or_else(|| self.vectors().next());
for value in self.vectors() {
if value.dims() != [length] {
return Err(LBFGSShardError::Shape("native history/direction/gradient length"));
}
if let Some(reference) = reference {
if value.dtype() != reference.dtype() {
return Err(LBFGSShardError::DType);
}
if value.device() != reference.device() {
return Err(LBFGSShardError::Device);
}
}
}
Ok(())
}
pub fn try_into_shard(
self,
layout: &LBFGSShardLayout,
rank: u32,
world: u32,
) -> Result<LBFGSShardedState<B>, LBFGSShardError> {
layout.validate(rank, world)?;
self.validate_vectors(layout.global_len()?, None)?;
let interval = layout.range(rank)?;
let slice = |value: Tensor<B, 1>| {
if world == 1 { value } else { value.slice(interval.clone()) }
};
let state = Self {
history_s: self.history_s.into_iter().map(slice).collect(),
history_y: self.history_y.into_iter().map(slice).collect(),
d: self.d.map(slice),
t: self.t,
prev_flat_grad: self.prev_flat_grad.map(slice),
prev_loss: self.prev_loss,
g_iter: self.g_iter,
};
LBFGSShardedState::from_local_state(state, layout, rank, world)
}
}
#[derive(Clone)]
pub struct LBFGSShardedState<B: Backend> {
version: u32,
rank: u32,
layout: LBFGSShardLayout,
state: LBFGSState<B>,
}
impl<B: Backend> LBFGSShardedState<B> {
pub fn from_local_state(
state: LBFGSState<B>,
layout: &LBFGSShardLayout,
rank: u32,
world: u32,
) -> Result<Self, LBFGSShardError> {
layout.validate(rank, world)?;
state.validate_vectors(layout.lengths[rank as usize], None)?;
Ok(Self { version: 1, rank, layout: layout.clone(), state })
}
pub fn state(&self) -> &LBFGSState<B> { &self.state }
pub fn rank(&self) -> u32 { self.rank }
pub fn layout(&self) -> &LBFGSShardLayout { &self.layout }
pub fn validate_placement(
&self,
layout: &LBFGSShardLayout,
rank: u32,
world: u32,
) -> Result<(), LBFGSShardError> {
layout.validate(rank, world)?;
if self.version != 1 || self.rank != rank || &self.layout != layout {
return Err(LBFGSShardError::Record);
}
self.state.validate_vectors(layout.lengths[rank as usize], None)
}
pub fn repartition_from_shards(
sources: &[Self],
destination: &LBFGSShardLayout,
rank: u32,
world: u32,
) -> Result<Self, LBFGSShardError> {
destination.validate(rank, world)?;
let first = sources.first()
.ok_or(LBFGSShardError::Layout("complete source checkpoint set is empty"))?;
let source_world = u32::try_from(sources.len())
.map_err(|_| LBFGSShardError::Layout("source rank count overflows"))?;
first.layout.validate(0, source_world)?;
if first.layout.global_len()? != destination.global_len()? {
return Err(LBFGSShardError::Shape("destination changes complete vector length"));
}
let reference = first.state.vectors().next();
for (source_rank, source) in sources.iter().enumerate() {
source.validate_placement(&first.layout, source_rank as u32, source_world)?;
if source.state.g_iter != first.state.g_iter
|| source.state.t.map(f64::to_bits) != first.state.t.map(f64::to_bits)
|| source.state.prev_loss.map(f64::to_bits) != first.state.prev_loss.map(f64::to_bits)
|| source.state.history_s.len() != first.state.history_s.len()
|| source.state.d.is_some() != first.state.d.is_some()
|| source.state.prev_flat_grad.is_some() != first.state.prev_flat_grad.is_some()
{
return Err(LBFGSShardError::Record);
}
source.state.validate_vectors(first.layout.lengths[source_rank], reference)?;
}
let target_interval = destination.range(rank)?;
let join = |select: &dyn Fn(&LBFGSState<B>) -> Tensor<B, 1>| {
let mut pieces = Vec::new();
for (source_rank, source) in sources.iter().enumerate() {
let source_interval = first.layout.range(source_rank as u32)?;
let start = target_interval.start.max(source_interval.start);
let end = target_interval.end.min(source_interval.end);
if start < end {
let value = select(&source.state);
let piece = if start == source_interval.start && end == source_interval.end {
value
} else {
value.slice(start - source_interval.start..end - source_interval.start)
};
pieces.push(piece);
}
}
let value = if pieces.len() == 1 {
pieces.pop().expect("one actual overlapping vector piece")
} else {
Tensor::cat(pieces, 0)
};
Ok::<_, LBFGSShardError>(value)
};
let history_s = (0..first.state.history_s.len())
.map(|index| join(&|state| state.history_s[index].clone()))
.collect::<Result<Vec<_>, _>>()?;
let history_y = (0..first.state.history_y.len())
.map(|index| join(&|state| state.history_y[index].clone()))
.collect::<Result<Vec<_>, _>>()?;
let d = if first.state.d.is_some() {
Some(join(&|state| state.d.as_ref().expect("validated direction presence").clone())?)
} else { None };
let prev_flat_grad = if first.state.prev_flat_grad.is_some() {
Some(join(&|state| state.prev_flat_grad.as_ref().expect("validated gradient presence").clone())?)
} else { None };
Self::from_local_state(
LBFGSState {
history_s,
history_y,
d,
t: first.state.t,
prev_flat_grad,
prev_loss: first.state.prev_loss,
g_iter: first.state.g_iter,
},
destination, rank, world,
)
}
pub fn to_device(mut self, device: &B::Device) -> Self {
self.state = self.state.to_device(device);
self
}
}
impl<B: Backend> Record<B> for LBFGSShardedState<B> {
type Item<S: PrecisionSettings> = (
u32, u32, Vec<usize>, <LBFGSState<B> as Record<B>>::Item<S>,
);
fn into_item<S: PrecisionSettings>(self) -> Self::Item<S> {
(self.version, self.rank, self.layout.lengths, self.state.into_item::<S>())
}
fn from_item<S: PrecisionSettings>(item: Self::Item<S>, device: &B::Device) -> Self {
Self {
version: item.0,
rank: item.1,
layout: LBFGSShardLayout::new(item.2),
state: LBFGSState::<B>::from_item::<S>(item.3, device),
}
}
}
pub(super) struct ShardedReductions<'a, C> {
pub(super) communicator: &'a C,
}
impl<C> ShardedReductions<'_, C> {
fn sum<B: Backend>(
&self,
value: Tensor<B, 1>,
) -> Result<Tensor<B, 1>, LBFGSShardedError<C::Error>>
where C: BroadcastTensorCollective<B> {
let dtype = value.dtype();
let device = value.device();
let output = self.communicator.all_reduce_sum(value.into_primitive().tensor())
.map_err(LBFGSShardedError::Collective)?;
let output = Tensor::<B, 1>::from_primitive(TensorPrimitive::Float(output));
if output.dims() != [1] {
return Err(LBFGSShardError::Shape("scalar all-reduce result").into());
}
if output.dtype() != dtype { return Err(LBFGSShardError::DType.into()); }
if output.device() != device { return Err(LBFGSShardError::Device.into()); }
Ok(output)
}
}
impl<B: Backend, C: BroadcastTensorCollective<B>> VectorReductions<B> for ShardedReductions<'_, C> {
type Error = LBFGSShardedError<C::Error>;
fn dot(
&mut self,
lhs: &Tensor<B, 1>,
rhs: &Tensor<B, 1>,
) -> Result<Tensor<B, 1>, Self::Error> {
self.sum(lhs.clone().dot(rhs.clone()))
}
fn sum_abs(&mut self, value: &Tensor<B, 1>) -> Result<f64, Self::Error> {
Ok(self.sum(value.clone().abs().sum())?.into_scalar().to_f64())
}
fn max_abs(&mut self, value: &Tensor<B, 1>) -> Result<f64, Self::Error> {
let local = value.clone().abs().max();
let dtype = local.dtype();
let device = local.device();
let output = self.communicator.all_gather_float(local.into_primitive().tensor())
.map_err(LBFGSShardedError::Collective)?;
let output = Tensor::<B, 1>::from_primitive(TensorPrimitive::Float(output));
if output.dims() != [self.communicator.world_size() as usize] {
return Err(LBFGSShardError::Shape("scalar maxima gather result").into());
}
if output.dtype() != dtype { return Err(LBFGSShardError::DType.into()); }
if output.device() != device { return Err(LBFGSShardError::Device.into()); }
Ok(output.max().into_scalar().to_f64())
}
}
impl<B: AutodiffBackend> LBFGS<B> {
pub fn to_sharded_record<C: BroadcastTensorCollective<B::InnerBackend>>(
&self,
layout: &LBFGSShardLayout,
communicator: &C,
) -> Result<LBFGSShardedState<B::InnerBackend>, LBFGSShardError> {
LBFGSShardedState::from_local_state(
self.state.clone(), layout, communicator.rank(), communicator.world_size(),
)
}
pub fn load_sharded_record<C: BroadcastTensorCollective<B::InnerBackend>>(
mut self,
record: LBFGSShardedState<B::InnerBackend>,
layout: &LBFGSShardLayout,
communicator: &C,
) -> Result<Self, LBFGSShardError> {
record.validate_placement(layout, communicator.rank(), communicator.world_size())?;
self.state = record.state;
Ok(self)
}
pub fn step_sharded<M, F, C>(
&mut self,
lr: LearningRate,
module: M,
mut closure: F,
layout: &LBFGSShardLayout,
communicator: &C,
) -> Result<(M, f64), LBFGSShardedError<C::Error>>
where
M: AutodiffModule<B> + Clone,
F: FnMut(M) -> (f64, GradientsParams),
C: BroadcastTensorCollective<B::InnerBackend>,
{
self.try_step_sharded(lr, module, |model| Ok(closure(model)), layout, communicator)
}
pub fn try_step_sharded<M, F, C>(
&mut self,
lr: LearningRate,
module: M,
closure: F,
layout: &LBFGSShardLayout,
communicator: &C,
) -> Result<(M, f64), LBFGSShardedError<C::Error>>
where
M: AutodiffModule<B> + Clone,
F: FnMut(M) -> Result<(f64, GradientsParams), LBFGSShardedError<C::Error>>,
C: BroadcastTensorCollective<B::InnerBackend>,
{
layout.validate(communicator.rank(), communicator.world_size())?;
let parameters = flatten_params_inner::<B, M>(&module)
.ok_or(LBFGSShardError::Shape("rank has no trainable parameter vector"))?;
let length = layout.lengths[communicator.rank() as usize];
if parameters.dims() != [length] {
return Err(LBFGSShardError::Shape("rank-local unique parameter vector length").into());
}
self.state.validate_vectors(length, Some(¶meters))?;
self.try_step_with_reductions(
lr, module, closure, &mut ShardedReductions { communicator }, Some(parameters), false,
)
.map(|(model, loss, _)| (model, loss))
}
}