详解 Linear 层、Softmax、LayerNorm 的梯度核实现,配合 FP16/BF16 混合精度训练与完整训练循环,附 8GB+ VRAM 硬件要求。
欢迎回来!在第一部分中,我们从向量加法到平铺 GEMM 构建了基础,最终用 HIP(CUDA/ROCm)组装了一个完整的 Transformer 块前向传播。但是一个没有梯度的前向传播只是一个非常昂贵的随机数生成器。
训练 LLM 需要反向传播(backpropagation)、优化器(如 AdamW)以及残酷的内存管理,才能将数十亿参数塞进显存。在第二部分中,我们将实现缺失的部分:每一层的梯度计算、权重更新步骤、混合精度训练(FP16/BF16)以及一个功能完整的训练循环。
读完本部分后,你将理解:
前置要求:完成第一部分,熟练掌握链式法则,以及至少 8GB 显存的 GPU 以便本地跟进。
Linear 层的前向传播是 Y = X * W + b(为简单起见忽略偏置)。在反向传播中,我们收到 dY(损失关于输出的梯度),必须计算:
数学上:dX = dY * W^T;dW = X^T * dY
我们可以复用之前用过的完全相同的 hipblasSgemm(或 rocBLAS),只需转置矩阵即可。
// Assuming d_Y is [B, S, D] and W is [D, D]
// Compute d_X = d_Y * W^T (Matrix multiply)
float alpha = 1.0f, beta = 0.0f;
hipblasSgemm(handle, HIPBLAS_OP_N, HIPBLAS_OP_T,
S, D, D, // M, N, K
&alpha,
d_Y, S, // matrix A (dY)
W, D, // matrix B (W) transposed internally
&beta,
d_X, S); // matrix C (dX)
// Compute d_W = X^T * d_Y
hipblasSgemm(handle, HIPBLAS_OP_T, HIPBLAS_OP_N,
D, D, S,
&alpha,
X, S, // matrix A (X) transposed
d_Y, S, // matrix B (dY)
&beta,
d_W, D); // matrix C (dW)
对于多头注意力,由于 softmax 和 Q*K^T 乘法,梯度会更加复杂,但原理相同:前向传播中的每个矩阵乘法对应反向传播中的两个矩阵乘法。
虽然 GEMM 处理线性部分,但我们需要自定义内核来处理非线性部分。
设 P 为前向 softmax 的输出(概率)。反向传播根据 dY 计算 dX。公式为:dX_i = P_i * (dY_i - sum(P_j * dY_j))
__global__ void softmax_backward_kernel(const float* d_Y, const float* P,
float* d_X, int rows, int cols) {
int row = blockIdx.x * blockDim.x + threadIdx.x;
if (row >= rows) return;
// First, compute dot product of P and dY for this row
float dot = 0.0f;
for (int j = 0; j < cols; ++j) {
dot += P[row * cols + j] * d_Y[row * cols + j];
}
// Second, compute dX = P * (dY - dot)
for (int j = 0; j < cols; ++j) {
int idx = row * cols + j;
d_X[idx] = P[idx] * (d_Y[idx] - dot);
}
}
LayerNorm 需要计算关于输入 X 以及缩放/偏置参数 gamma 和 beta 的梯度。我们在这里不写完整内核以节省篇幅,但模式是:
专业提示:在前向传播中将均值和 inv_std(标准差的倒数)存储在一个小缓冲区中。这样可以避免在反向传播时重新计算它们。
在计算了 dW 和 db 之后,我们需要更新模型权重。AdamW 是 LLM 的标准优化器。它为每个参数维护两个指数移动平均:m(动量)和 v(方差)。
我们不 为每个参数更新启动单独的内核,而是编写一个融合内核,在一次传递中更新所有内容。这最大限度地减少了内核启动开销并最大化内存带宽。
__global__ void adamw_update_kernel(float* W, float* dW, float* m, float* v,
int num_params, float lr, float beta1,
float beta2, float eps, float weight_decay,
int step) {
int idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx >= num_params) return;
// Bias correction factors
float bias_correction1 = 1.0f - powf(beta1, step);
float bias_correction2 = 1.0f - powf(beta2, step);
// Update biased first moment estimate
m[idx] = beta1 * m[idx] + (1.0f - beta1) * dW[idx];
// Update biased second raw moment estimate
v[idx] = beta2 * v[idx] + (1.0f - beta2) * dW[idx] * dW[idx];
// Compute bias-corrected estimates
float m_hat = m[idx] / bias_correction1;
float v_hat = v[idx] / bias_correction2;
// Update weight (with weight decay)
W[idx] = W[idx] - lr * (m_hat / (sqrtf(v_hat) + eps) + weight_decay * W[idx]);
// Optionally reset gradient to zero for next iteration (or do it async)
dW[idx] = 0.0f;
}
在这个内核上启动 num_params 个线程(例如 10 亿个浮点数)。这是非常高效的,并利用了 GPU 的大规模并行能力。
现代 GPU(NVIDIA Ampere+ 和 AMD CDNA+)有专门的 FP16/BF16 矩阵乘法硬件,有效地将吞吐量翻倍。要用混合精度训练 LLM:
以下是我们如何修改前向传播:
// Use half type in HIP
#include <hip/hip_fp16.h>
__global__ void cast_and_scale_gradients(half* dW_half, float* dW_float,
float scale, int n) {
int idx = threadIdx.x + blockIdx.x * blockDim.x;
if (idx < n) dW_float[idx] = (float)dW_half[idx] * scale;
}
注意:rocBLAS/cuBLAS 支持 FP16 的 hipblasHgemm。使用 hipblasGemmEx 选择计算类型(FP32 用于累加)以获得最高精度。
LLM 训练中最大的瓶颈是内存。为反向传播存储每个激活值对于 7B+ 参数的模型是不可能的。
激活检查点(或梯度检查点)只保存特定层(例如每 4 个 Transformer 块)的输入,并在反向传播期间重新计算中间激活值。
实现策略:
这大致将内存占用减半,但增加了约 30-40% 的计算量。对于大型 LLM,这是必需的。
// During Forward
if (layer_idx % checkpoint_interval == 0) {
hipMemcpy(checkpoint_buffer + offset, d_X, size, hipMemcpyDeviceToDevice);
}
// During Backward
if (layer_idx % checkpoint_interval == 0) {
// 1. Load X from checkpoint
// 2. Run Forward pass of this layer (discard output, keep activations)
// 3. Run Backward pass using recomputed activations
}
现在,我们将所有内容整合在一起。在纯 HIP/C++ 中一次训练迭代如下:
void train_step() {
// 1. Copy batch from CPU to GPU (async)
hipMemcpyAsync(d_input, h_input, batch_bytes, hipMemcpyHostToDevice, stream);
// 2. Forward Pass
forward_transformer(d_input, d_output, ...); // Uses FP16 for matmuls
// 3. Compute Loss (Cross Entropy)
float loss = compute_loss_kernel(d_output, d_labels);
// 4. Loss Scaling
scale_loss_kernel<<<...>>>(d_loss_scaled, loss, loss_scale);
// 5. Backward Pass (traverses graph in reverse order)
backward_transformer(d_output, d_input, ...); // Computes FP16 gradients
// 6. Unscale Gradients (cast to FP32)
cast_and_scale_gradients<<<...>>>(dW_weights, dW_fp32, 1.0f/loss_scale, n);
// 7. Apply Gradient Clipping (optional, to prevent exploding gradients)
float norm = compute_l2_norm_kernel(dW_fp32, n);
if (norm > max_norm) scale_gradients(dW_fp32, max_norm / norm, n);
// 8. Optimizer Step (AdamW on FP32 master weights)
adamw_update_kernel<<<(n+255)/256, 256>>>(W_fp32, dW_fp32, m, v, ...);
// 9. Copy updated FP32 weights back to FP16 for next forward pass
cast_fp32_to_fp16_kernel<<<...>>>(W_fp16, W_fp32, n);
// 10. Synchronize Stream
hipStreamSynchronize(stream);
}
你无法优化你无法衡量的东西。使用这些工具:
关注这些指标:
如果你已经读到这里,恭喜!你刚刚从头构建了一个现代 LLM 训练器的架构骨架——涵盖了 CUDA/ROCm 上的 GEMM 优化、自定义反向内核、AdamW 优化器、混合精度以及节省内存的检查点。
显然,PyTorch 和 JAX 等框架透明地处理了所有这些,并添加了分布式训练(FSDP、ZeRO 和 all-reduce),这些我们还没有涉及。但理解这些底层原语让你成为 GPU 计算的大师。无论你是在 NVIDIA H100 还是 AMD MI300X 上运行,你现在都清楚地知道 loss.backward() 和 optimizer.step() 在底层做了什么。
第三部分预告?我们将深入研究多 GPU 分布式训练——使用 NCCL/RCCL 实现 All-Reduce、Ring-AllReduce 和分片策略(ZeRO 阶段)。
在那之前,祝你内核编码愉快,愿你的占用率高高、warp 分歧少少!
关于反向传播、损失缩放或检查点有问题吗?在下方评论区留言!我会阅读每一条并详细回复。
下一部分见!