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:
@@ -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]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user