diff --git a/crates/ml/src/cuda_pipeline/experience_kernels.cu b/crates/ml/src/cuda_pipeline/experience_kernels.cu index 14faf31ed..0e3257b2c 100644 --- a/crates/ml/src/cuda_pipeline/experience_kernels.cu +++ b/crates/ml/src/cuda_pipeline/experience_kernels.cu @@ -4651,3 +4651,107 @@ extern "C" __global__ void homeostatic_regularizer( penalties[k] = pen; } + +/* ───────── Risk-budget branch (5th DQN head) ───────── */ + +extern "C" __global__ void risk_budget_forward( + const float* __restrict__ h_s2, + const float* __restrict__ w_risk_fc, + const float* __restrict__ b_risk_fc, + const float* __restrict__ w_risk_out, + const float* __restrict__ b_risk_out, + float* __restrict__ risk_hidden, + float* __restrict__ risk_budget_out, + int B, int SH2, int AH +) { + int i = blockIdx.x * blockDim.x + threadIdx.x; + if (i >= B) return; + + const float* h = h_s2 + (long long)i * SH2; + float* h_risk = risk_hidden + (long long)i * AH; + + for (int j = 0; j < AH; j++) { + float val = b_risk_fc[j]; + for (int k = 0; k < SH2; k++) { + val += w_risk_fc[(long long)j * SH2 + k] * h[k]; + } + h_risk[j] = fmaxf(val, 0.0f); + } + + float raw = b_risk_out[0]; + for (int j = 0; j < AH; j++) { + raw += w_risk_out[j] * h_risk[j]; + } + risk_budget_out[i] = 1.0f / (1.0f + expf(-raw)); +} + +extern "C" __global__ void apply_risk_budget( + const float* __restrict__ risk_budget, + float* __restrict__ q_values, + float* __restrict__ cvar_alpha_buf, + float* __restrict__ commit_lambda_buf, + int B, + int b0_size, + int b1_size, + int b2_size, + int b3_size +) { + int i = blockIdx.x * blockDim.x + threadIdx.x; + if (i >= B) return; + + float R = risk_budget[i]; + int total_actions = b0_size + b1_size + b2_size + b3_size; + + /* Scale magnitude Q-values: Small×1, Half×sqrt(R), Full×R */ + float* q_mag = q_values + (long long)i * total_actions + b0_size; + /* q_mag[0] = Small: unchanged */ + if (b1_size > 1) q_mag[1] *= sqrtf(fmaxf(R, 1e-6f)); + if (b1_size > 2) q_mag[2] *= fmaxf(R, 1e-6f); + + cvar_alpha_buf[i] = 0.1f + 0.4f * R; + commit_lambda_buf[i] = 0.01f * (1.0f - R); +} + +extern "C" __global__ void risk_budget_backward( + const float* __restrict__ h_s2, + const float* __restrict__ risk_hidden, + const float* __restrict__ risk_budget, + const float* __restrict__ d_q_mag, + const float* __restrict__ q_mag_pre, + const float* __restrict__ w_risk_out, + float* __restrict__ d_w_risk_fc, + float* __restrict__ d_b_risk_fc, + float* __restrict__ d_w_risk_out, + float* __restrict__ d_b_risk_out, + int B, int SH2, int AH, int b1_size +) { + int i = blockIdx.x * blockDim.x + threadIdx.x; + if (i >= B) return; + + float R = risk_budget[i]; + + float d_R = 0.0f; + const float* dq = d_q_mag + (long long)i * b1_size; + const float* qp = q_mag_pre + (long long)i * b1_size; + if (b1_size > 1) d_R += dq[1] * qp[1] * 0.5f / fmaxf(sqrtf(fmaxf(R, 1e-6f)), 1e-6f); + if (b1_size > 2) d_R += dq[2] * qp[2]; + + float d_raw = d_R * R * (1.0f - R); + + const float* h_risk = risk_hidden + (long long)i * AH; + atomicAdd(d_b_risk_out, d_raw); + for (int j = 0; j < AH; j++) { + atomicAdd(&d_w_risk_out[j], d_raw * h_risk[j]); + } + + const float* h = h_s2 + (long long)i * SH2; + for (int j = 0; j < AH; j++) { + float d_hj = d_raw * w_risk_out[j]; + if (h_risk[j] <= 0.0f) d_hj = 0.0f; + + atomicAdd(&d_b_risk_fc[j], d_hj); + for (int k = 0; k < SH2; k++) { + atomicAdd(&d_w_risk_fc[(long long)j * SH2 + k], d_hj * h[k]); + } + } +}