feat(rl): IntegratedTrainer checkpoint save/load

Binary format: magic + version + step + sections. Sections include
ISV (585 f32), 26 device buffers (DQN/policy/value/IQN/FRD/NoisyNet/
outcome weights + targets), 20 AdamW optimizer states, and encoder
(delegated to CfcTrunk). All device transfers use mapped-pinned DtoD.
load_checkpoint invalidates CUDA graphs to force re-capture.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-27 20:29:26 +02:00
parent a1277af6c7
commit 66ec7f75f4

View File

@@ -7719,6 +7719,371 @@ impl IntegratedTrainer {
}
Ok(probs)
}
// ── Checkpoint save / load ──────────────────────────────────────────
/// Persist all trainable weights, ISV state, and AdamW optimizer state
/// to a binary checkpoint file.
///
/// Binary format:
/// ```text
/// [magic: u32 = 0x464F5843] // "FOXC"
/// [version: u32 = 1]
/// [step: u64]
/// [n_sections: u32]
/// For each section:
/// [name_len: u16][name: UTF-8 bytes][data_len: u64][data: raw bytes]
/// ```
///
/// The encoder (perception trunk) is saved to a sibling file with
/// `.encoder` extension via `PerceptionTrainer::save_checkpoint`.
///
/// All device→host transfers use mapped-pinned staging + DtoD per
/// `feedback_no_htod_htoh_only_mapped_pinned`.
pub fn save_checkpoint(&self, path: &std::path::Path, step: u64) -> Result<()> {
use std::io::Write;
const MAGIC: u32 = 0x464F5843; // "FOXC"
const VERSION: u32 = 1;
// ── Encoder checkpoint (own format / sibling file) ───────────
let encoder_path = path.with_extension("encoder");
self.perception
.save_checkpoint(&encoder_path)
.context("save encoder checkpoint")?;
// ── Collect device-buffer sections ───────────────────────────
// (name, CUdeviceptr, len_in_f32s)
let device_sections: Vec<(&str, u64, usize)> = vec![
("dqn_w", self.dqn_head.w_d.raw_ptr(), self.dqn_head.w_d.len()),
("dqn_b", self.dqn_head.b_d.raw_ptr(), self.dqn_head.b_d.len()),
("dqn_w_target", self.dqn_head.w_target_d.raw_ptr(), self.dqn_head.w_target_d.len()),
("dqn_b_target", self.dqn_head.b_target_d.raw_ptr(), self.dqn_head.b_target_d.len()),
("policy_w", self.policy_head.w_d.raw_ptr(), self.policy_head.w_d.len()),
("policy_b", self.policy_head.b_d.raw_ptr(), self.policy_head.b_d.len()),
("value_w", self.value_head.w_d.raw_ptr(), self.value_head.w_d.len()),
("value_b", self.value_head.b_d.raw_ptr(), self.value_head.b_d.len()),
("iqn_w_embed", self.iqn_head.w_embed_d.raw_ptr(), self.iqn_head.w_embed_d.len()),
("iqn_b_embed", self.iqn_head.b_embed_d.raw_ptr(), self.iqn_head.b_embed_d.len()),
("iqn_w_out", self.iqn_head.w_out_d.raw_ptr(), self.iqn_head.w_out_d.len()),
("iqn_b_out", self.iqn_head.b_out_d.raw_ptr(), self.iqn_head.b_out_d.len()),
("iqn_w_embed_target", self.iqn_head.w_embed_target_d.raw_ptr(), self.iqn_head.w_embed_target_d.len()),
("iqn_b_embed_target", self.iqn_head.b_embed_target_d.raw_ptr(), self.iqn_head.b_embed_target_d.len()),
("iqn_w_out_target", self.iqn_head.w_out_target_d.raw_ptr(), self.iqn_head.w_out_target_d.len()),
("iqn_b_out_target", self.iqn_head.b_out_target_d.raw_ptr(), self.iqn_head.b_out_target_d.len()),
("frd_w1", self.frd_head.w1_d.raw_ptr(), self.frd_head.w1_d.len()),
("frd_b1", self.frd_head.b1_d.raw_ptr(), self.frd_head.b1_d.len()),
("frd_w2", self.frd_head.w2_d.raw_ptr(), self.frd_head.w2_d.len()),
("frd_b2", self.frd_head.b2_d.raw_ptr(), self.frd_head.b2_d.len()),
("noisy_mu_w", self.noisy_exploration.mu_w.raw_ptr(), self.noisy_exploration.mu_w.len()),
("noisy_sigma_w", self.noisy_exploration.sigma_w.raw_ptr(), self.noisy_exploration.sigma_w.len()),
("noisy_mu_b", self.noisy_exploration.mu_b.raw_ptr(), self.noisy_exploration.mu_b.len()),
("noisy_sigma_b", self.noisy_exploration.sigma_b.raw_ptr(), self.noisy_exploration.sigma_b.len()),
("outcome_w", self.outcome_head.w_d.raw_ptr(), self.outcome_head.w_d.len()),
("outcome_b", self.outcome_head.b_d.raw_ptr(), self.outcome_head.b_d.len()),
];
// Find the largest device buffer to size the staging area.
let max_len = device_sections
.iter()
.map(|(_, _, n)| *n)
.max()
.unwrap_or(0)
.max(RL_SLOTS_END);
let staging = unsafe { MappedF32Buffer::new(max_len) }
.map_err(|e| anyhow::anyhow!("checkpoint staging alloc: {e}"))?;
// ISV (1) + device buffers (26) + adam (20) = 47 sections.
let n_sections: u32 = 1 + device_sections.len() as u32 + 20;
let mut file = std::io::BufWriter::new(
std::fs::File::create(path)
.with_context(|| format!("create checkpoint {}", path.display()))?,
);
// ── Header ──────────────────────────────────────────────────
file.write_all(&MAGIC.to_le_bytes())?;
file.write_all(&VERSION.to_le_bytes())?;
file.write_all(&step.to_le_bytes())?;
file.write_all(&n_sections.to_le_bytes())?;
// Helper: write a section header (name_len + name + data_len).
let write_section_header =
|w: &mut std::io::BufWriter<std::fs::File>, name: &str, data_len: u64| -> Result<()> {
let name_bytes = name.as_bytes();
w.write_all(&(name_bytes.len() as u16).to_le_bytes())?;
w.write_all(name_bytes)?;
w.write_all(&data_len.to_le_bytes())?;
Ok(())
};
// ── Section 1: ISV (already host-resident via mapped-pinned) ─
{
let isv_len = RL_SLOTS_END;
let isv_bytes = (isv_len * std::mem::size_of::<f32>()) as u64;
write_section_header(&mut file, "isv", isv_bytes)?;
// Volatile reads for each element.
for i in 0..isv_len {
let val = unsafe { std::ptr::read_volatile(self.isv_mapped.host_ptr.add(i)) };
file.write_all(&val.to_le_bytes())?;
}
}
// ── Device-buffer sections (DtoD → staging → file) ──────────
let raw_s = self.raw_stream;
for (name, src_ptr, n_f32) in &device_sections {
let byte_len = (*n_f32 * std::mem::size_of::<f32>()) as u64;
write_section_header(&mut file, name, byte_len)?;
if *n_f32 > 0 {
unsafe {
raw_memcpy_dtod_async(
staging.dev_ptr,
*src_ptr,
*n_f32 * std::mem::size_of::<f32>(),
raw_s,
)
.map_err(|e| anyhow::anyhow!("checkpoint save DtoD {name}: {e:?}"))?;
raw_stream_sync(raw_s)
.map_err(|e| anyhow::anyhow!("checkpoint save sync {name}: {e:?}"))?;
}
let host_data = unsafe {
std::slice::from_raw_parts(staging.host_ptr, *n_f32)
};
file.write_all(bytemuck::cast_slice(host_data))?;
}
}
// ── AdamW sections ──────────────────────────────────────────
let adam_pairs: Vec<(&str, &AdamW)> = vec![
("adam_dqn_w", &self.dqn_w_adam),
("adam_dqn_b", &self.dqn_b_adam),
("adam_iqn_w_embed", &self.iqn_w_embed_adam),
("adam_iqn_b_embed", &self.iqn_b_embed_adam),
("adam_iqn_w_out", &self.iqn_w_out_adam),
("adam_iqn_b_out", &self.iqn_b_out_adam),
("adam_policy_w", &self.policy_w_adam),
("adam_policy_b", &self.policy_b_adam),
("adam_value_w", &self.value_w_adam),
("adam_value_b", &self.value_b_adam),
("adam_frd_w1", &self.frd_w1_adam),
("adam_frd_b1", &self.frd_b1_adam),
("adam_frd_w2", &self.frd_w2_adam),
("adam_frd_b2", &self.frd_b2_adam),
("adam_noisy_mu_w", &self.noisy_mu_w_adam),
("adam_noisy_sigma_w", &self.noisy_sigma_w_adam),
("adam_noisy_mu_b", &self.noisy_mu_b_adam),
("adam_noisy_sigma_b", &self.noisy_sigma_b_adam),
("adam_outcome_w", &self.outcome_w_adam),
("adam_outcome_b", &self.outcome_b_adam),
];
for (name, adam) in &adam_pairs {
// Serialize to a temp buffer to measure length, then write header + data.
let mut buf = Vec::new();
adam.save(&mut buf)
.with_context(|| format!("AdamW save {name}"))?;
write_section_header(&mut file, name, buf.len() as u64)?;
file.write_all(&buf)?;
}
file.flush()?;
Ok(())
}
/// Restore all trainable weights, ISV state, and AdamW optimizer state
/// from a binary checkpoint file. Returns the step number embedded in
/// the checkpoint header.
///
/// Unknown section names are silently skipped for forward compatibility.
///
/// All host→device transfers use mapped-pinned staging + DtoD per
/// `feedback_no_htod_htoh_only_mapped_pinned`.
pub fn load_checkpoint(&mut self, dev: &MlDevice, path: &std::path::Path) -> Result<u64> {
use std::io::Read;
const MAGIC: u32 = 0x464F5843;
const VERSION: u32 = 1;
// ── Encoder checkpoint (sibling file) ────────────────────────
let encoder_path = path.with_extension("encoder");
if encoder_path.exists() {
let trunk_cfg = crate::cfc::trunk::CfcConfig {
n_in: crate::cfc::snap_features::ENCODER_INPUT_DIM,
n_hid: HIDDEN_DIM,
cfc_n_in: HIDDEN_DIM,
mamba2_state_dim: self.perception.config().mamba2_state_dim,
n_batch: self.perception.config().n_batch,
seq_len: self.perception.config().seq_len,
};
let new_trunk =
crate::cfc::trunk::CfcTrunk::load_checkpoint(dev, &trunk_cfg, &encoder_path)
.context("load encoder checkpoint")?;
self.perception.trunk = new_trunk;
}
let mut file = std::io::BufReader::new(
std::fs::File::open(path)
.with_context(|| format!("open checkpoint {}", path.display()))?,
);
// ── Header ──────────────────────────────────────────────────
let mut buf4 = [0u8; 4];
let mut buf8 = [0u8; 8];
let mut buf2 = [0u8; 2];
file.read_exact(&mut buf4)?;
let magic = u32::from_le_bytes(buf4);
anyhow::ensure!(magic == MAGIC, "bad checkpoint magic: {magic:#010x}");
file.read_exact(&mut buf4)?;
let version = u32::from_le_bytes(buf4);
anyhow::ensure!(version == VERSION, "unsupported checkpoint version: {version}");
file.read_exact(&mut buf8)?;
let step = u64::from_le_bytes(buf8);
file.read_exact(&mut buf4)?;
let n_sections = u32::from_le_bytes(buf4);
// Build a lookup from section name to (device ptr, f32 count) for
// device-buffer restoration.
let device_lookup: std::collections::HashMap<&str, (u64, usize)> = [
("dqn_w", (self.dqn_head.w_d.raw_ptr(), self.dqn_head.w_d.len())),
("dqn_b", (self.dqn_head.b_d.raw_ptr(), self.dqn_head.b_d.len())),
("dqn_w_target", (self.dqn_head.w_target_d.raw_ptr(), self.dqn_head.w_target_d.len())),
("dqn_b_target", (self.dqn_head.b_target_d.raw_ptr(), self.dqn_head.b_target_d.len())),
("policy_w", (self.policy_head.w_d.raw_ptr(), self.policy_head.w_d.len())),
("policy_b", (self.policy_head.b_d.raw_ptr(), self.policy_head.b_d.len())),
("value_w", (self.value_head.w_d.raw_ptr(), self.value_head.w_d.len())),
("value_b", (self.value_head.b_d.raw_ptr(), self.value_head.b_d.len())),
("iqn_w_embed", (self.iqn_head.w_embed_d.raw_ptr(), self.iqn_head.w_embed_d.len())),
("iqn_b_embed", (self.iqn_head.b_embed_d.raw_ptr(), self.iqn_head.b_embed_d.len())),
("iqn_w_out", (self.iqn_head.w_out_d.raw_ptr(), self.iqn_head.w_out_d.len())),
("iqn_b_out", (self.iqn_head.b_out_d.raw_ptr(), self.iqn_head.b_out_d.len())),
("iqn_w_embed_target", (self.iqn_head.w_embed_target_d.raw_ptr(), self.iqn_head.w_embed_target_d.len())),
("iqn_b_embed_target", (self.iqn_head.b_embed_target_d.raw_ptr(), self.iqn_head.b_embed_target_d.len())),
("iqn_w_out_target", (self.iqn_head.w_out_target_d.raw_ptr(), self.iqn_head.w_out_target_d.len())),
("iqn_b_out_target", (self.iqn_head.b_out_target_d.raw_ptr(), self.iqn_head.b_out_target_d.len())),
("frd_w1", (self.frd_head.w1_d.raw_ptr(), self.frd_head.w1_d.len())),
("frd_b1", (self.frd_head.b1_d.raw_ptr(), self.frd_head.b1_d.len())),
("frd_w2", (self.frd_head.w2_d.raw_ptr(), self.frd_head.w2_d.len())),
("frd_b2", (self.frd_head.b2_d.raw_ptr(), self.frd_head.b2_d.len())),
("noisy_mu_w", (self.noisy_exploration.mu_w.raw_ptr(), self.noisy_exploration.mu_w.len())),
("noisy_sigma_w", (self.noisy_exploration.sigma_w.raw_ptr(), self.noisy_exploration.sigma_w.len())),
("noisy_mu_b", (self.noisy_exploration.mu_b.raw_ptr(), self.noisy_exploration.mu_b.len())),
("noisy_sigma_b", (self.noisy_exploration.sigma_b.raw_ptr(), self.noisy_exploration.sigma_b.len())),
("outcome_w", (self.outcome_head.w_d.raw_ptr(), self.outcome_head.w_d.len())),
("outcome_b", (self.outcome_head.b_d.raw_ptr(), self.outcome_head.b_d.len())),
]
.into_iter()
.collect();
// Staging buffer sized to the largest device section.
let max_staging = device_lookup.values().map(|(_, n)| *n).max().unwrap_or(0).max(RL_SLOTS_END);
let staging = unsafe { MappedF32Buffer::new(max_staging) }
.map_err(|e| anyhow::anyhow!("checkpoint load staging alloc: {e}"))?;
let raw_s = self.raw_stream;
for _ in 0..n_sections {
// Read section header.
file.read_exact(&mut buf2)?;
let name_len = u16::from_le_bytes(buf2) as usize;
let mut name_buf = vec![0u8; name_len];
file.read_exact(&mut name_buf)?;
let name = String::from_utf8(name_buf)
.context("checkpoint section name not valid UTF-8")?;
file.read_exact(&mut buf8)?;
let data_len = u64::from_le_bytes(buf8) as usize;
if name == "isv" {
// ISV: write directly to isv_mapped host_ptr via
// write_volatile, then DtoD host→device is unnecessary
// because isv_mapped IS the device memory (mapped-pinned).
let isv_n = RL_SLOTS_END;
anyhow::ensure!(
data_len == isv_n * std::mem::size_of::<f32>(),
"ISV section size mismatch: file {data_len} vs expected {}",
isv_n * std::mem::size_of::<f32>()
);
for i in 0..isv_n {
let mut vbuf = [0u8; 4];
file.read_exact(&mut vbuf)?;
let val = f32::from_le_bytes(vbuf);
unsafe {
std::ptr::write_volatile(self.isv_mapped.host_ptr.add(i), val);
}
}
} else if let Some(&(dst_ptr, expected_n)) = device_lookup.get(name.as_str()) {
// Device buffer: read file → staging host_ptr → DtoD → device.
let expected_bytes = expected_n * std::mem::size_of::<f32>();
anyhow::ensure!(
data_len == expected_bytes,
"section '{name}' size mismatch: file {data_len} vs expected {expected_bytes}"
);
if expected_n > 0 {
let host_slice = unsafe {
std::slice::from_raw_parts_mut(staging.host_ptr, expected_n)
};
file.read_exact(bytemuck::cast_slice_mut(host_slice))?;
unsafe {
raw_memcpy_dtod_async(
dst_ptr,
staging.dev_ptr,
expected_bytes,
raw_s,
)
.map_err(|e| anyhow::anyhow!("checkpoint load DtoD {name}: {e:?}"))?;
raw_stream_sync(raw_s)
.map_err(|e| anyhow::anyhow!("checkpoint load sync {name}: {e:?}"))?;
}
}
} else if name.starts_with("adam_") {
// AdamW section: dispatch to the matching instance's load().
let mut section_data = vec![0u8; data_len];
file.read_exact(&mut section_data)?;
let mut cursor = std::io::Cursor::new(&section_data[..]);
match name.as_str() {
"adam_dqn_w" => self.dqn_w_adam.load(&mut cursor)?,
"adam_dqn_b" => self.dqn_b_adam.load(&mut cursor)?,
"adam_iqn_w_embed" => self.iqn_w_embed_adam.load(&mut cursor)?,
"adam_iqn_b_embed" => self.iqn_b_embed_adam.load(&mut cursor)?,
"adam_iqn_w_out" => self.iqn_w_out_adam.load(&mut cursor)?,
"adam_iqn_b_out" => self.iqn_b_out_adam.load(&mut cursor)?,
"adam_policy_w" => self.policy_w_adam.load(&mut cursor)?,
"adam_policy_b" => self.policy_b_adam.load(&mut cursor)?,
"adam_value_w" => self.value_w_adam.load(&mut cursor)?,
"adam_value_b" => self.value_b_adam.load(&mut cursor)?,
"adam_frd_w1" => self.frd_w1_adam.load(&mut cursor)?,
"adam_frd_b1" => self.frd_b1_adam.load(&mut cursor)?,
"adam_frd_w2" => self.frd_w2_adam.load(&mut cursor)?,
"adam_frd_b2" => self.frd_b2_adam.load(&mut cursor)?,
"adam_noisy_mu_w" => self.noisy_mu_w_adam.load(&mut cursor)?,
"adam_noisy_sigma_w" => self.noisy_sigma_w_adam.load(&mut cursor)?,
"adam_noisy_mu_b" => self.noisy_mu_b_adam.load(&mut cursor)?,
"adam_noisy_sigma_b" => self.noisy_sigma_b_adam.load(&mut cursor)?,
"adam_outcome_w" => self.outcome_w_adam.load(&mut cursor)?,
"adam_outcome_b" => self.outcome_b_adam.load(&mut cursor)?,
_ => { /* unknown adam section — skip for forward compat */ }
}
} else {
// Unknown section — skip for forward compatibility.
let mut discard = vec![0u8; data_len];
file.read_exact(&mut discard)?;
}
}
// Invalidate captured CUDA graphs — weight addresses haven't moved
// but values changed, so any graph that baked in the old data needs
// re-capture on the next step.
self.prefill_graph = None;
self.postfill_graph = None;
self.reward_graph = None;
self.training_graph = None;
self.graph_warmup_done = false;
self.training_warmup_done = false;
Ok(step)
}
}
/// Free-function entry point for the `grad_h_accumulate_scaled` kernel.