perf(rl): replace ALL cudarc memcpy/sync with raw CUDA API in hot path

15 diag staging DtoD + 1 FRD labels DtoD + diag sync all converted
to raw_memcpy_dtod_async / raw_stream_sync. Zero cudarc driver calls
in steady-state step loop.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-26 14:20:30 +02:00
parent f1093e9a07
commit a9c417a79e
2 changed files with 52 additions and 39 deletions

View File

@@ -94,7 +94,7 @@ use ml_alpha::rl::isv_slots::{
RL_OUTCOME_AUX_LAMBDA_INDEX,
};
use ml_alpha::trainer::diag_staging::DiagStaging;
use cudarc::driver::DevicePtrMut;
use ml_alpha::trainer::raw_launch::raw_memcpy_dtod_async;
use ml_alpha::trainer::integrated::{
IntegratedStepStats, IntegratedTrainer, IntegratedTrainerConfig,
};
@@ -555,13 +555,12 @@ fn main() -> Result<()> {
let nbytes = frd_labels_bh.len() * std::mem::size_of::<i32>();
unsafe {
let s = trainer.stream.cu_stream();
let (dev_dst, _g) = trainer.frd_labels_d.device_ptr_mut(&trainer.stream);
cudarc::driver::result::memcpy_dtod_async(
dev_dst,
raw_memcpy_dtod_async(
trainer.frd_labels_d.raw_ptr(),
trainer.frd_labels_staging.dev_ptr,
nbytes,
s,
).context("frd_labels DtoD async")?;
).map_err(|e| anyhow::anyhow!("frd_labels DtoD async: {:?}", e))?;
}
}

View File

