diff --git a/crates/ml-alpha/src/trainer/integrated.rs b/crates/ml-alpha/src/trainer/integrated.rs index 003420968..c46fdb99a 100644 --- a/crates/ml-alpha/src/trainer/integrated.rs +++ b/crates/ml-alpha/src/trainer/integrated.rs @@ -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, 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::()) 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::()) 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::(), + 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 { + 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::(), + "ISV section size mismatch: file {data_len} vs expected {}", + isv_n * std::mem::size_of::() + ); + 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::(); + 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.