Files
foxhunt/docs/plans/2026-03-03-bf16-training.md
2026-03-04 18:36:28 +01:00

18 KiB

BF16 Training Implementation Plan

For Claude: REQUIRED SUB-SKILL: Use superpowers:executing-plans to implement this plan task-by-task.

Goal: Switch all 10 ML models from FP32 to BF16 training on Ampere+ GPUs with dynamic detection and zero casts in the training hot path.

Architecture: Add a single training_dtype(device) -> DType function that returns BF16 on Ampere+ CUDA, F32 elsewhere. Thread it through all ~150 VarBuilder sites, training tensor creation, GPU replay buffer, and checkpoint loading. Loss stays F32 (single scalar cast). CUDA experience collector kernels untouched (stay f32).

Tech Stack: Candle (v0.9.1 git pin), half crate (2.6.0), cudarc, safetensors


Phase 1: Core Infrastructure

Task 1: Add training_dtype() function

Files:

  • Modify: crates/ml/src/dqn/mixed_precision.rs

Step 1: Add the public function after detect_from_gpu_name() (~line 180)

/// Returns the optimal training DType for the given device.
/// Ampere+ CUDA GPUs → BF16 (tensor core acceleration).
/// Everything else (CPU, older GPUs) → F32.
pub fn training_dtype(device: &candle_core::Device) -> candle_core::DType {
    match device {
        candle_core::Device::Cuda(_) => {
            // Check if GPU supports BF16 natively
            if let Some(config) = detect_from_gpu_name_auto() {
                match config.dtype {
                    DTypeSelection::BF16 => candle_core::DType::BF16,
                    DTypeSelection::F16 => candle_core::DType::F32, // F16 needs loss scaling, stay F32
                }
            } else {
                candle_core::DType::F32
            }
        }
        _ => candle_core::DType::F32,
    }
}

Note: detect_from_gpu_name_auto() already exists at line 182 — it reads the GPU name from the CUDA device and calls detect_from_gpu_name(). Reuse it.

Step 2: Re-export from the module's public API

Ensure training_dtype is accessible as crate::dqn::mixed_precision::training_dtype. Check the module's pub use or mod visibility.

Step 3: Build check

SQLX_OFFLINE=true cargo check -p ml 2>&1 | head -20

Step 4: Commit

git add crates/ml/src/dqn/mixed_precision.rs
git commit -m "feat(ml): add training_dtype() for dynamic BF16 detection"

Task 2: Fix Mamba2 scalar_tensor BF16 rejection

Files:

  • Modify: crates/ml/src/mamba/mod.rs:605

Step 1: Change the match arm at line 605

Replace the BF16/F16 rejection:

// Before (line 605):
DType::F8E4M3 | DType::U8 | DType::U32 | DType::I64 | DType::BF16 | DType::F16 => {
    Err(MLError::ModelError(format!(
        "Unsupported dtype: {:?}",
        dtype
    )))
},

// After:
DType::BF16 | DType::F16 => {
    // Create in F32 then cast — half types can't be created directly from f64
    Tensor::new(&[value as f32], device)?
        .to_dtype(dtype)?
        .reshape(())?
        .ok_or_else(|| MLError::ModelError("scalar reshape failed".into()))
},
DType::F8E4M3 | DType::U8 | DType::U32 | DType::I64 => {
    Err(MLError::ModelError(format!(
        "Unsupported dtype: {:?}",
        dtype
    )))
},

Note: Check the exact return type — scalar_tensor may return Result<Tensor, MLError>. Adjust the reshape/return accordingly. The key is: create as f32, cast to target dtype, return scalar.

Step 2: Build check

SQLX_OFFLINE=true cargo check -p ml 2>&1 | head -20

Step 3: Commit

git add crates/ml/src/mamba/mod.rs
git commit -m "fix(ml): allow BF16/F16 in Mamba2 scalar_tensor helper"

Phase 2: DQN Module (largest surface area)

Task 3: VarBuilder sites — DQN core networks

Change DType::F32training_dtype(&device) (or training_dtype(device) if device is a reference) in all VarBuilder::from_varmap calls across the DQN module.

