jgrusewski
75f83888ca
feat(bf16): mixed-precision kernels, f32 IS-weights, CUTLASS padding, fast_isnan
Mixed-precision loss/grad kernels:
- MSE + C51 loss: float softmax/projection/TD-error (prevents bf16 exp overflow)
- MSE + C51 grad: float arithmetic + bf16 range clamp before atomicAdd
- Shared memory: float (4 bytes/elem) for numerically stable reductions
- Bias kernels: float add+clamp ±500 (prevents bf16 Inf cascade between layers)
- Noisy bias kernel: same float clamping
fast_isnan/fast_isinf (ROOT CAUSE FIX):
- nvcc --use_fast_math implies --no-nans → isnan()/isinf() compiled to false
- ALL NaN guards across ALL kernels were dead code
- Added bit-pattern IEEE 754 checks to common_device_functions.cuh
- Replaced isnan/isinf in 7 kernel files (21 occurrences)
- ml-dqn build.rs: all kernels now get common header (no more standalone)
f32 PER IS-weights:
- GpuBatchSlices.weights: CudaSlice<u16> → CudaSlice<f32>
- GpuBatch.weights: GpuTensor → CudaSlice<f32>
- Loss/grad kernel signatures: const __nv_bfloat16* → const float*
- Upload path: separate f32 memcpy instead of bf16 staging
- Eliminates bf16 overflow in IS-weight storage
CUTLASS padding:
- pad32() helper: round up to next multiple of 32
- 6 value-logit buffers: pad32(num_atoms) (51 → 64)
- 6 branch-logit buffers: +32*3 padding per branch
895/895 unit tests, 8/9 smoke tests pass.
50-epoch convergence: NaN at step ~100-200 — backward pass produces NaN
gradients within the CUDA graph replay (same atomic execution as Adam).
Root cause: bf16 backward GemmEx inputs can overflow. Needs mixed-precision
backward pass (same pattern as loss kernels) or f32 gradient output buffers.
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-03-28 21:40:09 +01:00
..
2026-03-13 10:18:35 +01:00
2026-03-27 19:52:44 +01:00
2026-03-14 11:35:15 +01:00
2026-03-13 10:18:35 +01:00
2026-03-13 10:18:35 +01:00
2026-03-28 21:40:09 +01:00
2026-03-13 10:18:35 +01:00
2026-03-13 10:18:35 +01:00
2026-03-13 10:18:35 +01:00
2026-03-28 12:13:45 +01:00
2026-03-13 10:18:35 +01:00
2026-03-28 21:40:09 +01:00
2026-03-28 01:51:45 +01:00
2026-03-28 01:51:45 +01:00
2026-03-15 11:59:31 +01:00
2026-03-27 00:33:05 +01:00
2026-03-19 00:39:03 +01:00
2026-03-13 10:18:35 +01:00
2026-03-13 10:18:35 +01:00
2026-03-28 13:11:47 +01:00
2026-03-16 21:01:28 +01:00
2026-03-13 10:18:35 +01:00
2026-03-13 10:18:35 +01:00
2026-03-13 10:18:35 +01:00
2026-03-28 01:51:45 +01:00
2026-03-13 10:18:35 +01:00
2026-03-13 10:18:35 +01:00
2026-03-13 10:18:35 +01:00
2026-03-18 00:53:47 +01:00
2026-03-14 11:35:15 +01:00
2026-03-15 11:59:31 +01:00