Register/Shared Memory Tiling
๊ฐ์
Register/Shared Memory Tiling์ GPU์์ GEMM(Generic Matrix Multiply) ์ฐ์ฐ์ ๊ทนํ ์ฑ๋ฅ์ผ๋ก ์ํํ๊ธฐ ์ํ ๋ฉ๋ชจ๋ฆฌ ๊ณ์ธต ๊ธฐ๋ฐ ํ์ผ๋ง ๊ธฐ๋ฒ์ด๋ค. Square Tiling์ด ์บ์ ์์ค์์ ํ์ผ์ ๋ถํ ํ๋ ๊ฐ๋ ์ด๋ผ๋ฉด, Register/Shared Memory Tiling์ ๊ธ๋ก๋ฒ ๋ฉ๋ชจ๋ฆฌ โ ๊ณต์ ๋ฉ๋ชจ๋ฆฌ(SMEM) โ ๋ ์ง์คํฐ๊น์ง ์ธ ๊ณ์ธต์ ๊ฑธ์ณ ๋ฐ์ดํฐ ํ๋ฆ์ ์ฒด๊ณ์ ์ผ๋ก ๊ด๋ฆฌํ์ฌ ์ฐ์ฐ ์ฌ์ฌ์ฉ(Operand Reuse)์ ๊ทน๋ํํ๋ค.
์ด ๊ธฐ๋ฒ์ ํต์ฌ์ ๊ฐ ์ค๋ ๋๊ฐ ๋จ์ผ ์ถ๋ ฅ ์์๊ฐ ์๋ TMรTN ํฌ๊ธฐ์ ์์ ๋ถ๋ถ ํ๋ ฌ(sub-matrix)์ ๊ณ์ฐํ๋ "์์ํฐ ํ๋ก๋ํธ(Outer Product)" ํจํด์ด๋ค. ๊ณต์ ๋ฉ๋ชจ๋ฆฌ์์ ํ์ผ์ ๋ก๋ํ ํ, ๊ฐ ์ค๋ ๋๊ฐ ์์ ์ ๋ ์ง์คํฐ์ A์ ํ ์กฐ๊ฐ๊ณผ B์ ์ด ์กฐ๊ฐ์ ์ ์ฅํ๊ณ , TMรTN ๋์ฐ๊ธฐ์์ ์ธ์ ์ฐ์ฐ์ ์ํํ๋ค. ์ด๋ฅผ ํตํด ์ค๋ ๋๋น ์ฐ์ฐ ๋ฐ๋(Arithmetic Intensity)๋ฅผ ๊ทน์ ์ผ๋ก ๋์ด๊ณ , ๋ฉ๋ชจ๋ฆฌ ๋์ญํญ ๋ณ๋ชฉ์ ํด์ํ๋ค.
cuBLAS, CUTLASS, Triton ๋ฑ ๋ชจ๋ ๊ณ ์ฑ๋ฅ GEMM ๋ผ์ด๋ธ๋ฌ๋ฆฌ๊ฐ ์ด ํ์ผ๋ง ์ ๋ต์ ๊ธฐ๋ณธ์ผ๋ก ์ฑํํ๊ณ ์์ผ๋ฉฐ, NVIDIA A100/H100์์ FP16 ํ ์ ์ฝ์ด ๊ธฐ์ค 90% ์ด์์ ํผํฌ ์ฑ๋ฅ์ ๋ฌ์ฑํ ์ ์๊ฒ ํ๋ ํต์ฌ ์ต์ ํ ๊ธฐ๋ฒ์ด๋ค.
ํต์ฌ ๊ฐ๋
๋ฉ๋ชจ๋ฆฌ ๊ณ์ธต๊ณผ ํ์ผ๋ง
GPU ๋ฉ๋ชจ๋ฆฌ ๊ณ์ธต์ ์ธ ๊ฐ์ง ์ฃผ์ ์์ค์ผ๋ก ๊ตฌ์ฑ๋๋ค:
| ๊ณ์ธต | ํฌ๊ธฐ (A100 ๊ธฐ์ค) | ๋์ญํญ | ์ ๊ทผ ์ง์ฐ | ํ์ผ๋ง ์ญํ |
|---|---|---|---|---|
| ๊ธ๋ก๋ฒ ๋ฉ๋ชจ๋ฆฌ (HBM) | 80 GB | 2 TB/s | ์๋ฐฑ ์ฌ์ดํด | ์ ์ฒด ํ๋ ฌ ์ ์ฅ, ํ์ผ ๋ก๋ ์์ค |
| ๊ณต์ ๋ฉ๋ชจ๋ฆฌ (SMEM) | 192 KB/SM | ~19 TB/s (์จ์นฉ) | ~20 ์ฌ์ดํด | ๋ธ๋ก ๊ฐ ๋ฐ์ดํฐ ๊ณต์ , ํ์ผ ์บ์ |
| ๋ ์ง์คํฐ ํ์ผ | 256 KB/SM | ๋ฌด์ ํ | 0 ์ฌ์ดํด | ์ค๋ ๋๋ณ ์ฐ์ฐ ๋ฐ์ดํฐ ์ ์ฅ |
๊ฐ ๊ณ์ธต๋ณ ํ์ผ ํฌ๊ธฐ ํ๋ผ๋ฏธํฐ:
- BM, BN: ๊ณต์ ๋ฉ๋ชจ๋ฆฌ์ ๋ก๋ํ A์ B์ ํ์ผ ํฌ๊ธฐ (๋์ด ร ๋๋น)
- BK: K ์ฐจ์ ๋ฐฉํฅ์ผ๋ก ํ ๋ฒ์ ์ฒ๋ฆฌํ ๊น์ด
- TM, TN: ๊ฐ ์ค๋ ๋๊ฐ ๋ ์ง์คํฐ์์ ๊ณ์ฐํ๋ ์ถ๋ ฅ ์๋ธ ๋งคํธ๋ฆญ์ค ํฌ๊ธฐ
์์ํฐ ํ๋ก๋ํธ GEMM
๊ธฐ์กด Square Tiling์ด ํ ์ค๋ ๋์ ๋ํด ํ ๋ฒ์ ํ๋์ ์์๋ฅผ ๊ณ์ฐํ๋ ๋ฐ๋ฉด, Register Tiling์ ์์ํฐ ํ๋ก๋ํธ ๋ฐฉ์์ผ๋ก ์๋ํ๋ค:
for k = 0 to K step BK:
A_tile[BMรBK]๋ฅผ ๊ณต์ ๋ฉ๋ชจ๋ฆฌ๋ก ๋ก๋
B_tile[BKรBN]๋ฅผ ๊ณต์ ๋ฉ๋ชจ๋ฆฌ๋ก ๋ก๋
for t = 0 to BK:
a = A_tile[:, t] โ BM๊ฐ ์์๋ฅผ ๋ ์ง์คํฐ๋ก ๋ก๋
b = B_tile[t, :] โ BN๊ฐ ์์๋ฅผ ๋ ์ง์ค๋ฒ๋ก ๋ก๋
C_reg[i][j] += a[i] ร b[j] (์ธ์ ์ฐ์ฐ)
๊ฐ k ๋ฐ๋ณต์์ BM๊ฐ์ A ์์์ BN๊ฐ์ B ์์๋ฅผ ๋ก๋ํ์ฌ TMรTN ํฌ๊ธฐ์ ๋์ฐ๊ธฐ๋ฅผ ์ ๋ฐ์ดํธํ๋ค. ์ด๋ฅผ ํตํด ์ฐ์ฐ ๋ ๋ฉ๋ชจ๋ฆฌ ๋น์จ(Compute-to-Memory Ratio)์ด BMรBN/BK๋ก ์ฆ๊ฐํ๋ค.
๊ณต์ ๋ฉ๋ชจ๋ฆฌ ํ์ผ ๋ก๋ฉ
๋ธ๋ก์ ๋ชจ๋ ์ค๋ ๋๊ฐ ํ๋ ฅ์ ์ผ๋ก ๊ธ๋ก๋ฒ ๋ฉ๋ชจ๋ฆฌ์์ ๊ณต์ ๋ฉ๋ชจ๋ฆฌ๋ก ํ์ผ์ ๋ก๋ํ๋ค:
// ํ์ผ ๋ก๋ฉ ์ฝ๋๋ค์ด์
int thread_id = threadIdx.x + threadIdx.y * blockDim.x;
int total_threads = blockDim.x * blockDim.y;
// A[BMรBK] ํ์ผ์ total_threads๋ก ๊ท ๋ฑ ๋ถํ ๋ก๋ฉ
for (int i = thread_id; i < BM * BK; i += total_threads) {
int row = i / BK;
int col = i % BK;
smem_A[row][col] = global_A[block_row * BM + row][k_tile * BK + col];
}
// B[BKรBN] ํ์ผ๋ ๋์ผํ๊ฒ ๋ก๋ฉ
for (int i = thread_id; i < BK * BN; i += total_threads) {
int row = i / BN;
int col = i % BN;
smem_B[row][col] = global_B[k_tile * BK + row][block_col * BN + col];
}
__syncthreads();
๋๋ธ ๋ฒํผ๋ง
๊ณต์ ๋ฉ๋ชจ๋ฆฌ ํ์ผ ๋ก๋ฉ๊ณผ ์ฐ์ฐ์ ๊ฒน์ณ ์คํํ์ฌ ์ง์ฐ ์๊ฐ์ ์จ๊ธฐ๋ ๊ธฐ๋ฒ์ด๋ค:
- ๋ฒํผ 0: ํ์ฌ ํ์ผ ์ฐ์ฐ ์ํ
- ๋ฒํผ 1: ๋ค์ ํ์ผ์ ๊ธ๋ก๋ฒ ๋ฉ๋ชจ๋ฆฌ์์ ๋ก๋ฉ
- ๋ค์ ๋ฐ๋ณต์์ ๋ฒํผ ๊ต์ฒด
์ด๋ฅผ ํตํด ๋ฉ๋ชจ๋ฆฌ ๋ก๋ฉ ์ง์ฐ์ ์ฐ์ฐ์ผ๋ก ์จ๊ธธ ์ ์๋ค.
๋น๊ต/๋ถ์
ํ์ผ๋ง ์ ๋ต ๋น๊ต
| ์ ๋ต | ๋ฐ์ดํฐ ์ฌ์ฌ์ฉ ์์น | ์ค๋ ๋๋น ์ถ๋ ฅ | ๋ฉ๋ชจ๋ฆฌ ์ ์ฝ | ์ฑ๋ฅ ํน์ฑ |
|---|---|---|---|---|
| Square Tiling | L1/L2 ์บ์ | 1๊ฐ | ์บ์ ํฌ๊ธฐ | ์ค๊ฐ ์ฑ๋ฅ |
| Shared Memory Tiling | ๊ณต์ ๋ฉ๋ชจ๋ฆฌ | 1๊ฐ | SMEM ํฌ๊ธฐ | ๋์ ์ฑ๋ฅ |
| Register Tiling | ๋ ์ง์คํฐ | TMรTN๊ฐ | ๋ ์ง์คํฐ ์ | ์ต๊ณ ์ฑ๋ฅ |
| Register + SMEM Tiling | ๊ณต์ ๋ฉ๋ชจ๋ฆฌ + ๋ ์ง์คํฐ | TMรTN๊ฐ | SMEM + ๋ ์ง์คํฐ | ์ต์ ์ฑ๋ฅ |
๋ฐ์ดํฐํ๋ก์ฐ ํจํด ๋น๊ต
| ํจํด | ์ค๋ช | ๋ ์ง์คํฐ ์ฌ์ฉ | SMEM ์ฌ์ฉ | ํน์ง |
|---|---|---|---|---|
| Output Stationary | ๊ฒฐ๊ณผ ๋์ฐ, ์ ๋ ฅ ์คํธ๋ฆฌ๋ฐ | TMรTN ๋์ฐ๊ธฐ | ์ต์ | ๋ฉ๋ชจ๋ฆฌ ์ต์ํ |
| Input Stationary | ์ ๋ ฅ ํ์ผ ๋ก๋, ์ถ๋ ฅ ์ฌ์ฌ์ฉ | TM+TN ์ ๋ ฅ ์กฐ๊ฐ | BMรBK + BKรBN | ์ ๋ ฅ ์ฌ์ฌ์ฉ ๊ทน๋ํ |
| Weight Stationary | ๊ฐ์ค์น ํ์ผ ๊ณ ์ | TM+TN ๊ฐ์ค์น ์กฐ๊ฐ | BMรBK | ์ถ๋ก ์ ์ ํฉ |
| Outer Product | ็ปๅ ธ GEMM ํ์ผ๋ง | TM+TN + TMรTN | BMรBK + BKรBN | ๋ฒ์ฉ ์ต์ ํ |
ํ๋ผ๋ฏธํฐ ํ๋ ์ํฅ
| ํ๋ผ๋ฏธํฐ | ์ฆ๊ฐ ์ ํจ๊ณผ | ๊ฐ์ ์ ํจ๊ณผ | ์ต์ ๊ฐ ์์ (A100 FP16) |
|---|---|---|---|
| BM | A ๋ฐ์ดํฐ ์ฌ์ฌ์ฉ ์ฆ๊ฐ | SMEM ๋ถ์กฑ, ์ค๋ ๋ ์ ๊ฐ์ | 128 |
| BN | B ๋ฐ์ดํฐ ์ฌ์ฌ์ฉ ์ฆ๊ฐ | SMEM ๋ถ์กฑ, ์ค๋ ๋ ์ ๊ฐ์ | 128 |
| BK | ๋ฐ๋ณต๋น ์ฐ์ฐ๋ ์ฆ๊ฐ | ์ฌ์ฌ์ฉ ๊ฐ์ | 16~32 |
| TM | ์ค๋ ๋๋น ์ถ๋ ฅ ์ฆ๊ฐ | ๋ ์ง์คํฐ ๋ถ์กฑ | 8 |
| TN | ์ค๋ ๋๋น ์ถ๋ ฅ ์ฆ๊ฐ | ๋ ์ง์คํฐ ๋ถ์กฑ | 8 |
์ ์ ์จ(Occupancy) ๋ถ์
๋ ์ง์คํฐ/์ค๋ ๋ = TM + TN + TMรTN (์
๋ ฅ ์กฐ๊ฐ + ๋์ฐ๊ธฐ)
๋ธ๋ก๋น ์ค๋ ๋ = (BM/TM) ร (BN/TN)
SMEM/๋ธ๋ก = BMรBKร4 + BKรBNร4 (FP32 ๊ธฐ์ค)
์์ (BM=128, BN=128, BK=16, TM=8, TN=8):
- ๋ ์ง์คํฐ/์ค๋ ๋: 8 + 8 + 64 = 80๊ฐ
- ๋ธ๋ก๋น ์ค๋ ๋: 16 ร 16 = 256๊ฐ
- SMEM/๋ธ๋ก: 128ร16ร4 + 16ร128ร4 = 16KB
- SM๋น ๋ธ๋ก ์: ์ ํ์ (๋ ์ง์คํฐ 256KB ๊ธฐ์ค ์ต๋ 40๋ธ๋ก)
๋์ ์๋ฆฌ
์ ์ฒด ์๊ณ ๋ฆฌ์ฆ ๊ตฌ์กฐ
// BM=128, BN=128, BK=16, TM=8, TN=8 ๊ธฐ์ค GEMM ์ปค๋
__global__ void gemm_tiled(float* A, float* B, float* C, int M, int N, int K) {
__shared__ float smem_A[BM][BK]; // 128ร16 = 2KB
__shared__ float smem_B[BK][BN]; // 16ร128 = 2KB
// ๊ฐ ์ค๋ ๋์ TMรTN ์ถ๋ ฅ ํ์ผ
float C_reg[TM][TN] = {0.0f};
float A_reg[TM];
float B_reg[TN];
int bx = blockIdx.x;
int by = blockIdx.y;
// 1. ํ์ผ ๋ฃจํ: K ์ฐจ์์ BK ๋จ์๋ก ์ํ
for (int k_tile = 0; k_tile < K; k_tile += BK) {
// 2. ํ๋ ฅ์ ํ์ผ ๋ก๋ฉ (๊ธ๋ก๋ฒ โ ๊ณต์ ๋ฉ๋ชจ๋ฆฌ)
load_tile_to_smem(A, smem_A, by, k_tile);
load_tile_to_smem(B, smem_B, k_tile, bx);
__syncthreads();
// 3. ๋ด๋ถ ๋ฃจํ: BK ๊น์ด๋งํผ ์ธ์ ์ฐ์ฐ
for (int k = 0; k < BK; k++) {
// A ์กฐ๊ฐ ๋ก๋: smem_A์์ TM๊ฐ ์์
for (int i = 0; i < TM; i++) {
A_reg[i] = smem_A[threadRow * TM + i][k];
}
// B ์กฐ๊ฐ ๋ก๋: smem_B์์ TN๊ฐ ์์
for (int j = 0; j < TN; j++) {
B_reg[j] = smem_B[k][threadCol * TN + j];
}
// ์ธ์ ์ฐ์ฐ: TMรTN ๋์ฐ๊ธฐ ์
๋ฐ์ดํธ
for (int i = 0; i < TM; i++) {
for (int j = 0; j < TN; j++) {
C_reg[i][j] += A_reg[i] * B_reg[j];
}
}
}
__syncthreads();
}
// 4. ๊ฒฐ๊ณผ ์ ์ฅ (๋ ์ง์คํฐ โ ๊ธ๋ก๋ฒ ๋ฉ๋ชจ๋ฆฌ)
for (int i = 0; i < TM; i++) {
for (int j = 0; j < TN; j++) {
C[(by * BM + threadRow * TM + i) * N + bx * BN + threadCol * TN + j]
= C_reg[i][j];
}
}
}
๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ํจํด ๋ถ์
๋ธ๋ก๋น ๊ธ๋ก๋ฒ ๋ฉ๋ชจ๋ฆฌ ๋ก๋ฉ๋:
- A: BMรK = 128รK ์์
- B: KรBN = Kร128 ์์
- ์ด: 256K ์์
๋ธ๋ก๋น ์ฐ์ฐ๋:
- BMรBNรK = 128ร128รK = 16384K FLOP
์ฐ์ฐ ๋ ๋ฉ๋ชจ๋ฆฌ ๋น์จ:
FLOP/Byte = 16384K / (256K ร 4) = 16 FLOP/Byte
Square Tiling ๋๋น 16๋ฐฐ ํฅ์๋ ์ฐ์ฐ ๋ฐ๋๋ฅผ ๋ฌ์ฑํ๋ค.
๋ ์ง์คํฐ ๋ ๋ฒจ ์ธ์ ์ฐ์ฐ ์์ธ
k=0: A_reg[0..7] ร B_reg[0..7] โ C_reg[0..7][0..7] (64 FLOP)
k=1: A_reg[0..7] ร B_reg[0..7] โ C_reg[0..7][0..7] (64 FLOP)
...
k=15: A_reg[0..7] ร B_reg[0..7] โ C_reg[0..7][0..7] (64 FLOP)
์ด: 16 ร 64 = 1024 FLOP per thread
๊ฐ ์ค๋ ๋๊ฐ BK=16๋ฒ์ ์ธ์ ์ฐ์ฐ์ผ๋ก TMรTN=64๊ฐ ์ถ๋ ฅ ์์๋ฅผ ๋์ฐํ๋ค.
์ปจํฐ์ด ๋ฒํผ๋ง ๊ตฌํ
__shared__ float smem_A[2][BM][BK]; // ๋ ๊ฐ์ A ํ์ผ ๋ฒํผ
__shared__ float smem_B[2][BK][BN]; // ๋ ๊ฐ์ B ํ์ผ ๋ฒํผ
// ์ฒซ ๋ฒ์งธ ํ์ผ ๋ฏธ๋ฆฌ ๋ก๋ฉ
load_tile_to_smem(A, smem_A[0], by, 0);
load_tile_to_smem(B, smem_B[0], 0, bx);
__syncthreads();
for (int k_tile = 0; k_tile < K; k_tile += BK) {
int curr = (k_tile / BK) % 2;
int next = 1 - curr;
// ํ์ฌ ๋ฒํผ์์ ์ฐ์ฐ
compute_outer_product(smem_A[curr], smem_B[curr], C_reg);
// ๋ค์ ํ์ผ์ ๋ค๋ฅธ ๋ฒํผ๋ก ๋ฏธ๋ฆฌ ๋ก๋ฉ
if (k_tile + BK < K) {
load_tile_to_smem(A, smem_A[next], by, k_tile + BK);
load_tile_to_smem(B, smem_B[next], k_tile + BK, bx);
}
__syncthreads();
}
์ฅ๋จ์
์ฅ์
- ๊ทน๋ํ๋ ์ฐ์ฐ ๋ฐ๋: TMรTN ํฌ๊ธฐ์ ์ธ์ ์ฐ์ฐ์ ํตํด ์ค๋ ๋๋น FLOP์ ๋ํญ ์ฆ๊ฐ์์ผ ๋ฉ๋ชจ๋ฆฌ ๋์ญํญ ๋ณ๋ชฉ์ ์ต์ํํ๋ค
- ๊ณ์ธต์ ๋ฐ์ดํฐ ์ฌ์ฌ์ฉ: ๊ธ๋ก๋ฒ โ ๊ณต์ ๋ฉ๋ชจ๋ฆฌ โ ๋ ์ง์คํฐ ์ธ ๊ณ์ธต์์ ๊ฐ๊ฐ BMรBN/BK, BM/BK, 1๋ฐฐ์ ๋ฐ์ดํฐ ์ฌ์ฌ์ฉ์ ๋ฌ์ฑํ๋ค
- ํ ์ ์ฝ์ด์์ ์์ฐ์ค๋ฌ์ด ๊ฒฐํฉ: NVIDIA Tensor Core์ WMMA/MMA ๋ช ๋ น๊ณผ ํ์ผ ๊ตฌ์กฐ๊ฐ ํธํ๋์ด ํผํฉ ์ ๋ฐ๋ ์ฐ์ฐ์์ ์ถ๊ฐ์ ์ธ ์ฑ๋ฅ ํฅ์์ ์ ๊ณตํ๋ค
- ์ ์ฐํ ํ๋ผ๋ฏธํฐ ์กฐ์ : BM/BN/BK/TM/TN์ ํ๋์จ์ด ํน์ฑ์ ๋ง๊ฒ ์กฐ์ ํ์ฌ ๋ค์ํ ์ํคํ ์ฒ์ ์ต์ ํํ ์ ์๋ค
- ๋์ ์ค์ฉ์ฑ: cuBLAS, CUTLASS, Triton ๋ฑ ๊ฒ์ฆ๋ ๊ตฌํ์ฒด๊ฐ ์กด์ฌํ์ฌ ํ๋ก๋์ ํ๊ฒฝ์์ ์ฆ์ ํ์ฉ ๊ฐ๋ฅํ๋ค
๋จ์
- ๊ตฌํ ๋ณต์ก๋: ๋ฉ๋ชจ๋ฆฌ ๊ณ์ธต ๊ฐ ๋๊ธฐํ, ๋ ์ง์คํฐ ํ ๋น, ํ์ผ ๋ก๋ฉ ์์ ๋ฑ ๊ณ ๋ คํ ์์๊ฐ ๋ง์ ๋๋ฒ๊น ์ด ์ด๋ ต๋ค
- ๋ ์ง์คํฐ ์๋ฐ(Register Pressure): TMรTN ๋์ฐ๊ธฐ๊ฐ ๋ง์ ๋ ์ง์คํฐ๋ฅผ ์๋ชจํ์ฌ, ์ ์ ์จ(Occupancy)์ด ๋ฎ์์ง ์ ์๋ค. ์ค๋ ๋ ์ ๊ฐ์๋ ๋ณ๋ ฌ์ฑ์ ์ ํ์ํจ๋ค
- ๊ณต์ ๋ฉ๋ชจ๋ฆฌ ์ ์ฝ: SMEM ํฌ๊ธฐ๊ฐ ์ ํ์ ์ด๋ฏ๋ก ํ์ผ ํฌ๊ธฐ๋ฅผ ๊ณผ๋ํ๊ฒ ํค์ธ ์ ์์ผ๋ฉฐ, ๋ฑ ํฌ ์ถฉ๋(Bank Conflict) ํํผ๋ฅผ ์ํ ํจ๋ฉ์ด ํ์ํ๋ค
- ๋น์ ํ ํ๋ ฌ ์ฒ๋ฆฌ: ํ๋ ฌ ์ฐจ์์ด ํ์ผ ํฌ๊ธฐ๋ก ๋๋์ด๋จ์ด์ง์ง ์์ผ๋ฉด ํจ๋ฉ์ด ๋ฐ์ํ์ฌ ์ฐ์ฐ ๋ญ๋น์ ๋ณต์ก๋ ์ฆ๊ฐ๋ฅผ ์ด๋ํ๋ค
- ํ๋์จ์ด ์์กด์ฑ: ์ต์ ํ์ผ ํฌ๊ธฐ๋ GPU ์ํคํ ์ฒ(Ampere, Hopper ๋ฑ)์ ๋ฐ๋ผ ๋ฌ๋ผ์ง๋ฏ๋ก ์ด์์ฑ์ด ๋จ์ด์ง๋ค
๊ด๋ จ ๊ธฐ์
ํ๋ ์์ํฌ ๋ฐ ๋ผ์ด๋ธ๋ฌ๋ฆฌ
- NVIDIA cuBLAS: ์ํคํ ์ฒ๋ณ ์ต์ ํ๋ GEMM ์ปค๋ ์งํฉ์ผ๋ก, Register/Shared Memory Tiling์ ์๋ ์ ํํ์ฌ ์ฌ์ฉ
- NVIDIA CUTLASS: ํ ํ๋ฆฟ ๊ธฐ๋ฐ GEMM ๋ผ์ด๋ธ๋ฌ๋ฆฌ๋ก, BM/BN/BK/TM/TN ํ๋ผ๋ฏธํฐ๋ฅผ ์ปดํ์ผ ์์ ์ ์ค์ ํ์ฌ ๋ค์ํ ํ์ผ๋ง ์ ๋ต ๊ตฌํ
- OpenAI Triton: ๊ณ ์์ค ์ธ์ด๋ก ํ์ผ๋ง์ ๊ธฐ์ ํ๋ฉด ์ปดํ์ผ๋ฌ๊ฐ ์๋์ผ๋ก SMEM/๋ ์ง์คํฐ ํ ๋น ๋ฐ ๋๊ธฐํ๋ฅผ ์์ฑ
- FlashAttention: ํ์ผ๋ง์ ์ดํ ์ ์ฐ์ฐ์ ์ ์ฉํ์ฌ ๋ฉ๋ชจ๋ฆฌ O(Nยฒ) โ O(N)์ผ๋ก ๊ฐ์
- XGemm/CK (Composable Kernel): AMD GPU์ฉ ํ์ผ๋ง GEMM ๋ผ์ด๋ธ๋ฌ๋ฆฌ
๊ด๋ จ ๊ฐ๋
- Bank Conflict: ๊ณต์ ๋ฉ๋ชจ๋ฆฌ์ 32๊ฐ ๋ฑ ํฌ์ ๋ํ ๋์ ์ ๊ทผ ์ถฉ๋. ํจ๋ฉ(+4) ๋๋ ์ค์์ฆ(Swizzle)๋ก ํด๊ฒฐ
- Warp-level GEMM: ์ํ ๋ด 32๊ฐ ์ค๋ ๋๊ฐ ํ๋ ฅ์ ์ผ๋ก ํฐ ํ์ผ์ ์ฒ๋ฆฌํ๋ ๊ธฐ๋ฒ
- Tensor Core: NVIDIA GPU์ ๋งคํธ๋ฆญ์ค ๊ณฑ์ ์ ์ฉ ์ ๋. WMMA/MMA ์ธํฐํ์ด์ค๋ก ์ ๊ทผ
- Pipeline Prefetching: ๋ค์ ํ์ผ์ ๋น๋๊ธฐ๋ก ๋ก๋ฉํ์ฌ ํ์ฌ ์ฐ์ฐ๊ณผ ๊ฒน์น๊ฒ ํ๋ ๊ธฐ๋ฒ
- Occupancy: SM๋น ํ์ฑ ์ค๋ ๋ ์. ๋ ์ง์คํฐ/SMEM ์ฌ์ฉ๋์ ๋ฐ๋น๋ก
์ฐธ๊ณ ๋ฌธํ
- "Anatomy of High-Performance Many-Threaded Matrix Multiplication" - Soliman & Rodriguez (2019)
- "CUTLASS: CUDA Templates for Linear Algebra Subroutines" - NVIDIA
- "Triton: An Intermediate Language and Compiler for Tiled Neural Network Computations" - Tillet et al. (2019)
- "Optimizing Sgemm on NVIDIA GPUs" - NVIDIA Developer Blog
- "FlashAttention: Fast and Memory-Efficient Exact Attention" - Dao et al. (2022)
- "Efficient Memory Management for Large Language Model Serving with PagedAttention" - Kwon et al. (2023)
ํต์ฌ ์ ๋ฆฌ
- Register/Shared Memory Tiling์ ๊ธ๋ก๋ฒ ๋ฉ๋ชจ๋ฆฌ โ ๊ณต์ ๋ฉ๋ชจ๋ฆฌ โ ๋ ์ง์คํฐ ์ธ ๊ณ์ธต์์ ํ์ผ์ ์ฒด๊ณ์ ์ผ๋ก ๊ด๋ฆฌํ๋ GPU GEMM ์ต์ ํ ๊ธฐ๋ฒ์ด๋ค
- ๊ฐ ์ค๋ ๋๊ฐ TMรTN ํฌ๊ธฐ์ ์ถ๋ ฅ ์๋ธ ๋งคํธ๋ฆญ์ค๋ฅผ ์ธ์ (Outer Product) ๋ฐฉ์์ผ๋ก ๊ณ์ฐํ์ฌ, ์ค๋ ๋๋น ์ฐ์ฐ ๋ฐ๋๋ฅผ ๊ทน๋ํํ๋ค
- ๊ณต์ ๋ฉ๋ชจ๋ฆฌ ํ์ผ์ ๋ธ๋ก ๋ด ๋ชจ๋ ์ค๋ ๋๊ฐ ํ๋ ฅ์ ์ผ๋ก ๋ก๋ฉํ์ฌ ๋ฐ์ดํฐ ์ฌ์ฌ์ฉ์ ๋ฌ์ฑํ๊ณ , ๋ ์ง์คํฐ ํ์ผ์ ์ค๋ ๋ ๋ด๋ถ์์ ์ฐ์ฐ์ ์ํํ๋ค
- ๋๋ธ ๋ฒํผ๋ง์ ํตํด ํ์ผ ๋ก๋ฉ๊ณผ ์ฐ์ฐ์ ๊ฒน์ณ ์คํํ์ฌ ๋ฉ๋ชจ๋ฆฌ ์ง์ฐ ์๊ฐ์ ์จ๊ธธ ์ ์๋ค
- cuBLAS, CUTLASS, Triton ๋ฑ ๋ชจ๋ ๊ณ ์ฑ๋ฅ GEMM ๋ผ์ด๋ธ๋ฌ๋ฆฌ๊ฐ ์ด ์ ๋ต์ ๊ธฐ๋ณธ์ผ๋ก ์ฑํํ๋ฉฐ, NVIDIA ํ ์ ์ฝ์ด์ ๊ฒฐํฉํ์ฌ ํผํฌ ์ฑ๋ฅ์ 90% ์ด์์ ๋ฌ์ฑํ ์ ์๋ค