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:
@@ -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(§ion_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.
|
||||
|
||||
Reference in New Issue
Block a user