Files (all need the same mechanical change):

  • crates/ml/src/dqn/dqn.rs:669
  • crates/ml/src/dqn/network.rs:256,261,301,363
  • crates/ml/src/dqn/agent.rs:366,371,597
  • crates/ml/src/dqn/dueling.rs:139
  • crates/ml/src/dqn/distributional_dueling.rs:155
  • crates/ml/src/dqn/quantile_regression.rs:90
  • crates/ml/src/dqn/curiosity.rs:38
  • crates/ml/src/dqn/factored_q_network.rs:72
  • crates/ml/src/dqn/rainbow_agent.rs:69,72
  • crates/ml/src/dqn/rainbow_network.rs:429,450

Pattern for each site:

// Before:
VarBuilder::from_varmap(&vars, DType::F32, &device)
// After:
VarBuilder::from_varmap(&vars, training_dtype(&device), &device)

Add use crate::dqn::mixed_precision::training_dtype; at the top of each file that doesn't already import it.

Step 1: Apply the change to all files listed above. Use replace_all where DType::F32 appears only in VarBuilder contexts. Where DType::F32 also appears in non-VarBuilder contexts (tensor creation, loss), change only the VarBuilder lines.

Step 2: Build check

SQLX_OFFLINE=true cargo check -p ml 2>&1 | head -30

Step 3: Commit

git add crates/ml/src/dqn/
git commit -m "feat(ml): BF16 VarBuilder for DQN core networks"

Task 4: VarBuilder sites — DQN layer modules

Same pattern for the layer-level modules that have many VarBuilder sites:

Files:

  • crates/ml/src/dqn/noisy_layers.rs:314,324,346,391,431,447,487
  • crates/ml/src/dqn/residual.rs:178,196,222,249,276,301,330,353
  • crates/ml/src/dqn/attention.rs:465,479,509,549,583,614
  • crates/ml/src/dqn/spectral_norm.rs:240,254,275,307,330,355,376
  • crates/ml/src/dqn/rmsnorm.rs:239,254,269,309,353,358,406,410,458,484

Same mechanical change. These files likely have DType::F32 ONLY in VarBuilder contexts, so replace_all may be safe. Verify by reading each file first.

Step 1: Apply changes. Step 2: Build check. Step 3: Commit

git add crates/ml/src/dqn/
git commit -m "feat(ml): BF16 VarBuilder for DQN layers (noisy, residual, attention, spectral, rmsnorm)"

Task 5: GPU replay buffer — BF16 states

Files:

  • Modify: crates/ml/src/cuda_pipeline/gpu_replay_buffer.rs:74-75

Step 1: Change states/next_states allocation to use dynamic dtype

// Line 74-75, change:
let states = Tensor::zeros(&[cap, sdim], DType::F32, device)?;
let next_states = Tensor::zeros(&[cap, sdim], DType::F32, device)?;
// To:
let dtype = training_dtype(device);
let states = Tensor::zeros(&[cap, sdim], dtype, device)?;
let next_states = Tensor::zeros(&[cap, sdim], dtype, device)?;

Keep rewards, dones, priorities as DType::F32. Keep actions as DType::U32.

Step 2: Verify insert_batch() callers cast correctly

The insert_batch() at line 167 uses slice_scatter which requires matching dtypes. The caller (DQN trainer) builds state tensors from f32 experience data. Add a .to_dtype(self.states.dtype())? cast on the incoming states/next_states args inside insert_batch():

// Inside insert_batch(), before slice_scatter:
let states = states.to_dtype(self.states.dtype())?;
let next_states = next_states.to_dtype(self.next_states.dtype())?;

This is the ONE cast at the data ingestion boundary. After this, all sample() returns match the buffer dtype (BF16 on Ampere+).

Step 3: Build check. Step 4: Commit

git add crates/ml/src/cuda_pipeline/gpu_replay_buffer.rs
git commit -m "feat(ml): BF16 states in GPU replay buffer (50% VRAM savings)"

Task 6: GPU weights extraction — handle BF16 weights

Files:

  • Modify: crates/ml/src/cuda_pipeline/gpu_weights.rs:179,204

