17 tasks across 4 phases: - Phase 1: 10 tasks converting 35 CUDA kernel internals to native BF16 (bf16_* wrappers, native arithmetic, no __bfloat162float casts) - Phase 2: Delete GpuTensor + nvrtc dependency (replace with raw CudaSlice) - Phase 3: Fix Rust compilation errors at host boundaries - Phase 4: Wire cublasGemmEx + test Includes lessons learned from failed bulk-sed approach and explicit conversion rules for every agent. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
396 lines
15 KiB
Markdown
396 lines
15 KiB
Markdown
# BF16 Native Completion — Zero Casts, Zero nvrtc, Zero GpuTensor
|
||
|
||
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||
|
||
**Goal:** Complete the full BF16 conversion: native BF16 arithmetic in ALL CUDA kernels (no __bfloat162float casts), delete GpuTensor + nvrtc dependency, fix all Rust compilation errors, wire cublasGemmEx.
|
||
|
||
**Architecture:** Each CUDA kernel is converted individually by a dedicated agent that reads the full file and rewrites float locals/shared memory to `__nv_bfloat16` using bf16_* wrappers from `common_device_functions.cuh`. GpuTensor is deleted and replaced with raw `CudaSlice<half::bf16>` operations. The nvrtc cudarc feature is removed entirely — all kernels precompiled as cubins.
|
||
|
||
**Tech Stack:** Rust, CUDA (`__nv_bfloat16`, `cuda_bf16.h`), cudarc 0.19.3 (`CudaSlice<half::bf16>`, `f16` feature), cuBLAS (`cublasGemmEx`, `CUDA_R_16BF`)
|
||
|
||
---
|
||
|
||
## ⛔ LESSONS LEARNED — DO NOT REPEAT
|
||
|
||
| What failed | Why | Correct approach |
|
||
|---|---|---|
|
||
| Bulk `sed` on CUDA kernel internals | Mangles parentheses, creates `float * __nv_bfloat16` ambiguity | Agent reads FULL file, converts MANUALLY |
|
||
| Changing types without fixing arithmetic | `__nv_bfloat16` has no implicit conversion to/from `float` | Change BOTH pointer types AND local variable types |
|
||
| Leaving `float` shared memory with BF16 buffers | `__shfl_xor_sync` and arithmetic mix BF16 and float | Shared memory → `__nv_bfloat16`, use `bf16_warp_sum()` |
|
||
| Agents creating conversion layers | `f32_to_bf16_kernel` per-step = conversion layer | Data born as BF16 from source, NO per-step conversion |
|
||
| Incremental kernel conversion | Forward produces BF16, loss expects F32 → mismatch | ALL kernels converted atomically |
|
||
|
||
## BF16 Conversion Rules (for EVERY agent)
|
||
|
||
```
|
||
POINTERS: float* → __nv_bfloat16* (already done at parameter level)
|
||
LOCALS: float val → __nv_bfloat16 val (when loading from BF16 buffers)
|
||
SHARED MEM: __shared__ float → __shared__ __nv_bfloat16
|
||
ARITHMETIC: +, -, *, /, >, < → native (SM80+)
|
||
MATH: sqrtf() → bf16_sqrt() (from common header)
|
||
expf() → bf16_exp()
|
||
logf() → bf16_log()
|
||
powf() → bf16_pow()
|
||
fabsf() → bf16_fabs()
|
||
fmaxf() → bf16_fmax()
|
||
fminf() → bf16_fmin()
|
||
cosf() → bf16_cos()
|
||
rsqrtf()→ bf16(rsqrtf(__bfloat162float(x))) (no bf16 rsqrt HW)
|
||
tanhf() → bf16(tanhf(__bfloat162float(x)))
|
||
SHUFFLE: __shfl_xor_sync() → bf16_shfl_xor()
|
||
warp sum → bf16_warp_sum()
|
||
warp max → bf16_warp_max()
|
||
ATOMIC: atomicAdd() → atomicAddBF16() (from common header)
|
||
LITERALS: 0.0f → bf16_zero()
|
||
1.0f → bf16_one()
|
||
other → bf16(literal)
|
||
SCALARS: float lr, float gamma → KEEP as float, convert at use: bf16(lr)
|
||
INTEGERS: int → KEEP as int
|
||
EXCEPTIONS: total_loss → float* (single scalar, atomicAdd monitoring)
|
||
grad_norm_buf → float* (single scalar)
|
||
```
|
||
|
||
---
|
||
|
||
## Current State (commit f2304d68)
|
||
|
||
| Component | Status |
|
||
|---|---|
|
||
| Rust types (`CudaSlice<f32>` → `CudaSlice<half::bf16>`) | ✅ Done across ml + ml-core |
|
||
| CUDA kernel parameter types (`float*` → `__nv_bfloat16*`) | ✅ Done, 36/37 compile |
|
||
| CUDA kernel INTERNAL code (float locals → __nv_bfloat16) | ❌ 35 files have casts |
|
||
| dt_kernels.cu compilation | ❌ Fails (last kernel) |
|
||
| Common header BF16 wrappers | ✅ Complete (bf16_sqrt, bf16_warp_sum, etc.) |
|
||
| ml-core Rust errors | ❌ 45 errors (22 nvrtc + 23 BF16 boundary) |
|
||
| ml crate Rust errors | ❌ 23 errors |
|
||
| cublasGemmEx wired | ❌ forward_online_bf16 exists but not wired |
|
||
| GpuTensor | ❌ Still exists, wraps CudaSlice, uses nvrtc |
|
||
|
||
---
|
||
|
||
## Phase 1: CUDA Kernels — Native BF16 Internals
|
||
|
||
### Task 1: Fix dt_kernels.cu (the one failing kernel)
|
||
|
||
**Complexity:** HIGH (24 casts, 14 kernels, 723 lines, shared memory, warp shuffles, layernorm, softmax, GELU)
|
||
|
||
**Files:**
|
||
- Modify: `crates/ml/src/cuda_pipeline/dt_kernels.cu`
|
||
|
||
- [ ] **Step 1: Read the ENTIRE file (723 lines)**
|
||
- [ ] **Step 2: Convert ALL float output parameters to __nv_bfloat16***
|
||
|
||
The following output parameters are still `float*` and must change:
|
||
- `dt_embed_kernel`: `output` → `__nv_bfloat16*`
|
||
- `dt_qkv_projection_kernel`: `Q_out`, `K_out`, `V_out` → `__nv_bfloat16*`
|
||
- `dt_causal_attention_kernel`: `output` → `__nv_bfloat16*`
|
||
- `dt_layernorm_kernel`: `output` → `__nv_bfloat16*`
|
||
- `dt_ffn_kernel`: `output` → `__nv_bfloat16*`
|
||
- `dt_action_head_kernel`: `logits` → `__nv_bfloat16*`
|
||
- `dt_return_to_go_kernel`: `returns_to_go` → `__nv_bfloat16*`
|
||
- `dt_build_trajectories_kernel`: `trajectories` → `__nv_bfloat16*`
|
||
- `dt_compute_rewards_actions_kernel`: `rewards` → `__nv_bfloat16*`
|
||
- Keep `total_loss` and `per_sample_loss` as `float*` (loss monitoring scalars)
|
||
- Keep `d_logits`, `dW`, `db`, `d_input` as `__nv_bfloat16*` (already converted)
|
||
|
||
- [ ] **Step 3: Convert ALL float locals to __nv_bfloat16**
|
||
- [ ] **Step 4: Convert shared memory to __nv_bfloat16**
|
||
|
||
`extern __shared__ float shmem[]` → `extern __shared__ __nv_bfloat16 shmem[]`
|
||
All `float*` pointers into shmem → `__nv_bfloat16*`
|
||
Warp reduction: use `bf16_warp_sum()` / `bf16_warp_max()`
|
||
|
||
- [ ] **Step 5: Replace math functions with bf16_* wrappers**
|
||
|
||
`expf()` → `bf16_exp()`, `logf()` → `bf16_log()`, `rsqrtf((float)N)` → `bf16(rsqrtf((float)N))`, `tanhf()` → `bf16(tanhf(__bfloat162float(x)))` (no bf16 tanh wrapper — add if needed), `fmaxf/fminf` → `bf16_fmax/bf16_fmin`
|
||
|
||
- [ ] **Step 6: Fix float literals**
|
||
|
||
`0.0f` → `bf16_zero()`, `1.0f` → `bf16_one()`, `-1e30f` → `bf16(-1e30f)`, `0.7978845608f` → `bf16(0.7978845608f)`, `0.044715f` → `bf16(0.044715f)`, `1e-8f` → `bf16(1e-8f)`
|
||
|
||
- [ ] **Step 7: Verify compilation**
|
||
|
||
Run: `SQLX_OFFLINE=true cargo check -p ml 2>&1 | grep "nvcc failed" | wc -l`
|
||
Expected: 0 (all 37 kernels compile)
|
||
|
||
- [ ] **Step 8: Commit**
|
||
|
||
```bash
|
||
git commit -m "feat(bf16): dt_kernels.cu native BF16 — all 37 CUDA kernels compile"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 2: Convert simple kernels to native BF16 (batch of 12)
|
||
|
||
**Complexity:** LOW (1-5 casts each, small kernels)
|
||
|
||
**Files (12 kernels, one agent):**
|
||
- `ema_kernel.cu` (3 casts)
|
||
- `relu_mask_kernel.cu` (1 cast)
|
||
- `per_update_kernel.cu` (2 casts)
|
||
- `iqn_cvar_kernel.cu` (2 casts)
|
||
- `trade_stats_kernel.cu` (1 cast)
|
||
- `monitoring_kernel.cu` (1 cast)
|
||
- `backward_kernels.cu` (3 casts)
|
||
- `backtest_forward_supervised_kernel.cu` (1 cast)
|
||
- `backtest_forward_ppo_kernel.cu` (1 cast)
|
||
- `statistics_kernel.cu` (3 casts)
|
||
- `curiosity_training_kernel.cu` (2 casts)
|
||
- `her_relabel_kernel.cu` (3 casts)
|
||
|
||
For each file:
|
||
- [ ] Read file, find all `__bfloat162float()` / `__float2bfloat16()` casts
|
||
- [ ] Change the local variable from `float` to `__nv_bfloat16`
|
||
- [ ] Remove the cast wrapper
|
||
- [ ] Replace math functions with bf16_* wrappers where needed
|
||
- [ ] Fix float literals → bf16()
|
||
|
||
- [ ] **Verify:** `SQLX_OFFLINE=true cargo check -p ml 2>&1 | grep "nvcc failed" | wc -l` = 0
|
||
- [ ] **Commit**
|
||
|
||
---
|
||
|
||
### Task 3: Convert medium kernels (batch of 8)
|
||
|
||
**Complexity:** MEDIUM (5-17 casts each)
|
||
|
||
**Files (8 kernels, one agent):**
|
||
- `bias_kernels.cu` (7 casts)
|
||
- `c51_grad_kernel.cu` (5 casts)
|
||
- `mse_grad_kernel.cu` (6 casts)
|
||
- `expected_q_kernel.cu` (5 casts)
|
||
- `nstep_kernel.cu` (5 casts)
|
||
- `signal_adapter_kernel.cu` (5 casts)
|
||
- `cql_grad_kernel.cu` (7 casts)
|
||
- `q_stats_kernel.cu` (7 casts)
|
||
|
||
Same conversion rules as Task 2.
|
||
|
||
- [ ] **Verify + Commit**
|
||
|
||
---
|
||
|
||
### Task 4: Convert training_guard_kernel.cu + epsilon_greedy_kernel.cu
|
||
|
||
**Complexity:** MEDIUM (8 + 11 casts, warp shuffles in training_guard)
|
||
|
||
- [ ] training_guard: warp reductions → `bf16_warp_sum()`
|
||
- [ ] epsilon_greedy: branching action selection with bonus arrays
|
||
- [ ] **Verify + Commit**
|
||
|
||
---
|
||
|
||
### Task 5: Convert c51_loss_kernel.cu + mse_loss_kernel.cu
|
||
|
||
**Complexity:** HIGH (14 + 16 casts, shared memory softmax, block reductions)
|
||
|
||
These are the core loss kernels. Shared memory stores softmax intermediates.
|
||
|
||
- [ ] Change `extern __shared__ float shmem[]` → `extern __shared__ __nv_bfloat16 shmem[]`
|
||
- [ ] Block reduction helpers (block_reduce_sum, block_reduce_max) → use `bf16_warp_sum`/`bf16_warp_max`
|
||
- [ ] log_softmax: `bf16_log()`, `bf16_exp()` for transcendentals
|
||
- [ ] Keep `total_loss` as `float*` (atomicAdd monitoring scalar)
|
||
- [ ] **Verify + Commit**
|
||
|
||
---
|
||
|
||
### Task 6: Convert attention_kernel.cu + attention_backward_kernel.cu + ensemble_kernels.cu
|
||
|
||
**Complexity:** HIGH (12 + 38 + 19 casts)
|
||
|
||
attention_backward has the most casts outside the complex kernels. These have:
|
||
- Multi-head attention QKV projections
|
||
- Softmax with temperature scaling
|
||
- Adam-style weight updates in attention_backward
|
||
- KL divergence in ensemble
|
||
|
||
- [ ] **Verify + Commit**
|
||
|
||
---
|
||
|
||
### Task 7: Convert backtest_gather_kernel.cu + backtest_env_kernel.cu + backtest_metrics_kernel.cu
|
||
|
||
**Complexity:** MEDIUM-HIGH (17 + 4 + 3 casts)
|
||
|
||
backtest_gather is the largest — loads states from BF16 market data.
|
||
|
||
- [ ] **Verify + Commit**
|
||
|
||
---
|
||
|
||
### Task 8: Convert experience_kernels.cu (dedicated agent)
|
||
|
||
**Complexity:** VERY HIGH (49 casts, 30 kernel parameters, portfolio simulation, trade physics)
|
||
|
||
The most complex kernel. Uses trade_physics.cuh functions. Portfolio state read/write. Reward computation. Episode reset.
|
||
|
||
- [ ] trade_physics.cuh functions: verify ALL converted (parameters + locals)
|
||
- [ ] env_step kernel: portfolio state → BF16 reads/writes
|
||
- [ ] experience gathering: state/reward output → BF16
|
||
- [ ] **Verify + Commit**
|
||
|
||
---
|
||
|
||
### Task 9: Convert dqn_utility_kernels.cu (dedicated agent)
|
||
|
||
**Complexity:** HIGH (31 casts, Adam optimizer, spectral norm, gradient operations)
|
||
|
||
- [ ] Adam kernel: params/grads/m/v → `__nv_bfloat16`, arithmetic → native
|
||
- [ ] Gradient norm: BF16 input, warp reduction → `bf16_warp_sum()`
|
||
- [ ] SAXPY/clip kernels: BF16 operands
|
||
- [ ] Spectral norm: BF16 matvec with shared memory
|
||
- [ ] `powf(beta1, t)` → `bf16_pow(bf16(beta1), bf16((float)t))`
|
||
- [ ] **Verify + Commit**
|
||
|
||
---
|
||
|
||
### Task 10: Convert iql_value_kernel.cu + iqn_dual_head_kernel.cu (dedicated agent)
|
||
|
||
**Complexity:** VERY HIGH (43 + 29 casts)
|
||
|
||
IQL has SiLU activation, IQN has cosine embedding and quantile Huber loss.
|
||
|
||
- [ ] SiLU: `x * sigmoid(x)` → `x * bf16_one() / (bf16_one() + bf16_exp(-x))`
|
||
- [ ] Cosine embedding: `bf16_cos(bf16(pi) * bf16((float)(d+1)) * tau)`
|
||
- [ ] Quantile Huber: piecewise with `bf16_fabs()`
|
||
- [ ] **Verify + Commit**
|
||
|
||
---
|
||
|
||
## Phase 2: Delete GpuTensor + nvrtc
|
||
|
||
### Task 11: Replace GpuTensor usages with raw CudaSlice
|
||
|
||
GpuTensor methods used (from analysis):
|
||
- `from_host(&[f32], shape, stream)` → `stream.clone_htod(&bf16_data)?`
|
||
- `from_vec(vec, shape, stream)` → same
|
||
- `zeros(shape, stream)` → `stream.alloc_zeros::<half::bf16>(n)?`
|
||
- `cat(&[tensors], dim, stream)` → sequential `cuMemcpyDtoDAsync`
|
||
- `full(shape, val, stream)` → custom fill kernel or `memset` for zero
|
||
- `scalar(val, stream)` → `stream.clone_htod(&[half::bf16::from_f32(val)])?`
|
||
- `randn(shape, stream)` → GPU RNG kernel (already exists in dqn_utility_kernels)
|
||
- `.data()` → direct CudaSlice reference (already a thin wrapper)
|
||
- `.shape()` → tracked separately as `(usize, usize)` or config-derived
|
||
- `.numel()` → `slice.len()`
|
||
- `.to_host(stream)` → `stream.memcpy_dtoh(&slice, &mut host)?`
|
||
|
||
**Files to modify (~20 files):**
|
||
- `crates/ml/src/trainers/dqn/fused_training.rs` — HER batch operations
|
||
- `crates/ml/src/trainers/dqn/trainer/train_step.rs` — batch upload
|
||
- `crates/ml/src/trainers/dqn/config.rs` — type references
|
||
- `crates/ml/src/inference.rs` — inference path
|
||
- All other files with GpuTensor references
|
||
|
||
- [ ] **Step 1:** Create `bf16_cat` and `bf16_fill` helper functions in gpu_weights.rs or a new bf16_ops.rs
|
||
- [ ] **Step 2:** Replace ALL `GpuTensor::xxx` calls with raw CudaSlice equivalents
|
||
- [ ] **Step 3:** Remove `GpuTensor` import from all files
|
||
- [ ] **Step 4:** Verify compilation
|
||
- [ ] **Commit**
|
||
|
||
---
|
||
|
||
### Task 12: Delete ml-core cuda_autograd + nvrtc
|
||
|
||
- [ ] **Step 1:** Remove `cudarc` nvrtc feature from ml-core/Cargo.toml
|
||
- [ ] **Step 2:** Delete `crates/ml-core/src/cuda_compile.rs` (runtime kernel compilation)
|
||
- [ ] **Step 3:** Delete or gut `crates/ml-core/src/cuda_autograd/` modules that depend on nvrtc:
|
||
- `elementwise.rs`, `reductions.rs`, `loss.rs`, `linear.rs`, `activations.rs`, `optimizer.rs`
|
||
- Keep `gpu_tensor.rs` temporarily if other crates reference GpuTensor type
|
||
- [ ] **Step 4:** Fix all compilation errors
|
||
- [ ] **Step 5:** If GpuTensor type is still referenced, create a thin `GpuTensor = CudaSlice<half::bf16>` type alias
|
||
- [ ] **Commit**
|
||
|
||
---
|
||
|
||
## Phase 3: Rust Compilation Fixes
|
||
|
||
### Task 13: Fix ml-core Rust errors
|
||
|
||
After removing nvrtc (22 errors gone), fix remaining 23 BF16 boundary errors:
|
||
|
||
- `Vec<f32>` → `Vec<half::bf16>` at host upload boundaries
|
||
- `[f32; N]` → `[half::bf16; N]` for fixed-size host arrays
|
||
- `clone_htod(&[f32])` → `clone_htod(&[half::bf16])`
|
||
- `memcpy_dtoh` → download to `Vec<half::bf16>`, convert to f32 for display
|
||
|
||
- [ ] **Verify:** `SQLX_OFFLINE=true cargo check -p ml-core` = 0 errors
|
||
- [ ] **Commit**
|
||
|
||
---
|
||
|
||
### Task 14: Fix ml crate Rust errors
|
||
|
||
~23 remaining errors in trainers, config, tests:
|
||
|
||
- Host data upload paths (feature vectors, rewards from replay buffer)
|
||
- Scalar readback paths (loss, grad_norm for monitoring)
|
||
- Test helpers (alloc_dueling, alloc_branching)
|
||
|
||
- [ ] **Verify:** `SQLX_OFFLINE=true cargo check -p ml` = 0 errors
|
||
- [ ] **Commit**
|
||
|
||
---
|
||
|
||
## Phase 4: Wire cuBLAS + Test
|
||
|
||
### Task 15: Wire cublasGemmEx BF16 forward path
|
||
|
||
`forward_online_bf16()` already exists in batched_forward.rs. Wire it into `launch_cublas_forward()` in gpu_dqn_trainer.rs.
|
||
|
||
- [ ] Change `launch_cublas_forward` to call `forward_online_bf16` with BF16 weight pointers
|
||
- [ ] Change `forward_target` and `forward_online_next` similarly
|
||
- [ ] **Verify + Commit**
|
||
|
||
---
|
||
|
||
### Task 16: Wire cublasGemmEx BF16 backward path
|
||
|
||
Add `gemmex_bf16` to batched_backward.rs (same pattern as forward).
|
||
|
||
- [ ] dW computation: BF16 × BF16 → BF16
|
||
- [ ] dX computation: BF16 × BF16 → BF16
|
||
- [ ] **Verify + Commit**
|
||
|
||
---
|
||
|
||
### Task 17: Full test suite validation
|
||
|
||
- [ ] `SQLX_OFFLINE=true cargo test -p ml --lib -- dqn`
|
||
- [ ] `SQLX_OFFLINE=true cargo test -p ml-dqn --lib`
|
||
- [ ] `SQLX_OFFLINE=true cargo test -p ml --lib -- gradient_budget`
|
||
- [ ] Update test helpers to construct `half::bf16` data natively (no Vec<f32>)
|
||
- [ ] **Final commit + push**
|
||
|
||
---
|
||
|
||
## Execution Dependencies
|
||
|
||
```
|
||
Phase 1 (CUDA kernels — can parallelize):
|
||
Task 1 (dt_kernels) ─┐
|
||
Task 2 (12 simple) ─┤
|
||
Task 3 (8 medium) ─┤→ ALL kernels compile natively
|
||
Task 4 (guard+epsilon) ─┤
|
||
Task 5 (c51+mse loss) ─┤
|
||
Task 6 (attn+ensemble) ─┤
|
||
Task 7 (backtest) ─┤
|
||
Task 8 (experience) ─┤
|
||
Task 9 (dqn_utility) ─┤
|
||
Task 10 (iql+iqn) ─┘
|
||
|
||
Phase 2 (GpuTensor removal):
|
||
Task 11 (replace usages) → Task 12 (delete infrastructure)
|
||
|
||
Phase 3 (Rust fixes):
|
||
Task 13 (ml-core) → Task 14 (ml)
|
||
|
||
Phase 4 (wire + test):
|
||
Task 15 (forward) → Task 16 (backward) → Task 17 (tests)
|
||
```
|
||
|
||
Phase 1 tasks are ALL independent — can run as parallel agents.
|
||
Phase 2 depends on Phase 1 (kernels must compile).
|
||
Phase 3 depends on Phase 2 (nvrtc removal changes error count).
|
||
Phase 4 depends on Phase 3 (Rust must compile).
|