在 H100/MI300X 上实现极致性能调优:IO-aware Flash Attention、Transformer Engine FP8 量化、多算子融合。
欢迎来到最终章!在第一部分,我们学会了走路(GEMM)。在第二部分,我们学会了奔跑(反向传播与 AdamW)。在第三部分,我们学会了在数千块 GPU 上飞行(分布式训练)。但如果你还留在这里,说明你并不满足于「能跑」的代码。你追求的是极致。你要从 H100 或 MI300X 中榨干每一滴 FLOPS。
在这一第四也就是最后一部分,我们不再把 GPU 当作通用处理器,而是把它当作一台内存受限的机器来对待。我们将实现以下内容:
Flash Attention —— 让 100k+ 上下文窗口成为可能的 IO 感知内核。
FP8 量化(Transformer Engine)—— 使用 8 位浮点获得 2 倍加速。
内核融合 —— 将 LayerNorm、残差连接和 Dropout 融合进单次传递。
分布式 Checkpointing —— 保存大规模分片模型而不撑爆文件系统。
到本部分结束时,你的自定义训练循环将媲美 PyTorch 2.0 + DeepSpeed 的性能。你将成为真正的 GPU 巫师。
前置要求:掌握前三个部分的内容、拥有计算能力 8.9+ 的 GPU(Ada Lovelace/Hopper)或 AMD CDNA 3(MI300)以原生运行 FP8,以及对极致性能的渴望。
标准 attention 先计算 S = Q * K^T(将 S 存入 HBM),再读取 S 计算 softmax,再读取 softmax 与 V 相乘。这意味着要多次往返慢速全局内存(HBM)。Flash Attention 通过使用分块(Tiling)和在线 softmax 数学将整个 attention 过程融合进单个内核来解决这个问题。
核心技巧:在线 Softmax
不再存储完整的注意力矩阵 S(形状 [seq_len, seq_len]),而是用适合放入共享内存(SRAM)的小块来处理 Q、K、V。我们动态重新计算 softmax 归一化。
以下是一个简化的 Flash Attention 内核骨架(单头,省略了 causal mask):
#define TILE_SIZE 32 // Blocks of 32x32 in shared memory
__global__ void flash_attention_kernel(const float* Q, const float* K, const float* V,
float* O, int N, int d_head) {
// Shared memory tiles
__shared__ float Q_tile[TILE_SIZE][d_head];
__shared__ float K_tile[TILE_SIZE][d_head];
__shared__ float V_tile[TILE_SIZE][d_head];
int tx = threadIdx.x, ty = threadIdx.y;
int batch_idx = blockIdx.z; // assuming batch is in z-dim
int q_start = blockIdx.y * TILE_SIZE;
int kv_start = blockIdx.x * TILE_SIZE;
// Load Q tile into shared memory (per block)
for (int i = ty; i < d_head; i += blockDim.y) {
Q_tile[tx][i] = Q[((batch_idx * N) + q_start + tx) * d_head + i];
}
__syncthreads();
// Local accumulators for attention output
float out_acc[d_head] = {0.0f};
float l = 0.0f; // normalization sum
float m = -INFINITY; // running maximum
// Loop over all KV blocks
for (int kv_block = 0; kv_block < N / TILE_SIZE; ++kv_block) {
// Load K and V tiles into shared memory (load from global)
load_kv_tile(K, V, kv_block, ...);
__syncthreads();
// Compute Q * K^T for this tile (results in registers)
float scores[TILE_SIZE] = {0.0f};
for (int i = 0; i < TILE_SIZE; ++i) {
float sum = 0.0f;
for (int k = 0; k < d_head; ++k) {
sum += Q_tile[ty][k] * K_tile[i][k];
}
scores[i] = sum * rsqrtf((float)d_head);
}
// Online softmax update (for this tile)
for (int i = 0; i < TILE_SIZE; ++i) {
float score = scores[i];
float new_m = fmaxf(m, score);
float exp_diff = expf(m - new_m);
float exp_score = expf(score - new_m);
// Correct previous accumulated values
for (int k = 0; k < d_head; ++k) out_acc[k] *= exp_diff;
l = l * exp_diff + exp_score;
m = new_m;
// Accumulate weighted V
for (int k = 0; k < d_head; ++k) {
out_acc[k] += exp_score * V_tile[i][k];
}
}
__syncthreads();
}
// Final normalization and write to global O
for (int k = 0; k < d_head; ++k) {
int idx = ((batch_idx * N) + q_start + ty) * d_head + k;
O[idx] = out_acc[k] / l;
}
}
注意:这比第一部分中的朴素 attention 要快得多,因为我们从不把 S 写入 HBM。在实践中你会使用 Flash Attention 2 或 3(加入了 warp 级并行),但这个内核教你的是exact原理。
FP8(8 位浮点)相比 FP16 将带宽和计算吞吐量提高了一倍。NVIDIA 的 Transformer Engine 和 AMD 的 FP8 支持使用每张量(per-tensor)缩放因子来避免溢出。
我们定义一个自定义半精度类型(或在 AMD 上使用 __hip_fp8,在 CUDA 上使用 __nv_fp8)。以下是一个将张量向下转换到 FP8 的转换内核:
// Simplified FP8 cast with dynamic scale
__global__ void cast_fp32_to_fp8_kernel(const float* input, __nv_fp8_e4m3* output,
float scale, int n) {
int idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx >= n) return;
float val = input[idx] * scale; // scale up to use FP8 range
// Saturate to FP8 max range (roughly ~448 for E4M3)
val = fminf(fmaxf(val, -448.0f), 448.0f);
output[idx] = __nv_fp8_e4m3(val);
}
在训练期间,我们维护一个最大值历史记录来动态调整每层的缩放因子。对于 MatMul,在主要 GEMM 之前将 Q、K、V 转换为 FP8,并在 FP32 中累积结果(使用 hipblasGemmEx 和计算类型 HIPBLAS_COMPUTE_32F)。
在标准 transformer block 中,我们做:
X = Attention(X) + X(残差相加)→ 写入 HBM。
X = LayerNorm(X) → 从 HBM 读取,写入 HBM。
X = Dropout(FFN(X)) + X → 读写 HBM。
这些都是内存受限操作。我们可以将它们融合进一个内核,只读取一次和写入一次。以下是残差 + LayerNorm 路径的融合内核:
__global__ void fused_residual_layernorm_kernel(float* X, const float* Attn_Out,
float* gamma, float* beta,
int rows, int cols) {
extern __shared__ float sdata[]; // shared memory for reduction
int row = blockIdx.x;
int tid = threadIdx.x;
float* x_row = X + row * cols;
float* attn_row = Attn_Out + row * cols;
// 1. Compute residual add and mean/variance in registers
float sum = 0.0f, sq_sum = 0.0f;
for (int i = tid; i < cols; i += blockDim.x) {
float val = x_row[i] + attn_row[i]; // residual connection
x_row[i] = val; // store temporarily (we will normalize in-place)
sum += val;
sq_sum += val * val;
}
// Shared memory reduction (omitted for brevity, use warp shuffle)
float mean = sum / cols;
float variance = sq_sum / cols - mean * mean;
float inv_std = rsqrtf(variance + 1e-5f);
// 2. Apply LayerNorm in-place and write back
for (int i = tid; i < cols; i += blockDim.x) {
float normalized = (x_row[i] - mean) * inv_std;
x_row[i] = normalized * gamma[i] + beta[i];
}
}
这个内核完成了 3 次独立 CUDA/HIP 调用的工作,节省了两次完整的全局内存读写。这对于模型中非 GEMM 部分大约是 30% 的加速!
当在第三部分使用 ZeRO-3 训练时,每块 GPU 只持有权重的一部分。如果你天真地把所有权重收集到 rank 0 来保存一个文件,你很可能会耗尽内存并制造一个巨大的 I/O 瓶颈。
相反,我们实现并行 Checkpointing。每块 GPU 同时将其本地分片保存到文件系统(例如 model_shard_rank_0.bin、model_shard_rank_1.bin)。
void save_checkpoint_sharded(void* local_weights, size_t local_size, int rank, const char* base_dir) {
char filename[256];
snprintf(filename, sizeof(filename), "%s/checkpoint_step_%d_rank_%d.bin", base_dir, step, rank);
// Use buffered, asynchronous file writes to not stall the GPU
// 1. Copy from device to pinned host memory (async)
float* host_buffer;
hipHostMalloc(&host_buffer, local_size, hipHostMallocDefault);
hipMemcpyAsync(host_buffer, local_weights, local_size, hipMemcpyDeviceToHost, stream);
hipStreamSynchronize(stream);
// 2. Write to disk using standard fwrite or POSIX (on a separate thread ideally)
FILE* fp = fopen(filename, "wb");
fwrite(host_buffer, 1, local_size, fp);
fclose(fp);
// 3. Also save a metadata file (json) containing the world_size, shapes, and dtypes
hipHostFree(host_buffer);
}
加载时,我们只需反向操作:每个 rank 加载自己的 .bin 文件。这随 GPU 数量线性扩展,完全消除了「rank 0 瓶颈」。
慢速 CPU 数据加载器会让你的巨型 GPU 集群饥饿。我们使用双缓冲(Double Buffering)和固定内存(Pinned Memory):
// Initialize two buffers
float* h_buffers[2];
float* d_buffers[2];
hipHostMalloc(&h_buffers[0], batch_size, hipHostMallocWriteCombined);
hipHostMalloc(&h_buffers[1], batch_size, hipHostMallocWriteCombined);
hipMalloc(&d_buffers[0], batch_size);
hipMalloc(&d_buffers[1], batch_size);
int current = 0;
for (int step = 0; step < total_steps; ++step) {
// While GPU processes buffer 'current', CPU loads buffer '1-current'
std::thread loader(load_next_batch, h_buffers[1-current]);
hipMemcpyAsync(d_buffers[current], h_buffers[current], batch_size,
hipMemcpyHostToDevice, compute_stream);
// ... Run Forward/Backward using d_buffers[current] ...
loader.join();
current = 1 - current;
}
这确保了你的 GPU 计算内核永不等待数据。这是生产系统中的标准做法,但在自定义 C++ 训练器中经常被忽视。
就这样。我们从第一部分的 vec_add 起步,现在已经构建了一个生产级、多 GPU、FP8 混合精度的训练器,具有融合内核和全异步操作。你已经实现了 Flash Attention 逻辑,通过 Ring-AllReduce 在节点间扩展梯度,用 ZeRO 分片内存,以及用内核融合优化内存传输。
编写这样一个自定义训练器是一项艰巨的任务——这就是 PyTorch 等框架存在的原因。但下次你看到 torch.compile 警告、DeepSpeed 配置文件或 Flash Attention 导入时,你不会再看到黑魔法。你会看到我们刚刚一起构建的exact C++/HIP 逻辑。
你已经正式从「GPU 用户」毕业成为「GPU 架构师」。
最后给你一个挑战:尝试将 FP8 转换与 Flash Attention 内核结合起来。它会坏掉,会令人沮丧,但当你修复它时,你将拥有一个比 99% 开源实现更快的内核。
如果你想要第五部分(虽然我以为这是终点!),我们可以探索图编译(静态计算图)或 CPU 卸载(当你的 VRAM 用尽时)。请在评论区告诉我!
从我硅心的最深处感谢你加入这段旅程。保持你的内核分块整齐,你的 wavefront 饱满,你的 HBM 带宽饱和。
下次见,在 LLM 的另一侧相见!🚀