Hierarchical Tiling
๊ฐ์
Hierarchical Tiling์ GEMM(General Matrix Multiply) ์ฐ์ฐ์์ ๋ฉ๋ชจ๋ฆฌ ๊ณ์ธต ๊ตฌ์กฐ์ ๋ง์ถฐ ํ๋ ฌ์ ์ฌ๋ฌ ๋จ๊ณ๋ก ๋ถํ ํ๋ ๊ธฐ๋ฒ์ด๋ค. ๋จ์ Square Tiling์ด ๋จ์ผ ์บ์ ๋ ๋ฒจ์ ๊ธฐ์ค์ผ๋ก ํ์ผ ํฌ๊ธฐ๋ฅผ ๊ฒฐ์ ํ๋ ๋ฐ๋ฉด, Hierarchical Tiling์ ๋ ์ง์คํฐ, Shared Memory, L1/L2 ์บ์, DRAM ๋ฑ ๊ฐ ๋ฉ๋ชจ๋ฆฌ ๊ณ์ธต์ ๋์ญํญยท๋ ์ดํด์ยท์ฉ๋ ํน์ฑ์ ์ต์ ํ๋ ํ์ผ ํฌ๊ธฐ๋ฅผ ์ค๊ณํ์ฌ ๋ฐ์ดํฐ ์ฌ์ฌ์ฉ์ ๊ทน๋ํํ๊ณ ๋ณ๋ ฌ์ฑ์ ์์ ํ ํ์ฉํ๋ค.
GPU์ ๋ฉ๋ชจ๋ฆฌ ๊ณ์ธต์ ์๋ฐฑ GB/s ๋์ญํญ์ Global Memory(HBM)์์ ์ TB/s์ Shared Memory, ๊ทธ๋ฆฌ๊ณ ์ KB์ ๋ ์ง์คํฐ ํ์ผ๊น์ง ์์ง์ ๊ตฌ์กฐ๋ฅผ ๊ฐ๋๋ค. Hierarchical Tiling์ ์ด ๊ฐ ๊ณ์ธต์์ ๋ฐ์ดํฐ๊ฐ ํ ๋ฒ ๋ก๋๋๋ฉด ํด๋น ๊ณ์ธต ๋ด์์ ์ต๋ํ ์ฌ์ฌ์ฉ๋๋๋ก ํ์ผ ํฌ๊ธฐ๋ฅผ ์กฐ์ ํจ์ผ๋ก์จ, ๊ณ์ธต ๊ฐ ๋ฐ์ดํฐ ์ด๋ ๋น์ฉ์ ์ต์ํํ๋ค. NVIDIA CUTLASS, cuBLAS ๋ฑ ๊ณ ์ฑ๋ฅ GEMM ๋ผ์ด๋ธ๋ฌ๋ฆฌ์ ํต์ฌ ์ต์ ํ ๊ธฐ๋ฒ์ผ๋ก, H100/A100 GPU์์ ์ด๋ก ์ ํผํฌ ์ฑ๋ฅ์ 90% ์ด์์ ๋ฌ์ฑํ๋ ๋ฐ ๊ธฐ์ฌํ๋ค.
ํต์ฌ ๊ฐ๋
๋ฉ๋ชจ๋ฆฌ ๊ณ์ธต๊ณผ ํ์ผ ๋ถํ
Hierarchical Tiling์ ํต์ฌ์ ๊ฐ ๋ฉ๋ชจ๋ฆฌ ๊ณ์ธต์ ๋ง๋ ํ์ผ ํฌ๊ธฐ๋ฅผ ์ค๊ณํ๋ ๊ฒ์ด๋ค:
| ๋ฉ๋ชจ๋ฆฌ ๊ณ์ธต | CUTLASS ์ฉ์ด | ํ์ผ ๋จ์ | ์ญํ |
|---|---|---|---|
| Global Memory (DRAM) | Threadblock Cluster | CtaTileM ร CtaTileN ร CtaTileK | SM ๊ฐ ๋ณ๋ ฌ์ฑ, ๋์ฉ๋ ๋ฐ์ดํฐ ๋ก๋ |
| Shared Memory (SMEM) | Threadblock/Warp | WarpTileM ร WarpTileN ร WarpTileK | SM ๋ด ๋ณ๋ ฌ์ฑ, ๋ฐ์ดํฐ ๊ณต์ |
| Registers (REG) | Thread/Atom | MmaM ร MmaN ร MmaK | ํ๋์จ์ด ๋ช ๋ น์ด ์์ค ์ฐ์ฐ |
๊ฐ ๊ณ์ธต์ ํ์ผ์ ์์ ๊ณ์ธต์ ํ์ผ์ ์๋ธ ๋๋ฐ์ด๋(subdivide)ํ์ฌ ์ ์๋๋ค. ์๋ฅผ ๋ค์ด, ํ๋์ Threadblock(CtaTile)์ ์ฌ๋ฌ ๊ฐ์ Warp Tile๋ก ๋ถํ ๋๊ณ , ๊ฐ Warp Tile์ ๋ค์ ์ฌ๋ฌ ๊ฐ์ MMA Atom์ผ๋ก ๋ถํ ๋๋ค.
Threadblock Tiling (Global โ Shared Memory)
Threadblock Tiling์ Global Memory์์ Shared Memory๋ก์ ํจ์จ์ ๋ฐ์ดํฐ ๋ก๋๋ฅผ ๋ด๋นํ๋ค:
C = A ร B (MรK) ร (KรN) โ (MรN)
for cta_n = 0 to N step CtaTileN:
for cta_m = 0 to M step CtaTileM:
for cta_k = 0 to K step CtaTileK:
// CTA(Threadblock)๊ฐ ๋ด๋นํ๋ MรN ํ์ผ
// CtaTileK ๋งํผ K์ถ์ ์ํํ๋ฉฐ ๋ฐ์ดํฐ ๋ก๋
Load A[cta_m:cta_m+CtaTileM, cta_k:cta_k+CtaTileK] โ SMEM
Load B[cta_k:cta_k+CtaTileK, cta_n:cta_n+CtaTileN] โ SMEM
// Shared Memory์์ MMA ์ฐ์ฐ
ํ์ผ ํ๋ผ๋ฏธํฐ (CUTLASS ๊ธฐ์ค):
| ํ๋ผ๋ฏธํฐ | ์๋ฏธ | ์์ (FP16 Tensor Core) |
|---|---|---|
| CtaTileM | Threadblock์ M ์ฐจ์ ํ์ผ ํฌ๊ธฐ | 128, 256 |
| CtaTileN | Threadblock์ N ์ฐจ์ ํ์ผ ํฌ๊ธฐ | 128, 256 |
| CtaTileK | K์ถ ํ์ผ ํฌ๊ธฐ (๋ฐ๋ณต ๋จ์) | 8, 32 |
์ฑ๋ฅ ํธ๋ ์ด๋์คํ:
- ํฐ ํ์ผ: Global Memory ๋ก๋ ํ์ ๊ฐ์ (์ข์), Register/Shared Memory ์๋น ์ฆ๊ฐ (๋์จ)
- ์์ ํ์ผ: ๋ก๋ ํ์ ์ฆ๊ฐ (๋์จ), ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋ ๊ฐ์ (์ข์)
Warp Tiling (Shared Memory โ Register)
Warp Tiling์ Shared Memory์์ Register๋ก์ ํจ์จ์ ๋ฐ์ดํฐ ๋ก๋์ Tensor Core ํ์ฉ์ ๋ด๋นํ๋ค:
// Threadblock ๋ด๋ถ์์ ๊ฐ ์ํ๊ฐ ๋ด๋นํ๋ ์๋ธํ์ผ
for warp_n = 0 to CtaTileN step WarpTileN:
for warp_m = 0 to CtaTileM step WarpTileM:
for warp_k = 0 to CtaTileK step WarpTileK:
// K์ถ ์ ์ฒด๋ฅผ unrollํ์ฌ ์ฐ์ฐ
Load SMEM โ Register (WarpTileM ร WarpTileK + WarpTileN ร WarpTileK)
MMA(Register A, Register B, Register C)
ํ์ผ ํ๋ผ๋ฏธํฐ:
| ํ๋ผ๋ฏธํฐ | ์๋ฏธ | ์์ |
|---|---|---|
| WarpTileM | ์ํ์ M ์ฐจ์ ํ์ผ ํฌ๊ธฐ | 64, 128 |
| WarpTileN | ์ํ์ N ์ฐจ์ ํ์ผ ํฌ๊ธฐ | 64, 128 |
| WarpTileK | ์ํ์ K ์ฐจ์ ํ์ผ ํฌ๊ธฐ | 32 |
์ค์ ๊ณ ๋ ค์ฌํญ:
- Shared Memory bank conflict ์์ด์ผ ํจ
- ์ํ ๋ด ๋ฐ์ดํฐ ์ค๋ณต ๋ก๋๋ฅผ ์ต์ํ
- Tensor Core ์ฌ์ฉ ์ MMA instruction ํฌ๊ธฐ์ ๋ง์ถฐ์ผ ํจ
Thread/Atom Tiling (Register ๋ด MMA)
Thread/Atom Tiling์ ์ค์ Tensor Core/CUDA Core MMA ์ฐ์ฐ์ ์ ์ถ๋ ฅ์ ๊ด๋ฆฌํ๋ค:
// Warp ๋ด๋ถ์์ MMA instruction ์์ค ์ฐ์ฐ
for mma_k = 0 to WarpTileK step MmaK:
for mma_n = 0 to WarpTileN step MmaN:
for mma_m = 0 to WarpTileM step MmaM:
mma_instruction(d, a, b, c) // Tensor Core ์ฐ์ฐ
ํ์ผ ํ๋ผ๋ฏธํฐ (Tensor Core ๊ธฐ์ค):
| ํ๋ผ๋ฏธํฐ | ์๋ฏธ | ์์ (FP16) |
|---|---|---|
| MmaM | MMA instruction์ M ์ฐจ์ | 16 |
| MmaN | MMA instruction์ N ์ฐจ์ | 8 |
| MmaK | MMA instruction์ K ์ฐจ์ | 16 |
๋ ์ง์คํฐ ์ฌ์ฉ๋ ๊ณ์ฐ:
๋ ์ง์คํฐ ์ โ (WarpTileM ร WarpTileK + WarpTileN ร WarpTileK + WarpTileM ร WarpTileN) ร bytes_per_element / 4
โ ๋๋ฌด ํฐ WarpTile์ ๋ ์ง์คํฐ ๋ถ์กฑ์ผ๋ก occupancy ๊ฐ์
CUTLASS์ 5๊ณ์ธต ๋ถํด ๊ตฌ์กฐ
CUTLASS 3.x๋ 5๊ณ์ธต API๋ก GEMM์ ๋ถํดํ๋ค:
| ๋ ๋ฒจ | CUTLASS ์ฉ์ด | ๋ณ๋ ฌ์ฑ | ๋ฉ๋ชจ๋ฆฌ ๊ณ์ธต |
|---|---|---|---|
| Device | GemmUniversalAdapter | ํธ์คํธ ์ธํฐํ์ด์ค | ์ปค๋ ๋ฐ์นญ |
| Kernel | GemmUniversal | SM ๊ฐ ๋ณ๋ ฌ์ฑ (Grid) | Global Memory |
| Collective | CollectiveMma | SM ๋ด ๋ณ๋ ฌ์ฑ (Threadblock Cluster) | Shared Memory |
| Tiled MMA/Copy | TiledMma | Warp ๋ด ๋ณ๋ ฌ์ฑ | Register |
| Atom | Mma_Atom, Copy_Atom | ์ค๋ ๋ ๋ด ILP | Register |
๋น๊ต/๋ถ์
ํ์ผ ํฌ๊ธฐ ํธ๋ ์ด๋์คํ
| ๊ณ ๋ ค์ฌํญ | ํฐ ํ์ผ | ์์ ํ์ผ |
|---|---|---|
| Global Memory ๋ก๋ | ์ ์ (โ) | ๋ง์ (โ) |
| Shared Memory ์ฌ์ฉ | ๋ง์ (โ) | ์ ์ (โ) |
| Register ์ฌ์ฉ | ๋ง์ (โ) | ์ ์ (โ) |
| Occupancy | ๋ฎ์ (โ) | ๋์ (โ) |
| ๋ฐ์ดํฐ ์ฌ์ฌ์ฉ | ๋์ (โ) | ๋ฎ์ (โ) |
| ๋ฉ๋ชจ๋ฆฌ ๋ ์ดํด์ ์๋ | ์ด๋ ค์ | ์ฌ์ |
| ๊ฒฝ๊ณ ํ์ผ ์ค๋ฒํค๋ | ํผ | ์์ |
ํ์ผ ํฌ๊ธฐ ๊ฒฐ์ ๊ณต์
Shared Memory ๊ธฐ์ค:
Stage ํฌ๊ธฐ = (CtaTileM ร CtaTileK + CtaTileN ร CtaTileK) ร element_size
์ด Shared Memory = Stage ํฌ๊ธฐ ร num_stages
โ SM๋น Shared Memory ํ๋ (96KB~228KB) ์ด๋ด์ด์ผ ํจ
Register ๊ธฐ์ค:
๋์ ๊ธฐ(accumulator) ๋ ์ง์คํฐ = CtaTileM ร CtaTileN / (warp ์) / 32
โ ์ด ๋ ์ง์คํฐ ์ฌ์ฉ๋ โค 255 (SM๋น)
โ occupancy = min(255/๋ ์ง์คํฐ์, SM๋น ์ต๋ ์ํ์)
์ผ๋ฐ์ ์ธ ํ์ผ ํฌ๊ธฐ ์์ (H100/A100)
| ํ๋ผ๋ฏธํฐ | FP16 Tensor Core | FP32 CUDA Core |
|---|---|---|
| CtaTileM ร N ร K | 256ร128ร32 | 128ร128ร8 |
| WarpTileM ร N ร K | 64ร64ร32 | 64ร64ร8 |
| MmaM ร N ร K | 16ร8ร16 (mma.sync) | 1ร1ร1 (simt) |
| Stages (Pipelining) | 3~5 | 2 |
| Warps per CTA | 4ร2ร1 = 8 | 4ร2ร1 = 8 |
๋์ ์๋ฆฌ
Pipelining (์ํํธ์จ์ด ํ์ดํ๋ผ์ด๋)
๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ์ง์ฐ์๊ฐ(memory latency)์ ์จ๊ธฐ๊ธฐ ์ํด ๋๋ธ ๋ฒํผ๋ง์ ์ฌ์ฉํ๋ค:
Stage 1: [Load SMEM_0 from Global] [MMA on SMEM_1] [Store results]
Stage 2: [Load SMEM_1 from Global] [MMA on SMEM_0] [Store results]
...
- Threadblock-scoped Shared Memory tiles: 2-Stage ์ด์์ ๋๋ธ ๋ฒํผ๋ง
- ํ ํ์ผ์ ํ์ฌ MMA ์ฐ์ฐ ์ฌ์ฉ
- ๋ค๋ฅธ ํ์ผ์ ๋ค์ ๋ฐ๋ณต์ ์ํ Global Memory ๋ก๋
- Warp-scoped Matrix Fragments: 2-Stage ๋ ์ง์คํฐ ๋๋ธ ๋ฒํผ๋ง
Threadblock Rasterization
L2 Cache ํ์ฉ๋๋ฅผ ๊ทน๋ํํ๊ธฐ ์ํด ์ฐ์์ ์ผ๋ก ๋ฐ์นญ๋๋ Threadblock์ด ์ธ์ ํ GEMM ์์ญ์ ๋ด๋นํ๋๋ก ๋งคํํ๋ค:
- ๊ฐ์ SMEM ํ์ผ์ ๋์์ ์ ๊ทผํ ํ๋ฅ ์ฆ๊ฐ โ L2 cache hit์จ ํฅ์
threadblock_swizzle.h์์ ๋ค์ํ ์ค์ผ์ค๋ง ํจ์ ์ ๊ณต
Parallelized Reductions
Split-K (Threadblock ๊ฐ ๋ณ๋ ฌ์ฑ):
- M, N์ด ์๊ณ K๊ฐ ํด ๋ ์ ํจ
- K ์ฐจ์์ ์ฌ๋ฌ Threadblock์ผ๋ก ๋ถํ โ ๋ณ๋ ฌ ์ฐ์ฐ ํ Reducer
- 2-์ปค๋ ํ์: Partitioned-K GEMM + Batched Reduction
Sliced-K (Warp ๊ฐ ๋ณ๋ ฌ์ฑ):
- Threadblock ๋ด์์ K ์ฐจ์์ ์ฌ๋ฌ warp๋ก ๋ถํ
- Split-K๋ณด๋ค ์ ์ ์ค๋ฒํค๋
- warp ๊ฐ ๋ถ๋ถ ํฉ(partial sum)์ ๋ง์ง๋ง์ reduction
Hopper Warp Specialization
Hopper ์ํคํ ์ฒ(SM90)์์ ๋์ ๋ Warp Specialization์ Producer-Consumer ๋ถ๋ฆฌ๋ก ๋น๋๊ธฐ ํ์ดํ๋ผ์ธ์ ๊ตฌํํ๋ค:
- Producer Warps: TMA(Tensor Memory Accelerator)๋ก Global โ Shared Memory ๋น๋๊ธฐ ๋ก๋
- Consumer Warps: Shared Memory์์ Tensor Core MMA ์ฐ์ฐ
- Persistent Kernel: Grid ํฌ๊ธฐ๋ฅผ SM ์๋งํผ๋ง ๋ฐ์นญํ์ฌ ์ปค๋ ๋ฐ์นญ ์ค๋ฒํค๋๋ฅผ amortize
- Tile Scheduler: ๋ฐ์นญ๋ Threadblock์ ๋์ ์ผ๋ก ํ์ผ ํ ๋น
์ฅ๋จ์
์ฅ์
- ๋ฉ๋ชจ๋ฆฌ ๋์ญํญ ์ต์ ํ: ๊ฐ ๊ณ์ธต์ ํน์ฑ์ ๋ง๋ ํ์ผ ํฌ๊ธฐ๋ก ๋ฐ์ดํฐ ์ด๋ ๋น์ฉ ์ต์ํ
- ๋ณ๋ ฌ์ฑ ๊ทน๋ํ: Grid/Threadblock/Warp ์์ค์ ๊ณ์ธต์ ๋ณ๋ ฌ์ฑ ํ์ฉ
- ** occupancy ์กฐ์ :** ํ์ผ ํฌ๊ธฐ ์กฐ์ ์ผ๋ก ๋ ์ง์คํฐ/Shared Memory ์ฌ์ฉ๋ ์ ์ด
- ํ๋์จ์ด ํ์ฉ: Tensor Core, TMA, Warp Specialization ๋ฑ ์ต์ ํ๋์จ์ด ๊ธฐ๋ฅ ํ์ฉ
- ๋ฒ์ฉ์ฑ: FP16, BF16, FP8, INT8 ๋ฑ ๋ค์ํ ์ ๋ฐ๋ ์ง์
๋จ์
- ์ค๊ณ ๋ณต์ก์ฑ: ๋ค์์ ํ์ผ ํ๋ผ๋ฏธํฐ๋ฅผ ํ๋ํด์ผ ํจ
- ํ๋์จ์ด ์์กด: ์ํคํ ์ฒ๋ณ ์ต์ ํ์ผ ํฌ๊ธฐ๊ฐ ๋ค๋ฆ
- ์ปค๋ ๋ฐ์นญ ์ค๋ฒํค๋: Persistent Kernel์ด ์๋ ๊ฒฝ์ฐ ๋นํจ์จ์
- ๊ฒฝ๊ณ ํ์ผ ์ฒ๋ฆฌ: ๋ฌธ์ ํฌ๊ธฐ๊ฐ ํ์ผ ํฌ๊ธฐ๋ก ๋๋์ด ๋จ์ด์ง์ง ์์ผ๋ฉด ๋ถํ์ํ ์ฐ์ฐ ๋ฐ์
๊ด๋ จ ๊ธฐ์
- CUTLASS: NVIDIA ๊ณ ์ฑ๋ฅ ์ ํ๋์ CUDA C++ ํ ํ๋ฆฟ ๋ผ์ด๋ธ๋ฌ๋ฆฌ (Volta~Blackwell ์ง์)
- cuBLAS: NVIDIA ๊ธฐ๋ณธ ์ ํ๋์ ๋ผ์ด๋ธ๋ฌ๋ฆฌ (๋ด๋ถ์ ์ผ๋ก CUTLASS ๊ธฐ๋ฐ ์ปค๋ ์ฌ์ฉ)
- cuBLASXt: Multi-GPU + CPU ํ์ด๋ธ๋ฆฌ๋ ํ์ผ๋ง
- cuBLASDx: On-Chip GEMM (Device-Side GEMM)
- Tensor Core: NVIDIA GPU์ ๋งคํธ๋ฆญ์ค ์ฐ์ฐ ๊ฐ์๊ธฐ (MMA instruction)
- TMA (Tensor Memory Accelerator): Hopper ์ํคํ ์ฒ์ ๋น๋๊ธฐ ๋ฉ๋ชจ๋ฆฌ ๋ก๋ ์์ง
- Software Pipelining: ๋ฉ๋ชจ๋ฆฌ ๋ ์ดํด์ ์จ๊ธฐ๊ธฐ ์ํ ๋๋ธ ๋ฒํผ๋ง ๊ธฐ๋ฒ
ํต์ฌ ์ ๋ฆฌ
Hierarchical Tiling์ ๋ฉ๋ชจ๋ฆฌ ๊ณ์ธต ๊ตฌ์กฐ์ ๋ง์ถฐ ํ๋ ฌ์ ๋ค๋จ๊ณ๋ก ๋ถํ ํ๋ GEMM ์ต์ ํ ๊ธฐ๋ฒ์ผ๋ก, Threadblock(Warp)/Thread/Atom ์์ค์์ ๊ฐ๊ฐ ์ต์ ์ ํ์ผ ํฌ๊ธฐ๋ฅผ ์ค๊ณํ์ฌ ๋ฐ์ดํฐ ์ฌ์ฌ์ฉ์ ๊ทน๋ํํ๋ค. NVIDIA CUTLASS์ 5๊ณ์ธต ๋ถํด ๊ตฌ์กฐ(Device/Kernel/Collective/TiledMMA/Atom)๋ ์ด ๊ธฐ๋ฒ์ ํ์ค์ ์ธ ๊ตฌํ ํ๋ ์์ํฌ๋ก, Hopper ์ํคํ ์ฒ์ Warp Specialization๊ณผ ๊ฒฐํฉ๋์ด ์ด๋ก ์ ํผํฌ ์ฑ๋ฅ์ 90% ์ด์์ ๋ฌ์ฑํ๋ค. ํ์ผ ํฌ๊ธฐ๋ Shared Memory ์ฉ๋, Register ์, occupancy ๊ฐ์ ํธ๋ ์ด๋์คํ๋ฅผ ํตํด ๊ฒฐ์ ๋๋ฉฐ, Pipelining๊ณผ Threadblock Rasterization๊ณผ ๊ฐ์ ๊ธฐ๋ฒ์ด ๊ฒฐํฉ๋์ด ์ต์ ์ ์ฑ๋ฅ์ ์ ๊ณตํ๋ค.