Step 1: Cast to F32 before extraction

The CUDA experience collector kernel expects f32 weights. When model weights are BF16, cast before extracting:

// In extract_one() at line 179, change:
.to_vec1::<f32>()
// To:
.to_dtype(candle_core::DType::F32)?
.to_vec1::<f32>()

Same for sync_one() at line 204. This is the boundary cast from BF16 model weights → f32 CUDA kernel. Happens once per epoch during experience collection, not in the training hot path.

Step 2: Build check. Step 3: Commit

git add crates/ml/src/cuda_pipeline/gpu_weights.rs
git commit -m "feat(ml): handle BF16 weights in GPU weight extraction"

Task 7: GPU data pre-upload — BF16 feature tensors

Files:

  • Modify: crates/ml/src/cuda_pipeline/mod.rs:132,135,346,353

Step 1: Cast feature/target uploads to training dtype

In DqnGpuData::upload() (line 132):

// After creating the tensor from f32 data, cast:
let features = Tensor::from_vec(flat_features, (num_bars, feature_dim), device)?
    .to_dtype(training_dtype(device))?;
let targets = Tensor::from_vec(flat_targets, (num_bars, target_dim), device)?
    .to_dtype(training_dtype(device))?;

Same pattern for GpuBufferPool::upload_dqn (lines 346, 353) — cast after from_slice.

For PPO PpoGpuData::upload() (line 418) — same cast.

This is the data ingestion boundary cast. All downstream build_batch_states() and bar_features() calls return BF16 directly.

Step 2: Build check. Step 3: Commit

git add crates/ml/src/cuda_pipeline/mod.rs
git commit -m "feat(ml): BF16 GPU data pre-upload for DQN and PPO"

Task 8: DQN training tensors — CPU replay buffer path

Files:

  • Modify: crates/ml/src/dqn/dqn.rscompute_loss_internal()

The CPU replay buffer path creates training batch tensors from Vec<f32>. These need to match the model's weight dtype.

Step 1: Cast batch tensors at creation

At lines 1540-1569, after each Tensor::from_vec:

// States/next_states — cast to model dtype for matmul compatibility
let states_tensor = Tensor::from_vec(states, (batch_size, self.config.state_dim), device)?
    .to_dtype(training_dtype(device))?;
let next_states_tensor = Tensor::from_vec(next_states, (batch_size, self.config.state_dim), device)?
    .to_dtype(training_dtype(device))?;

Actions stay U32. Rewards/dones/importance-weights stay F32 — they're used in loss math, not matmuls.

For the GPU replay buffer path, states already come back as BF16 from sample() (Task 5), so no change needed there.

Step 2: Verify loss stays F32

The distributional loss at line 1650/1658 already has to_dtype(DType::F32) enforcement. Keep as-is.

Step 3: Build check. Step 4: Commit

git add crates/ml/src/dqn/dqn.rs
git commit -m "feat(ml): BF16 training tensors in DQN compute_loss"

Task 9: DQN trainer auxiliary tensors

Files:

  • Modify: crates/ml/src/trainers/dqn/trainer.rs:1015,1130,2134,2141,3123,3647

Same pattern — cast state tensors used in select_actions_batch, curiosity, and Q-value logging to training_dtype(&self.device):

let tensor = Tensor::from_vec(states, shape, &self.device)?
    .to_dtype(training_dtype(&self.device))?;

