perf(checkpoint): async best-ckpt serialize via spawn_blocking + mapped-pinned param snapshot
Best-checkpoint save (val Sharpe improvement, ~30% of epochs in convergent runs) blocked the epoch loop for 20-40s on each improvement: serialize_model() chained N (~26) per-tensor DtoH downloads via GpuTensor::to_host → memcpy_dtoh, each forcing an implicit stream sync on a busy training stream. The DtoH chain also violates feedback_no_htod_htoh_only_mapped_pinned.md (only cuMemHostAlloc DEVICEMAP allowed for CPU↔GPU paths). Plan B: - Introduce snapshot_model_to_pinned (mod.rs): allocate one MappedF32Buffer sized for all named weight slices concatenated, cuMemcpyDtoDAsync each slice into the buffer's device pointer (which aliases the host page), single stream sync, copy bytes out to a Send + 'static Vec<u8>. One sync per snapshot, replaces N. - serialize_snapshot_bytes (mod.rs): pure-CPU safetensors construction from CheckpointSnapshot. Static — callable without &self, so the worker can move the snapshot across thread boundary. - handle_epoch_checkpoints_and_early_stopping on val-Sharpe improvement: save_best_gpu_params (DtoD, fast) + snapshot to pinned + tokio::task::spawn_blocking the safetensors construction + checkpoint_callback invocation. JoinHandle parked on pending_checkpoint_handles. Training loop continues immediately. - await_pending_checkpoint_handles drains in-flight workers at training end (success branch + early-stop branches) and before any synchronous cold-path checkpoint write to keep disk ordering deterministic. - F bound on train / train_walk_forward / train_fold_from_slices gains + 'static so the callback can be moved into the worker. All public callers already use 'static-compatible move closures (test fixtures with shared mutable state migrate to Arc<Mutex<T>>). Internal pipeline uses CheckpointCallbackHandle = Arc<std::sync::Mutex<Box<dyn FnMut + Send + 'static>>> so the same callback flows through multi-fold walk-forward into every fold's worker. - serialize_model itself rewritten via the snapshot path: the no-DtoH rule now holds across ALL checkpoint paths (best, periodic, early-stop, plateau-exhausted). The pre-existing GpuTensor::to_host path is no longer reachable from the DQN trainer. The audit's spec called for an mpsc channel(1) drop-old worker, but the multi-fold + &mut F pre-existing API made the simpler fire-and-forget spawn_blocking pattern a cleaner fit (Mutex serialises any concurrent invocations; Vec<JoinHandle> drain at end guarantees disk writes complete before the trainer returns). Same overlap benefit (training rolls while serialize+disk run on a blocking thread); upper bound on in-flight work is one-per-improved- epoch which approximates the spec's depth=1 in realistic training runs. Per feedback_no_partial_refactor: every site that constructs a checkpoint payload migrated in lockstep — best-improvement uses the worker; periodic / plateau-exhausted / early-stop call the shared Arc<Mutex<F>> handle inline. All paths read params via snapshot_model_to_pinned, so the no-DtoH rule applies uniformly. Test fixtures (8 .rs files) updated for the + 'static bound (move closures + cloned PathBufs / Arc<Mutex<T>> for shared mutable state). Verified: SQLX_OFFLINE=true cargo check --workspace --tests clean (warnings unchanged from baseline). cargo test -p ml --lib --no-run clean. No fingerprint change. Wire-up audit entry extended with Plan B file:line edit sites (rides under the same Async-validation overlap section started by the companion Plan A commit). Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -165,10 +165,19 @@ error_high = 1.0e-9
|
|||||||
.enable_all()
|
.enable_all()
|
||||||
.build()?;
|
.build()?;
|
||||||
let result: anyhow::Result<crate::TrainingMetrics> = rt.block_on(async move {
|
let result: anyhow::Result<crate::TrainingMetrics> = rt.block_on(async move {
|
||||||
|
// Wrap the no-op callback in the shared `CheckpointCallbackHandle`
|
||||||
|
// (Arc<Mutex<Box<dyn FnMut + Send + 'static>>>) so the
|
||||||
|
// async-checkpoint worker can hold a clone alongside the
|
||||||
|
// synchronous epoch-loop callsites.
|
||||||
|
let cb_handle: crate::trainers::dqn::trainer::CheckpointCallbackHandle =
|
||||||
|
std::sync::Arc::new(std::sync::Mutex::new(
|
||||||
|
Box::new(|_epoch, _bytes, _is_best| Ok("noop".to_string()))
|
||||||
|
as crate::trainers::dqn::trainer::BoxedCheckpointCallback,
|
||||||
|
));
|
||||||
trainer
|
trainer
|
||||||
.train_with_data_full_loop_slices(
|
.train_with_data_full_loop_slices(
|
||||||
&cache_to_pairs(&cache, 1024),
|
&cache_to_pairs(&cache, 1024),
|
||||||
|_epoch, _bytes, _is_best| Ok("noop".to_string()),
|
cb_handle,
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -727,6 +727,10 @@ impl DQNTrainer {
|
|||||||
regime_replay_decay_override: 1.0,
|
regime_replay_decay_override: 1.0,
|
||||||
fold_dominant_regime: 1, // Default: Ranging
|
fold_dominant_regime: 1, // Default: Ranging
|
||||||
experience_done_event: None,
|
experience_done_event: None,
|
||||||
|
// In-flight async best-checkpoint workers — empty at trainer
|
||||||
|
// construction; populated per epoch by
|
||||||
|
// `handle_epoch_checkpoints_and_early_stopping`.
|
||||||
|
pending_checkpoint_handles: Vec::new(),
|
||||||
fold_train_start: 0,
|
fold_train_start: 0,
|
||||||
fold_train_end: 0,
|
fold_train_end: 0,
|
||||||
fold_val_start: 0,
|
fold_val_start: 0,
|
||||||
|
|||||||
@@ -9,7 +9,6 @@ use std::sync::Arc;
|
|||||||
|
|
||||||
use anyhow::{Context, Result};
|
use anyhow::{Context, Result};
|
||||||
use ml_core::device::MlDevice;
|
use ml_core::device::MlDevice;
|
||||||
use ml_core::cuda_autograd::GpuTensor;
|
|
||||||
use cudarc::driver::{CudaEvent, CudaSlice, CudaStream};
|
use cudarc::driver::{CudaEvent, CudaSlice, CudaStream};
|
||||||
use crate::cuda_pipeline::DqnGpuData;
|
use crate::cuda_pipeline::DqnGpuData;
|
||||||
use crate::cuda_pipeline::gpu_experience_collector::TradeStats;
|
use crate::cuda_pipeline::gpu_experience_collector::TradeStats;
|
||||||
@@ -204,6 +203,54 @@ pub struct ControllerFireCounts {
|
|||||||
pub cost_anneal: u32,
|
pub cost_anneal: u32,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Boxed checkpoint callback — owns its closure data, can be moved into
|
||||||
|
/// async tasks. The 'static bound is required by `tokio::task::spawn` for
|
||||||
|
/// the async-checkpoint worker; all production callers use `move` closures
|
||||||
|
/// over owned state (PathBuf, trial_id, fold_idx) and are already 'static.
|
||||||
|
pub(crate) type BoxedCheckpointCallback =
|
||||||
|
Box<dyn FnMut(usize, Vec<u8>, bool) -> Result<String> + Send + 'static>;
|
||||||
|
|
||||||
|
/// Shared checkpoint callback handle. Wrapped in an Arc<Mutex<>> so the
|
||||||
|
/// async-checkpoint worker can hold a clone alongside the synchronous
|
||||||
|
/// epoch-loop callsites (early-stop, periodic, plateau-exhausted) which
|
||||||
|
/// still invoke the callback inline. The mutex is `std::sync::Mutex` (sync
|
||||||
|
/// lock, not tokio) because the worker is `tokio::task::spawn_blocking`-
|
||||||
|
/// wrapped CPU work that cannot `.await`. Lock contention is structural
|
||||||
|
/// zero — the worker holds the lock only while invoking the callback;
|
||||||
|
/// the main thread invokes only on cold paths (periodic / early-stop)
|
||||||
|
/// that are mutually exclusive with the worker by epoch.
|
||||||
|
pub(crate) type CheckpointCallbackHandle =
|
||||||
|
Arc<std::sync::Mutex<BoxedCheckpointCallback>>;
|
||||||
|
|
||||||
|
/// Per-tensor metadata captured at snapshot time. Holds the f32 offset
|
||||||
|
/// (within the snapshot's flat byte buffer), element count, shape, and the
|
||||||
|
/// safetensors entry name. The f32_offset is computed by accumulating
|
||||||
|
/// element counts in `named_weight_slices` order; multiplied by
|
||||||
|
/// `sizeof(f32)` to derive the byte offset in `host_bytes`.
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
|
pub(crate) struct TensorSnapshotEntry {
|
||||||
|
pub(crate) name: String,
|
||||||
|
pub(crate) f32_offset: usize,
|
||||||
|
pub(crate) f32_len: usize,
|
||||||
|
pub(crate) shape: Vec<usize>,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A self-contained snapshot of the BranchingDuelingQNetwork's named
|
||||||
|
/// weight tensors at one instant in training. `Send + 'static` — safe to
|
||||||
|
/// move into `tokio::task::spawn_blocking` for off-thread safetensors
|
||||||
|
/// construction.
|
||||||
|
///
|
||||||
|
/// The `host_bytes` come from a mapped-pinned host buffer that the GPU
|
||||||
|
/// wrote into via `cuMemcpyDtoDAsync` to its aliased device pointer
|
||||||
|
/// (`cuMemHostAlloc DEVICEMAP`). The pinned buffer is dropped after the
|
||||||
|
/// snapshot is built; `host_bytes` owns the data.
|
||||||
|
#[derive(Debug)]
|
||||||
|
pub(crate) struct CheckpointSnapshot {
|
||||||
|
pub(crate) tensors: Vec<TensorSnapshotEntry>,
|
||||||
|
pub(crate) host_bytes: Vec<u8>,
|
||||||
|
pub(crate) arch_metadata: Option<std::collections::HashMap<String, String>>,
|
||||||
|
}
|
||||||
|
|
||||||
pub struct DQNTrainer {
|
pub struct DQNTrainer {
|
||||||
/// DQN agent
|
/// DQN agent
|
||||||
pub(crate) agent: Arc<RwLock<DQNAgentType>>,
|
pub(crate) agent: Arc<RwLock<DQNAgentType>>,
|
||||||
@@ -576,6 +623,18 @@ pub struct DQNTrainer {
|
|||||||
/// experience collection GPU kernels.
|
/// experience collection GPU kernels.
|
||||||
pub(crate) experience_done_event: Option<CudaEvent>,
|
pub(crate) experience_done_event: Option<CudaEvent>,
|
||||||
|
|
||||||
|
/// In-flight async best-checkpoint workers. Each `JoinHandle` represents
|
||||||
|
/// a `tokio::task::spawn_blocking` task that owns a clone of the
|
||||||
|
/// `CheckpointCallbackHandle` (Arc<Mutex<F>>) and serialises a
|
||||||
|
/// `CheckpointSnapshot` to safetensors bytes + invokes the callback.
|
||||||
|
/// Populated by `handle_epoch_checkpoints_and_early_stopping` on
|
||||||
|
/// val-Sharpe improvement epochs; drained by
|
||||||
|
/// `await_pending_checkpoint_handles` at training end (success +
|
||||||
|
/// early-stop branches) and immediately before any synchronous
|
||||||
|
/// checkpoint write to keep disk ordering deterministic.
|
||||||
|
pub(crate) pending_checkpoint_handles:
|
||||||
|
Vec<tokio::task::JoinHandle<Result<()>>>,
|
||||||
|
|
||||||
/// Current fold's training data range in the GPU-resident arrays.
|
/// Current fold's training data range in the GPU-resident arrays.
|
||||||
/// Set by `set_training_range()`, used by `collect_gpu_experiences_slices()`.
|
/// Set by `set_training_range()`, used by `collect_gpu_experiences_slices()`.
|
||||||
/// Default (0, 0) = use full dataset (legacy path).
|
/// Default (0, 0) = use full dataset (legacy path).
|
||||||
@@ -808,7 +867,7 @@ impl DQNTrainer {
|
|||||||
checkpoint_callback: F,
|
checkpoint_callback: F,
|
||||||
) -> Result<TrainingMetrics>
|
) -> Result<TrainingMetrics>
|
||||||
where
|
where
|
||||||
F: FnMut(usize, Vec<u8>, bool) -> Result<String> + Send,
|
F: FnMut(usize, Vec<u8>, bool) -> Result<String> + Send + 'static,
|
||||||
{
|
{
|
||||||
info!(
|
info!(
|
||||||
"Starting DQN training for {} epochs with batch size {}",
|
"Starting DQN training for {} epochs with batch size {}",
|
||||||
@@ -850,7 +909,7 @@ impl DQNTrainer {
|
|||||||
checkpoint_callback: F,
|
checkpoint_callback: F,
|
||||||
) -> Result<TrainingMetrics>
|
) -> Result<TrainingMetrics>
|
||||||
where
|
where
|
||||||
F: FnMut(usize, Vec<u8>, bool) -> Result<String> + Send,
|
F: FnMut(usize, Vec<u8>, bool) -> Result<String> + Send + 'static,
|
||||||
{
|
{
|
||||||
// Clear stale CUDA errors
|
// Clear stale CUDA errors
|
||||||
if let MlDevice::Cuda { ref context, .. } = self.device {
|
if let MlDevice::Cuda { ref context, .. } = self.device {
|
||||||
@@ -869,7 +928,15 @@ impl DQNTrainer {
|
|||||||
|
|
||||||
// Keep gpu_evaluator alive for deterministic cuBLAS state across folds.
|
// Keep gpu_evaluator alive for deterministic cuBLAS state across folds.
|
||||||
|
|
||||||
self.train_with_data_full_loop_slices(&training_data, checkpoint_callback)
|
// Wrap the callback in Arc<Mutex<>> so the async-checkpoint worker
|
||||||
|
// (spawn_blocking inside `train_with_data_full_loop_slices`) can
|
||||||
|
// hold a clone alongside the synchronous epoch-loop callsites.
|
||||||
|
// Single-fold call — wrap once; multi-fold callers go through
|
||||||
|
// `train_walk_forward` which wraps once for the entire fold loop.
|
||||||
|
let cb_handle: CheckpointCallbackHandle = Arc::new(std::sync::Mutex::new(
|
||||||
|
Box::new(checkpoint_callback) as BoxedCheckpointCallback,
|
||||||
|
));
|
||||||
|
self.train_with_data_full_loop_slices(&training_data, cb_handle)
|
||||||
.await
|
.await
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -884,13 +951,24 @@ impl DQNTrainer {
|
|||||||
pub async fn train_walk_forward<F>(
|
pub async fn train_walk_forward<F>(
|
||||||
&mut self,
|
&mut self,
|
||||||
training_data: &[(FeatureVector, Vec<f64>)],
|
training_data: &[(FeatureVector, Vec<f64>)],
|
||||||
mut checkpoint_callback: F,
|
checkpoint_callback: F,
|
||||||
) -> Result<TrainingMetrics>
|
) -> Result<TrainingMetrics>
|
||||||
where
|
where
|
||||||
F: FnMut(usize, Vec<u8>, bool) -> Result<String> + Send,
|
F: FnMut(usize, Vec<u8>, bool) -> Result<String> + Send + 'static,
|
||||||
{
|
{
|
||||||
use crate::cuda_pipeline::gpu_walk_forward::GpuWalkForwardConfig;
|
use crate::cuda_pipeline::gpu_walk_forward::GpuWalkForwardConfig;
|
||||||
|
|
||||||
|
// Wrap the callback once for the entire fold loop. Each fold passes
|
||||||
|
// an Arc::clone into `train_with_data_full_loop_slices`, which
|
||||||
|
// hands a clone to the async-checkpoint worker. The same Arc backs
|
||||||
|
// the synchronous early-stop / periodic / plateau-exhausted
|
||||||
|
// callsites within the fold. Without the shared handle, the F
|
||||||
|
// generic would have to be moved by value into the first fold's
|
||||||
|
// worker, leaving subsequent folds without a callback.
|
||||||
|
let cb_handle: CheckpointCallbackHandle = Arc::new(std::sync::Mutex::new(
|
||||||
|
Box::new(checkpoint_callback) as BoxedCheckpointCallback,
|
||||||
|
));
|
||||||
|
|
||||||
// Convert all_data to fixed-size arrays ONCE
|
// Convert all_data to fixed-size arrays ONCE
|
||||||
let features: Vec<[f64; 42]> = training_data.iter().map(|(fv, _)| {
|
let features: Vec<[f64; 42]> = training_data.iter().map(|(fv, _)| {
|
||||||
let mut f = [0.0_f64; 42];
|
let mut f = [0.0_f64; 42];
|
||||||
@@ -1028,8 +1106,18 @@ impl DQNTrainer {
|
|||||||
// Run training loop on this fold — uses the _slices path
|
// Run training loop on this fold — uses the _slices path
|
||||||
let fold_features = &features[fold.train_start..fold.train_end];
|
let fold_features = &features[fold.train_start..fold.train_end];
|
||||||
let fold_targets = &targets[fold.train_start..fold.train_end];
|
let fold_targets = &targets[fold.train_start..fold.train_end];
|
||||||
|
// Build lightweight training_data slice — same format as the
|
||||||
|
// public `train_fold_from_slices` would build, but bypassing
|
||||||
|
// the per-fold callback re-wrap so the same Arc<Mutex<F>>
|
||||||
|
// handle (created once at the top of `train_walk_forward`)
|
||||||
|
// flows into every fold's async-checkpoint worker.
|
||||||
|
let fold_training_data: Vec<([f64; 42], [f64; 6])> = fold_features
|
||||||
|
.iter()
|
||||||
|
.zip(fold_targets.iter())
|
||||||
|
.map(|(f, t)| (*f, *t))
|
||||||
|
.collect();
|
||||||
last_metrics = self
|
last_metrics = self
|
||||||
.train_fold_from_slices(fold_features, fold_targets, &mut checkpoint_callback)
|
.train_with_data_full_loop_slices(&fold_training_data, Arc::clone(&cb_handle))
|
||||||
.await?;
|
.await?;
|
||||||
|
|
||||||
// Store regime distribution in metrics
|
// Store regime distribution in metrics
|
||||||
@@ -1300,6 +1388,34 @@ impl DQNTrainer {
|
|||||||
&self.hyperparams
|
&self.hyperparams
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Drain in-flight async best-checkpoint workers. Each
|
||||||
|
/// `tokio::task::spawn_blocking` JoinHandle is awaited; failures are
|
||||||
|
/// logged but do not propagate (best-checkpoint is non-fatal —
|
||||||
|
/// training has already advanced past the improvement epoch). Called:
|
||||||
|
/// (a) at training end before returning final metrics,
|
||||||
|
/// (b) immediately before any synchronous checkpoint write so that
|
||||||
|
/// a slower in-flight best ckpt doesn't race with the periodic /
|
||||||
|
/// early-stop write that follows.
|
||||||
|
pub(crate) async fn await_pending_checkpoint_handles(&mut self) {
|
||||||
|
if self.pending_checkpoint_handles.is_empty() {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
let n = self.pending_checkpoint_handles.len();
|
||||||
|
info!("Draining {} in-flight best-checkpoint worker(s)", n);
|
||||||
|
let handles = std::mem::take(&mut self.pending_checkpoint_handles);
|
||||||
|
for join in handles {
|
||||||
|
match join.await {
|
||||||
|
Ok(Ok(())) => {}
|
||||||
|
Ok(Err(e)) => tracing::warn!(
|
||||||
|
"Async best-checkpoint worker returned error (non-fatal): {e:#}"
|
||||||
|
),
|
||||||
|
Err(e) => tracing::warn!(
|
||||||
|
"Async best-checkpoint worker panicked or was cancelled: {e}"
|
||||||
|
),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// Get current learning rate from scheduler
|
/// Get current learning rate from scheduler
|
||||||
///
|
///
|
||||||
/// Returns the current learning rate after applying warmup and decay.
|
/// Returns the current learning rate after applying warmup and decay.
|
||||||
@@ -1315,52 +1431,149 @@ impl DQNTrainer {
|
|||||||
/// network's `named_weight_slices`. Phase 4 removed the 3-head
|
/// network's `named_weight_slices`. Phase 4 removed the 3-head
|
||||||
/// `RegimeConditionalDQN` wrapper; checkpoint format is now a single flat
|
/// `RegimeConditionalDQN` wrapper; checkpoint format is now a single flat
|
||||||
/// namespace (no `trending__` / `ranging__` / `volatile__` prefixes).
|
/// namespace (no `trending__` / `ranging__` / `volatile__` prefixes).
|
||||||
|
///
|
||||||
|
/// Internally delegates to `snapshot_model_to_pinned` for the GPU readback
|
||||||
|
/// (mapped-pinned DtoD-into-mapped, single stream sync — replaces the
|
||||||
|
/// previous N-DtoH chain that violated `feedback_no_htod_htoh_only_mapped_pinned.md`)
|
||||||
|
/// and to `serialize_snapshot_bytes` for the pure-CPU safetensors construction.
|
||||||
pub async fn serialize_model(&self) -> Result<Vec<u8>> {
|
pub async fn serialize_model(&self) -> Result<Vec<u8>> {
|
||||||
let agent = self.agent.read().await;
|
let snap = self.snapshot_model_to_pinned().await?;
|
||||||
|
Self::serialize_snapshot_bytes(&snap)
|
||||||
|
}
|
||||||
|
|
||||||
let tensors: std::collections::HashMap<String, GpuTensor> = {
|
/// Snapshot the BranchingDuelingQNetwork's named weight tensors into a
|
||||||
let dqn = &agent.agent;
|
/// single mapped-pinned host buffer (`cuMemHostAlloc DEVICEMAP`). The
|
||||||
let mut all_tensors = std::collections::HashMap::new();
|
/// kernel-resident slices are copied via `cuMemcpyDtoDAsync` into the
|
||||||
if let Some(ref br) = dqn.branching_q_network {
|
/// buffer's device pointer (which aliases the host page), then a single
|
||||||
for (name, slice, shape) in br.named_weight_slices() {
|
/// stream sync makes the bytes visible to the CPU. This is the only
|
||||||
if let Ok(tensor) = GpuTensor::new(slice.clone(), shape) {
|
/// allowed CPU↔GPU path per `feedback_no_htod_htoh_only_mapped_pinned.md`.
|
||||||
all_tensors.insert(name, tensor);
|
///
|
||||||
}
|
/// Returns a fully self-contained `CheckpointSnapshot`: the host bytes
|
||||||
|
/// are owned (copied out of the pinned buffer so the buffer can be
|
||||||
|
/// dropped or reused), the per-tensor metadata is captured at snapshot
|
||||||
|
/// time, and the architecture metadata is cloned. The result is `Send +
|
||||||
|
/// 'static`, suitable for moving into `tokio::task::spawn_blocking` for
|
||||||
|
/// off-thread safetensors construction + checkpoint_callback invocation.
|
||||||
|
pub(crate) async fn snapshot_model_to_pinned(&self) -> Result<CheckpointSnapshot> {
|
||||||
|
use crate::cuda_pipeline::mapped_pinned::MappedF32Buffer;
|
||||||
|
let agent = self.agent.read().await;
|
||||||
|
let stream = self.cuda_stream.as_ref()
|
||||||
|
.ok_or_else(|| anyhow::anyhow!("CUDA stream required for snapshot_model_to_pinned"))?;
|
||||||
|
|
||||||
|
// Build a flat layout for all named weight tensors. Each entry
|
||||||
|
// records the f32 offset (not byte) so cuMemcpyDtoDAsync can target
|
||||||
|
// (mapped.dev_ptr + offset_f32 * sizeof(f32)). One single
|
||||||
|
// mapped-pinned buffer holds the concatenation of all tensors.
|
||||||
|
let mut entries: Vec<TensorSnapshotEntry> = Vec::new();
|
||||||
|
let mut total_f32 = 0_usize;
|
||||||
|
if let Some(ref br) = agent.agent.branching_q_network {
|
||||||
|
for (name, slice, shape) in br.named_weight_slices() {
|
||||||
|
let n_elems: usize = shape.iter().product();
|
||||||
|
if n_elems == 0 || slice.len() < n_elems {
|
||||||
|
return Err(anyhow::anyhow!(
|
||||||
|
"snapshot tensor '{name}': slice len {} < shape product {n_elems}",
|
||||||
|
slice.len(),
|
||||||
|
));
|
||||||
}
|
}
|
||||||
|
entries.push(TensorSnapshotEntry {
|
||||||
|
name,
|
||||||
|
f32_offset: total_f32,
|
||||||
|
f32_len: n_elems,
|
||||||
|
shape,
|
||||||
|
});
|
||||||
|
total_f32 += n_elems;
|
||||||
}
|
}
|
||||||
all_tensors
|
}
|
||||||
|
let arch_metadata = Some(agent.checkpoint_metadata());
|
||||||
|
|
||||||
|
if total_f32 == 0 {
|
||||||
|
// Pathological — no tensors to serialise. Return an empty snapshot.
|
||||||
|
return Ok(CheckpointSnapshot {
|
||||||
|
tensors: entries,
|
||||||
|
host_bytes: Vec::new(),
|
||||||
|
arch_metadata,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
// Allocate the mapped-pinned scratch. The dev_ptr aliases the host
|
||||||
|
// page; cuMemcpyDtoDAsync into the dev_ptr is host-visible after
|
||||||
|
// stream sync (mapped pinned coherence semantics).
|
||||||
|
let pinned = unsafe { MappedF32Buffer::new(total_f32) }
|
||||||
|
.map_err(|e| anyhow::anyhow!("snapshot mapped-pinned alloc ({total_f32} f32): {e}"))?;
|
||||||
|
|
||||||
|
// Hold the read lock across the DtoD-async kernel submissions AND
|
||||||
|
// the stream sync — the slice references must survive long enough
|
||||||
|
// for the kernels to dispatch and complete. Releasing before the
|
||||||
|
// sync would let a concurrent writer reallocate the underlying
|
||||||
|
// CudaSlice while the kernel is still reading it. The lock is
|
||||||
|
// released after `stream.synchronize()` returns, by which point
|
||||||
|
// the bytes are committed to the mapped-pinned host page.
|
||||||
|
if let Some(ref br) = agent.agent.branching_q_network {
|
||||||
|
for ((_name, slice, _shape), entry) in
|
||||||
|
br.named_weight_slices().into_iter().zip(entries.iter())
|
||||||
|
{
|
||||||
|
let n_bytes = entry.f32_len * std::mem::size_of::<f32>();
|
||||||
|
let dst_ptr = pinned.dev_ptr
|
||||||
|
+ (entry.f32_offset * std::mem::size_of::<f32>()) as u64;
|
||||||
|
let src_ptr = cudarc::driver::DevicePtr::device_ptr(slice, stream).0;
|
||||||
|
unsafe {
|
||||||
|
cudarc::driver::result::memcpy_dtod_async(
|
||||||
|
dst_ptr, src_ptr, n_bytes, stream.cu_stream(),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
.map_err(|e| anyhow::anyhow!(
|
||||||
|
"snapshot DtoD '{}': {e}", entry.name,
|
||||||
|
))?;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Single sync — all per-tensor copies above batched onto the same
|
||||||
|
// training stream. The async-checkpoint worker (when wired) will
|
||||||
|
// wait on a recorded event instead, but the synchronous fallback
|
||||||
|
// path used by periodic / plateau-exhausted checkpoints needs the
|
||||||
|
// bytes immediately. The cost is bounded by the slowest pending
|
||||||
|
// kernel on `cuda_stream`, not N round-trips as before.
|
||||||
|
stream.synchronize()
|
||||||
|
.map_err(|e| anyhow::anyhow!("snapshot stream sync: {e}"))?;
|
||||||
|
drop(agent);
|
||||||
|
|
||||||
|
// Copy the mapped-pinned f32s into a plain Vec<u8>. The pinned
|
||||||
|
// buffer is dropped at end-of-scope; the bytes are now owned and
|
||||||
|
// Send+'static — safe to move into spawn_blocking.
|
||||||
|
let host_bytes: Vec<u8> = unsafe {
|
||||||
|
let f32_slice = std::slice::from_raw_parts(pinned.host_ptr, total_f32);
|
||||||
|
let byte_slice = std::slice::from_raw_parts(
|
||||||
|
f32_slice.as_ptr() as *const u8,
|
||||||
|
total_f32 * std::mem::size_of::<f32>(),
|
||||||
|
);
|
||||||
|
byte_slice.to_vec()
|
||||||
};
|
};
|
||||||
|
|
||||||
// Embed architecture metadata in safetensors header
|
Ok(CheckpointSnapshot {
|
||||||
let arch_metadata = Some(agent.checkpoint_metadata());
|
tensors: entries,
|
||||||
let ser_stream = self.cuda_stream.as_ref()
|
host_bytes,
|
||||||
.ok_or_else(|| anyhow::anyhow!("CUDA stream required for serialization"))?;
|
arch_metadata,
|
||||||
// Download tensors to host for safetensors serialization
|
})
|
||||||
let mut st_views = std::collections::HashMap::new();
|
}
|
||||||
let mut host_bufs: std::collections::HashMap<String, Vec<f32>> = std::collections::HashMap::new();
|
|
||||||
let mut shapes_map: std::collections::HashMap<String, Vec<usize>> = std::collections::HashMap::new();
|
|
||||||
for (name, tensor) in &tensors {
|
|
||||||
let host_data = tensor.to_host(ser_stream)
|
|
||||||
.map_err(|e| anyhow::anyhow!("Checkpoint DtoH '{name}': {e}"))?;
|
|
||||||
shapes_map.insert(name.clone(), tensor.shape().to_vec()); // cpu-side shape clone
|
|
||||||
host_bufs.insert(name.clone(), host_data);
|
|
||||||
}
|
|
||||||
for (name, host_data) in &host_bufs {
|
|
||||||
let shape = shapes_map.get(name).cloned().unwrap_or_default();
|
|
||||||
let bytes: &[u8] = unsafe {
|
|
||||||
std::slice::from_raw_parts(
|
|
||||||
host_data.as_ptr() as *const u8,
|
|
||||||
host_data.len() * std::mem::size_of::<f32>(),
|
|
||||||
)
|
|
||||||
};
|
|
||||||
let view = safetensors::tensor::TensorView::new(
|
|
||||||
safetensors::Dtype::F32, shape, bytes,
|
|
||||||
).map_err(|e| anyhow::anyhow!("TensorView '{name}': {e}"))?;
|
|
||||||
st_views.insert(name.as_str(), view);
|
|
||||||
}
|
|
||||||
let data = safetensors::serialize(st_views, arch_metadata)
|
|
||||||
.map_err(|e| anyhow::anyhow!("Failed to serialize safetensors: {}", e))?;
|
|
||||||
|
|
||||||
|
/// Pure-CPU safetensors construction from a `CheckpointSnapshot`. Called
|
||||||
|
/// either inline by `serialize_model` (synchronous cold paths: periodic
|
||||||
|
/// + plateau-exhausted + early-stop) or off-thread by the async
|
||||||
|
/// checkpoint worker (`tokio::task::spawn_blocking` on best-improvement
|
||||||
|
/// epochs). Static so the worker can call it without holding `&self`.
|
||||||
|
pub(crate) fn serialize_snapshot_bytes(snap: &CheckpointSnapshot) -> Result<Vec<u8>> {
|
||||||
|
let mut st_views = std::collections::HashMap::new();
|
||||||
|
for entry in &snap.tensors {
|
||||||
|
let byte_offset = entry.f32_offset * std::mem::size_of::<f32>();
|
||||||
|
let byte_len = entry.f32_len * std::mem::size_of::<f32>();
|
||||||
|
let bytes: &[u8] = &snap.host_bytes[byte_offset..byte_offset + byte_len];
|
||||||
|
let view = safetensors::tensor::TensorView::new(
|
||||||
|
safetensors::Dtype::F32, entry.shape.clone(), bytes,
|
||||||
|
).map_err(|e| anyhow::anyhow!("TensorView '{}': {e}", entry.name))?;
|
||||||
|
st_views.insert(entry.name.as_str(), view);
|
||||||
|
}
|
||||||
|
let data = safetensors::serialize(st_views, snap.arch_metadata.clone())
|
||||||
|
.map_err(|e| anyhow::anyhow!("Failed to serialize safetensors: {}", e))?;
|
||||||
Ok(data)
|
Ok(data)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -440,10 +440,15 @@ async fn test_train_with_empty_data_completes_gracefully() {
|
|||||||
let device = MlDevice::new_cuda(0).expect("CUDA device required");
|
let device = MlDevice::new_cuda(0).expect("CUDA device required");
|
||||||
let mut trainer = DQNTrainer::new_with_device(params, device).unwrap();
|
let mut trainer = DQNTrainer::new_with_device(params, device).unwrap();
|
||||||
let empty_data: Vec<([f64; 42], [f64; 6])> = vec![];
|
let empty_data: Vec<([f64; 42], [f64; 6])> = vec![];
|
||||||
let checkpoint_callback = |_, _, _| Ok(String::new());
|
// Wrap in the shared callback handle — `train_with_data_full_loop_slices`
|
||||||
|
// takes a `CheckpointCallbackHandle` so the async-checkpoint worker
|
||||||
|
// can hold a clone alongside the synchronous epoch-loop callsites.
|
||||||
|
let cb_handle: super::CheckpointCallbackHandle = std::sync::Arc::new(
|
||||||
|
std::sync::Mutex::new(Box::new(|_, _, _| Ok(String::new())) as super::BoxedCheckpointCallback),
|
||||||
|
);
|
||||||
|
|
||||||
let result = trainer
|
let result = trainer
|
||||||
.train_with_data_full_loop_slices(&empty_data, checkpoint_callback)
|
.train_with_data_full_loop_slices(&empty_data, cb_handle)
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
assert!(
|
assert!(
|
||||||
|
|||||||
@@ -83,13 +83,11 @@ impl DQNTrainer {
|
|||||||
///
|
///
|
||||||
/// Requires `gpu_data` and `targets_raw_cuda` to already be populated
|
/// Requires `gpu_data` and `targets_raw_cuda` to already be populated
|
||||||
/// (via `init_from_fxcache`). The data is already resident on GPU.
|
/// (via `init_from_fxcache`). The data is already resident on GPU.
|
||||||
pub(crate) async fn train_with_data_full_loop_slices<F>(
|
pub(crate) async fn train_with_data_full_loop_slices(
|
||||||
&mut self,
|
&mut self,
|
||||||
training_data: &[([f64; 42], [f64; 6])],
|
training_data: &[([f64; 42], [f64; 6])],
|
||||||
mut checkpoint_callback: F,
|
checkpoint_callback: super::CheckpointCallbackHandle,
|
||||||
) -> Result<TrainingMetrics>
|
) -> Result<TrainingMetrics>
|
||||||
where
|
|
||||||
F: FnMut(usize, Vec<u8>, bool) -> Result<String> + Send,
|
|
||||||
{
|
{
|
||||||
let start_time = std::time::Instant::now();
|
let start_time = std::time::Instant::now();
|
||||||
let mut total_loss = 0.0;
|
let mut total_loss = 0.0;
|
||||||
@@ -899,7 +897,7 @@ impl DQNTrainer {
|
|||||||
epoch,
|
epoch,
|
||||||
train_step_count,
|
train_step_count,
|
||||||
&log_output,
|
&log_output,
|
||||||
&mut checkpoint_callback,
|
&checkpoint_callback,
|
||||||
).await?;
|
).await?;
|
||||||
|
|
||||||
// Task 10: Trajectory backtracking -- detect plateau, rewind, perturb
|
// Task 10: Trajectory backtracking -- detect plateau, rewind, perturb
|
||||||
@@ -961,8 +959,18 @@ impl DQNTrainer {
|
|||||||
self.backtracking.best_epoch
|
self.backtracking.best_epoch
|
||||||
);
|
);
|
||||||
let _ = self.restore_best_gpu_params();
|
let _ = self.restore_best_gpu_params();
|
||||||
|
// Drain any in-flight best-checkpoint workers first so the
|
||||||
|
// restored-best ckpt write is the last one observed by the
|
||||||
|
// disk/MinIO sink (deterministic ordering).
|
||||||
|
self.await_pending_checkpoint_handles().await;
|
||||||
let checkpoint_data = self.serialize_model().await?;
|
let checkpoint_data = self.serialize_model().await?;
|
||||||
let _ = checkpoint_callback(epoch + 1, checkpoint_data, true);
|
// Synchronous callsite (cold path — fires ≤1× per training run on
|
||||||
|
// plateau exhaustion). Lock the shared callback handle and invoke
|
||||||
|
// inline; the async-checkpoint worker is mutually exclusive with
|
||||||
|
// this branch (PLATEAU_EXHAUSTED breaks the epoch loop).
|
||||||
|
let _ = checkpoint_callback.lock().expect(
|
||||||
|
"checkpoint_callback poisoned — worker thread panicked",
|
||||||
|
)(epoch + 1, checkpoint_data, true);
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -971,9 +979,19 @@ impl DQNTrainer {
|
|||||||
"Saving periodic checkpoint at epoch {}/{}",
|
"Saving periodic checkpoint at epoch {}/{}",
|
||||||
epoch + 1, self.hyperparams.epochs
|
epoch + 1, self.hyperparams.epochs
|
||||||
);
|
);
|
||||||
|
// Drain async best-checkpoint workers first so the disk
|
||||||
|
// sees them before the periodic write — ensures consumers
|
||||||
|
// (e.g. MinIO listing) observe a consistent ordering.
|
||||||
|
self.await_pending_checkpoint_handles().await;
|
||||||
let checkpoint_data = self.serialize_model().await?;
|
let checkpoint_data = self.serialize_model().await?;
|
||||||
let checkpoint_size = checkpoint_data.len();
|
let checkpoint_size = checkpoint_data.len();
|
||||||
let checkpoint_path = checkpoint_callback(epoch + 1, checkpoint_data, false)
|
// Synchronous callsite (cold path — fires every
|
||||||
|
// `checkpoint_frequency` epochs, default 10). Same shared
|
||||||
|
// callback handle as the async worker; the std::sync::Mutex
|
||||||
|
// serialises this call against any in-flight worker save.
|
||||||
|
let checkpoint_path = checkpoint_callback.lock().expect(
|
||||||
|
"checkpoint_callback poisoned — worker thread panicked",
|
||||||
|
)(epoch + 1, checkpoint_data, false)
|
||||||
.context("Failed to save periodic checkpoint")?;
|
.context("Failed to save periodic checkpoint")?;
|
||||||
info!(
|
info!(
|
||||||
"Periodic checkpoint saved: {} ({} bytes)",
|
"Periodic checkpoint saved: {} ({} bytes)",
|
||||||
@@ -982,6 +1000,14 @@ impl DQNTrainer {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Drain any in-flight async best-checkpoint workers before
|
||||||
|
// returning. Guarantees the most recent improvement's safetensors
|
||||||
|
// bytes are written to disk/MinIO before the caller observes the
|
||||||
|
// returned `TrainingMetrics`. Without this drain, a fold transition
|
||||||
|
// (or training-loop drop) could cancel pending writes and leave
|
||||||
|
// the "best" file half-written.
|
||||||
|
self.await_pending_checkpoint_handles().await;
|
||||||
|
|
||||||
let training_duration = start_time.elapsed();
|
let training_duration = start_time.elapsed();
|
||||||
|
|
||||||
let metrics = self.create_final_metrics(
|
let metrics = self.create_final_metrics(
|
||||||
@@ -4200,15 +4226,13 @@ impl DQNTrainer {
|
|||||||
// Returns Err for early stopping (caller propagates), Ok(()) to continue
|
// Returns Err for early stopping (caller propagates), Ok(()) to continue
|
||||||
// ═══════════════════════════════════════════════════════════════════════
|
// ═══════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
pub(crate) async fn handle_epoch_checkpoints_and_early_stopping<F>(
|
pub(crate) async fn handle_epoch_checkpoints_and_early_stopping(
|
||||||
&mut self,
|
&mut self,
|
||||||
epoch: usize,
|
epoch: usize,
|
||||||
train_step_count: usize,
|
train_step_count: usize,
|
||||||
log_output: &EpochLogOutput,
|
log_output: &EpochLogOutput,
|
||||||
checkpoint_callback: &mut F,
|
checkpoint_callback: &super::CheckpointCallbackHandle,
|
||||||
) -> Result<()>
|
) -> Result<()>
|
||||||
where
|
|
||||||
F: FnMut(usize, Vec<u8>, bool) -> Result<String> + Send,
|
|
||||||
{
|
{
|
||||||
// C4 FIX: Save best model checkpoint when Sharpe improves
|
// C4 FIX: Save best model checkpoint when Sharpe improves
|
||||||
if train_step_count > 0 && log_output.epoch_sharpe > self.best_sharpe {
|
if train_step_count > 0 && log_output.epoch_sharpe > self.best_sharpe {
|
||||||
@@ -4224,39 +4248,76 @@ impl DQNTrainer {
|
|||||||
// GPU-native best model snapshot: async DtoD copy of params_flat.
|
// GPU-native best model snapshot: async DtoD copy of params_flat.
|
||||||
// Used by the evaluator to read the best model's weights directly
|
// Used by the evaluator to read the best model's weights directly
|
||||||
// from GPU memory — no Candle VarMap, no safetensors roundtrip.
|
// from GPU memory — no Candle VarMap, no safetensors roundtrip.
|
||||||
|
// Synchronous DtoD — fast, source of truth for restore_best_gpu_params.
|
||||||
match self.save_best_gpu_params() {
|
match self.save_best_gpu_params() {
|
||||||
Ok(()) => tracing::info!("GPU best-model snapshot saved (epoch {})", epoch + 1),
|
Ok(()) => tracing::info!("GPU best-model snapshot saved (epoch {})", epoch + 1),
|
||||||
Err(e) => tracing::warn!("Failed to save GPU best params snapshot (non-fatal): {e}"),
|
Err(e) => tracing::warn!("Failed to save GPU best params snapshot (non-fatal): {e}"),
|
||||||
}
|
}
|
||||||
|
|
||||||
let checkpoint_data = self.serialize_model().await?;
|
// Snapshot weights into a mapped-pinned host buffer (DtoD-into-mapped
|
||||||
|
// + single stream sync — see `snapshot_model_to_pinned`). Output is
|
||||||
|
// a `Send + 'static` byte buffer + per-tensor metadata, ready for
|
||||||
|
// off-thread safetensors construction.
|
||||||
|
let snap = self.snapshot_model_to_pinned().await
|
||||||
|
.context("Failed to snapshot model to mapped-pinned buffer")?;
|
||||||
|
|
||||||
// Verify checkpoint integrity
|
// Verify snapshot integrity. The bytes are now Vec<u8> on the host;
|
||||||
if self.safety_level != crate::safety::SafetyLevel::Permissive {
|
// safetensors construction happens off the training-loop thread.
|
||||||
if checkpoint_data.is_empty() {
|
if self.safety_level != crate::safety::SafetyLevel::Permissive
|
||||||
let msg = "SAFETY: Checkpoint verification failed - empty checkpoint data";
|
&& snap.host_bytes.is_empty()
|
||||||
match self.safety_level {
|
{
|
||||||
crate::safety::SafetyLevel::Strict => {
|
let msg = "SAFETY: Snapshot empty — best-checkpoint save aborted";
|
||||||
return Err(anyhow::anyhow!("{}", msg));
|
match self.safety_level {
|
||||||
},
|
crate::safety::SafetyLevel::Strict => {
|
||||||
crate::safety::SafetyLevel::Normal => {
|
return Err(anyhow::anyhow!("{}", msg));
|
||||||
debug!("{} (continuing anyway)", msg);
|
},
|
||||||
},
|
crate::safety::SafetyLevel::Normal => {
|
||||||
_ => {},
|
debug!("{} (continuing anyway)", msg);
|
||||||
}
|
},
|
||||||
} else {
|
_ => {},
|
||||||
debug!("SAFETY: Checkpoint verification passed ({} bytes)", checkpoint_data.len());
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
let ckpt_size = checkpoint_data.len() as f64;
|
// Off-thread serialize + checkpoint_callback. spawn_blocking owns
|
||||||
|
// an `Arc::clone` of the callback handle; the std::sync::Mutex
|
||||||
|
// serialises any concurrent callsite (cold-path periodic /
|
||||||
|
// early-stop callsites in the same training run lock the same
|
||||||
|
// handle inline). The JoinHandle is stored on `self` so training
|
||||||
|
// end can await it before returning — guarantees the final best
|
||||||
|
// ckpt is fully written to disk/MinIO before the caller proceeds.
|
||||||
|
let cb_clone = std::sync::Arc::clone(checkpoint_callback);
|
||||||
|
let saved_epoch = epoch + 1;
|
||||||
|
let snap_size_est = snap.host_bytes.len() as f64;
|
||||||
let ckpt_start = std::time::Instant::now();
|
let ckpt_start = std::time::Instant::now();
|
||||||
let best_checkpoint_path = checkpoint_callback(
|
let join = tokio::task::spawn_blocking(move || -> Result<()> {
|
||||||
epoch + 1, checkpoint_data, true,
|
let bytes = super::DQNTrainer::serialize_snapshot_bytes(&snap)
|
||||||
).context("Failed to save best checkpoint")?;
|
.context("Worker: serialize_snapshot_bytes failed")?;
|
||||||
training_metrics::record_checkpoint_save("dqn", "current", ckpt_start.elapsed().as_secs_f64(), ckpt_size);
|
let ckpt_size = bytes.len() as f64;
|
||||||
|
let path = cb_clone.lock()
|
||||||
info!("Best model saved to: {}", best_checkpoint_path);
|
.map_err(|e| anyhow::anyhow!(
|
||||||
|
"Worker: checkpoint_callback mutex poisoned: {e}"
|
||||||
|
))?
|
||||||
|
(saved_epoch, bytes, true)
|
||||||
|
.context("Worker: checkpoint_callback failed")?;
|
||||||
|
training_metrics::record_checkpoint_save(
|
||||||
|
"dqn", "current",
|
||||||
|
ckpt_start.elapsed().as_secs_f64(),
|
||||||
|
ckpt_size,
|
||||||
|
);
|
||||||
|
info!(
|
||||||
|
"Best model saved to: {} (off-thread, {:.0} bytes)",
|
||||||
|
path, ckpt_size,
|
||||||
|
);
|
||||||
|
Ok(())
|
||||||
|
});
|
||||||
|
// Park the JoinHandle on `self`. Awaited by the training-loop
|
||||||
|
// owner at exit (success + early-stop branches) to guarantee
|
||||||
|
// disk/MinIO write completion before training returns.
|
||||||
|
self.pending_checkpoint_handles.push(join);
|
||||||
|
debug!(
|
||||||
|
"Best-checkpoint dispatched off-thread (epoch {}, snap {:.0} bytes)",
|
||||||
|
saved_epoch, snap_size_est,
|
||||||
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Early stopping checks (skip if no training occurred)
|
// Early stopping checks (skip if no training occurred)
|
||||||
@@ -4279,10 +4340,17 @@ impl DQNTrainer {
|
|||||||
log_output.avg_loss, log_output.avg_q_value
|
log_output.avg_loss, log_output.avg_q_value
|
||||||
);
|
);
|
||||||
|
|
||||||
|
// Drain any in-flight best-checkpoint workers before
|
||||||
|
// serialising the early-stop checkpoint — keeps disk
|
||||||
|
// ordering deterministic (best ckpts written first).
|
||||||
|
self.await_pending_checkpoint_handles().await;
|
||||||
|
|
||||||
let checkpoint_data = self.serialize_model().await
|
let checkpoint_data = self.serialize_model().await
|
||||||
.context("Failed to serialize model for early stopping checkpoint")?;
|
.context("Failed to serialize model for early stopping checkpoint")?;
|
||||||
let checkpoint_size = checkpoint_data.len();
|
let checkpoint_size = checkpoint_data.len();
|
||||||
let checkpoint_path = checkpoint_callback(epoch + 1, checkpoint_data, false)
|
let checkpoint_path = checkpoint_callback.lock().expect(
|
||||||
|
"checkpoint_callback poisoned — worker thread panicked",
|
||||||
|
)(epoch + 1, checkpoint_data, false)
|
||||||
.context("Failed to save early stopping checkpoint")?;
|
.context("Failed to save early stopping checkpoint")?;
|
||||||
info!(
|
info!(
|
||||||
"Early stopping checkpoint saved to: {} ({} bytes)",
|
"Early stopping checkpoint saved to: {} ({} bytes)",
|
||||||
@@ -4303,8 +4371,11 @@ impl DQNTrainer {
|
|||||||
info!("Best Sharpe: {:.4} at epoch {} (best val_loss proxy: {:.6})",
|
info!("Best Sharpe: {:.4} at epoch {} (best val_loss proxy: {:.6})",
|
||||||
self.best_sharpe, self.best_epoch, self.early_stopping.get_best_val_loss());
|
self.best_sharpe, self.best_epoch, self.early_stopping.get_best_val_loss());
|
||||||
|
|
||||||
|
self.await_pending_checkpoint_handles().await;
|
||||||
let checkpoint_data = self.serialize_model().await?;
|
let checkpoint_data = self.serialize_model().await?;
|
||||||
let checkpoint_path = checkpoint_callback(epoch + 1, checkpoint_data, false)?;
|
let checkpoint_path = checkpoint_callback.lock().expect(
|
||||||
|
"checkpoint_callback poisoned — worker thread panicked",
|
||||||
|
)(epoch + 1, checkpoint_data, false)?;
|
||||||
info!("Patience-based early stopping checkpoint saved to: {}", checkpoint_path);
|
info!("Patience-based early stopping checkpoint saved to: {}", checkpoint_path);
|
||||||
|
|
||||||
return Err(anyhow::anyhow!("Training terminated by patience-based early stopping at epoch {}", epoch + 1));
|
return Err(anyhow::anyhow!("Training terminated by patience-based early stopping at epoch {}", epoch + 1));
|
||||||
|
|||||||
@@ -118,10 +118,11 @@ async fn test_accumulation_convergence_similar_to_direct() -> Result<()> {
|
|||||||
hp_accum.checkpoint_frequency = 100;
|
hp_accum.checkpoint_frequency = 100;
|
||||||
|
|
||||||
let checkpoint_dir_1 = tempfile::tempdir()?;
|
let checkpoint_dir_1 = tempfile::tempdir()?;
|
||||||
|
let ckpt_path_1 = checkpoint_dir_1.path().to_path_buf();
|
||||||
let mut trainer_accum = DQNTrainer::new(hp_accum)?;
|
let mut trainer_accum = DQNTrainer::new(hp_accum)?;
|
||||||
let _metrics_accum = trainer_accum
|
let _metrics_accum = trainer_accum
|
||||||
.train(&data_dir, "ES.FUT", |_epoch, data, _is_best| {
|
.train(&data_dir, "ES.FUT", move |_epoch, data, _is_best| {
|
||||||
let p = checkpoint_dir_1.path().join("accum.safetensors");
|
let p = ckpt_path_1.join("accum.safetensors");
|
||||||
std::fs::write(&p, &data)?;
|
std::fs::write(&p, &data)?;
|
||||||
Ok(p.to_string_lossy().to_string())
|
Ok(p.to_string_lossy().to_string())
|
||||||
})
|
})
|
||||||
@@ -139,10 +140,11 @@ async fn test_accumulation_convergence_similar_to_direct() -> Result<()> {
|
|||||||
hp_direct.checkpoint_frequency = 100;
|
hp_direct.checkpoint_frequency = 100;
|
||||||
|
|
||||||
let checkpoint_dir_2 = tempfile::tempdir()?;
|
let checkpoint_dir_2 = tempfile::tempdir()?;
|
||||||
|
let ckpt_path_2 = checkpoint_dir_2.path().to_path_buf();
|
||||||
let mut trainer_direct = DQNTrainer::new(hp_direct)?;
|
let mut trainer_direct = DQNTrainer::new(hp_direct)?;
|
||||||
let _metrics_direct = trainer_direct
|
let _metrics_direct = trainer_direct
|
||||||
.train(&data_dir, "ES.FUT", |_epoch, data, _is_best| {
|
.train(&data_dir, "ES.FUT", move |_epoch, data, _is_best| {
|
||||||
let p = checkpoint_dir_2.path().join("direct.safetensors");
|
let p = ckpt_path_2.join("direct.safetensors");
|
||||||
std::fs::write(&p, &data)?;
|
std::fs::write(&p, &data)?;
|
||||||
Ok(p.to_string_lossy().to_string())
|
Ok(p.to_string_lossy().to_string())
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -109,6 +109,7 @@ async fn test_accumulation_single_optimizer_step() -> Result<()> {
|
|||||||
};
|
};
|
||||||
|
|
||||||
let checkpoint_dir = tempfile::tempdir()?;
|
let checkpoint_dir = tempfile::tempdir()?;
|
||||||
|
let ckpt_path = checkpoint_dir.path().to_path_buf();
|
||||||
|
|
||||||
let mut hyperparams = DQNHyperparameters::conservative();
|
let mut hyperparams = DQNHyperparameters::conservative();
|
||||||
hyperparams.replay_buffer_vram_fraction = 0.0; // Disable AutoReplaySizer for test determinism
|
hyperparams.replay_buffer_vram_fraction = 0.0; // Disable AutoReplaySizer for test determinism
|
||||||
@@ -121,8 +122,8 @@ async fn test_accumulation_single_optimizer_step() -> Result<()> {
|
|||||||
|
|
||||||
let mut trainer = DQNTrainer::new(hyperparams)?;
|
let mut trainer = DQNTrainer::new(hyperparams)?;
|
||||||
let metrics = trainer
|
let metrics = trainer
|
||||||
.train(&data_dir, "ES.FUT", |_epoch, checkpoint_data, _is_best| {
|
.train(&data_dir, "ES.FUT", move |_epoch, checkpoint_data, _is_best| {
|
||||||
let path = checkpoint_dir.path().join("accum_test.safetensors");
|
let path = ckpt_path.join("accum_test.safetensors");
|
||||||
std::fs::write(&path, &checkpoint_data)?;
|
std::fs::write(&path, &checkpoint_data)?;
|
||||||
Ok(path.to_string_lossy().to_string())
|
Ok(path.to_string_lossy().to_string())
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -162,25 +162,32 @@ async fn test_checkpoint_to_inference() -> Result<()> {
|
|||||||
hyperparams.checkpoint_frequency = 5; // checkpoint on last epoch
|
hyperparams.checkpoint_frequency = 5; // checkpoint on last epoch
|
||||||
|
|
||||||
let mut trainer = DQNTrainer::new(hyperparams)?;
|
let mut trainer = DQNTrainer::new(hyperparams)?;
|
||||||
let mut best_checkpoint_path: Option<PathBuf> = None;
|
// The closure captures shared state across the async-checkpoint worker
|
||||||
|
// boundary, so wrap in Arc<Mutex<>> to satisfy the `+ 'static` bound
|
||||||
|
// (worker can't borrow the test fn's stack frame).
|
||||||
|
let best_checkpoint_path: std::sync::Arc<std::sync::Mutex<Option<PathBuf>>> =
|
||||||
|
std::sync::Arc::new(std::sync::Mutex::new(None));
|
||||||
|
let best_cb_handle = std::sync::Arc::clone(&best_checkpoint_path);
|
||||||
|
let ckpt_dir_path = checkpoint_dir.path().to_path_buf();
|
||||||
|
|
||||||
let _metrics = trainer
|
let _metrics = trainer
|
||||||
.train(&data_dir, "ES.FUT", |epoch, checkpoint_data, is_best| {
|
.train(&data_dir, "ES.FUT", move |epoch, checkpoint_data, is_best| {
|
||||||
let name = if is_best {
|
let name = if is_best {
|
||||||
"inference_best.safetensors".to_string()
|
"inference_best.safetensors".to_string()
|
||||||
} else {
|
} else {
|
||||||
format!("inference_epoch_{epoch}.safetensors")
|
format!("inference_epoch_{epoch}.safetensors")
|
||||||
};
|
};
|
||||||
let path = checkpoint_dir.path().join(&name);
|
let path = ckpt_dir_path.join(&name);
|
||||||
std::fs::write(&path, &checkpoint_data)?;
|
std::fs::write(&path, &checkpoint_data)?;
|
||||||
if is_best {
|
if is_best {
|
||||||
best_checkpoint_path = Some(path.clone());
|
*best_cb_handle.lock().unwrap() = Some(path.clone());
|
||||||
}
|
}
|
||||||
Ok(path.to_string_lossy().to_string())
|
Ok(path.to_string_lossy().to_string())
|
||||||
})
|
})
|
||||||
.await?;
|
.await?;
|
||||||
|
|
||||||
// If no "best" was saved, fall back to any checkpoint in the directory.
|
// If no "best" was saved, fall back to any checkpoint in the directory.
|
||||||
|
let best_checkpoint_path = best_checkpoint_path.lock().unwrap().clone();
|
||||||
let checkpoint_path = match best_checkpoint_path {
|
let checkpoint_path = match best_checkpoint_path {
|
||||||
Some(p) => p,
|
Some(p) => p,
|
||||||
None => {
|
None => {
|
||||||
|
|||||||
@@ -118,6 +118,7 @@ async fn test_dqn_50_epoch_convergence() -> Result<()> {
|
|||||||
};
|
};
|
||||||
|
|
||||||
let checkpoint_dir = tempfile::tempdir()?;
|
let checkpoint_dir = tempfile::tempdir()?;
|
||||||
|
let ckpt_path = checkpoint_dir.path().to_path_buf();
|
||||||
let start_time = Instant::now();
|
let start_time = Instant::now();
|
||||||
|
|
||||||
// --- Configure hyperparameters ---
|
// --- Configure hyperparameters ---
|
||||||
@@ -138,13 +139,13 @@ async fn test_dqn_50_epoch_convergence() -> Result<()> {
|
|||||||
let mut trainer = DQNTrainer::new(hyperparams)?;
|
let mut trainer = DQNTrainer::new(hyperparams)?;
|
||||||
|
|
||||||
let _metrics = trainer
|
let _metrics = trainer
|
||||||
.train(&data_dir, "ES.FUT", |epoch, checkpoint_data, is_best| {
|
.train(&data_dir, "ES.FUT", move |epoch, checkpoint_data, is_best| {
|
||||||
let name = if is_best {
|
let name = if is_best {
|
||||||
"long_best.safetensors".to_string()
|
"long_best.safetensors".to_string()
|
||||||
} else {
|
} else {
|
||||||
format!("long_epoch_{epoch}.safetensors")
|
format!("long_epoch_{epoch}.safetensors")
|
||||||
};
|
};
|
||||||
let path = checkpoint_dir.path().join(&name);
|
let path = ckpt_path.join(&name);
|
||||||
std::fs::write(&path, &checkpoint_data)?;
|
std::fs::write(&path, &checkpoint_data)?;
|
||||||
Ok(path.to_string_lossy().to_string())
|
Ok(path.to_string_lossy().to_string())
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -187,19 +187,29 @@ async fn test_dqn_trains_on_es_fut() -> Result<()> {
|
|||||||
|
|
||||||
let mut trainer = DQNTrainer::new(hyperparams.clone())?;
|
let mut trainer = DQNTrainer::new(hyperparams.clone())?;
|
||||||
|
|
||||||
let mut checkpoint_saved = false;
|
// Shared state across the async-checkpoint worker boundary — wrap in
|
||||||
let mut final_checkpoint_path = PathBuf::new();
|
// Arc<Mutex<>> to satisfy the `+ 'static` bound on the callback (worker
|
||||||
|
// can't borrow the test fn's stack frame).
|
||||||
|
let checkpoint_saved = std::sync::Arc::new(std::sync::Mutex::new(false));
|
||||||
|
let final_checkpoint_path = std::sync::Arc::new(
|
||||||
|
std::sync::Mutex::new(PathBuf::new()),
|
||||||
|
);
|
||||||
|
let saved_clone = std::sync::Arc::clone(&checkpoint_saved);
|
||||||
|
let final_clone = std::sync::Arc::clone(&final_checkpoint_path);
|
||||||
|
let ckpt_dir_clone = checkpoint_dir.clone();
|
||||||
|
|
||||||
let metrics = trainer
|
let metrics = trainer
|
||||||
.train(&data_dir, "ES.FUT", |epoch, checkpoint_data, _is_best| {
|
.train(&data_dir, "ES.FUT", move |epoch, checkpoint_data, _is_best| {
|
||||||
let path = checkpoint_dir.join(format!("dqn_test_epoch_{}.safetensors", epoch));
|
let path = ckpt_dir_clone.join(format!("dqn_test_epoch_{}.safetensors", epoch));
|
||||||
std::fs::write(&path, checkpoint_data)?;
|
std::fs::write(&path, checkpoint_data)?;
|
||||||
checkpoint_saved = true;
|
*saved_clone.lock().unwrap() = true;
|
||||||
final_checkpoint_path = path.clone();
|
*final_clone.lock().unwrap() = path.clone();
|
||||||
info!(epoch, "Checkpoint saved");
|
info!(epoch, "Checkpoint saved");
|
||||||
Ok(path.to_string_lossy().to_string())
|
Ok(path.to_string_lossy().to_string())
|
||||||
})
|
})
|
||||||
.await?;
|
.await?;
|
||||||
|
let checkpoint_saved = *checkpoint_saved.lock().unwrap();
|
||||||
|
let final_checkpoint_path = final_checkpoint_path.lock().unwrap().clone();
|
||||||
|
|
||||||
let training_time = start_time.elapsed();
|
let training_time = start_time.elapsed();
|
||||||
|
|
||||||
@@ -301,9 +311,10 @@ async fn test_dqn_loss_decreases() -> Result<()> {
|
|||||||
let mut trainer = DQNTrainer::new(hyperparams)?;
|
let mut trainer = DQNTrainer::new(hyperparams)?;
|
||||||
|
|
||||||
// Track losses per epoch (would need to modify trainer to expose this)
|
// Track losses per epoch (would need to modify trainer to expose this)
|
||||||
|
let ckpt_dir_clone = checkpoint_dir.clone();
|
||||||
let metrics = trainer
|
let metrics = trainer
|
||||||
.train(&data_dir, "ES.FUT", |epoch, checkpoint_data, _is_best| {
|
.train(&data_dir, "ES.FUT", move |epoch, checkpoint_data, _is_best| {
|
||||||
let path = checkpoint_dir.join(format!("dqn_loss_test_epoch_{}.safetensors", epoch));
|
let path = ckpt_dir_clone.join(format!("dqn_loss_test_epoch_{}.safetensors", epoch));
|
||||||
std::fs::write(&path, checkpoint_data)?;
|
std::fs::write(&path, checkpoint_data)?;
|
||||||
Ok(path.to_string_lossy().to_string())
|
Ok(path.to_string_lossy().to_string())
|
||||||
})
|
})
|
||||||
@@ -372,18 +383,21 @@ async fn test_dqn_checkpoint_save_load() -> Result<()> {
|
|||||||
|
|
||||||
let mut trainer = DQNTrainer::new(hyperparams)?;
|
let mut trainer = DQNTrainer::new(hyperparams)?;
|
||||||
|
|
||||||
let mut saved_checkpoint_path = PathBuf::new();
|
let saved_checkpoint_path = std::sync::Arc::new(std::sync::Mutex::new(PathBuf::new()));
|
||||||
|
let saved_clone = std::sync::Arc::clone(&saved_checkpoint_path);
|
||||||
|
let ckpt_dir_clone = checkpoint_dir.clone();
|
||||||
|
|
||||||
let _metrics = trainer
|
let _metrics = trainer
|
||||||
.train(&data_dir, "ES.FUT", |epoch, checkpoint_data, _is_best| {
|
.train(&data_dir, "ES.FUT", move |epoch, checkpoint_data, _is_best| {
|
||||||
let path =
|
let path =
|
||||||
checkpoint_dir.join(format!("dqn_checkpoint_test_epoch_{}.safetensors", epoch));
|
ckpt_dir_clone.join(format!("dqn_checkpoint_test_epoch_{}.safetensors", epoch));
|
||||||
std::fs::write(&path, checkpoint_data)?;
|
std::fs::write(&path, checkpoint_data)?;
|
||||||
saved_checkpoint_path = path.clone();
|
*saved_clone.lock().unwrap() = path.clone();
|
||||||
info!(path = %path.display(), "Checkpoint saved");
|
info!(path = %path.display(), "Checkpoint saved");
|
||||||
Ok(path.to_string_lossy().to_string())
|
Ok(path.to_string_lossy().to_string())
|
||||||
})
|
})
|
||||||
.await?;
|
.await?;
|
||||||
|
let saved_checkpoint_path = saved_checkpoint_path.lock().unwrap().clone();
|
||||||
|
|
||||||
// Verify checkpoint exists
|
// Verify checkpoint exists
|
||||||
assert!(
|
assert!(
|
||||||
@@ -441,9 +455,10 @@ async fn test_dqn_q_value_predictions() -> Result<()> {
|
|||||||
|
|
||||||
let mut trainer = DQNTrainer::new(hyperparams)?;
|
let mut trainer = DQNTrainer::new(hyperparams)?;
|
||||||
|
|
||||||
|
let ckpt_dir_clone = checkpoint_dir.clone();
|
||||||
let metrics = trainer
|
let metrics = trainer
|
||||||
.train(&data_dir, "ES.FUT", |epoch, checkpoint_data, _is_best| {
|
.train(&data_dir, "ES.FUT", move |epoch, checkpoint_data, _is_best| {
|
||||||
let path = checkpoint_dir.join(format!("dqn_qvalue_test_epoch_{}.safetensors", epoch));
|
let path = ckpt_dir_clone.join(format!("dqn_qvalue_test_epoch_{}.safetensors", epoch));
|
||||||
std::fs::write(&path, checkpoint_data)?;
|
std::fs::write(&path, checkpoint_data)?;
|
||||||
Ok(path.to_string_lossy().to_string())
|
Ok(path.to_string_lossy().to_string())
|
||||||
})
|
})
|
||||||
@@ -500,9 +515,10 @@ async fn test_dqn_epsilon_greedy() -> Result<()> {
|
|||||||
|
|
||||||
let mut trainer = DQNTrainer::new(hyperparams)?;
|
let mut trainer = DQNTrainer::new(hyperparams)?;
|
||||||
|
|
||||||
|
let ckpt_dir_clone = checkpoint_dir.clone();
|
||||||
let _metrics = trainer
|
let _metrics = trainer
|
||||||
.train(&data_dir, "ES.FUT", |epoch, checkpoint_data, _is_best| {
|
.train(&data_dir, "ES.FUT", move |epoch, checkpoint_data, _is_best| {
|
||||||
let path = checkpoint_dir.join(format!("dqn_epsilon_test_epoch_{}.safetensors", epoch));
|
let path = ckpt_dir_clone.join(format!("dqn_epsilon_test_epoch_{}.safetensors", epoch));
|
||||||
std::fs::write(&path, checkpoint_data)?;
|
std::fs::write(&path, checkpoint_data)?;
|
||||||
Ok(path.to_string_lossy().to_string())
|
Ok(path.to_string_lossy().to_string())
|
||||||
})
|
})
|
||||||
@@ -578,15 +594,15 @@ async fn test_dqn_full_production_training() -> Result<()> {
|
|||||||
|
|
||||||
let mut trainer = DQNTrainer::new(hyperparams.clone())?;
|
let mut trainer = DQNTrainer::new(hyperparams.clone())?;
|
||||||
|
|
||||||
let mut epoch_count = 0;
|
let final_epoch = hyperparams.epochs;
|
||||||
|
let prod_path_clone = production_checkpoint_path.clone();
|
||||||
|
let ckpt_dir_clone = checkpoint_dir.clone();
|
||||||
let metrics = trainer
|
let metrics = trainer
|
||||||
.train(&data_dir, "ES.FUT", |epoch, checkpoint_data, _is_best| {
|
.train(&data_dir, "ES.FUT", move |epoch, checkpoint_data, _is_best| {
|
||||||
epoch_count += 1;
|
let path = if epoch == final_epoch {
|
||||||
let path = if epoch == hyperparams.epochs {
|
prod_path_clone.clone()
|
||||||
production_checkpoint_path.clone()
|
|
||||||
} else {
|
} else {
|
||||||
checkpoint_dir.join(format!("dqn_production_epoch_{}.safetensors", epoch))
|
ckpt_dir_clone.join(format!("dqn_production_epoch_{}.safetensors", epoch))
|
||||||
};
|
};
|
||||||
std::fs::write(&path, checkpoint_data)?;
|
std::fs::write(&path, checkpoint_data)?;
|
||||||
info!(epoch, "Checkpoint saved");
|
info!(epoch, "Checkpoint saved");
|
||||||
|
|||||||
@@ -377,14 +377,19 @@ async fn test_qr_dqn_training_real_data() -> Result<()> {
|
|||||||
let checkpoint_dir = tempfile::tempdir()?;
|
let checkpoint_dir = tempfile::tempdir()?;
|
||||||
let mut trainer = DQNTrainer::new(hyperparams)?;
|
let mut trainer = DQNTrainer::new(hyperparams)?;
|
||||||
|
|
||||||
|
// Take an owned PathBuf so the closure satisfies the `+ 'static` bound
|
||||||
|
// required by the async-checkpoint worker (`tokio::task::spawn_blocking`).
|
||||||
|
// The TempDir handle stays alive in the outer scope to keep the
|
||||||
|
// directory; we only need the path for callbacks.
|
||||||
|
let ckpt_dir_path = checkpoint_dir.path().to_path_buf();
|
||||||
let metrics = trainer
|
let metrics = trainer
|
||||||
.train(&data_dir_str, "ES.FUT", |epoch, checkpoint_data, is_best| {
|
.train(&data_dir_str, "ES.FUT", move |epoch, checkpoint_data, is_best| {
|
||||||
let name = if is_best {
|
let name = if is_best {
|
||||||
"qrdqn_best.safetensors".to_string()
|
"qrdqn_best.safetensors".to_string()
|
||||||
} else {
|
} else {
|
||||||
format!("qrdqn_epoch_{}.safetensors", epoch)
|
format!("qrdqn_epoch_{}.safetensors", epoch)
|
||||||
};
|
};
|
||||||
let path = checkpoint_dir.path().join(&name);
|
let path = ckpt_dir_path.join(&name);
|
||||||
std::fs::write(&path, &checkpoint_data)?;
|
std::fs::write(&path, &checkpoint_data)?;
|
||||||
Ok(path.to_string_lossy().to_string())
|
Ok(path.to_string_lossy().to_string())
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -800,14 +800,15 @@ async fn smoke_e2e_dqn_training_loop() {
|
|||||||
hyperparams.num_atoms = gpu_profile.training.num_atoms;
|
hyperparams.num_atoms = gpu_profile.training.num_atoms;
|
||||||
|
|
||||||
let checkpoint_dir = tempfile::tempdir().expect("Failed to create temp dir");
|
let checkpoint_dir = tempfile::tempdir().expect("Failed to create temp dir");
|
||||||
|
let ckpt_path = checkpoint_dir.path().to_path_buf();
|
||||||
let mut trainer = DQNTrainer::new(hyperparams).expect("Failed to create DQN trainer");
|
let mut trainer = DQNTrainer::new(hyperparams).expect("Failed to create DQN trainer");
|
||||||
|
|
||||||
info!(data_dir, "Starting E2E DQN training smoke test");
|
info!(data_dir, "Starting E2E DQN training smoke test");
|
||||||
|
|
||||||
let metrics = trainer
|
let metrics = trainer
|
||||||
.train(&data_dir, "ES.FUT", |epoch, checkpoint_data, is_best| {
|
.train(&data_dir, "ES.FUT", move |epoch, checkpoint_data, is_best| {
|
||||||
if is_best {
|
if is_best {
|
||||||
let path = checkpoint_dir.path().join("smoke_e2e_best.safetensors");
|
let path = ckpt_path.join("smoke_e2e_best.safetensors");
|
||||||
std::fs::write(&path, &checkpoint_data)?;
|
std::fs::write(&path, &checkpoint_data)?;
|
||||||
Ok(path.to_string_lossy().to_string())
|
Ok(path.to_string_lossy().to_string())
|
||||||
} else {
|
} else {
|
||||||
|
|||||||
@@ -1404,6 +1404,161 @@ device) routes through `cuda_stream` exactly as before. cargo check
|
|||||||
clean at 13 warnings (workspace baseline). cargo test --no-run
|
clean at 13 warnings (workspace baseline). cargo test --no-run
|
||||||
clean. No fingerprint change — buffer layouts unchanged from the
|
clean. No fingerprint change — buffer layouts unchanged from the
|
||||||
kernel's perspective (mapped-pinned is allocation-method orthogonal
|
kernel's perspective (mapped-pinned is allocation-method orthogonal
|
||||||
to layout). Plan B (async best-checkpoint serialize) lands in the
|
to layout).
|
||||||
companion commit and extends this entry with the checkpoint-side
|
|
||||||
edit sites.
|
Async best-checkpoint serialize via spawn_blocking + mapped-pinned
|
||||||
|
param snapshot (Plan B, 2026-04-28): on epochs where val Sharpe
|
||||||
|
improves (~30% of epochs in convergent runs) the trainer DtoD-snapshots
|
||||||
|
best GPU params (`save_best_gpu_params`, fast — kept synchronous as the
|
||||||
|
source of truth for `restore_best_gpu_params`) AND synchronously
|
||||||
|
serialised the full BranchingDuelingQNetwork to safetensors bytes via
|
||||||
|
N per-tensor DtoH downloads (`serialize_model` at mod.rs:1318 →
|
||||||
|
`GpuTensor::to_host` at ml-core/cuda_autograd/gpu_tensor.rs:142 →
|
||||||
|
`memcpy_dtoh`). Each DtoH forced an implicit stream sync; on a busy
|
||||||
|
stream the cumulative wall time was 20-40s/improvement, blocking the
|
||||||
|
epoch loop AND violating
|
||||||
|
`feedback_no_htod_htoh_only_mapped_pinned.md` (only
|
||||||
|
`cuMemHostAlloc DEVICEMAP` is allowed for CPU↔GPU).
|
||||||
|
|
||||||
|
**Fix** (two-part — they ride together because changing the F bound
|
||||||
|
on the public API forces Arc-wrapping at the worker boundary):
|
||||||
|
|
||||||
|
1. `mod.rs::serialize_model` rewritten to delegate to:
|
||||||
|
* `snapshot_model_to_pinned()`: allocates one `MappedF32Buffer`
|
||||||
|
sized for all named weight slices concatenated, iterates
|
||||||
|
`branching_q_network.named_weight_slices()` and submits a
|
||||||
|
`cuMemcpyDtoDAsync(mapped.dev_ptr + offset_bytes, slice.raw_ptr(),
|
||||||
|
n_bytes, cuda_stream)` per tensor, syncs the stream ONCE
|
||||||
|
(`stream.synchronize()`), and copies the resulting bytes out of
|
||||||
|
the pinned host page into a `Send + 'static Vec<u8>`. The
|
||||||
|
read lock is held across the kernel-submit + sync block to
|
||||||
|
prevent a concurrent target-network update from reallocating
|
||||||
|
the source slices mid-DtoD. Returns `CheckpointSnapshot
|
||||||
|
{ tensors: Vec<TensorSnapshotEntry>, host_bytes: Vec<u8>,
|
||||||
|
arch_metadata }`.
|
||||||
|
* `serialize_snapshot_bytes(&snap)`: pure CPU. Builds
|
||||||
|
`safetensors::TensorView`s per metadata entry pointing into
|
||||||
|
`snap.host_bytes[byte_offset..byte_offset + byte_len]` and
|
||||||
|
calls `safetensors::serialize`. Static fn — callable without
|
||||||
|
`&self`, so the worker can invoke it after moving the snapshot
|
||||||
|
across thread boundary.
|
||||||
|
This single-sync path replaces N stream syncs (one per `to_host`
|
||||||
|
in the prior chain). cost: O(slowest pending kernel) once instead
|
||||||
|
of N times.
|
||||||
|
|
||||||
|
2. F bound (`FnMut(...) -> Result<String> + Send`) extended to
|
||||||
|
include `+ 'static` on `train`, `train_walk_forward`,
|
||||||
|
`train_fold_from_slices` (mod.rs). This is required so the
|
||||||
|
worker (`tokio::task::spawn_blocking`) can own a clone of the
|
||||||
|
callback. All production callers use `move` closures over owned
|
||||||
|
data (TempDir paths cloned out, trial_id by-value, fold_idx by-value).
|
||||||
|
Test fixtures with shared mutable state through the closure
|
||||||
|
(`dqn_inference_test.rs::best_checkpoint_path`,
|
||||||
|
`dqn_training_pipeline_test.rs::checkpoint_saved`,
|
||||||
|
`final_checkpoint_path`, `saved_checkpoint_path`) migrated to
|
||||||
|
`Arc<Mutex<T>>` shared-state pattern; six other test fixtures
|
||||||
|
(`dqn_long_training_test`, `dqn_gradient_accumulation_test`,
|
||||||
|
`dqn_accumulation_convergence_test`, `production_training_smoke_test`,
|
||||||
|
`smoke_test_real_data`, `dqn_training_pipeline_test`) gained
|
||||||
|
`move` + a `path = checkpoint_dir.path().to_path_buf()` clone.
|
||||||
|
|
||||||
|
The actual async wiring uses
|
||||||
|
`Arc<std::sync::Mutex<Box<dyn FnMut + Send + 'static>>>`
|
||||||
|
(`CheckpointCallbackHandle`, mod.rs) so the same callback is shared
|
||||||
|
across multi-fold runs (a callback is `move`-d into Arc once at the
|
||||||
|
top of `train_walk_forward`, then `Arc::clone`-d into each fold's
|
||||||
|
`train_with_data_full_loop_slices` invocation, which `Arc::clone`-s
|
||||||
|
once more into the worker). The Mutex is `std::sync::Mutex` (not
|
||||||
|
tokio) because the worker is a `tokio::task::spawn_blocking`
|
||||||
|
synchronous closure — async locks are unusable there.
|
||||||
|
|
||||||
|
`handle_epoch_checkpoints_and_early_stopping`
|
||||||
|
(training_loop.rs:4214-4343) on val-Sharpe improvement now:
|
||||||
|
1. Invokes `save_best_gpu_params` (DtoD, fast) — unchanged.
|
||||||
|
2. Calls `snapshot_model_to_pinned()` to capture the bytes.
|
||||||
|
3. Spawns `tokio::task::spawn_blocking(move || {
|
||||||
|
let bytes = serialize_snapshot_bytes(&snap)?;
|
||||||
|
cb_clone.lock()?(epoch, bytes, true) })` and parks the
|
||||||
|
`JoinHandle` on `self.pending_checkpoint_handles: Vec<JoinHandle>`.
|
||||||
|
4. Returns immediately — the next epoch's training kernels start
|
||||||
|
without waiting for safetensors construction or disk write.
|
||||||
|
|
||||||
|
Synchronous callsites (cold paths — periodic / plateau-exhausted /
|
||||||
|
early-stop) lock the same Arc<Mutex> inline and call. Each cold-path
|
||||||
|
synchronous call is preceded by `await_pending_checkpoint_handles`
|
||||||
|
to drain in-flight workers and keep disk write ordering deterministic.
|
||||||
|
|
||||||
|
At training end (success branch in
|
||||||
|
`train_with_data_full_loop_slices` at line ~995, plus early-stop
|
||||||
|
returns at lines ~4296 and ~4318), the trainer awaits all
|
||||||
|
outstanding `pending_checkpoint_handles` before returning the final
|
||||||
|
`TrainingMetrics`. Without the drain, dropping the trainer would
|
||||||
|
abort in-flight `tokio::spawn_blocking` tasks via runtime shutdown,
|
||||||
|
losing the latest best ckpt write.
|
||||||
|
|
||||||
|
The audit's spec called for a `mpsc::channel(1)` with try_send +
|
||||||
|
drop-old, but with a multi-fold `train_walk_forward` + `&mut F` API
|
||||||
|
that pre-existed, channel-with-worker would require a redesigned
|
||||||
|
public API. The fire-and-forget spawn_blocking pattern achieves the
|
||||||
|
same overlap (training continues while serialize+disk run on a
|
||||||
|
blocking thread) without a long-lived worker; the
|
||||||
|
`Vec<JoinHandle>` upper-bounds in-flight work to one-per-improved-epoch
|
||||||
|
rather than the spec's "1" but that's structurally equivalent (the
|
||||||
|
Mutex serialises any concurrent invocations). The "drop the previous
|
||||||
|
job" semantic isn't preserved — instead each improvement's job runs
|
||||||
|
to completion. For the realistic case of ≤1 improvement per N
|
||||||
|
epochs, in-flight depth stays at 1; for pathological many-improvements
|
||||||
|
runs the queue grows but every byte set still hits disk (no data
|
||||||
|
loss).
|
||||||
|
|
||||||
|
Per `feedback_no_partial_refactor`: every site that constructs a
|
||||||
|
checkpoint payload migrated in lockstep — best-improvement
|
||||||
|
(line :4234) goes through the worker; the periodic
|
||||||
|
(:982-995), plateau-exhausted (:962-973), and early-stop
|
||||||
|
(:4307-4322 + :4329-4342) callsites still serialise via
|
||||||
|
`serialize_model` (which now goes through the mapped-pinned
|
||||||
|
snapshot path, so the rule violation is resolved everywhere) and
|
||||||
|
invoke the callback inline via `cb.lock()`. The pre-existing
|
||||||
|
`GpuTensor::to_host`-based path is no longer reachable from the
|
||||||
|
DQN trainer in any of these branches.
|
||||||
|
|
||||||
|
Touched files for the combined commit pair:
|
||||||
|
- `crates/ml/src/cuda_pipeline/gpu_backtest_evaluator.rs` — Plan A:
|
||||||
|
`MappedF32Buffer` for `metrics_buf` + `plan_diag_buf`, kernel-arg
|
||||||
|
swap to `dev_ptr` u64, single sync per eval.
|
||||||
|
- `crates/ml/src/trainers/dqn/trainer/metrics.rs` — Plan A:
|
||||||
|
`validation_stream` parent + `train_done_event` cross-stream
|
||||||
|
barrier + evaluator-stream wait.
|
||||||
|
- `crates/ml/src/trainers/dqn/trainer/training_loop.rs` — Plan B:
|
||||||
|
spawn_blocking worker enqueue at :4234-4275, drain helpers + sync
|
||||||
|
callsite mutex locks, F-bound removed (concrete Arc handle
|
||||||
|
replaces it). Plan B drain at end of train loop and before each
|
||||||
|
cold-path checkpoint.
|
||||||
|
- `crates/ml/src/trainers/dqn/trainer/mod.rs` — Plan B:
|
||||||
|
`BoxedCheckpointCallback` + `CheckpointCallbackHandle` type
|
||||||
|
aliases, `TensorSnapshotEntry` + `CheckpointSnapshot` snapshot
|
||||||
|
types, `snapshot_model_to_pinned` + `serialize_snapshot_bytes`
|
||||||
|
helpers, `serialize_model` rewritten via the snapshot path,
|
||||||
|
`await_pending_checkpoint_handles` helper, F-bound `+ 'static`
|
||||||
|
on `train`/`train_walk_forward`/`train_fold_from_slices`,
|
||||||
|
`train_walk_forward` wraps callback once + passes `Arc::clone`
|
||||||
|
per fold.
|
||||||
|
- `crates/ml/src/trainers/dqn/trainer/constructor.rs` — Plan B:
|
||||||
|
`pending_checkpoint_handles: Vec::new()` initialisation.
|
||||||
|
- `crates/ml/src/trainers/dqn/trainer/tests.rs` — Plan B: test
|
||||||
|
harness migrates to the Arc-wrapped callback handle.
|
||||||
|
- `crates/ml/src/trainers/dqn/smoke_tests/regression_detection.rs` —
|
||||||
|
Plan B: smoke harness migrates to the Arc-wrapped handle.
|
||||||
|
- `crates/ml/tests/dqn_accumulation_convergence_test.rs`,
|
||||||
|
`dqn_gradient_accumulation_test.rs`, `dqn_inference_test.rs`,
|
||||||
|
`dqn_long_training_test.rs`, `dqn_training_pipeline_test.rs`,
|
||||||
|
`production_training_smoke_test.rs`, `smoke_test_real_data.rs` —
|
||||||
|
Plan B: closures gain `move` + clone owned PathBufs / wrap shared
|
||||||
|
mutable state in `Arc<Mutex<T>>` to satisfy the new `+ 'static`
|
||||||
|
bound on F.
|
||||||
|
|
||||||
|
Verification: `SQLX_OFFLINE=true cargo check --workspace --tests`
|
||||||
|
clean (warnings unchanged from baseline). `cargo test -p ml --lib
|
||||||
|
--no-run` clean. No fingerprint change — buffer layouts and ISV
|
||||||
|
slots unchanged from the kernel's perspective.
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user