feat: risk_budget_forward + apply_risk_budget + risk_budget_backward CUDA kernels

5th branch: h_s2 → ReLU hidden → sigmoid R ∈ (0,1).
apply_risk_budget: scales magnitude Q-values (Full×R, Half×sqrt(R)),
produces per-sample CVaR alpha and commitment lambda.
Backward: chain rule through sigmoid → ReLU → FC weights via atomicAdd.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-16 01:41:33 +02:00
parent 8b53fe25c7
commit be3dd47fbf

View File

@@ -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]);
}
}
}