โšก AI Optimization

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 ๋ฉ”๋ชจ๋ฆฌ ๊ณ„์ธต ๊ตฌ์กฐ

๋ฉ”๋ชจ๋ฆฌ ๊ณ„์ธต๊ณผ ํƒ€์ผ ๋ถ„ํ• 

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 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๊ณผ ๊ฐ™์€ ๊ธฐ๋ฒ•์ด ๊ฒฐํ•ฉ๋˜์–ด ์ตœ์ ์˜ ์„ฑ๋Šฅ์„ ์ œ๊ณตํ•œ๋‹ค.