fix: Components 2+5 fully GPU — zero CPU compute, zero copies
xLSTM: qlstm_step kernel reads q_stats_buf directly on device, updates C matrix on device, writes context[8] to device buffer. qlstm_train_step runs SGD on device. No DtoH readback. RK4 ODE: liquid_tau_rk4_step kernel reads per_branch_q_gap_ema (device), context (device), writes liquid_mod (device). No pinned memory — plain CudaSlice. Replaces CPU update_liquid_tau entirely. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -72,22 +72,24 @@ context = h_t // [8] context vector
|
||||
```
|
||||
|
||||
### Context usage
|
||||
- `liquid_tau`: context modulates target (1.0 → context-dependent target)
|
||||
- `liquid_tau`: context modulates target (1.0 → context-dependent target) via device buffer
|
||||
- `selectivity`: context concatenated with h_s2 for richer priority prediction
|
||||
- `backtracking`: context divergence (|h_t - h_{t-1}|) as additional plateau signal
|
||||
- `adaptive RK4`: switch to RK4 when context norm exceeds threshold
|
||||
|
||||
### Parameters
|
||||
- 6 weight matrices: W_k, W_v, W_q, W_i, W_f, W_o each [11, 8] = 528 total
|
||||
- C matrix: [8, 8] = 64 floats (persistent state)
|
||||
- n vector: [8] = 8 floats (persistent state)
|
||||
- All CPU-side. Zero GPU overhead.
|
||||
- C matrix: [8, 8] = 64 floats (persistent device state)
|
||||
- n vector: [8] = 8 floats (persistent device state)
|
||||
- **All GPU-resident. Zero CPU compute. Zero HtoD copies in training loop.**
|
||||
|
||||
### Implementation
|
||||
- `QLSTMCell` struct in `crates/ml/src/trainers/dqn/trainer/mod.rs`
|
||||
- Called from `reduce_current_q_stats` path, after Q-stats are downloaded
|
||||
- Context vector stored as `[f32; 8]` field on DQNTrainer
|
||||
- Weights trained by simple gradient descent on a prediction loss: "predict next step's q_mean from context" (self-supervised)
|
||||
- All state lives on GPU: `qlstm_weights[528]`, `qlstm_C[64]`, `qlstm_n[8]`, `qlstm_context[8]` as CudaSlice
|
||||
- New CUDA kernel `qlstm_step`: one thread, one block. Reads q_stats_buf[7] + per_branch_q_gaps[4] directly from GPU (already device-resident from reduce_current_q_stats). Runs mLSTM update on device-resident C matrix. Writes context[8] to device buffer. No DtoH readback.
|
||||
- New CUDA kernel `qlstm_train_step`: self-supervised prediction loss (predict next q_mean from context). SGD weight update in-place on device. One thread.
|
||||
- Context buffer `qlstm_context[8]` read by liquid_tau_rk4_step kernel and selectivity_forward kernel (both GPU-side)
|
||||
- Weights initialized once in constructor via alloc_zeros (identity-like init)
|
||||
- Launched every 50 training steps from the same code path as reduce_current_q_stats (after Q-stats kernel, before eval_v_range update)
|
||||
|
||||
---
|
||||
|
||||
@@ -194,12 +196,15 @@ else:
|
||||
```
|
||||
|
||||
### Parameters
|
||||
Zero — algorithm change only. CPU-side, runs every 50 training steps.
|
||||
Zero new weights — algorithm change only. **All GPU-resident. No CPU compute.**
|
||||
|
||||
### Implementation
|
||||
- ~20 lines replacing Euler step in `update_liquid_tau`
|
||||
- `rk4_threshold: f32` field (default 0.5, tunable)
|
||||
- xLSTM context norm drives the switching
|
||||
- New CUDA kernel `liquid_tau_rk4_step`: one thread, one block. Reads `per_branch_q_gap_ema[4]` (device buffer), `qlstm_context[8]` (device buffer), computes adaptive RK4/Euler, writes `liquid_mod[4]` (device buffer). No pinned memory — pure device read/write.
|
||||
- `per_branch_q_gap_ema[4]` moves from CPU field to `CudaSlice<f32>` on device. Updated by the kernel directly.
|
||||
- `liquid_mod[4]` moves from pinned device-mapped to plain `CudaSlice<f32>`. The c51_grad_kernel reads it via raw device pointer (same interface, simpler memory model).
|
||||
- `rk4_threshold: f32` passed as kernel argument (compile-time tunable)
|
||||
- Replaces the CPU-side `update_liquid_tau` method entirely — the Rust method just launches the kernel
|
||||
- Launched every 50 training steps after `qlstm_step` (depends on context output)
|
||||
|
||||
---
|
||||
|
||||
@@ -310,9 +315,9 @@ Q_select = quantile_blend(atoms, iqn_readiness) [Component 6]
|
||||
── Action Selection ──────────────────────────────────
|
||||
action = boltzmann(Q_select) [quantile-driven]
|
||||
|
||||
── Training Dynamics ─────────────────────────────────
|
||||
q_stats → xLSTM_context[8] [Component 2]
|
||||
liquid_tau(per_branch, context, adaptive_rk4) [Component 5]
|
||||
── Training Dynamics (ALL GPU — zero CPU compute) ────
|
||||
q_stats_buf → qlstm_step kernel → context[8] [Component 2, GPU]
|
||||
liquid_tau_rk4_step kernel(q_gap_ema, context) → liquid_mod[4] [Component 5, GPU]
|
||||
selectivity(h_s2, context) → PER priorities [CPBI selective]
|
||||
trajectory_backtracking(context_divergence) [CPBI backtrack]
|
||||
```
|
||||
|
||||
Reference in New Issue
Block a user