use std::path::{Path, PathBuf};
use std::sync::{Arc, Mutex};
use crate::distributed::ModelSchema;
use crate::distributed::controller::{DTYPE_F32, RoundFrame, TensorPayload};
use crate::nn::checkpoint::{LoadReport, RawCheckpointEntry, dtype_tag, save_checkpoint_from_raw_file};
use crate::nn::{Buffer, Module, Parameter};
use crate::tensor::{DType, Result, TensorError};
pub(crate) fn consensus_param_key(i: usize) -> String {
format!("p{i}")
}
pub(crate) fn consensus_buffer_key(j: usize) -> String {
format!("b{j}")
}
pub fn load_consensus_checkpoint<M: Module + ?Sized>(model: &M, path: &str) -> Result<LoadReport> {
let params: Vec<(String, Parameter)> = model
.parameters()
.into_iter()
.enumerate()
.map(|(i, p)| (consensus_param_key(i), p))
.collect();
let buffers: Vec<(String, Buffer)> = model
.buffers()
.into_iter()
.enumerate()
.map(|(j, b)| (consensus_buffer_key(j), b))
.collect();
crate::nn::load_checkpoint_file(path, ¶ms, &buffers, None)
}
pub struct CheckpointForge {
schema: Option<ModelSchema>,
inner: Mutex<ForgeState>,
}
#[derive(Default)]
struct ForgeState {
pending_path: Option<PathBuf>,
accumulated: Vec<TensorPayload>,
pending_outer: Option<Vec<TensorPayload>>,
}
impl CheckpointForge {
pub fn new(schema: Option<ModelSchema>) -> Arc<Self> {
Arc::new(CheckpointForge {
schema,
inner: Mutex::new(ForgeState::default()),
})
}
pub fn can_write_model(&self) -> bool {
self.schema.is_some()
}
pub fn arm(&self, model_path: PathBuf) {
let mut st = self.inner.lock().expect("checkpoint forge mutex poisoned");
st.pending_path = Some(model_path);
st.accumulated.clear();
st.pending_outer = None;
}
pub fn is_armed(&self) -> bool {
self.inner
.lock()
.expect("checkpoint forge mutex poisoned")
.pending_path
.is_some()
}
pub fn stash_outer_momentum(&self, payloads: Vec<TensorPayload>) {
let mut st = self.inner.lock().expect("checkpoint forge mutex poisoned");
st.pending_outer = Some(payloads);
}
pub fn accumulate(&self, frame: RoundFrame) {
let Some(schema) = self.schema.as_ref() else {
return; };
let want = schema.tensor_count();
let (path, payloads, outer) = {
let mut st = self.inner.lock().expect("checkpoint forge mutex poisoned");
if st.pending_path.is_none() {
return; }
st.accumulated.extend(frame.tensors);
if st.accumulated.len() < want {
return; }
let path = st.pending_path.take().expect("armed checked above");
let payloads = std::mem::take(&mut st.accumulated);
let outer = st.pending_outer.take();
(path, payloads, outer)
};
if payloads.len() != want {
eprintln!(
"flodl ddp: consensus checkpoint accumulated {} tensors but schema \
expects {} ({} params + {} buffers); .fdl skipped for {}",
payloads.len(),
want,
schema.param_names.len(),
schema.buffer_names.len(),
path.display(),
);
return;
}
let outer = outer.filter(|o| {
let ok = o.len() == schema.param_names.len();
if !ok {
eprintln!(
"flodl ddp: outer-momentum has {} tensors but schema expects \
{} params; .outer.fdl skipped for {}",
o.len(),
schema.param_names.len(),
path.display(),
);
}
ok
});
let schema = schema.clone();
let spawn = std::thread::Builder::new()
.name("flodl-ckpt-writer".to_string())
.spawn(move || {
if let Err(e) = write_consensus_fdl(&schema, &payloads, &path) {
eprintln!(
"flodl ddp: consensus checkpoint write to {} failed: {e}",
path.display(),
);
}
if let Some(outer_payloads) = outer {
let outer_path = path.with_extension("outer.fdl");
if let Err(e) = write_outer_momentum_fdl(&outer_payloads, &outer_path) {
eprintln!(
"flodl ddp: outer-momentum checkpoint write to {} failed: {e}",
outer_path.display(),
);
}
}
});
if let Err(e) = spawn {
eprintln!("flodl ddp: failed to spawn checkpoint writer thread: {e}");
}
}
}
fn write_consensus_fdl(schema: &ModelSchema, payloads: &[TensorPayload], path: &Path) -> Result<()> {
if payloads.len() != schema.tensor_count() {
return Err(TensorError::new(&format!(
"checkpoint_forge: accumulated {} tensors but schema expects \
{} ({} params + {} buffers) — schema/accumulation mismatch",
payloads.len(),
schema.tensor_count(),
schema.param_names.len(),
schema.buffer_names.len(),
)));
}
let shapes: Vec<Vec<i64>> = payloads
.iter()
.map(|p| p.shape.iter().map(|&d| d as i64).collect())
.collect();
let param_count = schema.param_names.len();
let keys: Vec<String> = (0..payloads.len())
.map(|i| {
if i < param_count {
consensus_param_key(i)
} else {
consensus_buffer_key(i - param_count)
}
})
.collect();
let upcast: Vec<Option<Vec<u8>>> = payloads
.iter()
.map(|p| {
if p.dtype == DTYPE_F32 {
Ok(None)
} else {
let vals = crate::distributed::controller::payload_to_f32(p)?;
Ok(Some(
crate::distributed::controller::f32_slice_to_payload_bytes(
&vals, DTYPE_F32,
)?,
))
}
})
.collect::<Result<_>>()?;
let mut entries = Vec::with_capacity(payloads.len());
for (i, p) in payloads.iter().enumerate() {
entries.push(RawCheckpointEntry {
name: keys[i].as_str(),
shape: &shapes[i],
dtype_tag: dtype_tag(DType::Float32),
raw: upcast[i].as_deref().unwrap_or(&p.bytes),
});
}
let path_str = path.to_str().ok_or_else(|| {
TensorError::new(&format!(
"checkpoint_forge: non-utf8 checkpoint path {}",
path.display(),
))
})?;
let tmp = format!("{path_str}.tmp");
save_checkpoint_from_raw_file(&tmp, &entries, None)?;
std::fs::rename(&tmp, path_str).map_err(|e| {
TensorError::new(&format!(
"checkpoint_forge: atomic rename {tmp} -> {path_str} failed: {e}"
))
})?;
Ok(())
}
fn write_outer_momentum_fdl(payloads: &[TensorPayload], path: &Path) -> Result<()> {
let shapes: Vec<Vec<i64>> = payloads
.iter()
.map(|p| p.shape.iter().map(|&d| d as i64).collect())
.collect();
let keys: Vec<String> = (0..payloads.len()).map(consensus_param_key).collect();
let mut entries = Vec::with_capacity(payloads.len());
for (i, p) in payloads.iter().enumerate() {
if p.dtype != DTYPE_F32 {
return Err(TensorError::new(&format!(
"checkpoint_forge: outer payload[{i}] dtype {} not supported (v1 f32 only)",
p.dtype,
)));
}
entries.push(RawCheckpointEntry {
name: keys[i].as_str(),
shape: &shapes[i],
dtype_tag: dtype_tag(DType::Float32),
raw: &p.bytes,
});
}
let path_str = path.to_str().ok_or_else(|| {
TensorError::new(&format!(
"checkpoint_forge: non-utf8 outer-momentum path {}",
path.display(),
))
})?;
let tmp = format!("{path_str}.tmp");
save_checkpoint_from_raw_file(&tmp, &entries, None)?;
std::fs::rename(&tmp, path_str).map_err(|e| {
TensorError::new(&format!(
"checkpoint_forge: atomic rename {tmp} -> {path_str} failed: {e}"
))
})?;
Ok(())
}
pub fn load_outer_momentum<M: Module + ?Sized>(
model: &M,
path: &str,
) -> Result<Vec<crate::tensor::Tensor>> {
let params = model.parameters();
let mut targets: Vec<(String, Parameter)> = Vec::with_capacity(params.len());
for (i, p) in params.iter().enumerate() {
let zeros = crate::tensor::Tensor::zeros_like(&p.variable.data())?;
targets.push((consensus_param_key(i), Parameter::new(zeros, "outer_momentum")));
}
crate::nn::load_checkpoint_file(path, &targets, &[], None)?;
Ok(targets.iter().map(|(_, p)| p.variable.data()).collect())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::distributed::cpu_reduce::tensors_to_round_frame;
use crate::tensor::{Device, Tensor};
fn cpu_tensor(vals: &[f32], shape: &[i64]) -> Tensor {
Tensor::from_f32(vals, shape, Device::CPU).unwrap()
}
#[test]
fn accumulate_params_then_buffers_round_trips_through_load() {
let schema = ModelSchema {
param_names: vec!["w".to_string(), "b".to_string()],
buffer_names: vec!["running_mean".to_string()],
};
let forge = CheckpointForge::new(Some(schema));
assert!(forge.can_write_model());
let w = cpu_tensor(&[1.0, 2.0, 3.0, 4.0], &[2, 2]);
let b = cpu_tensor(&[5.0, 6.0], &[2]);
let rm = cpu_tensor(&[7.0, 8.0], &[2]);
let params_frame = tensors_to_round_frame(&[&w, &b], DTYPE_F32).unwrap();
let buffers_frame = tensors_to_round_frame(&[&rm], DTYPE_F32).unwrap();
let dir = std::env::temp_dir().join(format!("flodl_forge_{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("consensus.fdl");
forge.arm(path.clone());
forge.accumulate(params_frame); assert!(!path.exists(), "no write before all model frames arrive");
forge.accumulate(buffers_frame);
let mut found = false;
for _ in 0..200 {
if path.exists() {
found = true;
break;
}
std::thread::sleep(std::time::Duration::from_millis(5));
}
assert!(found, "completed accumulation produced the .fdl");
use crate::nn::{Buffer, Parameter};
let tw = Parameter::new(cpu_tensor(&[0.0; 4], &[2, 2]), "w");
let tb = Parameter::new(cpu_tensor(&[0.0; 2], &[2]), "b");
let trm = Buffer::new(cpu_tensor(&[0.0; 2], &[2]), "running_mean");
crate::nn::load_checkpoint_file(
path.to_str().unwrap(),
&[("p0".to_string(), tw.clone()), ("p1".to_string(), tb.clone())],
&[("b0".to_string(), trm.clone())],
None,
)
.unwrap();
assert_eq!(tw.variable.data().to_f32_vec().unwrap(), vec![1.0, 2.0, 3.0, 4.0]);
assert_eq!(tb.variable.data().to_f32_vec().unwrap(), vec![5.0, 6.0]);
assert_eq!(trm.get().to_f32_vec().unwrap(), vec![7.0, 8.0]);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn bf16_frames_write_f32_checkpoint() {
use crate::distributed::controller::DTYPE_BF16;
let schema = ModelSchema {
param_names: vec!["w".to_string()],
buffer_names: vec![],
};
let forge = CheckpointForge::new(Some(schema));
let w = cpu_tensor(&[1.5, -2.0, 0.25, 42.0], &[4]);
let frame = tensors_to_round_frame(&[&w], DTYPE_BF16).unwrap();
assert_eq!(frame.tensors[0].dtype, DTYPE_BF16);
let dir =
std::env::temp_dir().join(format!("flodl_forge_bf16_{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("consensus.fdl");
forge.arm(path.clone());
forge.accumulate(frame);
let mut found = false;
for _ in 0..200 {
if path.exists() {
found = true;
break;
}
std::thread::sleep(std::time::Duration::from_millis(5));
}
assert!(found, "bf16 accumulation produced the .fdl");
use crate::nn::Parameter;
let tw = Parameter::new(cpu_tensor(&[0.0; 4], &[4]), "w");
crate::nn::load_checkpoint_file(
path.to_str().unwrap(),
&[("p0".to_string(), tw.clone())],
&[],
None,
)
.unwrap();
let loaded = tw.variable.data();
assert_eq!(loaded.dtype(), crate::tensor::DType::Float32);
assert_eq!(loaded.to_f32_vec().unwrap(), vec![1.5, -2.0, 0.25, 42.0]);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn outer_momentum_writes_sidecar_and_round_trips() {
let schema = ModelSchema {
param_names: vec!["w".to_string(), "b".to_string()],
buffer_names: vec!["running_mean".to_string()],
};
let forge = CheckpointForge::new(Some(schema));
let w = cpu_tensor(&[1.0, 2.0, 3.0, 4.0], &[2, 2]);
let b = cpu_tensor(&[5.0, 6.0], &[2]);
let rm = cpu_tensor(&[7.0, 8.0], &[2]);
let mw = cpu_tensor(&[0.1, 0.2, 0.3, 0.4], &[2, 2]);
let mb = cpu_tensor(&[0.5, 0.6], &[2]);
let dir = std::env::temp_dir().join(format!("flodl_forge_outer_{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("consensus.fdl");
forge.arm(path.clone());
assert!(forge.is_armed());
forge.stash_outer_momentum(tensors_to_round_frame(&[&mw, &mb], DTYPE_F32).unwrap().tensors);
forge.accumulate(tensors_to_round_frame(&[&w, &b], DTYPE_F32).unwrap());
forge.accumulate(tensors_to_round_frame(&[&rm], DTYPE_F32).unwrap());
let outer_path = path.with_extension("outer.fdl");
let mut found = false;
for _ in 0..200 {
if path.exists() && outer_path.exists() {
found = true;
break;
}
std::thread::sleep(std::time::Duration::from_millis(5));
}
assert!(found, "both .fdl and .outer.fdl written");
use crate::nn::Parameter;
let t0 = Parameter::new(cpu_tensor(&[0.0; 4], &[2, 2]), "m0");
let t1 = Parameter::new(cpu_tensor(&[0.0; 2], &[2]), "m1");
crate::nn::load_checkpoint_file(
outer_path.to_str().unwrap(),
&[("p0".to_string(), t0.clone()), ("p1".to_string(), t1.clone())],
&[],
None,
)
.unwrap();
assert_eq!(t0.variable.data().to_f32_vec().unwrap(), vec![0.1, 0.2, 0.3, 0.4]);
assert_eq!(t1.variable.data().to_f32_vec().unwrap(), vec![0.5, 0.6]);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn no_outer_momentum_writes_no_sidecar() {
let schema = ModelSchema {
param_names: vec!["w".to_string()],
buffer_names: vec![],
};
let forge = CheckpointForge::new(Some(schema));
let w = cpu_tensor(&[1.0, 2.0], &[2]);
let dir = std::env::temp_dir().join(format!("flodl_forge_noouter_{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("c.fdl");
forge.arm(path.clone());
forge.accumulate(tensors_to_round_frame(&[&w], DTYPE_F32).unwrap());
let outer_path = path.with_extension("outer.fdl");
let mut model_found = false;
for _ in 0..200 {
if path.exists() {
model_found = true;
break;
}
std::thread::sleep(std::time::Duration::from_millis(5));
}
assert!(model_found, ".fdl written");
std::thread::sleep(std::time::Duration::from_millis(20));
assert!(!outer_path.exists(), "no .outer.fdl when no momentum stashed");
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn unarmed_accumulate_is_noop() {
let schema = ModelSchema {
param_names: vec!["w".to_string()],
buffer_names: vec![],
};
let forge = CheckpointForge::new(Some(schema));
let w = cpu_tensor(&[1.0, 2.0], &[2]);
forge.accumulate(tensors_to_round_frame(&[&w], DTYPE_F32).unwrap());
}
#[test]
fn arm_clears_partial_accumulation() {
let schema = ModelSchema {
param_names: vec!["w".to_string(), "b".to_string()],
buffer_names: vec![],
};
let forge = CheckpointForge::new(Some(schema));
let stale = cpu_tensor(&[9.0], &[1]);
let w = cpu_tensor(&[1.0, 2.0], &[2]);
let b = cpu_tensor(&[3.0], &[1]);
let dir = std::env::temp_dir().join(format!("flodl_forge_arm_{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("c.fdl");
forge.arm(path.clone());
forge.accumulate(tensors_to_round_frame(&[&stale], DTYPE_F32).unwrap()); forge.arm(path.clone());
forge.accumulate(tensors_to_round_frame(&[&w], DTYPE_F32).unwrap()); forge.accumulate(tensors_to_round_frame(&[&b], DTYPE_F32).unwrap());
let mut found = false;
for _ in 0..200 {
if path.exists() {
found = true;
break;
}
std::thread::sleep(std::time::Duration::from_millis(5));
}
assert!(found, "re-armed accumulation produced the .fdl");
use crate::nn::Parameter;
let tw = Parameter::new(cpu_tensor(&[0.0; 2], &[2]), "w");
let tb = Parameter::new(cpu_tensor(&[0.0; 1], &[1]), "b");
crate::nn::load_checkpoint_file(
path.to_str().unwrap(),
&[("p0".to_string(), tw.clone()), ("p1".to_string(), tb.clone())],
&[],
None,
)
.unwrap();
assert_eq!(tw.variable.data().to_f32_vec().unwrap(), vec![1.0, 2.0]);
assert_eq!(tb.variable.data().to_f32_vec().unwrap(), vec![3.0]);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn write_consensus_fdl_rejects_tensor_count_mismatch() {
let schema = ModelSchema {
param_names: vec!["w".to_string()],
buffer_names: vec![],
};
let a = cpu_tensor(&[1.0], &[1]);
let b = cpu_tensor(&[2.0], &[1]);
let frame = tensors_to_round_frame(&[&a, &b], DTYPE_F32).unwrap();
let path = std::env::temp_dir().join("flodl_forge_mismatch.fdl");
let err = write_consensus_fdl(&schema, &frame.tensors, &path).unwrap_err();
assert!(err.to_string().contains("mismatch"), "got: {err}");
}
}