use std::sync::Arc;
use crate::distributed::controller::{RoundFrame, RoundKind};
use crate::distributed::cpu_reduce::{round_frame_to_tensors, tensors_to_round_frame};
use crate::tensor::{Result, Tensor};
pub type OuterOptimizerFactory =
Arc<dyn Fn() -> Box<dyn OuterOptimizer> + Send + Sync>;
pub trait OuterOptimizer: Send {
fn outer_step(
&mut self,
prev_global: &[Tensor],
consensus: &[Tensor],
) -> Result<Vec<Tensor>>;
fn checkpoint_state(&self) -> Option<Vec<Tensor>> {
None
}
fn load_checkpoint_state(&mut self, _state: Vec<Tensor>) -> Result<()> {
Ok(())
}
fn resets_inner(&self) -> bool {
false
}
}
#[derive(Debug, Default, Clone, Copy)]
pub struct OuterAvg;
impl OuterOptimizer for OuterAvg {
fn outer_step(
&mut self,
_prev_global: &[Tensor],
consensus: &[Tensor],
) -> Result<Vec<Tensor>> {
Ok(consensus.to_vec())
}
}
pub struct SlowMomentum {
lr: f64,
mu: f64,
velocity: Vec<Tensor>,
}
impl SlowMomentum {
pub fn new(lr: f64, mu: f64) -> Self {
SlowMomentum { lr, mu, velocity: Vec::new() }
}
}
impl OuterOptimizer for SlowMomentum {
fn outer_step(
&mut self,
prev_global: &[Tensor],
consensus: &[Tensor],
) -> Result<Vec<Tensor>> {
let n = consensus.len();
let fresh = self.velocity.len() != n;
let mut new_global = Vec::with_capacity(n);
let mut new_velocity = Vec::with_capacity(n);
for i in 0..n {
let g = prev_global[i].sub(&consensus[i])?;
let v = if fresh {
g
} else {
self.velocity[i].mul_scalar(self.mu)?.add(&g)?
};
let step = v.mul_scalar(self.lr)?;
new_global.push(prev_global[i].sub(&step)?);
new_velocity.push(v);
}
self.velocity = new_velocity;
Ok(new_global)
}
fn checkpoint_state(&self) -> Option<Vec<Tensor>> {
if self.velocity.is_empty() {
None
} else {
Some(self.velocity.clone())
}
}
fn load_checkpoint_state(&mut self, state: Vec<Tensor>) -> Result<()> {
self.velocity = state;
Ok(())
}
}
pub struct NesterovMomentum {
lr: f64,
mu: f64,
velocity: Vec<Tensor>,
}
impl NesterovMomentum {
pub fn new(lr: f64, mu: f64) -> Self {
NesterovMomentum { lr, mu, velocity: Vec::new() }
}
}
impl OuterOptimizer for NesterovMomentum {
fn outer_step(
&mut self,
prev_global: &[Tensor],
consensus: &[Tensor],
) -> Result<Vec<Tensor>> {
let n = consensus.len();
let fresh = self.velocity.len() != n;
let mut new_global = Vec::with_capacity(n);
let mut new_velocity = Vec::with_capacity(n);
for i in 0..n {
let g = prev_global[i].sub(&consensus[i])?;
let v = if fresh {
g.copy()?
} else {
self.velocity[i].mul_scalar(self.mu)?.add(&g)?
};
let look_ahead = v.mul_scalar(self.mu)?.add(&g)?;
let step = look_ahead.mul_scalar(self.lr)?;
new_global.push(prev_global[i].sub(&step)?);
new_velocity.push(v);
}
self.velocity = new_velocity;
Ok(new_global)
}
fn checkpoint_state(&self) -> Option<Vec<Tensor>> {
if self.velocity.is_empty() {
None
} else {
Some(self.velocity.clone())
}
}
fn load_checkpoint_state(&mut self, state: Vec<Tensor>) -> Result<()> {
self.velocity = state;
Ok(())
}
fn resets_inner(&self) -> bool {
true
}
}
pub struct OuterStepper {
opt: Box<dyn OuterOptimizer>,
prev_global: Option<Vec<Tensor>>,
seen_params_this_window: bool,
}
impl OuterStepper {
pub fn new(opt: Box<dyn OuterOptimizer>) -> Self {
OuterStepper {
opt,
prev_global: None,
seen_params_this_window: false,
}
}
pub fn process_frame(&mut self, frame: RoundFrame) -> Result<RoundFrame> {
match frame.kind {
RoundKind::Control => {
self.seen_params_this_window = false;
Ok(frame)
}
RoundKind::Model if !self.seen_params_this_window => {
if !crate::distributed::realized_work::is_realized(frame.weight) {
return Ok(frame);
}
self.seen_params_this_window = true;
let weight = frame.weight;
let wire_dtype = frame
.tensors
.first()
.map_or(crate::distributed::controller::DTYPE_F32, |t| t.dtype);
let consensus = round_frame_to_tensors(&frame)?;
let prev = self.prev_global.take().unwrap_or_else(|| consensus.clone());
let new_global = self.opt.outer_step(&prev, &consensus)?;
self.prev_global = Some(new_global.clone());
let refs: Vec<&Tensor> = new_global.iter().collect();
let mut stepped = tensors_to_round_frame(&refs, wire_dtype)?;
stepped.weight = weight;
Ok(stepped)
}
RoundKind::Model => Ok(frame),
}
}
pub fn checkpoint_state(&self) -> Option<Vec<Tensor>> {
self.opt.checkpoint_state()
}
pub fn load_checkpoint_state(&mut self, state: Vec<Tensor>) -> Result<()> {
self.opt.load_checkpoint_state(state)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::distributed::controller::{DTYPE_BF16, DTYPE_F32};
use crate::tensor::{Device, test_device};
fn t(vals: &[f32], shape: &[i64]) -> Tensor {
Tensor::from_f32(vals, shape, test_device()).unwrap()
}
#[test]
fn outer_avg_is_identity() {
let mut opt = OuterAvg;
let prev = vec![t(&[1.0, 2.0, 3.0], &[3]), t(&[4.0, 5.0], &[2])];
let consensus = vec![t(&[10.0, 20.0, 30.0], &[3]), t(&[40.0, 50.0], &[2])];
let out = opt.outer_step(&prev, &consensus).unwrap();
assert_eq!(out.len(), consensus.len());
assert_eq!(out[0].to_f32_vec().unwrap(), vec![10.0, 20.0, 30.0]);
assert_eq!(out[1].to_f32_vec().unwrap(), vec![40.0, 50.0]);
}
#[test]
fn outer_avg_ignores_prev_global() {
let mut opt = OuterAvg;
let consensus = vec![t(&[7.0, 8.0], &[2])];
let bogus_prev = vec![t(&[0.0, 0.0], &[2])];
let out = opt.outer_step(&bogus_prev, &consensus).unwrap();
assert_eq!(out[0].to_f32_vec().unwrap(), vec![7.0, 8.0]);
let out2 = opt.outer_step(&consensus, &consensus).unwrap();
assert_eq!(out2[0].to_f32_vec().unwrap(), vec![7.0, 8.0]);
}
#[test]
fn stepper_with_outer_avg_is_byte_identical() {
let p0 = t(&[1.25, -2.5, 3.75], &[3]);
let p1 = t(&[4.0, 5.0], &[2]);
let buf = t(&[9.0, 8.0], &[2]);
let mut params_frame = tensors_to_round_frame(&[&p0, &p1], DTYPE_F32).unwrap();
params_frame.weight = 1.0; let buffers_frame = tensors_to_round_frame(&[&buf], DTYPE_F32).unwrap();
let mut control_frame = tensors_to_round_frame(&[&p0], DTYPE_F32).unwrap();
control_frame.kind = RoundKind::Control;
let mut stepper = OuterStepper::new(Box::new(OuterAvg));
let c_out = stepper.process_frame(control_frame.clone()).unwrap();
assert_eq!(c_out, control_frame, "Control frame must pass through");
let p_out = stepper.process_frame(params_frame.clone()).unwrap();
assert_eq!(p_out, params_frame, "params frame must be byte-identical under OuterAvg");
let b_out = stepper.process_frame(buffers_frame.clone()).unwrap();
assert_eq!(b_out, buffers_frame, "buffers frame must pass through");
stepper.process_frame(control_frame).unwrap();
let p_out2 = stepper.process_frame(params_frame.clone()).unwrap();
assert_eq!(p_out2, params_frame, "params still byte-identical second window");
}
#[test]
fn stepper_preserves_bf16_frame_dtype() {
let p = t(&[1.25, -2.5, 3.75], &[3]);
let mut params_frame = tensors_to_round_frame(&[&p], DTYPE_BF16).unwrap();
params_frame.weight = 1.0;
assert_eq!(params_frame.tensors[0].dtype, DTYPE_BF16);
let mut stepper = OuterStepper::new(Box::new(OuterAvg));
let out = stepper.process_frame(params_frame.clone()).unwrap();
assert_eq!(out, params_frame, "OuterAvg identity is byte-exact in bf16 too");
}
#[test]
fn slow_momentum_heavy_ball_math() {
let mut opt = SlowMomentum::new(0.5, 0.9);
let w1 = opt
.outer_step(&[t(&[1.0, 2.0], &[2])], &[t(&[1.0, 2.0], &[2])])
.unwrap();
assert_eq!(w1[0].to_f32_vec().unwrap(), vec![1.0, 2.0]);
let w2 = opt
.outer_step(&[t(&[1.0, 2.0], &[2])], &[t(&[0.5, 1.0], &[2])])
.unwrap();
let w2v = w2[0].to_f32_vec().unwrap();
assert!((w2v[0] - 0.75).abs() < 1e-6 && (w2v[1] - 1.5).abs() < 1e-6, "got {w2v:?}");
let w3 = opt
.outer_step(&[t(&[0.75, 1.5], &[2])], &[t(&[0.7, 1.4], &[2])])
.unwrap();
let w3v = w3[0].to_f32_vec().unwrap();
assert!((w3v[0] - 0.5).abs() < 1e-6 && (w3v[1] - 1.0).abs() < 1e-6, "got {w3v:?}");
}
#[test]
fn stepper_slow_momentum_steps_params_only() {
let mut control = tensors_to_round_frame(&[&t(&[0.0], &[1])], DTYPE_F32).unwrap();
control.kind = RoundKind::Control;
let buffers = tensors_to_round_frame(&[&t(&[9.0, 8.0], &[2])], DTYPE_F32).unwrap();
let mut stepper = OuterStepper::new(Box::new(SlowMomentum::new(0.5, 0.9)));
stepper.process_frame(control.clone()).unwrap();
let mut p1 = tensors_to_round_frame(&[&t(&[2.0, 4.0], &[2])], DTYPE_F32).unwrap();
p1.weight = 1.0; let p1_out = stepper.process_frame(p1.clone()).unwrap();
assert_eq!(p1_out, p1, "first-window params unchanged (g=0)");
let b1_out = stepper.process_frame(buffers.clone()).unwrap();
assert_eq!(b1_out, buffers, "buffers pass through");
stepper.process_frame(control).unwrap();
let mut p2 = tensors_to_round_frame(&[&t(&[1.0, 2.0], &[2])], DTYPE_F32).unwrap();
p2.weight = 1.0;
let p2_out = stepper.process_frame(p2).unwrap();
let stepped = round_frame_to_tensors(&p2_out).unwrap()[0].to_f32_vec().unwrap();
assert!(
(stepped[0] - 1.5).abs() < 1e-6 && (stepped[1] - 3.0).abs() < 1e-6,
"second-window params stepped: got {stepped:?}"
);
let b2_out = stepper.process_frame(buffers.clone()).unwrap();
assert_eq!(b2_out, buffers, "buffers still pass through after a real step");
}
#[test]
fn outer_avg_checkpoint_state_is_none() {
let opt = OuterAvg;
assert!(opt.checkpoint_state().is_none());
}
#[test]
fn slow_momentum_checkpoint_round_trip_is_faithful() {
let mut warm = SlowMomentum::new(0.5, 0.9);
warm.outer_step(&[t(&[1.0, 2.0], &[2])], &[t(&[1.0, 2.0], &[2])]).unwrap();
warm.outer_step(&[t(&[1.0, 2.0], &[2])], &[t(&[0.5, 1.0], &[2])]).unwrap();
let saved = warm.checkpoint_state().expect("has momentum after stepping");
let mut resumed = SlowMomentum::new(0.5, 0.9);
resumed.load_checkpoint_state(saved).unwrap();
let prev = [t(&[0.75, 1.5], &[2])];
let cons = [t(&[0.7, 1.4], &[2])];
let a = warm.outer_step(&prev, &cons).unwrap()[0].to_f32_vec().unwrap();
let b = resumed.outer_step(&prev, &cons).unwrap()[0].to_f32_vec().unwrap();
for (i, (x, y)) in a.iter().zip(&b).enumerate() {
assert!((x - y).abs() < 1e-6, "param[{i}]: resumed {y} != warmed {x}");
}
let mut fresh = SlowMomentum::new(0.5, 0.9);
let c = fresh.outer_step(&prev, &cons).unwrap()[0].to_f32_vec().unwrap();
assert!(
(a[0] - c[0]).abs() > 1e-6 || (a[1] - c[1]).abs() > 1e-6,
"warmed and from-zero should differ, else the test is vacuous"
);
}
#[test]
fn nesterov_resets_inner_others_dont() {
assert!(NesterovMomentum::new(0.7, 0.9).resets_inner(), "DiLoCo resets inner");
assert!(!SlowMomentum::new(0.5, 0.7).resets_inner(), "SlowMo keeps inner continuous");
assert!(!OuterAvg.resets_inner(), "OuterAvg keeps inner continuous");
}
#[test]
fn nesterov_look_ahead_math() {
let mut opt = NesterovMomentum::new(0.5, 0.9);
let w1 = opt
.outer_step(&[t(&[1.0, 2.0], &[2])], &[t(&[1.0, 2.0], &[2])])
.unwrap();
assert_eq!(w1[0].to_f32_vec().unwrap(), vec![1.0, 2.0]);
let w2 = opt
.outer_step(&[t(&[1.0, 2.0], &[2])], &[t(&[0.5, 1.0], &[2])])
.unwrap();
let w2v = w2[0].to_f32_vec().unwrap();
assert!((w2v[0] - 0.525).abs() < 1e-6 && (w2v[1] - 1.05).abs() < 1e-6, "got {w2v:?}");
let w3 = opt
.outer_step(&[t(&[0.525, 1.05], &[2])], &[t(&[0.5, 1.0], &[2])])
.unwrap();
let w3v = w3[0].to_f32_vec().unwrap();
assert!((w3v[0] - 0.29875).abs() < 1e-5 && (w3v[1] - 0.5975).abs() < 1e-5, "got {w3v:?}");
}
#[test]
fn nesterov_checkpoint_round_trip() {
let mut warm = NesterovMomentum::new(0.5, 0.9);
warm.outer_step(&[t(&[1.0, 2.0], &[2])], &[t(&[1.0, 2.0], &[2])]).unwrap();
warm.outer_step(&[t(&[1.0, 2.0], &[2])], &[t(&[0.5, 1.0], &[2])]).unwrap();
let saved = warm.checkpoint_state().expect("has momentum");
let mut resumed = NesterovMomentum::new(0.5, 0.9);
resumed.load_checkpoint_state(saved).unwrap();
let prev = [t(&[0.525, 1.05], &[2])];
let cons = [t(&[0.5, 1.0], &[2])];
let a = warm.outer_step(&prev, &cons).unwrap()[0].to_f32_vec().unwrap();
let b = resumed.outer_step(&prev, &cons).unwrap()[0].to_f32_vec().unwrap();
for (x, y) in a.iter().zip(&b) {
assert!((x - y).abs() < 1e-6, "resumed {y} != warmed {x}");
}
}
#[test]
fn outer_avg_as_trait_object() {
let mut opt: Box<dyn OuterOptimizer> = Box::new(OuterAvg);
let c = vec![t(&[1.5], &[1])];
let out = opt.outer_step(&c, &c).unwrap();
assert_eq!(out[0].to_f32_vec().unwrap(), vec![1.5]);
let _ = Device::CPU;
}
}