These are not in the training hot path (they're action selection and logging), so the single cast is fine.

Step 1: Apply casts at listed lines. Step 2: Build check. Step 3: Commit

git add crates/ml/src/trainers/dqn/trainer.rs
git commit -m "feat(ml): BF16 auxiliary tensors in DQN trainer"

Phase 3: PPO Module

Task 10: VarBuilder sites — PPO networks

Files:

  • crates/ml/src/ppo/ppo.rs:301,549
  • crates/ml/src/ppo/lstm_networks.rs:52,267
  • crates/ml/src/ppo/continuous_policy.rs:83
  • crates/ml/src/ppo/flow_policy/mod.rs:124
  • crates/ml/src/ppo/flow_policy/coupling_layer.rs:260

Same pattern: DType::F32training_dtype(&device).

Also change checkpoint loading at lines 1795 and 1852:

// Before:
VarBuilder::from_mmaped_safetensors(&[path], DType::F32, &device)
// After:
VarBuilder::from_mmaped_safetensors(&[path], training_dtype(&device), &device)

Step 1: Apply all changes. Step 2: Build check. Step 3: Commit

git add crates/ml/src/ppo/
git commit -m "feat(ml): BF16 VarBuilder and checkpoints for PPO networks"

Task 11: PPO training tensors

Files:

  • Modify: crates/ml/src/trainers/ppo.rs:834,895,960,1014

Cast state tensors to training dtype for forward pass compatibility:

let states = Tensor::from_vec(state_floats, shape, &self.device)?
    .to_dtype(training_dtype(&self.device))?;

Lines 1156, 1346, 1348, 1372 (rewards, returns, values) — keep F32, these are loss/metric tensors not fed to the network.

Step 1: Apply casts to state tensors only. Step 2: Build check. Step 3: Commit

git add crates/ml/src/trainers/ppo.rs
git commit -m "feat(ml): BF16 training tensors in PPO trainer"

Phase 4: Supervised Models (8 models)

Task 12: TFT VarBuilder sites

Files:

  • crates/ml/src/tft/mod.rs:339
  • crates/ml/src/tft/quantized_grn.rs:295,315
  • crates/ml/src/tft/quantized_attention.rs:417
  • crates/ml/src/tft/quantized_lstm.rs:417,442
  • crates/ml/src/tft/quantized_vsn.rs:61,249
  • crates/ml/src/tft/varmap_quantization.rs:676,722

Same DType::F32training_dtype(&device) pattern.

Step 1: Apply, build, commit.

git commit -m "feat(ml): BF16 VarBuilder for TFT"

Task 13: Mamba2 VarBuilder site

Files:

  • crates/ml/src/mamba/mod.rs:631
  • crates/ml/src/mamba/ssd_layer.rs:556

Already fixed scalar_tensor in Task 2. Now change VarBuilder dtype.

Step 1: Apply, build, commit.

git commit -m "feat(ml): BF16 VarBuilder for Mamba2"

Task 14: Liquid/CfC VarBuilder sites

Files:

  • crates/ml/src/liquid/candle_cfc.rs:450,460,475,489,499,524,548,558,578,601,639 (11 sites)
  • crates/ml/src/liquid/adapter.rs:58
  • crates/ml/src/liquid/training.rs:505

Step 1: Apply, build, commit.

git commit -m "feat(ml): BF16 VarBuilder for Liquid/CfC"

Task 15: KAN VarBuilder sites

Files:

  • crates/ml/src/kan/layer.rs:138,150
  • crates/ml/src/kan/network.rs:94,107
  • crates/ml/src/kan/trainable.rs:47

Step 1: Apply, build, commit.

git commit -m "feat(ml): BF16 VarBuilder for KAN"

Task 16: xLSTM VarBuilder sites

Files:

  • crates/ml/src/xlstm/slstm.rs:136,148,164,172
  • crates/ml/src/xlstm/mlstm.rs:229,243,259,267,282
  • crates/ml/src/xlstm/block.rs:127,139,152,161
  • crates/ml/src/xlstm/network.rs:172,185,198,214,236,248
  • crates/ml/src/xlstm/trainable.rs:45

Step 1: Apply, build, commit.

git commit -m "feat(ml): BF16 VarBuilder for xLSTM"

Task 17: Diffusion VarBuilder sites

Files:

  • crates/ml/src/diffusion/sampler.rs:163,230
  • crates/ml/src/diffusion/denoiser.rs:233,244,258,270,285
  • crates/ml/src/diffusion/trainable.rs:45

Step 1: Apply, build, commit.

git commit -m "feat(ml): BF16 VarBuilder for Diffusion"

Task 18: TGGN + TLOB VarBuilder sites

Files:

  • crates/ml/src/tgnn/trainable_adapter.rs:87
  • crates/ml/src/tlob/trainable_adapter.rs:116
  • crates/ml/src/trainers/tlob.rs:210

Step 1: Apply, build, commit.

git commit -m "feat(ml): BF16 VarBuilder for TGGN and TLOB"

Phase 5: Remaining Sites

Task 19: Ensemble adapters + misc

Files:

  • crates/ml/src/ensemble/adapters/liquid.rs:46,62
  • crates/ml/src/ensemble/adapters/diffusion.rs:49,72
  • crates/ml/src/ensemble/adapters/kan.rs:44,60
  • crates/ml/src/ensemble/adapters/tlob.rs:100,123
  • crates/ml/src/ensemble/adapters/tggn.rs:77,96
  • crates/ml/src/ensemble/adapters/xlstm.rs:57,79
  • crates/ml/src/portfolio_transformer.rs:190
  • crates/ml/src/features/multi_timeframe.rs:269,560,576
  • crates/ml/src/trainers/online_learning.rs:581
  • crates/ml/src/explainability/integrated_gradients.rs:200,247,301

Step 1: Apply, build, commit.

git commit -m "feat(ml): BF16 VarBuilder for ensemble adapters and misc modules"

Phase 6: Validation

Task 20: Full workspace build and test

Step 1: Workspace build

SQLX_OFFLINE=true cargo check --workspace 2>&1 | tail -5

Expected: 0 errors. Fix any dtype mismatches — common issues:

  • expected F32 but got BF16 — a tensor created without the dtype cast feeding into a module that expects matched dtypes
  • cannot add BF16 and F32 — missing cast at a boundary

Step 2: Clippy

SQLX_OFFLINE=true cargo clippy --workspace -- -D warnings 2>&1 | tail -10

Step 3: ML crate tests

SQLX_OFFLINE=true cargo test -p ml --lib 2>&1 | tail -20

Expected: 2506+ tests pass. Tests run on CPU → F32 path, no behavior change.

Step 4: Commit if any fixes were needed

git commit -m "fix(ml): resolve BF16 dtype mismatches"

Task 21: BF16 integration test (optional — requires GPU)

Create a minimal integration test verifying BF16 training works end-to-end on CUDA:

Files:

  • Create: crates/ml/tests/bf16_training_integration.rs
//! Integration test: verify BF16 training on Ampere+ GPU
//! Run with: SQLX_OFFLINE=true cargo test -p ml --test bf16_training_integration

#[cfg(feature = "cuda")]
mod bf16_tests {
    use ml::dqn::mixed_precision::training_dtype;
    use candle_core::{Device, DType};

    #[test]
    fn test_training_dtype_returns_bf16_on_cuda() {
        if let Ok(device) = Device::new_cuda(0) {
            let dtype = training_dtype(&device);
            // On Ampere+ (L40S, H100), should be BF16
            // On older GPUs, F32 is fine too
            assert!(dtype == DType::BF16 || dtype == DType::F32);
        }
    }

    #[test]
    fn test_training_dtype_returns_f32_on_cpu() {
        let device = Device::Cpu;
        assert_eq!(training_dtype(&device), DType::F32);
    }
}

Real validation is the hyperopt run on L40S — compare trial Sharpe distributions.

Step 1: Create test, build, commit.

git commit -m "test(ml): add BF16 training dtype integration test"

Summary

Phase Tasks Sites Changed Commit Count
1: Infrastructure 1-2 training_dtype fn + mamba fix 2
2: DQN 3-9 ~80 VarBuilder + CUDA pipeline + training tensors 7
3: PPO 10-11 ~9 VarBuilder + training tensors + checkpoints 2
4: Supervised 12-18 ~60 VarBuilder across 8 models 7
5: Remaining 19 ~20 ensemble/misc sites 1
6: Validation 20-21 Build + test + integration test 2
Total 21 tasks ~150 sites ~21 commits

Risk Checkpoints

After Phase 2 (DQN complete): full workspace build must pass. DQN is the most complex module — if it compiles, the rest is mechanical.

After Phase 6: all 2506+ tests must pass on CPU. GPU validation via hyperopt run on L40S.