feat(ml-alpha): TFT GRN forward+backward kernels for multi-horizon heads (Phase 1.7a)
Per-horizon GRN structure (Lim et al. 2021 §3.3 adapted to scalar output): eta_2[k, m] = GELU(W1[k, m, :] @ h + b1[k, m]) # [HIDDEN] → [HEAD_MID] eta_1[k, m] = W2[k, m, :] @ eta_2[k, :] + b2[k, m] # [HEAD_MID] → [HEAD_MID] gate_lin[k] = W_gate[k, :] @ eta_1[k, :] + b_gate[k] # → scalar main[k] = W_main[k, :] @ eta_1[k, :] + b_main[k] # → scalar skip[k] = W_skip[k, :] @ h + b_skip[k] # [HIDDEN] → scalar logit[k] = skip[k] + sigmoid(gate_lin[k]) * main[k] p[k] = sigmoid(logit[k]) Gated residual lets each per-horizon head learn "linear vs deeper-transform" gating, matching the regime-conditional alpha pattern from pearl_snapshot_alpha_is_regime_conditional (~20% of book states carry the edge; spread-Q4 hits 75% acc, middle quintiles below chance). Backward chain rule covers all 10 parameter tensors + the trunk gradient (skip-path direct + main-path through W2→GELU→W1, lambda-scaled). Single-writer discipline (no atomicAdd per feedback_no_atomicadd.md): - Thread m owns row m of grad_w1 (col i in 0..HIDDEN), row m of grad_w2 (col m_in in 0..HEAD_MID), and column m of d_eta_2. - Threads 0..4 own per-horizon scalar grads (skip/gate/main biases). - Trunk grad_h tiles i over 2 strides of HEAD_MID for HIDDEN=128 coverage. Shared mem: ~6.5KB (s_a1 + s_z2 + s_d_eta1 + s_d_eta2 + s_d_z1 + scalars), well within 48KB limit. Existing 2-layer MLP kernels (Tasks 1.3/1.4) stay in the cubin as ablation baseline; the wired path becomes GRN once perception.rs lands. build.rs cache-bust → v7. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
@@ -20,10 +20,11 @@ const KERNELS: &[&str] = &[
|
||||
"layer_norm", // Phase 1: trunk pre-CfC normalisation
|
||||
];
|
||||
|
||||
// Cache bust v6 (2026-05-17): Phase 1 model capacity — LayerNorm
|
||||
// kernel added (cuda/layer_norm.cu) + multi_horizon_heads.cu will get
|
||||
// 2-layer MLP heads. Old cubins don't have the new symbols. Force
|
||||
// fresh nvcc compile on the cluster's /cargo-target PVC.
|
||||
// Cache bust v7 (2026-05-17): Phase 1.7 GRN heads — appended
|
||||
// multi_horizon_heads_grn_fwd_batched + _bwd_batched kernels to
|
||||
// cuda/multi_horizon_heads.cu (TFT gated residual: 2-layer MLP body +
|
||||
// GLU gate + skip-projection from trunk). Old cubins don't have the
|
||||
// new symbols. Force fresh nvcc compile.
|
||||
|
||||
fn main() {
|
||||
println!("cargo:rerun-if-changed=build.rs");
|
||||
|
||||
@@ -425,3 +425,334 @@ extern "C" __global__ void multi_horizon_heads_2layer_bwd_batched(
|
||||
__syncthreads();
|
||||
}
|
||||
}
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────
|
||||
// TFT Gated Residual Network (GRN) heads — Phase 1.7 (2026-05-17)
|
||||
// ─────────────────────────────────────────────────────────────────────
|
||||
//
|
||||
// Per-horizon GRN structure (Lim et al. 2021 §3.3, adapted to scalar
|
||||
// output per horizon):
|
||||
//
|
||||
// eta_2[k, m_out] = GELU(W1[k, m_out, i] @ h[i] + b1[k, m_out]) # [HIDDEN] → [HEAD_MID]
|
||||
// eta_1[k, m_out] = W2[k, m_out, m_in] @ eta_2[k, m_in] + b2[...] # [HEAD_MID] → [HEAD_MID]
|
||||
// gate_lin[k] = W_gate[k, m] @ eta_1[k, m] + b_gate[k] # [HEAD_MID] → scalar
|
||||
// main[k] = W_main[k, m] @ eta_1[k, m] + b_main[k]
|
||||
// skip[k] = W_skip[k, i] @ h[i] + b_skip[k] # [HIDDEN] → scalar
|
||||
// logit[k] = skip[k] + sigmoid(gate_lin[k]) * main[k]
|
||||
// p[k] = sigmoid(logit[k])
|
||||
//
|
||||
// Naming convention follows the existing 2-layer kernel: w1 = first
|
||||
// linear (HIDDEN → HEAD_MID, GELU-activated), w2 = second linear
|
||||
// (HEAD_MID → HEAD_MID, no activation; produces eta_1).
|
||||
//
|
||||
// Thread layout (block_dim = HEAD_MID = 64):
|
||||
// m = threadIdx.x -- owns "output_row" dim of (k, m_out=m) for w1/w2/eta_1/eta_2
|
||||
// for skip / gate / main scalars: thread tid<5 handles horizon k=tid
|
||||
|
||||
extern "C" __global__ void multi_horizon_heads_grn_fwd_batched(
|
||||
const float* __restrict__ w1, // [5, HEAD_MID, HIDDEN]
|
||||
const float* __restrict__ b1, // [5, HEAD_MID]
|
||||
const float* __restrict__ w2, // [5, HEAD_MID_out, HEAD_MID_in]
|
||||
const float* __restrict__ b2, // [5, HEAD_MID_out]
|
||||
const float* __restrict__ w_gate, // [5, HEAD_MID]
|
||||
const float* __restrict__ b_gate, // [5]
|
||||
const float* __restrict__ w_main, // [5, HEAD_MID]
|
||||
const float* __restrict__ b_main, // [5]
|
||||
const float* __restrict__ w_skip, // [5, HIDDEN]
|
||||
const float* __restrict__ b_skip, // [5]
|
||||
const float* __restrict__ h, // [B, HIDDEN]
|
||||
int n_batch,
|
||||
float* __restrict__ probs, // [B, 5]
|
||||
float* __restrict__ z1_out, // [B, 5, HEAD_MID] — pre-GELU z1 (= W1 @ h + b1)
|
||||
float* __restrict__ a1_out, // [B, 5, HEAD_MID] — post-GELU eta_2
|
||||
float* __restrict__ z2_out, // [B, 5, HEAD_MID] — eta_1 (= W2 @ eta_2 + b2)
|
||||
float* __restrict__ gate_logit_out, // [B, 5] — pre-sigmoid gate scalar
|
||||
float* __restrict__ main_out, // [B, 5]
|
||||
float* __restrict__ logit_out // [B, 5] — pre-final-sigmoid logit (skip + sigma(gate)*main)
|
||||
) {
|
||||
int b_idx = blockIdx.x;
|
||||
int m = threadIdx.x;
|
||||
if (b_idx >= n_batch || m >= HEAD_MID_H) return;
|
||||
|
||||
const float* h_row = h + (long long)b_idx * HIDDEN_H;
|
||||
|
||||
__shared__ float s_a1[N_HORIZONS_H * HEAD_MID_H]; // post-GELU eta_2
|
||||
__shared__ float s_z2[N_HORIZONS_H * HEAD_MID_H]; // eta_1
|
||||
|
||||
// Pass 1: z1[k, m] = W1[k, m, :] @ h + b1[k, m]; a1 = GELU(z1).
|
||||
// Thread m owns row m (output dim) for all k. Sequential over HIDDEN.
|
||||
#pragma unroll
|
||||
for (int k = 0; k < N_HORIZONS_H; ++k) {
|
||||
float z = b1[k * HEAD_MID_H + m];
|
||||
for (int i = 0; i < HIDDEN_H; ++i) {
|
||||
z += w1[((long long)(k * HEAD_MID_H) + m) * HIDDEN_H + i] * h_row[i];
|
||||
}
|
||||
const float a = gelu_act(z);
|
||||
s_a1[k * HEAD_MID_H + m] = a;
|
||||
z1_out[((long long)b_idx * N_HORIZONS_H + k) * HEAD_MID_H + m] = z;
|
||||
a1_out[((long long)b_idx * N_HORIZONS_H + k) * HEAD_MID_H + m] = a;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Pass 2: z2[k, m] = W2[k, m_out=m, m_in=n] @ a1[k, n] + b2[k, m].
|
||||
// Pure linear — no activation. Thread m owns output row m for all k.
|
||||
#pragma unroll
|
||||
for (int k = 0; k < N_HORIZONS_H; ++k) {
|
||||
float z = b2[k * HEAD_MID_H + m];
|
||||
for (int n = 0; n < HEAD_MID_H; ++n) {
|
||||
z += w2[((long long)(k * HEAD_MID_H) + m) * HEAD_MID_H + n]
|
||||
* s_a1[k * HEAD_MID_H + n];
|
||||
}
|
||||
s_z2[k * HEAD_MID_H + m] = z;
|
||||
z2_out[((long long)b_idx * N_HORIZONS_H + k) * HEAD_MID_H + m] = z;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Pass 3: per-horizon scalars. Threads 0..4 each handle one horizon.
|
||||
// gate_lin[k] = W_gate[k, :] @ z2[k, :] + b_gate[k]
|
||||
// main[k] = W_main[k, :] @ z2[k, :] + b_main[k]
|
||||
// skip[k] = W_skip[k, :] @ h + b_skip[k]
|
||||
// logit[k] = skip[k] + sigmoid(gate_lin[k]) * main[k]
|
||||
// p[k] = sigmoid(logit[k])
|
||||
if (m < N_HORIZONS_H) {
|
||||
const int k = m;
|
||||
float gl = b_gate[k];
|
||||
float mv = b_main[k];
|
||||
#pragma unroll
|
||||
for (int n = 0; n < HEAD_MID_H; ++n) {
|
||||
const float z2_kn = s_z2[k * HEAD_MID_H + n];
|
||||
gl += w_gate[k * HEAD_MID_H + n] * z2_kn;
|
||||
mv += w_main[k * HEAD_MID_H + n] * z2_kn;
|
||||
}
|
||||
float sk = b_skip[k];
|
||||
for (int i = 0; i < HIDDEN_H; ++i) {
|
||||
sk += w_skip[k * HIDDEN_H + i] * h_row[i];
|
||||
}
|
||||
const float sigmoid_gl = 1.0f / (1.0f + expf(-gl));
|
||||
const float logit = sk + sigmoid_gl * mv;
|
||||
const float p = 1.0f / (1.0f + expf(-logit));
|
||||
|
||||
gate_logit_out[(long long)b_idx * N_HORIZONS_H + k] = gl;
|
||||
main_out[(long long)b_idx * N_HORIZONS_H + k] = mv;
|
||||
logit_out[(long long)b_idx * N_HORIZONS_H + k] = logit;
|
||||
probs[(long long)b_idx * N_HORIZONS_H + k] = p;
|
||||
}
|
||||
}
|
||||
|
||||
// GRN backward — chain rule:
|
||||
//
|
||||
// d_logit = grad_probs * p * (1 - p)
|
||||
// d_skip = d_logit
|
||||
// d_main = d_logit * sigma(gate_lin)
|
||||
// d_gate_a = d_logit * main
|
||||
// d_gate_lin = d_gate_a * sigma(gate_lin) * (1 - sigma(gate_lin))
|
||||
//
|
||||
// Per-horizon scalar grads:
|
||||
// grad_w_skip[k, i] += d_skip[k] * h[i]
|
||||
// grad_b_skip[k] += d_skip[k]
|
||||
// grad_w_main[k, m] += d_main[k] * eta_1[k, m]
|
||||
// grad_b_main[k] += d_main[k]
|
||||
// grad_w_gate[k, m] += d_gate_lin[k] * eta_1[k, m]
|
||||
// grad_b_gate[k] += d_gate_lin[k]
|
||||
//
|
||||
// Through eta_1:
|
||||
// d_eta_1[k, m] = d_main * w_main[k, m] + d_gate_lin * w_gate[k, m]
|
||||
// grad_w2[k, m_out, m_in] += d_eta_1[k, m_out] * eta_2[k, m_in]
|
||||
// grad_b2[k, m_out] += d_eta_1[k, m_out]
|
||||
// d_eta_2[k, m_in] = sum_{m_out} d_eta_1[k, m_out] * w2[k, m_out, m_in]
|
||||
//
|
||||
// Through GELU + W1:
|
||||
// d_z1[k, m] = d_eta_2[k, m] * gelu_prime(z1[k, m])
|
||||
// grad_w1[k, m, i] += d_z1[k, m] * h[i]
|
||||
// grad_b1[k, m] += d_z1[k, m]
|
||||
//
|
||||
// Trunk gradient (lambda-scaled, sums two paths):
|
||||
// grad_h[i] = sum_k lambda[k] * (
|
||||
// d_skip[k] * w_skip[k, i] (skip path)
|
||||
// + sum_m d_z1[k, m] * w1[k, m, i] (main path through W2→W1)
|
||||
// ) + carry[i]
|
||||
//
|
||||
// Single-writer discipline:
|
||||
// Thread m (= m_out for w2, m for w1) owns row m of grad_w2 (writes
|
||||
// cols m_in for all m_in) and row m of grad_w1 (writes cols i for
|
||||
// all i). Thread m also owns col m of d_eta_2 (sums over m_out).
|
||||
// Thread tid (with tid < 5) owns horizon-scalar grads.
|
||||
// Trunk grad_h tiles i over 2 strides of HEAD_MID.
|
||||
|
||||
extern "C" __global__ void multi_horizon_heads_grn_bwd_batched(
|
||||
const float* __restrict__ w1, // [5, HEAD_MID, HIDDEN]
|
||||
const float* __restrict__ w2, // [5, HEAD_MID_out, HEAD_MID_in]
|
||||
const float* __restrict__ w_gate, // [5, HEAD_MID]
|
||||
const float* __restrict__ w_main, // [5, HEAD_MID]
|
||||
const float* __restrict__ w_skip, // [5, HIDDEN]
|
||||
const float* __restrict__ probs, // [B, 5]
|
||||
const float* __restrict__ grad_probs, // [B, 5]
|
||||
const float* __restrict__ z1, // [B, 5, HEAD_MID]
|
||||
const float* __restrict__ a1, // [B, 5, HEAD_MID]
|
||||
const float* __restrict__ z2, // [B, 5, HEAD_MID]
|
||||
const float* __restrict__ gate_logit, // [B, 5]
|
||||
const float* __restrict__ main_val, // [B, 5]
|
||||
const float* __restrict__ h, // [B, HIDDEN]
|
||||
const float* __restrict__ grad_h_carry, // [B, HIDDEN] (nullptr OK)
|
||||
const float* __restrict__ lambda, // [5]
|
||||
int n_batch,
|
||||
float* __restrict__ grad_w1, // [5, HEAD_MID, HIDDEN] (+=)
|
||||
float* __restrict__ grad_b1, // [5, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_w2, // [5, HEAD_MID_out, HEAD_MID_in] (+=)
|
||||
float* __restrict__ grad_b2, // [5, HEAD_MID_out] (+=)
|
||||
float* __restrict__ grad_w_gate, // [5, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_b_gate, // [5] (+=)
|
||||
float* __restrict__ grad_w_main, // [5, HEAD_MID] (+=)
|
||||
float* __restrict__ grad_b_main, // [5] (+=)
|
||||
float* __restrict__ grad_w_skip, // [5, HIDDEN] (+=)
|
||||
float* __restrict__ grad_b_skip, // [5] (+=)
|
||||
float* __restrict__ grad_h // [B, HIDDEN] (overwrite + carry)
|
||||
) {
|
||||
int tid = threadIdx.x;
|
||||
if (tid >= HEAD_MID_H) return;
|
||||
|
||||
__shared__ float s_lambda[N_HORIZONS_H];
|
||||
if (tid < N_HORIZONS_H) {
|
||||
const float l = lambda[tid];
|
||||
s_lambda[tid] = (l > 0.0f) ? l : 1.0f; // sentinel: zero buffer → 1.0
|
||||
}
|
||||
|
||||
__shared__ float s_d_skip[N_HORIZONS_H]; // = d_logit (skip-path grad)
|
||||
__shared__ float s_d_main[N_HORIZONS_H];
|
||||
__shared__ float s_d_gate_lin[N_HORIZONS_H];
|
||||
__shared__ float s_d_eta1[N_HORIZONS_H * HEAD_MID_H]; // d_eta_1 = d_z2
|
||||
__shared__ float s_d_eta2[N_HORIZONS_H * HEAD_MID_H]; // d_eta_2 = d_a1
|
||||
__shared__ float s_d_z1[N_HORIZONS_H * HEAD_MID_H];
|
||||
__syncthreads();
|
||||
|
||||
for (int bi = 0; bi < n_batch; ++bi) {
|
||||
const int m = tid;
|
||||
const float* h_row = h + (long long)bi * HIDDEN_H;
|
||||
|
||||
// Pass 1: per-horizon scalar grads. Threads 0..4 only.
|
||||
if (tid < N_HORIZONS_H) {
|
||||
const int k = tid;
|
||||
const float p = probs[(long long)bi * N_HORIZONS_H + k];
|
||||
const float dp = grad_probs[(long long)bi * N_HORIZONS_H + k];
|
||||
const float d_logit = dp * p * (1.0f - p);
|
||||
|
||||
const float gl = gate_logit[(long long)bi * N_HORIZONS_H + k];
|
||||
const float sigmoid_gl = 1.0f / (1.0f + expf(-gl));
|
||||
const float mv = main_val[(long long)bi * N_HORIZONS_H + k];
|
||||
|
||||
const float d_skip = d_logit;
|
||||
const float d_main = d_logit * sigmoid_gl;
|
||||
const float d_gate_a = d_logit * mv;
|
||||
const float d_gate_lin = d_gate_a * sigmoid_gl * (1.0f - sigmoid_gl);
|
||||
|
||||
s_d_skip[k] = d_skip;
|
||||
s_d_main[k] = d_main;
|
||||
s_d_gate_lin[k] = d_gate_lin;
|
||||
|
||||
grad_b_skip[k] += d_skip;
|
||||
grad_b_main[k] += d_main;
|
||||
grad_b_gate[k] += d_gate_lin;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Pass 2: per-(k, m_out=m) compute d_eta_1[k, m] and accumulate
|
||||
// grad_w_main[k, m] / grad_w_gate[k, m]. Single-writer per (k, m).
|
||||
#pragma unroll
|
||||
for (int k = 0; k < N_HORIZONS_H; ++k) {
|
||||
const float d_main_k = s_d_main[k];
|
||||
const float d_gate_lin_k = s_d_gate_lin[k];
|
||||
const float w_main_km = w_main[k * HEAD_MID_H + m];
|
||||
const float w_gate_km = w_gate[k * HEAD_MID_H + m];
|
||||
const float z2_km = z2[((long long)bi * N_HORIZONS_H + k) * HEAD_MID_H + m];
|
||||
|
||||
const float d_eta1_km = d_main_k * w_main_km + d_gate_lin_k * w_gate_km;
|
||||
s_d_eta1[k * HEAD_MID_H + m] = d_eta1_km;
|
||||
|
||||
grad_w_main[k * HEAD_MID_H + m] += d_main_k * z2_km;
|
||||
grad_w_gate[k * HEAD_MID_H + m] += d_gate_lin_k * z2_km;
|
||||
|
||||
// grad_b2[k, m] += d_eta_1[k, m] -- thread m sole writer of row m.
|
||||
grad_b2[k * HEAD_MID_H + m] += d_eta1_km;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Pass 3: through W2.
|
||||
// grad_w2[k, m_out=m, m_in] += d_eta_1[k, m] * eta_2[k, m_in]
|
||||
// thread m loops over m_in and is sole writer of row m.
|
||||
// d_eta_2[k, m_in=m] = sum_{m_out} d_eta_1[k, m_out] * W2[k, m_out, m_in=m]
|
||||
// thread m owns column m.
|
||||
#pragma unroll
|
||||
for (int k = 0; k < N_HORIZONS_H; ++k) {
|
||||
const float d_eta1_km = s_d_eta1[k * HEAD_MID_H + m];
|
||||
// grad_w2[k, m_out=m, m_in=n] += d_eta_1[k, m] * eta_2[k, n]
|
||||
for (int n = 0; n < HEAD_MID_H; ++n) {
|
||||
const float a1_kn = a1[((long long)bi * N_HORIZONS_H + k) * HEAD_MID_H + n];
|
||||
grad_w2[((long long)(k * HEAD_MID_H) + m) * HEAD_MID_H + n] += d_eta1_km * a1_kn;
|
||||
}
|
||||
|
||||
// d_eta_2[k, m_in=m] = sum_{m_out} d_eta_1[k, m_out] * w2[k, m_out, m]
|
||||
float d_eta2_km = 0.0f;
|
||||
for (int m_out = 0; m_out < HEAD_MID_H; ++m_out) {
|
||||
d_eta2_km += s_d_eta1[k * HEAD_MID_H + m_out]
|
||||
* w2[((long long)(k * HEAD_MID_H) + m_out) * HEAD_MID_H + m];
|
||||
}
|
||||
s_d_eta2[k * HEAD_MID_H + m] = d_eta2_km;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Pass 4: through GELU. Thread m owns col m of d_z1.
|
||||
// d_z1[k, m] = d_eta_2[k, m] * gelu_prime(z1[k, m])
|
||||
// grad_b1[k, m] += d_z1[k, m]
|
||||
#pragma unroll
|
||||
for (int k = 0; k < N_HORIZONS_H; ++k) {
|
||||
const float d_eta2_km = s_d_eta2[k * HEAD_MID_H + m];
|
||||
const float z1_km = z1[((long long)bi * N_HORIZONS_H + k) * HEAD_MID_H + m];
|
||||
const float d_z1_km = d_eta2_km * gelu_prime(z1_km);
|
||||
s_d_z1[k * HEAD_MID_H + m] = d_z1_km;
|
||||
grad_b1[k * HEAD_MID_H + m] += d_z1_km;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// Pass 5: grad_w1[k, m, i] += d_z1[k, m] * h[i]
|
||||
// Thread m owns row m for all k, i (HIDDEN cols sequential).
|
||||
for (int i = 0; i < HIDDEN_H; ++i) {
|
||||
const float h_bi_i = h_row[i];
|
||||
#pragma unroll
|
||||
for (int k = 0; k < N_HORIZONS_H; ++k) {
|
||||
grad_w1[((long long)(k * HEAD_MID_H) + m) * HIDDEN_H + i] +=
|
||||
s_d_z1[k * HEAD_MID_H + m] * h_bi_i;
|
||||
}
|
||||
}
|
||||
|
||||
// Pass 6: trunk grad_w_skip + grad_h.
|
||||
// grad_w_skip[k, i] += d_skip[k] * h[i]
|
||||
// grad_h[bi, i] = sum_k lambda[k] * (
|
||||
// d_skip[k] * w_skip[k, i] # skip path
|
||||
// + sum_m d_z1[k, m] * w1[k, m, i] # main path
|
||||
// ) + carry[i]
|
||||
//
|
||||
// Tile i over 2 strides of HEAD_MID. Thread tid handles i=tid and i=tid+HEAD_MID.
|
||||
#pragma unroll
|
||||
for (int tile = 0; tile < HIDDEN_H / HEAD_MID_H; ++tile) {
|
||||
const int i = tile * HEAD_MID_H + tid;
|
||||
const float h_bi_i = h_row[i];
|
||||
float acc_grad_h = 0.0f;
|
||||
#pragma unroll
|
||||
for (int k = 0; k < N_HORIZONS_H; ++k) {
|
||||
grad_w_skip[k * HIDDEN_H + i] += s_d_skip[k] * h_bi_i;
|
||||
|
||||
float contrib = s_d_skip[k] * w_skip[k * HIDDEN_H + i];
|
||||
#pragma unroll
|
||||
for (int mm = 0; mm < HEAD_MID_H; ++mm) {
|
||||
contrib += s_d_z1[k * HEAD_MID_H + mm]
|
||||
* w1[((long long)(k * HEAD_MID_H) + mm) * HIDDEN_H + i];
|
||||
}
|
||||
acc_grad_h += s_lambda[k] * contrib;
|
||||
}
|
||||
const float carry = (grad_h_carry != nullptr)
|
||||
? grad_h_carry[(long long)bi * HIDDEN_H + i] : 0.0f;
|
||||
grad_h[(long long)bi * HIDDEN_H + i] = acc_grad_h + carry;
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user