Skip to content

程式碼逐行解讀

LayerNorm 實作 (C)

Forward Pass

c
void layernorm_forward(float* out, float* mean, float* rstd,
                       float* inp, float* weight, float* bias,
                       int B, int T, int C) {
    float eps = 1e-5f;
    for (int b = 0; b < B; b++) {
        for (int t = 0; t < T; t++) {
            // 計算每個 (b,t) 位置的 C 維向量
            float* x = inp + b * T * C + t * C;

            // 步驟 1: 計算平均值
            float m = 0.0f;
            for (int i = 0; i < C; i++) m += x[i];
            m = m / C;

            // 步驟 2: 計算變異數
            float v = 0.0f;
            for (int i = 0; i < C; i++) {
                float xshift = x[i] - m;
                v += xshift * xshift;
            }
            v = v / C;

            // 步驟 3: 計算 reciprocal standard deviation
            float s = 1.0f / sqrtf(v + eps);

            // 步驟 4: 歸一化 + scale + shift
            float* out_bt = out + b * T * C + t * C;
            for (int i = 0; i < C; i++) {
                float n = s * (x[i] - m);
                out_bt[i] = n * weight[i] + bias[i];
            }

            // 快取 mean/rstd 供 backward 使用
            mean[b * T + t] = m;
            rstd[b * T + t] = s;
        }
    }
}

Backward Pass

c
void layernorm_backward(float* dinp, float* dweight, float* dbias,
                        float* dout, float* inp, float* weight,
                        float* mean, float* rstd,
                        int B, int T, int C) {
    for (int b = 0; b < B; b++) {
        for (int t = 0; t < T; t++) {
            float* dout_bt = dout + b * T * C + t * C;
            float* inp_bt = inp + b * T * C + t * C;
            float* dinp_bt = dinp + b * T * C + t * C;
            float mean_bt = mean[b * T + t];
            float rstd_bt = rstd[b * T + t];

            // 第一次 pass: 計算 reduce 值
            float dnorm_mean = 0.0f, dnorm_norm_mean = 0.0f;
            for (int i = 0; i < C; i++) {
                float norm_bti = (inp_bt[i] - mean_bt) * rstd_bt;
                float dnorm_i = weight[i] * dout_bt[i];
                dnorm_mean += dnorm_i;
                dnorm_norm_mean += dnorm_i * norm_bti;
            }
            dnorm_mean /= C;
            dnorm_norm_mean /= C;

            // 第二次 pass: 計算各項梯度
            for (int i = 0; i < C; i++) {
                float norm_bti = (inp_bt[i] - mean_bt) * rstd_bt;
                float dnorm_i = weight[i] * dout_bt[i];
                dbias[i] += dout_bt[i];
                dweight[i] += norm_bti * dout_bt[i];

                float dval = dnorm_i - dnorm_mean
                           - norm_bti * dnorm_norm_mean;
                dinp_bt[i] += dval * rstd_bt;
            }
        }
    }
}

CUDA LayerNorm Kernel

Kernel 3 (無 shared memory)

cuda
__global__ void layernorm_forward_kernel3(
    floatX* __restrict__ out, float* __restrict__ mean, float* __restrict__ rstd,
    const floatX*  __restrict__ inp, const floatX*  __restrict__ weight,
    const floatX* __restrict__ bias, int N, int C) {
    int lane_id = threadIdx.x % WARP_SIZE;
    int warp_id = threadIdx.x / WARP_SIZE;
    int idx = blockIdx.x * (blockDim.x / WARP_SIZE) + warp_id;
    if(idx >= N) return;

    const floatX* x = inp + idx * C;

    // warp reduce 計算 mean 和 rstd
    // 每個 thread 負責 C / WARP_SIZE 個元素
    float sum = 0.0f;
    for (int i = lane_id; i < C; i += WARP_SIZE)
        sum += (float)x[i];
    sum = warpReduceSum(sum);
    float m = sum / C;

    // ... 續 rstd 和歸一化 ...
}

Kernel 6 (有 shared memory)

較優化的版本,將 weight/bias 載入 shared memory,使用 x128 vectorized load/store,並用 __ldcs/__stcs streaming hints 減少快取汙染。

Transformer Block Forward Pass (train_gpt2.c)

c
for (int l = 0; l < L; l++) {
    // 前層的 residual 輸入
    residual = l == 0 ? acts.encoded : acts.residual3 + (l-1) * B * T * C;

    // Layer 1: LayerNorm → QKV → Attention → Projection → Residual
    layernorm_forward(l_ln1, l_ln1_mean, l_ln1_rstd, residual, l_ln1w, l_ln1b, B, T, C);
    matmul_forward(l_qkv, l_ln1, l_qkvw, l_qkvb, B, T, C, 3*C);
    attention_forward(l_atty, l_preatt, l_att, l_qkv, B, T, C, NH);
    matmul_forward(l_attproj, l_atty, l_attprojw, l_attprojb, B, T, C, C);
    residual_forward(l_residual2, residual, l_attproj, B*T*C);

    // Layer 2: LayerNorm → FC → GeLU → FC → Residual
    layernorm_forward(l_ln2, l_ln2_mean, l_ln2_rstd, l_residual2, l_ln2w, l_ln2b, B, T, C);
    matmul_forward(l_fch, l_ln2, l_fcw, l_fcb, B, T, C, 4*C);
    gelu_forward(l_fch_gelu, l_fch, B*T*4*C);
    matmul_forward(l_fcproj, l_fch_gelu, l_fcprojw, l_fcprojb, B, T, 4*C, C);
    residual_forward(l_residual3, l_residual2, l_fcproj, B*T*C);
}