use serde::{Deserialize, Serialize};
use crate::tensor::{Result, TensorError};
pub const CHECKPOINT_META_SCHEMA_VERSION: u32 = 4;
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct EpochCoverage {
pub epoch: usize,
pub total_samples: usize,
pub uncovered_ranges: Vec<(usize, usize)>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CoverageBlock {
pub seed: u64,
pub batch_size: usize,
pub per_epoch: Vec<EpochCoverage>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ElCheState {
pub anchor: usize,
pub anchor_rank: Option<usize>,
pub smoothed_ms_per_batch: Vec<f64>,
pub phase: crate::distributed::el_che::Phase,
pub calibration_count: u64,
#[serde(default)]
pub trend_history: Option<Vec<f64>>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ModelSchema {
pub param_names: Vec<String>,
pub buffer_names: Vec<String>,
}
impl ModelSchema {
pub fn from_module<M: crate::nn::Module + ?Sized>(model: &M) -> Self {
ModelSchema {
param_names: model
.parameters()
.iter()
.map(|p| p.name.clone())
.collect(),
buffer_names: model
.buffers()
.iter()
.map(|b| b.name.clone())
.collect(),
}
}
pub fn tensor_count(&self) -> usize {
self.param_names.len() + self.buffer_names.len()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
#[serde(rename_all = "snake_case")]
pub enum SaveReason {
GracefulShutdown,
MaxFailureExceeded,
SingleSurvivor,
AllRanksLost,
ReduceStall,
Checkpoint,
}
impl SaveReason {
pub const fn to_u8(self) -> u8 {
match self {
SaveReason::GracefulShutdown => 0,
SaveReason::MaxFailureExceeded => 1,
SaveReason::SingleSurvivor => 2,
SaveReason::AllRanksLost => 3,
SaveReason::ReduceStall => 4,
SaveReason::Checkpoint => 5,
}
}
pub fn from_u8(byte: u8) -> Option<Self> {
match byte {
0 => Some(SaveReason::GracefulShutdown),
1 => Some(SaveReason::MaxFailureExceeded),
2 => Some(SaveReason::SingleSurvivor),
3 => Some(SaveReason::AllRanksLost),
4 => Some(SaveReason::ReduceStall),
5 => Some(SaveReason::Checkpoint),
_ => None,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CheckpointMeta {
pub schema_version: u32,
pub epoch: usize,
pub global_step: usize,
pub sync_round: u64,
pub world_size_at_save: usize,
pub save_reason: SaveReason,
#[serde(default)]
pub elche_state: Option<ElCheState>,
#[serde(default)]
pub coverage: Option<CoverageBlock>,
}
impl CheckpointMeta {
pub fn new(
epoch: usize,
global_step: usize,
sync_round: u64,
world_size_at_save: usize,
save_reason: SaveReason,
) -> Self {
Self {
schema_version: CHECKPOINT_META_SCHEMA_VERSION,
epoch,
global_step,
sync_round,
world_size_at_save,
save_reason,
elche_state: None,
coverage: None,
}
}
pub fn with_elche_state(mut self, state: ElCheState) -> Self {
self.elche_state = Some(state);
self
}
pub fn with_coverage(mut self, coverage: CoverageBlock) -> Self {
self.coverage = Some(coverage);
self
}
pub fn sidecar_path(stem: &str) -> std::path::PathBuf {
let mut p = std::path::PathBuf::from(stem);
p.set_extension("meta.json");
p
}
pub fn write_to_file(&self, path: &std::path::Path) -> Result<()> {
let content = serde_json::to_string_pretty(self).map_err(|e| {
TensorError::new(&format!(
"CheckpointMeta: serialize JSON for {}: {e}",
path.display(),
))
})?;
let tmp = path.with_extension("json.tmp");
std::fs::write(&tmp, content).map_err(|e| {
TensorError::new(&format!(
"CheckpointMeta: write {}: {e}",
tmp.display(),
))
})?;
std::fs::rename(&tmp, path).map_err(|e| {
TensorError::new(&format!(
"CheckpointMeta: atomic rename {} -> {}: {e}",
tmp.display(),
path.display(),
))
})
}
pub fn read_from_file(path: &std::path::Path) -> Result<Self> {
let content = std::fs::read_to_string(path).map_err(|e| {
TensorError::new(&format!(
"CheckpointMeta: read {}: {e}",
path.display(),
))
})?;
let meta: Self = serde_json::from_str(&content).map_err(|e| {
TensorError::new(&format!(
"CheckpointMeta: parse JSON from {}: {e}",
path.display(),
))
})?;
if meta.schema_version > CHECKPOINT_META_SCHEMA_VERSION {
return Err(TensorError::new(&format!(
"CheckpointMeta: schema version {} in {} is newer than \
the version {} this binary supports",
meta.schema_version,
path.display(),
CHECKPOINT_META_SCHEMA_VERSION,
)));
}
Ok(meta)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RankDeathRecord {
pub schema_version: u32,
pub rank: usize,
pub world_size: usize,
pub error: String,
pub unix_time_secs: u64,
}
pub const RANK_DEATH_RECORD_SCHEMA_VERSION: u32 = 1;
impl RankDeathRecord {
pub fn new(rank: usize, world_size: usize, error: String) -> Self {
let unix_time_secs = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
RankDeathRecord {
schema_version: RANK_DEATH_RECORD_SCHEMA_VERSION,
rank,
world_size,
error,
unix_time_secs,
}
}
pub fn write_to_file(&self, path: &std::path::Path) -> Result<()> {
let content = serde_json::to_string_pretty(self).map_err(|e| {
TensorError::new(&format!(
"RankDeathRecord: serialize JSON for {}: {e}",
path.display(),
))
})?;
std::fs::write(path, content).map_err(|e| {
TensorError::new(&format!(
"RankDeathRecord: write {}: {e}",
path.display(),
))
})
}
}
pub struct CheckpointBundle;
impl CheckpointBundle {
pub fn model_path(stem: &str) -> std::path::PathBuf {
let mut p = std::path::PathBuf::from(stem);
p.set_extension("fdl");
p
}
pub fn optim_path(stem: &str) -> std::path::PathBuf {
let mut p = std::path::PathBuf::from(stem);
p.set_extension("optim");
p
}
pub fn meta_path(stem: &str) -> std::path::PathBuf {
CheckpointMeta::sidecar_path(stem)
}
pub fn config_sidecar_path(stem: &str) -> std::path::PathBuf {
let mut p = std::path::PathBuf::from(stem);
p.set_extension("config.json");
p
}
pub fn rank_death_path(stem: &str, rank: usize) -> std::path::PathBuf {
let mut p = std::path::PathBuf::from(stem);
p.set_extension("");
let name = format!(
"{}.rank{rank}.death.json",
p.file_name().and_then(|s| s.to_str()).unwrap_or("checkpoint"),
);
p.set_file_name(name);
p
}
}
#[cfg(test)]
mod tests {
use super::*;
fn temp_dir(label: &str) -> std::path::PathBuf {
let dir = std::env::temp_dir().join(format!(
"flodl_meta_{label}_{}",
std::process::id()
));
std::fs::create_dir_all(&dir).unwrap();
dir
}
#[test]
fn sidecar_path_from_stem_appends_meta_json() {
let p = CheckpointMeta::sidecar_path("/tmp/run/ckpt_final");
assert_eq!(
p,
std::path::PathBuf::from("/tmp/run/ckpt_final.meta.json")
);
}
#[test]
fn rank_death_path_is_per_rank_and_shares_stem() {
assert_eq!(
CheckpointBundle::rank_death_path("/tmp/run/ckpt_final", 2),
std::path::PathBuf::from("/tmp/run/ckpt_final.rank2.death.json"),
);
assert_eq!(
CheckpointBundle::rank_death_path("/tmp/run/ckpt.fdl", 0),
std::path::PathBuf::from("/tmp/run/ckpt.rank0.death.json"),
);
assert_ne!(
CheckpointBundle::rank_death_path("ckpt", 0),
CheckpointBundle::rank_death_path("ckpt", 1),
);
}
#[test]
fn rank_death_record_round_trips_json() {
let dir = temp_dir("death");
let path = dir.join("ckpt.rank1.death.json");
let rec = RankDeathRecord::new(1, 4, "boom: cuda OOM".to_string());
rec.write_to_file(&path).unwrap();
let raw = std::fs::read_to_string(&path).unwrap();
let back: RankDeathRecord = serde_json::from_str(&raw).unwrap();
assert_eq!(back.rank, 1);
assert_eq!(back.world_size, 4);
assert_eq!(back.error, "boom: cuda OOM");
assert_eq!(back.schema_version, RANK_DEATH_RECORD_SCHEMA_VERSION);
}
#[test]
fn sidecar_path_from_fdl_strips_extension() {
let p = CheckpointMeta::sidecar_path("/tmp/run/ckpt_final.fdl");
assert_eq!(
p,
std::path::PathBuf::from("/tmp/run/ckpt_final.meta.json")
);
}
#[test]
fn roundtrip_json_preserves_all_fields() {
let dir = temp_dir("roundtrip");
let path = dir.join("ckpt.meta.json");
let meta = CheckpointMeta::new(
3,
7500,
12,
4,
SaveReason::MaxFailureExceeded,
);
meta.write_to_file(&path).unwrap();
let loaded = CheckpointMeta::read_from_file(&path).unwrap();
assert_eq!(loaded.epoch, 3);
assert_eq!(loaded.global_step, 7500);
assert_eq!(loaded.sync_round, 12);
assert_eq!(loaded.world_size_at_save, 4);
assert_eq!(loaded.save_reason, SaveReason::MaxFailureExceeded);
assert_eq!(loaded.schema_version, CHECKPOINT_META_SCHEMA_VERSION);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn save_reason_serializes_snake_case() {
let meta = CheckpointMeta::new(0, 0, 0, 1, SaveReason::SingleSurvivor);
let json = serde_json::to_string(&meta).unwrap();
assert!(
json.contains("\"save_reason\":\"single_survivor\""),
"expected snake_case save_reason in JSON, got: {json}",
);
}
#[test]
fn bundle_paths_share_stem() {
let stem = "/tmp/run/ckpt_final";
assert_eq!(
CheckpointBundle::model_path(stem),
std::path::PathBuf::from("/tmp/run/ckpt_final.fdl"),
);
assert_eq!(
CheckpointBundle::optim_path(stem),
std::path::PathBuf::from("/tmp/run/ckpt_final.optim"),
);
assert_eq!(
CheckpointBundle::meta_path(stem),
std::path::PathBuf::from("/tmp/run/ckpt_final.meta.json"),
);
assert_eq!(
CheckpointBundle::config_sidecar_path(stem),
std::path::PathBuf::from("/tmp/run/ckpt_final.config.json"),
);
}
#[test]
fn bundle_paths_strip_existing_extension() {
let stem = "/tmp/run/ckpt.fdl";
assert_eq!(
CheckpointBundle::model_path(stem),
std::path::PathBuf::from("/tmp/run/ckpt.fdl"),
);
assert_eq!(
CheckpointBundle::optim_path(stem),
std::path::PathBuf::from("/tmp/run/ckpt.optim"),
);
assert_eq!(
CheckpointBundle::meta_path(stem),
std::path::PathBuf::from("/tmp/run/ckpt.meta.json"),
);
}
#[test]
fn save_reason_u8_roundtrip_for_all_variants() {
for r in [
SaveReason::GracefulShutdown,
SaveReason::MaxFailureExceeded,
SaveReason::SingleSurvivor,
SaveReason::AllRanksLost,
SaveReason::ReduceStall,
SaveReason::Checkpoint,
] {
let byte = r.to_u8();
assert_eq!(SaveReason::from_u8(byte), Some(r));
}
}
#[test]
fn save_reason_from_unknown_byte_is_none() {
assert_eq!(SaveReason::from_u8(99), None);
assert_eq!(SaveReason::from_u8(255), None);
}
#[test]
fn v1_file_loads_with_elche_state_defaulted_to_none() {
let dir = temp_dir("v1_forward_compat");
let path = dir.join("ckpt.meta.json");
let raw_json = r#"{
"schema_version": 1,
"epoch": 5,
"global_step": 10000,
"sync_round": 25,
"world_size_at_save": 2,
"save_reason": "graceful_shutdown"
}"#;
std::fs::write(&path, raw_json).unwrap();
let loaded = CheckpointMeta::read_from_file(&path).unwrap();
assert_eq!(loaded.schema_version, 1);
assert_eq!(loaded.epoch, 5);
assert_eq!(loaded.elche_state, None);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn roundtrip_preserves_elche_state() {
let dir = temp_dir("elche_roundtrip");
let path = dir.join("ckpt.meta.json");
let state = ElCheState {
anchor: 12,
anchor_rank: Some(1),
smoothed_ms_per_batch: vec![5.0, 2.5, 4.0],
phase: crate::distributed::el_che::Phase::Stable,
calibration_count: 42,
trend_history: Some(vec![0.01, 0.015, 0.02, 0.025, 0.03]),
};
let meta = CheckpointMeta::new(
3,
7500,
12,
3,
SaveReason::GracefulShutdown,
)
.with_elche_state(state.clone());
meta.write_to_file(&path).unwrap();
let loaded = CheckpointMeta::read_from_file(&path).unwrap();
assert_eq!(loaded.elche_state, Some(state));
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn v2_file_loads_with_trend_history_defaulted_to_none() {
let dir = temp_dir("v2_forward_compat");
let path = dir.join("ckpt.meta.json");
let raw_json = r#"{
"schema_version": 2,
"epoch": 5,
"global_step": 10000,
"sync_round": 25,
"world_size_at_save": 2,
"save_reason": "graceful_shutdown",
"elche_state": {
"anchor": 8,
"anchor_rank": 0,
"smoothed_ms_per_batch": [3.0, 5.0],
"phase": "stable",
"calibration_count": 17
}
}"#;
std::fs::write(&path, raw_json).unwrap();
let loaded = CheckpointMeta::read_from_file(&path).unwrap();
assert_eq!(loaded.schema_version, 2);
let elche = loaded.elche_state.expect("v2 file carries elche_state");
assert_eq!(elche.anchor, 8);
assert_eq!(elche.trend_history, None);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn elche_restore_from_state_roundtrip() {
use crate::distributed::ElChe;
let mut original = ElChe::new(3, 8).with_overhead_target(0.05);
for _ in 0..3 {
original.report_timing(&[5.0, 7.0, 6.0], &[8, 8, 8], 0.5);
}
let snap = original.to_state();
assert!(snap.anchor_rank.is_some(), "calibrated run elects an anchor");
assert!(
snap.smoothed_ms_per_batch.iter().any(|&v| v > 0.0),
"calibrated run has positive smoothed readings"
);
let mut restored = ElChe::new(3, 8).with_overhead_target(0.20);
restored.restore_from_state(&snap).unwrap();
assert_eq!(restored.to_state().anchor, snap.anchor);
assert_eq!(restored.to_state().anchor_rank, snap.anchor_rank);
assert_eq!(restored.to_state().phase, snap.phase);
assert_eq!(
restored.to_state().calibration_count,
snap.calibration_count
);
assert_eq!(
restored.to_state().smoothed_ms_per_batch,
snap.smoothed_ms_per_batch
);
assert!(restored.is_calibrated(), "restored from positive-reading snap");
}
#[test]
fn elche_restore_from_state_rejects_world_size_mismatch() {
use crate::distributed::ElChe;
let mut three_rank = ElChe::new(3, 8);
for _ in 0..3 {
three_rank.report_timing(&[5.0, 6.0, 7.0], &[8, 8, 8], 0.5);
}
let snap = three_rank.to_state();
let mut two_rank = ElChe::new(2, 8);
let err = two_rank.restore_from_state(&snap).unwrap_err();
assert!(
err.to_string().contains("world_size"),
"expected world_size mismatch error, got: {err}"
);
}
#[test]
fn coord_config_resume_from_meta_applies_all_fields() {
use crate::distributed::cluster_coordinator::ClusterCoordinatorConfig;
use crate::distributed::ddp::ElChe;
use crate::distributed::ddp_run::{ApplyPolicy, AverageBackend};
let elche_state = ElCheState {
anchor: 16,
anchor_rank: Some(2),
smoothed_ms_per_batch: vec![4.0, 5.0, 6.0],
phase: crate::distributed::el_che::Phase::Mature,
calibration_count: 99,
trend_history: Some(vec![0.01, 0.02, 0.03]),
};
let meta = CheckpointMeta::new(
7,
42_000,
201,
3,
SaveReason::GracefulShutdown,
)
.with_elche_state(elche_state.clone());
let config = ClusterCoordinatorConfig::new(
ApplyPolicy::Cadence,
AverageBackend::Cpu,
3,
ElChe::new(3, 8),
)
.resume_from_meta(&meta);
assert_eq!(config.start_epoch, 7);
assert_eq!(config.start_global_step, 42_000);
assert_eq!(config.start_avg_count, 201);
assert_eq!(config.start_elche_state, Some(elche_state));
}
#[test]
fn v3_file_loads_with_coverage_defaulted_to_none() {
let dir = temp_dir("v3_forward_compat");
let path = dir.join("ckpt.meta.json");
let raw_json = r#"{
"schema_version": 3,
"epoch": 5,
"global_step": 10000,
"sync_round": 25,
"world_size_at_save": 2,
"save_reason": "graceful_shutdown",
"elche_state": {
"anchor": 8,
"anchor_rank": 0,
"smoothed_ms_per_batch": [3.0, 5.0],
"phase": "stable",
"calibration_count": 17,
"trend_history": [0.01, 0.02, 0.03]
}
}"#;
std::fs::write(&path, raw_json).unwrap();
let loaded = CheckpointMeta::read_from_file(&path).unwrap();
assert_eq!(loaded.schema_version, 3);
assert!(loaded.elche_state.is_some());
assert_eq!(loaded.coverage, None);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn roundtrip_preserves_coverage() {
let dir = temp_dir("coverage_roundtrip");
let path = dir.join("ckpt.meta.json");
let coverage = CoverageBlock {
seed: 42,
batch_size: 32,
per_epoch: vec![
EpochCoverage {
epoch: 7,
total_samples: 1024,
uncovered_ranges: vec![(320, 64), (640, 384)],
},
EpochCoverage {
epoch: 8,
total_samples: 1024,
uncovered_ranges: vec![(0, 1024)],
},
],
};
let meta = CheckpointMeta::new(7, 7500, 12, 3, SaveReason::GracefulShutdown)
.with_coverage(coverage.clone());
meta.write_to_file(&path).unwrap();
let loaded = CheckpointMeta::read_from_file(&path).unwrap();
assert_eq!(loaded.coverage, Some(coverage));
assert_eq!(loaded.schema_version, CHECKPOINT_META_SCHEMA_VERSION);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn model_schema_from_module_captures_ordered_names() {
use crate::nn::Linear;
use crate::tensor::Device;
let lin = Linear::on_device(4, 2, Device::CPU).unwrap();
let schema = ModelSchema::from_module(&lin);
assert_eq!(
schema.param_names,
vec!["weight".to_string(), "bias".to_string()],
"Linear params in declaration order"
);
assert!(schema.buffer_names.is_empty(), "Linear has no buffers");
assert_eq!(schema.tensor_count(), 2);
}
#[test]
fn future_schema_version_rejected() {
let dir = temp_dir("future");
let path = dir.join("ckpt.meta.json");
let raw_json = r#"{
"schema_version": 99999,
"epoch": 0,
"global_step": 0,
"sync_round": 0,
"world_size_at_save": 1,
"save_reason": "graceful_shutdown"
}"#;
std::fs::write(&path, raw_json).unwrap();
let err = CheckpointMeta::read_from_file(&path).unwrap_err();
assert!(
err.to_string().contains("newer than"),
"expected schema-version error, got: {err}",
);
std::fs::remove_dir_all(&dir).ok();
}
}