Skip to content

概念卡片

GPU 程式設計

CUDA Kernel

定義: 在 GPU 上執行的函式,以 __global__ 修飾,由 CPU 端透過 <<<grid, block>>> 啟動 例子: layernorm_forward_kernel3<<<grid_size, block_size, 0, stream>>>(out, mean, rstd, inp, weight, bias, N, C);

Grid / Block / Thread 層次

  • Grid: 所有 thread block 的集合
  • Block: 同一 block 的 thread 可透過 shared memory 溝通
  • Thread: 最小執行單元,每個 thread 處理一部分資料
  • llm.c 中典型配置: block_size = 256,grid size 由 CEIL_DIV(N, block_y) 計算

Warp

  • 32 個 thread 組成一個 warp,是 GPU 排程的最小單位
  • warpReduceSum: 在 warp 內對所有 thread 的值做歸納求和
  • Warp 內 thread 的 divergent branches 會導致序列化

Shared Memory

  • Block 內所有 thread 共享的 L1 快取
  • 大小有限(典型 48KB-164KB 取決於 GPU)
  • llm.c 的 layernorm_forward_kernel6 將 weight/bias 載入 shared memory 以減少全域記憶體存取

Memory Coalescing

  • 相鄰 thread 存取相鄰記憶體位置,合併為單一大量記憶體交易
  • 通道維度 C 設為最內層維度(stride=1)以最大化合併存取
  • 使用 x128 vectorized load/store(一次載入 128-bit)

數值格式

FP32 (32-bit float)

  • 1 sign + 8 exponent + 23 mantissa
  • train_gpt2.c 使用;CUDA 版本在 Master Weights 使用

BF16 (bfloat16)

  • 1 sign + 8 exponent + 7 mantissa
  • 與 FP32 等價的指數範圍,但精度較低
  • llm.c 預設精度,以 floatX 型別抽象

TF32

  • NVIDIA Ampere (SM 8.0+) 的加速模式
  • CUBLAS_COMPUTE_32F_FAST_TF32 啟用
  • 在矩陣乘法中自動使用,不需更改程式碼

優化技術

Kernel Fusion

  • 將多個連續 kernel 合併為一個,減少記憶體往返
  • fused_residual_forward5: 將 residual add + layernorm 合成一個 kernel
  • cuDNN Flash Attention: 將 QKV projection 後的 attention 計算 fusion

tiling

  • 將大矩陣分割成 tile 處理,利用 shared memory 快取
  • 減少對全域記憶體的讀取次數

Streaming Store (__stcs)

  • 提示編譯器此資料不會立即被重複使用
  • 繞過 L1 快取直接寫入 L2/DRAM,減少快取汙染

訓練技術

Gradient Accumulation

  • 將大 batch 拆成多個 micro-batch
  • 每個 micro-batch 計算梯度後 += 累加,最後一次更新權重
  • 用於總 batch size 無法一次載入 GPU 記憶體的情況

ZeRO Optimization

  • 將優化器狀態(m, v)、梯度、參數分散到各 GPU
  • Stage 1: 僅切分優化器狀態
  • Stage 2: + 梯度切分
  • Stage 3: + 參數切分

Checkpointing (Activation Recomputation)

  • Forward 時捨棄部分激活值,backward 時重新計算
  • --recompute 1: 重新計算 GeLU
  • --recompute 2: + 重新計算 LayerNorm
  • 以額外計算換取記憶體空間

AdamW

python
m = beta1 * m + (1-beta1) * grad
v = beta2 * v + (1-beta2) * grad^2
m_hat = m / (1 - beta1^t)
v_hat = v / (1 - beta2^t)
param -= lr * (m_hat / sqrt(v_hat) + eps) + weight_decay * param)