Every to_vec()/memcpy_dtoh across 48 files audited and annotated: - ~100 false positives: Rust slice .to_vec() (cpu-side, never touches GPU) - ~25 gpu-exit: legitimate scalar readbacks (loss, grad_norm, epoch state) - ~20 test-only readbacks: gated by #[cfg(test)] scope - ~10 cpu-side uploads: .to_vec() before from_vec() GPU upload - ~3 checkpoint exports: export_to_host at epoch boundary Annotations use inline comments: // cpu-side, // gpu-exit:, // test-only Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
6.2 KiB
Full GPU Forward Passes — Eliminate ALL CPU from ML Models
For agentic workers: REQUIRED: Use superpowers:subagent-driven-development. Steps use checkbox (
- [ ]) syntax for tracking.
Goal: Every model forward pass, backward pass, loss computation, and weight update runs 100% on GPU. Zero to_host(), to_vec(), memcpy_dtoh in any production code path. The only CPU touch is the final scalar that exits the system.
Architecture: Each model's forward pass must use only: GpuTensor ops (CUDA elementwise kernels), GpuLinear::forward() (cuBLAS sgemm), ActivationKernels (CUDA), ReductionKernels (CUDA), StreamTensor free functions (gpu_matmul, gpu_add, etc.). No host downloads.
Current State: 134 GPU→CPU downloads across 6 models + infrastructure
| Model | CPU downloads | Key operations on CPU |
|---|---|---|
| Mamba2 | 23 | Selective scan, SSM state update, loss computation |
| TFT | 17 | Variable selection, LSTM encoder per-step, quantile loss |
| xLSTM | 12 | mLSTM attention (QKV gate computation), sLSTM forward |
| KAN | 10 | B-spline evaluation, network forward |
| Liquid CfC | 8 | CfC cell dynamics, training perturbation |
| Diffusion | 4 | Denoiser forward, sampling |
| ml-dqn infra | 19 | Replay buffer, target update debug, dead neuron check |
| ml-ppo infra | 5 | PPO update_gpu residual downloads, CudaLinear checkpoint |
| ml crate infra | 33 | Action selector, portfolio, state build, signal adapter |
Phase 1: ml-dqn Infrastructure (19 downloads → 0)
Task 1a: gpu_replay_buffer.rs (11 downloads)
bf16_slice_to_gpu_tensor/u32_slice_to_gpu_tensor: type conversion helpers — replace with GPU cast kernelsmax_priorityreadbacks (5 sites): replace with GPU atomicMax, keep single scalar on devicesample_indices: indices/weights download for CPU PER — keep as CudaSlice, pass to fused trainercs_total: prefix sum total — single scalar, use mapped pinned memory
Task 1b: target_update.rs (3 remaining downloads)
get_average_value: download for debug logging — use ReductionKernels::stats().meandebug_compare_params: download for assertion — use GPU subtract + abs + max
Task 1c: dqn.rs (2 remaining downloads)
- Dead neuron detection: download weight_mu — use GPU NaN/zero check kernel
- Gradient finite check: download params — use GPU NaN check kernel
Task 1d: branching.rs (2 downloads)
- Debug weight comparison: use GPU subtract + abs + max
Phase 2: ml-supervised Models (74 downloads → 0)
Task 2a: TFT (17 downloads)
Files: tft/mod.rs, tft/lstm_encoder.rs, tft/variable_selection.rs, tft/quantile_outputs.rs
Key rewrites:
lstm_encoder.rs:129— per-step LSTM: download input, sequential CPU loop → needs GPU LSTM kernel (already have CudaLSTM in ml-ppo)variable_selection.rs:64-142— download inputs, CPU softmax+gather → GPU softmax + GPU gatherquantile_outputs.rs:73-181— download predictions/targets, CPU quantile loss → GPU quantile loss kernelmod.rs:617,664— download static context, predictions → keep on GPU
Task 2b: Mamba2 (23 downloads)
Files: mamba/mod.rs, mamba/scan_algorithms.rs, mamba/selective_state.rs, mamba/loss.rs
Key rewrites:
selective_state.rs:233— download input for f32→f64 conversion + sequential scan → GPU scan kernelscan_algorithms.rs:103— download for parallel scan → GPU parallel prefix sum (already have kernel!)loss.rs:17-58— download predictions/targets for MSE → GPU LossKernels::mse()mod.rs— 15+ downloads throughout forward/backward → systematic rewrite to StreamTensor ops
Task 2c: xLSTM (12 downloads)
Files: xlstm/network.rs, xlstm/mlstm.rs, xlstm/slstm.rs
Key rewrites:
mlstm.rs:119-125— download Q,K,V,i,f,o,c for mLSTM attention → GPU matmul + elementwise opsnetwork.rs:94— download input for per-layer processing → keep as StreamTensorslstm.rs:141-142— download outputs for assertion → GPU comparison
Task 2d: KAN (10 downloads)
Files: kan/spline.rs, kan/layer.rs, kan/network.rs
Key rewrites:
spline.rs:148-188— B-spline evaluation downloads grid+input → needs GPU B-spline kernellayer.rs:121— download output for shape check → GPU shape validationnetwork.rs:119— download final output → keep on GPU until caller needs it
Task 2e: Liquid CfC (8 downloads)
Files: liquid/candle_cfc.rs, liquid/training.rs, liquid/cells.rs, liquid/network.rs
Key rewrites:
candle_cfc.rs:223— download input for CfC cell → GPU CfC kerneltraining.rs:529— download for perturbation gradient → GPU perturbationcells.rs:328— download for backbone → keep on GPUnetwork.rs:198— download activations → keep on GPU
Task 2f: Diffusion (4 downloads)
Files: diffusion/denoiser.rs, diffusion/sampler.rs
Key rewrites:
denoiser.rs:249-282— test downloads, minimalsampler.rs:194— download samples → keep on GPU
Phase 3: ml-ppo Infrastructure (5 downloads → 0)
ppo.rs:643,658— update_gpu still downloads actions + returns for per-element gather → GPU gather kernelcuda_nn/linear.rs:299,303— weight/bias download → checkpoint only, keep but markcuda_nn/tensor_util.rs:27— CudaVec to_vec → only for debug
Phase 4: ml Crate Infrastructure (33 downloads → minimal)
- Action selector: final action readback (1 u32) — keep (exits system)
- GPU monitoring: summary stats readback — single struct, keep
- Signal adapter: test-only downloads — move behind #[cfg(test)]
- PPO collector download_*(): debug methods — delete or gate
- Portfolio sim: old CPU path — delete (GPU path exists)
- State building: some test code downloads — gate
- Training guard: accumulator readback — single scalar, keep
- DQN trainer readback_buf: mapped pinned memory — correct pattern, keep
- Metrics/training_loop: single scalar readbacks — keep
Execution Strategy
Phase 1 (ml-dqn): 1 agent, ~2 hours Phase 2a-2f (ml-supervised): 3 agents parallel (TFT+Mamba, xLSTM+KAN, Liquid+Diffusion), ~4-6 hours each Phase 3 (ml-ppo): 1 agent, ~1 hour Phase 4 (ml infra): 1 agent, ~2 hours
Total: 6 agents, ~8-12 hours with 3 parallel.