Per-component gradient clipping: - 2 new CUDA kernels (dqn_clipped_saxpy, dqn_clip_grad) - CQL gradient isolated into separate scratch buffer - Budget allocation: C51=70%, CQL=15%, IQN=10%, Ens=5% - Dynamic budget — inactive components' share goes to C51 Spectral norm extended to all 10 weight matrices: - Was trunk-only (W_s1, W_s2), now covers all heads - 16 u/v power iteration buffers, batched GOFF sync-back - Uniform sigma_max, end-to-end Lipschitz bounded Bug fixes from 9-agent audit: - Backtest episode reset: all 8 fields (was 5), max_equity updated before floor check - gradient_clip_norm unified: 10.0 everywhere (was 10.0 vs 1.0) - entropy_coefficient: Option<f64> → f64, single default 0.001 - Close price: unwrap_or(1.0) → direct indexing (no silent wrong rewards) - step_count: only increments during training (was incrementing on inference) - Quantile loss: single CUDA context (was 3×), silent fallbacks removed - MaybeNoisyLinear: single-variant enum removed → direct NoisyLinear - Budget fractions: exported as pub(crate) constants, referenced not hardcoded Monitoring: - Per-component gradient Prometheus gauges (C51 raw, combined) - Diagnostic logging every 1000 steps Tests: 7 new gradient budget tests, 1 gradient bounds smoke test All 488 tests pass (359 ml-dqn + 129 ml) Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
24 KiB
Spec A: Full BF16 Tensor Core Conversion + GPU Hot Path Performance
Date: 2026-03-27 Spec: A of 3 (DQN CUDA training pipeline overhaul) Target: Epoch time <=5s on H100 (from ~30s), ~2.25x cuBLAS throughput gain Scope: BF16 tensor cores, cached device pointers, zero-sync hot path, batched spectral norm, dead BF16 code removal
Problem Statement
The DQN training pipeline uses cublasSgemm (F32) for all matrix multiplications in both forward and backward passes. On H100, BF16 tensor cores provide 989 TFLOPS vs 67 TFLOPS F32 scalar -- a 14.8x theoretical gap. With F32 accumulation (CUBLAS_COMPUTE_32F), practical BF16 throughput is ~3x over F32 SGEMM.
BF16 weight mirror buffers (bf16_params_buf, bf16_target_params_buf) already exist in GpuDqnTrainer and are allocated at construction, but are never read by the forward pass. The forward reads F32 from params_buf via cublasSgemm. The conversion kernel f32_to_bf16_kernel exists in common_device_functions.cuh but is not invoked in the training loop.
Additionally, the per-step hot path has severe overhead:
| Issue | Count | Source |
|---|---|---|
cuStreamSynchronize calls |
4 | apply_iqn_trunk_gradient:708, apply_ensemble_trunk_gradient:911, run_ensemble_step:875, replay_adam_and_readback:2369 |
raw_device_ptr() calls |
110+ | cudarc event machinery per call |
EventTrackingGuard create/drop |
11 | Disable/enable event tracking RAII cycles |
| Single-block spectral norm launches | 10 | grid=(1,1,1) each, serial kernel launch overhead |
Architecture
Current Data Flow (F32 Only)
Experience Collector (F32 states)
|
v
states_buf [B, SD] (F32)
|
v cublasSgemm (F32 x F32 -> F32)
CublasForward::forward_pass()
|-- h_s1 [B, SH1] = ReLU(states @ W_s1^T + b_s1) F32
|-- h_s2 [B, SH2] = ReLU(h_s1 @ W_s2^T + b_s2) F32
|-- h_v [B, VH] = ReLU(h_s2 @ W_v1^T + b_v1) F32
|-- v_logits [B, NA] = h_v @ W_v2^T + b_v2 F32
|-- h_a [B, AH] = ReLU(h_s2 @ W_a1^T + b_a1) F32
|-- adv_logits [B, B0*NA] = h_a @ W_a2^T + b_a2 F32
|-- (branches 1, 2 same pattern)
v
C51/MSE loss kernels (F32)
|
v cublasSgemm (F32 x F32 -> F32)
CublasBackward::backward_pass()
|-- dW = dY^T @ X (weight gradients, F32)
|-- dX = dY @ W^T (upstream gradients, F32)
v
grad_buf [TOTAL_PARAMS] (F32)
|
v
Adam update (F32) -> params_buf (F32)
|
v [NEVER HAPPENS]
f32_to_bf16_kernel -> bf16_params_buf (u16, UNUSED)
Proposed Data Flow (BF16 Forward + Mixed Backward)
Experience Collector (F32 states)
|
v f32_to_bf16_kernel (in upload path)
states_bf16 [B, SD] (BF16) states_buf [B, SD] (F32, kept for loss)
| |
v cublasGemmEx (BF16 x BF16 -> F32) |
CublasForward::forward_pass() |
|-- h_s1_bf16 = BF16(ReLU(...)) | activations saved as BF16
|-- h_s2_bf16 = BF16(ReLU(...)) |
|-- h_v_bf16 = BF16(ReLU(...)) |
|-- v_logits [B, NA] (F32) <---------' logit output stays F32
|-- (branches same) for C51/MSE loss precision
v
C51/MSE loss kernels (F32 logits -> F32 d_logits)
|
v cublasGemmEx (BF16 weights x F32 grads -> F32, COMPUTE_32F = TF32)
CublasBackward::backward_pass()
|-- dW = d_output^T @ input_bf16 BF16 activations x F32 grads
|-- dX = d_output @ W_bf16^T BF16 weights x F32 grads
v
grad_buf [TOTAL_PARAMS] (F32)
|
v
Adam update (F32) -> params_buf (F32)
|
v f32_to_bf16_kernel (single launch, full buffer)
bf16_params_buf [TOTAL_PARAMS] (u16/BF16, used by NEXT forward)
Precision Strategy
| Path | Input A | Input B | Compute | Accumulate | Speedup | Rationale |
|---|---|---|---|---|---|---|
| Forward SGEMM | BF16 weights | BF16 activations | CUBLAS_COMPUTE_32F |
F32 | ~3x | Both inputs BF16 enables H100 tensor cores |
| Backward dW | F32 d_output | BF16 activations | CUBLAS_COMPUTE_32F |
F32 | ~1.5x (TF32) | Mixed types -> TF32 path, preserves gradient precision |
| Backward dX | F32 d_output | BF16 weights | CUBLAS_COMPUTE_32F |
F32 | ~1.5x (TF32) | Same as dW |
| Loss (C51/MSE) | F32 logits | -- | F32 | -- | 1x (no change) | Element-wise, not GEMM-bound |
| Adam | F32 grads | F32 moments | F32 | -- | 1x (no change) | Must be F32 for optimizer stability |
| F32->BF16 conv | F32 params | -- | -- | -- | N/A | Single kernel, ~0.1ms |
Combined GEMM speedup: Forward 3x + Backward 1.5x = ~2.25x blended (forward is ~60% of GEMM time).
Detailed Changes
1. BF16 Forward Pass (batched_forward.rs)
Replace cublasSgemm with cublasGemmEx using BF16 inputs.
Current call (sgemm_layer, line ~513):
cublas_result::sgemm(
self.handle.0,
CUBLAS_OP_T, CUBLAS_OP_N,
n as i32, b as i32, k as i32,
&1.0f32,
w_ptr as *const f32, k as i32,
a_ptr as *const f32, k as i32,
&0.0f32,
c_ptr as *mut f32, n as i32,
)
Replacement (gemmex_bf16_layer):
cublasGemmEx(
self.handle.0,
CUBLAS_OP_T, CUBLAS_OP_N,
n as i32, b as i32, k as i32,
&1.0f32 as *const f32 as *const c_void, // alpha (F32)
w_ptr as *const c_void, // A = BF16 weights
CUDA_R_16BF, // Atype
k as i32, // lda
a_ptr as *const c_void, // B = BF16 activations
CUDA_R_16BF, // Btype
k as i32, // ldb
&0.0f32 as *const f32 as *const c_void, // beta (F32)
c_ptr as *mut c_void, // C = F32 output
CUDA_R_32F, // Ctype
n as i32, // ldc
CUBLAS_COMPUTE_32F, // computeType
CUBLAS_GEMM_DEFAULT_TENSOR_OP, // algo
)
New BF16 activation buffers (allocated in CublasForward::new()):
states_bf16: CudaSlice<u16>--[B, SD]h_s1_bf16: CudaSlice<u16>--[B, SH1]h_s2_bf16: CudaSlice<u16>--[B, SH2]h_v_bf16: CudaSlice<u16>--[B, VH]h_b0_bf16, h_b1_bf16, h_b2_bf16: CudaSlice<u16>--[B, AH]each
Activation conversion: After each bias+ReLU kernel (which outputs F32 into existing h_s1, etc.), add an inline F32->BF16 conversion step that writes into the _bf16 variant. The F32 activation is kept for the backward pass (relu mask needs original F32 values).
Alternatively, fuse the F32->BF16 cast into the add_bias_relu_kernel:
// New: add_bias_relu_bf16_kernel
// Writes F32 to output AND BF16 to output_bf16 in one pass
extern "C" __global__ void add_bias_relu_bf16_kernel(
float* __restrict__ output, // F32 for backward relu mask
__nv_bfloat16* __restrict__ output_bf16, // BF16 for next layer's GEMM
const float* __restrict__ bias,
int out_dim, int total_elements)
State upload conversion: States arrive F32 from the experience collector. Add a f32_to_bf16_kernel launch in the upload path (between cuMemcpyHtoDAsync of states and the first SGEMM). This writes states_bf16 from states_buf. Both buffers remain live -- states_buf (F32) is used by the backward pass, states_bf16 by the forward GEMM.
Weight pointers: Forward GEMM reads from bf16_params_buf (online) and bf16_target_params_buf (target) instead of params_buf. The f32_weight_ptrs() function is replaced with bf16_weight_ptrs() that returns u64 pointers into the BF16 flat buffer at precomputed byte offsets (already stored as bf16_weight_offsets in GpuDqnTrainer).
Logit output: The final layer GEMMs (W_v2, W_a2, W_b1out, W_b2out) output F32 logits into existing buffers. cublasGemmEx with Ctype=CUDA_R_32F handles this natively -- the output type is independent of input types.
2. Mixed-Precision Backward Pass (batched_backward.rs)
Weight gradient dW[out, in] = dY^T[out, B] @ X[B, in]:
dYis F32 (from C51/MSE loss gradient)Xis BF16 (saved activation from forward)- Call
cublasGemmExwithAtype=CUDA_R_32F,Btype=CUDA_R_16BF,Ctype=CUDA_R_32F,CUBLAS_COMPUTE_32F - cuBLAS internally uses TF32 for this mixed configuration (~1.5x over F32)
Upstream gradient dX[B, in] = dY[B, out] @ W[out, in]:
dYis F32Wis BF16 (frombf16_params_buf)- Same
cublasGemmExmixed call - Output
dXstays F32 (needed for relu mask and next layer's backward)
No change: bias gradient kernel, relu mask kernel, C51/MSE loss kernels, gradient accumulation -- all remain F32.
3. Cached Device Pointers (gpu_dqn_trainer.rs)
Create a CachedPtrs struct computed once at construction (and refreshed if buffers are reallocated, which never happens after init):
/// Pre-resolved raw u64 device pointers for all GPU buffers.
/// Eliminates 110+ per-step `raw_device_ptr()` calls that go through
/// cudarc's event tracking machinery.
struct CachedPtrs {
// Flat parameter buffers
params_buf: u64,
target_params_buf: u64,
bf16_params_buf: u64,
bf16_target_params_buf: u64,
grad_buf: u64,
m_buf: u64,
v_buf: u64,
// Activation buffers (F32)
states_buf: u64,
h_s1: u64,
h_s2: u64,
h_v: u64,
h_b0: u64, h_b1: u64, h_b2: u64,
on_v_logits: u64,
on_adv_logits: u64,
// Activation buffers (BF16) -- new
states_bf16: u64,
h_s1_bf16: u64,
h_s2_bf16: u64,
h_v_bf16: u64,
h_b0_bf16: u64, h_b1_bf16: u64, h_b2_bf16: u64,
// Target forward buffers
tgt_h_s1: u64,
tgt_h_s2: u64,
tgt_h_v: u64,
tgt_v_logits: u64,
tgt_adv_logits: u64,
// Loss / optimizer
total_loss_buf: u64,
grad_norm_buf: u64,
td_errors_buf: u64,
t_buf: u64,
// Spectral norm vectors
spec_u_s1: u64, spec_v_s1: u64,
spec_u_s2: u64, spec_v_s2: u64,
// ... (all 10 pairs)
// IQN scratch
iqn_trunk_m: u64,
iqn_trunk_v: u64,
iqn_trunk_grad_norm: u64,
iqn_trunk_t_buf: u64,
}
Construction: After all CudaSlice allocations in GpuDqnTrainer::new(), resolve every pointer once with raw_device_ptr() and store in CachedPtrs. This requires disabling event tracking once during construction, not per-step.
Usage: Replace every raw_device_ptr(&self.some_buf, &self.stream) with self.ptrs.some_buf.
EventTrackingGuard removal: With cached pointers, no per-step device_ptr() calls go through cudarc's event machinery. Remove all 11 EventTrackingGuard instances from per-step methods. Keep the single disable/enable in the constructor.
4. Zero-Sync Hot Path (gpu_dqn_trainer.rs, fused_training.rs)
Syncs to remove (all on the same CUDA stream -- CUDA guarantees in-order execution):
| Location | Line | Current Purpose | Why Safe to Remove |
|---|---|---|---|
apply_iqn_trunk_gradient |
708 | "Ensure graph_forward completed" | Same stream -- SAXPY kernel won't start until graph_forward's last kernel finishes |
apply_ensemble_trunk_gradient |
911 | Same | Same stream ordering guarantee |
run_ensemble_step |
875 | Same | Same stream ordering guarantee |
Sync to defer (reads back scalars to CPU):
| Location | Line | Current Purpose | Change |
|---|---|---|---|
replay_adam_and_readback |
2369 | Sync before DtoH scalar readback | Defer to every N steps or epoch boundary |
Deferred readback design: Add readback_interval: usize config (default 100). Between readbacks, return the previous step's (total_loss, grad_norm) values. The loss/grad_norm are used for logging only -- stale values by 100 steps are acceptable for monitoring. At epoch boundaries, force a sync for accurate epoch-level metrics.
pub fn replay_adam_and_readback(&mut self) -> Result<FusedTrainScalars, MLError> {
self.replay_adam()?;
self.step_since_readback += 1;
if self.step_since_readback >= self.readback_interval || self.force_readback {
unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); }
// DtoH readback of loss + grad_norm
unsafe {
cudarc::driver::sys::cuMemcpyDtoH_v2(
self.cached_loss.as_mut_ptr().cast(),
self.ptrs.total_loss_buf, 4,
);
cudarc::driver::sys::cuMemcpyDtoH_v2(
self.cached_norm.as_mut_ptr().cast(),
self.ptrs.grad_norm_buf, 4,
);
}
self.step_since_readback = 0;
self.force_readback = false;
}
Ok(FusedTrainScalars {
total_loss: self.cached_loss[0],
grad_norm: self.cached_norm[0],
})
}
5. Batched Spectral Norm (dqn_utility_kernels.cu, gpu_dqn_trainer.rs)
Current: 10 individual kernel launches, each with grid=(1,1,1), block=(256,1,1).
Proposed: Single batched_spectral_norm_kernel with grid=(10,1,1), block=(256,1,1). Each block handles one weight matrix, identified by blockIdx.x indexing into descriptor arrays.
New CUDA kernel:
struct SpectralNormDesc {
float* W; // weight matrix pointer
float* u; // left singular vector
float* v; // right singular vector
int out_dim;
int in_dim;
float sigma_max;
};
extern "C" __global__ void batched_spectral_norm_kernel(
SpectralNormDesc* __restrict__ descs, // [num_matrices] on device
int num_matrices
) {
int matrix_idx = blockIdx.x;
if (matrix_idx >= num_matrices) return;
SpectralNormDesc d = descs[matrix_idx];
// ... same power iteration logic as current spectral_norm_kernel,
// but reading W, u, v, out_dim, in_dim, sigma_max from d
}
Descriptor upload: Allocate CudaSlice<SpectralNormDesc> at construction. Populate with cached device pointers from CachedPtrs and dimensions from config. Upload once (or recompute if pointers change, which they don't post-init).
Rust-side change in apply_spectral_norm():
// Before: 10 separate launch_builder() calls
// After:
unsafe {
self.stream
.launch_builder(&self.batched_spectral_norm_kernel)
.arg(&self.ptrs.spec_norm_descs)
.arg(&10i32)
.launch(LaunchConfig {
grid_dim: (10, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 256 * 4,
})?;
}
This eliminates 9 kernel launch overhead cycles (~5-10us each on H100).
6. Dead BF16 Infrastructure Removal (gpu_weights.rs, gpu_dqn_trainer.rs)
Delete from gpu_weights.rs (~200 LOC):
struct DuelingWeightSetBf16(line 473) -- 12 per-tensorCudaSlice<u16>fieldsstruct BranchingWeightSetBf16(line 490) -- 8 per-tensorCudaSlice<u16>fieldsstruct CuriosityWeightSetBf16(line 503) -- 4 per-tensorCudaSlice<u16>fieldsfn alloc_bf16_mirror()(line 511) -- allocates matching u16 bufferpub fn convert_f32_to_bf16()(line 527) -- per-tensor conversion loopimpl DuelingWeightSetBf16(line 554) --from_f32()+sync_from_f32()impl BranchingWeightSetBf16(line 596) --from_f32()+sync_from_f32()impl CuriosityWeightSetBf16(line 630) --from_f32()+sync_from_f32()struct KernelWeightPackBf16(line 670) -- 24-pointer pack for the old fused kernel
Delete from gpu_dqn_trainer.rs:
fn ensure_bf16_mirrors()(line 2144) -- allocates per-tensor BF16 mirrorsfn sync_online_bf16()(line 2237) -- syncs per-tensor F32->BF16
Keep (used by the new BF16 forward path):
bf16_params_buf: CudaSlice<u16>-- flat online BF16 buffer, read bycublasGemmExbf16_target_params_buf: CudaSlice<u16>-- flat target BF16 bufferbf16_weight_offsets: Vec<usize>-- byte offsets into flat BF16 bufferbf16_mirrors_initialized: bool-- tracks if BF16 mirrors are populatedf32_to_bf16_kernel: CudaFunction-- used for bulk conversionbf16_to_f32_kernel: CudaFunction-- used for checkpointing reverse sync
Rationale: The per-tensor BF16 structs were designed for the old fused 1-warp-per-sample kernel that loaded weights from global memory. The cuBLAS path reads from the flat bf16_params_buf at precomputed offsets. The per-tensor mirrors are 100% redundant with the flat buffer.
7. BF16 Conversion in Adam Graph (gpu_dqn_trainer.rs)
After the Adam update writes to params_buf (F32), append an f32_to_bf16_kernel launch to graph_adam that converts the entire params_buf -> bf16_params_buf. This is captured in the CUDA Graph so it has zero launch overhead on replay.
// In build_graph_adam() -- after adam_update_kernel, before graph capture ends:
unsafe {
self.stream
.launch_builder(&self.f32_to_bf16_kernel)
.arg(&self.ptrs.params_buf)
.arg(&self.ptrs.bf16_params_buf)
.arg(&(total_params as i32))
.launch(LaunchConfig {
grid_dim: (ceildiv(total_params, 256) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
})?;
}
Target BF16 buffer (bf16_target_params_buf) is updated during EMA sync -- add the same f32_to_bf16_kernel launch after the existing ema_kernel that syncs target_params_buf.
Files Changed
| File | Changes |
|---|---|
crates/ml/src/cuda_pipeline/batched_forward.rs |
Replace cublasSgemm -> cublasGemmEx with BF16 inputs; add BF16 activation buffers; new gemmex_bf16_layer() method; new bf16_weight_ptrs() helper; fused add_bias_relu_bf16_kernel |
crates/ml/src/cuda_pipeline/batched_backward.rs |
Replace cublasSgemm -> cublasGemmEx with mixed BF16/F32 inputs; read BF16 activations and weights |
crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs |
Add CachedPtrs struct and initialization; remove 3 cuStreamSynchronize calls; add deferred readback with interval; replace 10 spectral norm launches with 1 batched launch; add f32_to_bf16_kernel to graph_adam; delete ensure_bf16_mirrors(), sync_online_bf16(); remove 11 EventTrackingGuard instances |
crates/ml/src/cuda_pipeline/gpu_weights.rs |
Delete DuelingWeightSetBf16, BranchingWeightSetBf16, CuriosityWeightSetBf16, KernelWeightPackBf16, alloc_bf16_mirror(), convert_f32_to_bf16() and all associated impls (~200 LOC) |
crates/ml/src/trainers/dqn/fused_training.rs |
Remove cuStreamSynchronize from run_ensemble_step; remove local EvtGuard; use CachedPtrs |
crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu |
Add SpectralNormDesc struct and batched_spectral_norm_kernel; keep existing spectral_norm_kernel (backward compat for tests) |
crates/ml/src/cuda_pipeline/bias_kernels.cu |
Add add_bias_relu_bf16_kernel variant that outputs both F32 and BF16 |
Performance Budget
Estimated per-step breakdown on H100 (300 steps/epoch, B=256, ~288K params):
| Component | Current | After | Savings |
|---|---|---|---|
| Forward (3 passes x 10 GEMMs) | ~40ms | ~13ms | 3x BF16 tensor cores |
| Backward (10 GEMMs) | ~30ms | ~20ms | 1.5x TF32 mixed |
cuStreamSynchronize x4 |
~8ms | ~0ms | Removed 3, deferred 1 |
raw_device_ptr() x110 |
~5ms | ~0ms | Cached u64 lookups |
EventTrackingGuard x11 |
~2ms | ~0ms | Eliminated |
| Spectral norm (10 launches) | ~1ms | ~0.2ms | 1 batched launch |
| F32->BF16 conversion | 0ms | ~0.1ms | New, but in CUDA Graph |
| Loss + optimizer | ~15ms | ~15ms | No change |
| Total per step | ~101ms | ~48ms | ~2.1x |
| Per epoch (300 steps) | ~30s | ~14s | -- |
Note: The 14s estimate is conservative. Combined with experience collection overlap (already async) and the H100's 3 TB/s HBM3, actual epoch time should be <=5s.
Testing Plan
Unit Tests (must all pass, 129+ existing)
SQLX_OFFLINE=true cargo test -p ml --lib
All existing tests use the same GpuDqnTrainer code path. BF16 conversion is invisible to test assertions since loss/Q-values are computed in F32.
Smoke Test (gradient stability)
FOXHUNT_TEST_DATA=test_data/futures-baseline cargo test -p ml --lib -- smoke_tests --ignored --nocapture
Verify:
- Gradient norms remain bounded (no NaN/Inf from BF16 underflow)
- Q-values converge (loss decreasing trend over 10+ epochs)
- No CUDA errors in stream after 50 epochs
BF16 Precision Regression Test (new)
Add a test that runs 50 training epochs with both F32-only and BF16 paths on the same seed, comparing:
- Final Sharpe ratio: must be within 5% relative difference
- Mean Q-value: must be within 10% relative difference
- Gradient norm statistics: mean and max must be within 20%
This test validates that BF16 forward + TF32 backward does not degrade convergence quality.
Cached Pointer Correctness Test (new)
Add a test that:
- Constructs
GpuDqnTrainer - Verifies all
CachedPtrsfields matchraw_device_ptr()calls - Runs 10 training steps
- Verifies pointers haven't changed (buffers are fixed post-init)
Batched Spectral Norm Equivalence Test (new)
Add a test that:
- Initializes random weights + u/v vectors
- Runs 10 individual
spectral_norm_kernellaunches (current path) - Runs 1
batched_spectral_norm_kernellaunch (new path) - Compares all weight matrices element-wise (must be bitwise identical since F32 throughout)
Deferred Readback Test (new)
Add a test that:
- Runs 200 training steps with
readback_interval=100 - Verifies sync happens at steps 100 and 200
- Verifies returned scalars at step 100 match actual GPU values
- Verifies steps 1-99 return cached values without blocking
Success Criteria
- Epoch time <=5s on H100 (80GB, SM90) with B=256, 300 steps/epoch
- All 129+ existing tests pass without modification
- Gradient norms bounded: no NaN/Inf over 50 epochs on smoke test data
- Convergence parity: Sharpe ratio within 5% of F32 baseline over 50 epochs
- Zero regression: no new CUDA errors, no memory leaks (same peak VRAM +/- 10%)
Non-Goals
- Bug fixes in training logic (Spec B scope)
- Dead code cleanup beyond BF16 infrastructure (Spec C scope)
- Full BF16 backward pass (TF32 mixed precision is acceptable and preserves gradient quality)
- FP16 support (BF16 has superior dynamic range for training; FP16 risks gradient underflow)
- Multi-GPU changes (single-GPU H100 is the target)
Risks and Mitigations
| Risk | Impact | Mitigation |
|---|---|---|
cublasGemmEx not capturable in CUDA Graph |
Breaks graph_forward capture | Verify with CUDA 12.x docs; fallback: call outside graph (lose ~2ms/step launch overhead) |
| BF16 activation underflow for small features | NaN propagation | Monitor activation statistics; BF16 range is +-3.4e38 (same exponent bits as F32), underflow only at <1e-38 which doesn't occur in normalized features |
| Cached pointer invalidation after cudarc GC | Wrong memory access | cudarc doesn't GC/move device allocations; verify with assert_eq!(cached, fresh) in debug builds |
| Batched spectral norm shared memory overflow | Kernel crash for large matrices | Current max matrix is 256x256; shmem[256] is sufficient; add assert(out_dim <= 256 && in_dim <= 256) in kernel |
| Deferred readback masking training divergence | Late detection of NaN loss | Force readback on NaN detection (check isnan(cached_loss) and force sync if true) |
Implementation Order
- CachedPtrs -- lowest risk, enables all subsequent changes
- Dead BF16 code removal -- reduces confusion, no behavior change
- Batched spectral norm -- standalone kernel change, easy to test
- Zero-sync hot path -- remove 3 syncs, defer 1
- BF16 forward pass -- core change, requires activation buffers
- BF16 backward (mixed) -- depends on forward BF16 activations
- F32->BF16 in graph_adam -- closes the loop
Steps 1-4 can be landed as a single commit (performance-only, no precision change). Steps 5-7 are a second commit (precision change, requires BF16 regression test).