@@ -24,11 +24,13 @@ use std::sync::Arc;
use anyhow::{Context, Result};
use cudarc::driver::CudaStream;
use cudarc::driver::sys::CUstream;
use ml_core::device::MlDevice;
use crate::pinned_mem::MappedF32Buffer;
use crate::rl::frd::FRD_OUT_DIM;
use crate::rl::isv_slots::RL_SLOTS_END;
use crate::trainer::raw_launch::{raw_memcpy_dtod_async, raw_stream_sync};
/// Double-buffered mapped-pinned staging for async diagnostic reads.
///
@@ -44,6 +46,9 @@ use crate::rl::isv_slots::RL_SLOTS_END;
/// via `f32::to_bits() as i32`.
pub struct DiagStaging {
diag_stream: Arc<CudaStream>,
/// Cached raw `CUstream` handle for the diag stream — avoids
/// cudarc's `bind_to_thread` overhead on every memcpy/sync call.
raw_diag_stream: CUstream,
bufs: [MappedF32Buffer; 2],
current: usize,
/// Total f32 count per buffer.
@@ -107,8 +112,10 @@ impl DiagStaging {
let buf1 = unsafe { MappedF32Buffer::new(n_floats) }
.map_err(|e| anyhow::anyhow!("DiagStaging buf1: {e}"))?;
let raw_diag_stream = diag_stream.cu_stream();
Ok(Self {
diag_stream,
raw_diag_stream,
bufs: [buf0, buf1],
current: 0,
n_floats,
@@ -167,144 +174,144 @@ impl DiagStaging {
frd_logits_ptr: u64,
) -> Result<()> {
let buf = &self.bufs[self.current];
let s = self.diag_stream.cu_stream();
let s = self.raw_diag_stream;
let b = self.b_size;
unsafe {
// ISV
cudarc::driver::result::memcpy_dtod_async(
raw_memcpy_dtod_async(
buf.dev_ptr,
isv_ptr,
self.isv_len * 4,
s,
)
.context("diag snapshot isv")?;
.map_err(|e| anyhow::anyhow!("diag snapshot isv: {:?}", e))?;
// rewards
cudarc::driver::result::memcpy_dtod_async(
raw_memcpy_dtod_async(
buf.dev_ptr + (self.rewards_offset * 4) as u64,
rewards_ptr,
b * 4,
s,
)
.context("diag snapshot rewards")?;
.map_err(|e| anyhow::anyhow!("diag snapshot rewards: {:?}", e))?;
// dones
cudarc::driver::result::memcpy_dtod_async(
raw_memcpy_dtod_async(
buf.dev_ptr + (self.dones_offset * 4) as u64,
dones_ptr,
b * 4,
s,
)
.context("diag snapshot dones")?;
.map_err(|e| anyhow::anyhow!("diag snapshot dones: {:?}", e))?;
// actions (i32 -- same byte width as f32, reinterpret on host)
cudarc::driver::result::memcpy_dtod_async(
raw_memcpy_dtod_async(
buf.dev_ptr + (self.actions_offset * 4) as u64,
actions_ptr,
b * 4,
s,
)
.context("diag snapshot actions")?;
.map_err(|e| anyhow::anyhow!("diag snapshot actions: {:?}", e))?;
// raw_rewards
cudarc::driver::result::memcpy_dtod_async(
raw_memcpy_dtod_async(
buf.dev_ptr + (self.raw_rewards_offset * 4) as u64,
raw_rewards_ptr,
b * 4,
s,
)
.context("diag snapshot raw_rewards")?;
.map_err(|e| anyhow::anyhow!("diag snapshot raw_rewards: {:?}", e))?;
// trade_duration
cudarc::driver::result::memcpy_dtod_async(
raw_memcpy_dtod_async(
buf.dev_ptr + (self.trade_duration_offset * 4) as u64,
trade_duration_ptr,
b * 4,
s,
)
.context("diag snapshot trade_duration")?;
.map_err(|e| anyhow::anyhow!("diag snapshot trade_duration: {:?}", e))?;
// outcome_ema
cudarc::driver::result::memcpy_dtod_async(
raw_memcpy_dtod_async(
buf.dev_ptr + (self.outcome_ema_offset * 4) as u64,
outcome_ema_ptr,
b * 4,
s,
)
.context("diag snapshot outcome_ema")?;
.map_err(|e| anyhow::anyhow!("diag snapshot outcome_ema: {:?}", e))?;
// position_lots (i32 -- raw bytes into f32 slot)
cudarc::driver::result::memcpy_dtod_async(
raw_memcpy_dtod_async(
buf.dev_ptr + (self.position_lots_offset * 4) as u64,
position_lots_ptr,
b * 4,
s,
)
.context("diag snapshot position_lots")?;
.map_err(|e| anyhow::anyhow!("diag snapshot position_lots: {:?}", e))?;
// pyramid_count (i32 -- raw bytes into f32 slot)
cudarc::driver::result::memcpy_dtod_async(
raw_memcpy_dtod_async(
buf.dev_ptr + (self.pyramid_count_offset * 4) as u64,
pyramid_count_ptr,
b * 4,
s,
)
.context("diag snapshot pyramid_count")?;
.map_err(|e| anyhow::anyhow!("diag snapshot pyramid_count: {:?}", e))?;
// unit_entry_price (f32, B x 4)
cudarc::driver::result::memcpy_dtod_async(
raw_memcpy_dtod_async(
buf.dev_ptr + (self.unit_entry_price_offset * 4) as u64,
unit_entry_price_ptr,
b * 4 * 4,
s,
)
.context("diag snapshot unit_entry_price")?;
.map_err(|e| anyhow::anyhow!("diag snapshot unit_entry_price: {:?}", e))?;
// unit_entry_step (i32, B x 4 -- raw bytes)
cudarc::driver::result::memcpy_dtod_async(
raw_memcpy_dtod_async(
buf.dev_ptr + (self.unit_entry_step_offset * 4) as u64,
unit_entry_step_ptr,
b * 4 * 4,
s,
)
.context("diag snapshot unit_entry_step")?;
.map_err(|e| anyhow::anyhow!("diag snapshot unit_entry_step: {:?}", e))?;
// unit_lots (i32, B x 4 -- raw bytes)
cudarc::driver::result::memcpy_dtod_async(
raw_memcpy_dtod_async(
buf.dev_ptr + (self.unit_lots_offset * 4) as u64,
unit_lots_ptr,
b * 4 * 4,
s,
)
.context("diag snapshot unit_lots")?;
.map_err(|e| anyhow::anyhow!("diag snapshot unit_lots: {:?}", e))?;
// unit_trail (f32, B x 4)
cudarc::driver::result::memcpy_dtod_async(
raw_memcpy_dtod_async(
buf.dev_ptr + (self.unit_trail_offset * 4) as u64,
unit_trail_ptr,
b * 4 * 4,
s,
)
.context("diag snapshot unit_trail")?;
.map_err(|e| anyhow::anyhow!("diag snapshot unit_trail: {:?}", e))?;
// close_unit_index (i32, B -- raw bytes)
cudarc::driver::result::memcpy_dtod_async(
raw_memcpy_dtod_async(
buf.dev_ptr + (self.close_unit_index_offset * 4) as u64,
close_unit_index_ptr,
b * 4,
s,
)
.context("diag snapshot close_unit_index")?;
.map_err(|e| anyhow::anyhow!("diag snapshot close_unit_index: {:?}", e))?;
// frd_logits (f32, B x FRD_OUT_DIM)
cudarc::driver::result::memcpy_dtod_async(
raw_memcpy_dtod_async(
buf.dev_ptr + (self.frd_logits_offset * 4) as u64,
frd_logits_ptr,
b * FRD_OUT_DIM * 4,
s,
)
.context("diag snapshot frd_logits")?;
.map_err(|e| anyhow::anyhow!("diag snapshot frd_logits: {:?}", e))?;
}
Ok(())
@@ -315,13 +322,20 @@ impl DiagStaging {
/// written to is now the "complete" buffer readable by the host, and the
/// other buffer becomes the new write target.
pub fn sync_and_swap(&mut self) -> Result<()> {
self.diag_stream
.synchronize()
.context("diag stream sync")?;
unsafe {
raw_stream_sync(self.raw_diag_stream)
.map_err(|e| anyhow::anyhow!("diag stream sync: {:?}", e))?;
}
self.current = 1 - self.current;
Ok(())
}
/// The owned `CudaStream` backing the raw handle. Kept alive for the
/// lifetime of `DiagStaging`; dropping it would invalidate `raw_diag_stream`.
pub fn diag_stream(&self) -> &CudaStream {
&self.diag_stream
}
/// Read the PREVIOUS step's ISV data (the buffer NOT currently targeted
/// for writes). Call AFTER `sync_and_swap`.
pub fn read_isv(&self) -> &[f32] {