Square Tiling ๊ธฐ์ด
๊ฐ์
Square Tiling์ ํ๋ ฌ ๊ณฑ์ (Matrix Multiplication) ์ฐ์ฐ์ ํจ์จ์ ์ผ๋ก ์ํํ๊ธฐ ์ํ ๊ธฐ๋ณธ์ ์ธ ์ต์ ํ ๊ธฐ๋ฒ์ด๋ค. ํฐ ํ๋ ฌ์ ์์ ์ ์ฌ๊ฐํ ํ์ผ(Tile)๋ก ๋ถํ ํ์ฌ ์บ์ ๋ฉ๋ชจ๋ฆฌ์ ๋ ์ง์คํฐ ํ์ผ์ locality๋ฅผ ๊ทน๋ํํ๊ณ , ์ฐ์ฐ ์ฌ์ฌ์ฉ(Operand Reuse)์ ํตํด ๋ฉ๋ชจ๋ฆฌ ๋์ญํญ ์ฌ์ฉ๋์ ์ค์ธ๋ค. ์ด ๊ธฐ๋ฒ์ GEMM(Generic Matrix Multiply) ์ปค๋์ ์ฑ๋ฅ์ ๊ฒฐ์ ํ๋ ํต์ฌ ์์๋ก, GPU, CPU, AI ๊ฐ์๊ธฐ ๋ฑ ๋ค์ํ ํ๋์จ์ด ํ๋ซํผ์์ ๊ด๋ฒ์ํ๊ฒ ์ฌ์ฉ๋๋ค.
Square Tiling์ ํต์ฌ ์์ด๋์ด๋ "์์ ํ์ผ์ ํ ๋ฒ์ ์ฒ๋ฆฌํ๋ฉด ์บ์์ ์ ์ฌ๋ ๋ฐ์ดํฐ๋ฅผ ์ฌ๋ฌ ๋ฒ ์ฌ์ฌ์ฉํ ์ ์๋ค"๋ ๊ฒ์ด๋ค. ์๋ฅผ ๋ค์ด, 1024ร1024 ํ๋ ฌ ๊ณฑ์ ์์ ๋จ์ ๋ฃจํ๋ O(Nยณ)์ ์บ์ ๋ฏธ์ค๋ฅผ ๋ฐ์์ํค์ง๋ง, 32ร32 ํ์ผ๋ก ๋ถํ ํ๋ฉด ์บ์ ๋ฏธ์ค๊ฐ O(Nยณ/Tยฒ)๋ก ๊ฐ์ํ๋ค.
ํต์ฌ ๊ฐ๋
ํ์ผ๋ง(Tiling)์ ์๋ฆฌ
ํ๋ ฌ ๊ณฑ์ C = A ร B์ ์ ์์ ๋ฐ๋ฅด๋ฉด, ๊ฒฐ๊ณผ ํ๋ ฌ์ ๊ฐ ์์ C[i][j]๋ ๋ค์๊ณผ ๊ฐ์ด ๊ณ์ฐ๋๋ค:
C[i][j] = ฮฃ (A[i][k] ร B[k][j]) (k = 0 ~ N-1)
๋จ์ ๊ตฌํ์์๋ A์ i๋ฒ์งธ ํ๊ณผ B์ j๋ฒ์งธ ์ด์ ๋งค๋ฒ ๋ฉ๋ชจ๋ฆฌ์์ ๋ก๋ํ์ง๋ง, Square Tiling์ ์ ์ฉํ๋ฉด:
- ํ๋ ฌ ๋ถํ : ํ๋ ฌ A, B๋ฅผ TรT ํฌ๊ธฐ์ ์ ์ฌ๊ฐํ ํ์ผ๋ก ๋ถํ
- ํ์ผ ๋จ์ ์ฐ์ฐ: ๊ฐ ํ์ผ ์(A ํ์ผ, B ํ์ผ)์ ๋ํด ๋ถ๋ถ ๊ณฑ์ ์ํ
- ๋์ฐ: ๋ถ๋ถ ๊ณฑ์ ๊ฒฐ๊ณผ๋ฅผ C ํ์ผ์ ๋์ฐ
- ๋ฐ๋ณต: ๋ชจ๋ ํ์ผ ์์ ๋ํด ์ ๊ณผ์ ์ ๋ฐ๋ณต
์บ์ ๋ ๋ฒจ๊ณผ ํ์ผ ํฌ๊ธฐ
ํ์ผ ํฌ๊ธฐ๋ ํ๋์จ์ด์ ๋ฉ๋ชจ๋ฆฌ ๊ณ์ธต ๊ตฌ์กฐ์ ๋ฐ๋ผ ๊ฒฐ์ ๋๋ค:
| ๋ฉ๋ชจ๋ฆฌ ๊ณ์ธต | ํฌ๊ธฐ | ๋ํ ์์ | ํ์ผ ํฌ๊ธฐ ์ํฅ |
|---|---|---|---|
| ๋ ์ง์คํฐ ํ์ผ | ์ KB | NVIDIA SM๋น 256KB | L1 ํ์ผ ํฌ๊ธฐ ์ ํ |
| L1 ์บ์ | 16~128 KB | A100: 192KB/SM | ์ต์ ํ์ผ ํฌ๊ธฐ ๊ฒฐ์ |
| L2 ์บ์ | ์ MB | A100: 40MB | ์ค๊ฐ ํ์ผ ํฌ๊ธฐ |
| ๋ฉ๋ชจ๋ฆฌ(DRAM) | ์ GB | A100: 80GB HBM2e | ์ ์ฒด ํ๋ ฌ ํฌ๊ธฐ |
์ต์ ํ์ผ ํฌ๊ธฐ ๊ณต์:
T = โ(์บ์ ํฌ๊ธฐ / (3 ร ์์ ํฌ๊ธฐ ร 2))
- ์์ ํฌ๊ธฐ: FP16์ด๋ฉด 2๋ฐ์ดํธ, FP32์ด๋ฉด 4๋ฐ์ดํธ
- 3: A, B, C ์ธ ํ๋ ฌ์ ํ์ผ์ด ๋์์ ์บ์์ ์์ด์ผ ํจ
- 2: A ํ์ผ๊ณผ B ํ์ผ ๋ ๊ฐ๊ฐ ํ์
๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ํจํด
Square Tiling์ ๋ฉ๋ชจ๋ฆฌ ์ ๊ทผ ํจํด์ ํฌ๊ฒ ๊ฐ์ ํ๋ค:
ํ์ผ๋ง ์ (๋จ์ ๋ฃจํ):
- A ํ: O(N)๋ฒ ๋ก๋
- B ์ด: O(N)๋ฒ ๋ก๋
- ์ด ์บ์ ๋ฏธ์ค: O(Nยณ)
ํ์ผ๋ง ํ:
- A ํ์ผ: O(T)๋ฒ ๋ก๋ โ O(N/T)๋ฒ ์ฌ์ฌ์ฉ
- B ํ์ผ: O(T)๋ฒ ๋ก๋ โ O(N/T)๋ฒ ์ฌ์ฌ์ฉ
- ์ด ์บ์ ๋ฏธ์ค: O(Nยณ/Tยฒ)
Cache-Oblivious ์๊ณ ๋ฆฌ์ฆ
Square Tiling์ ํ๋์จ์ด ์บ์ ํฌ๊ธฐ๋ฅผ ๋ช ์์ ์ผ๋ก ๊ณ ๋ คํ๋ Cache-Aware ์ ๊ทผ๋ฒ์ด๋ค. ๋ฐ๋ฉด, Cache-Oblivious ์๊ณ ๋ฆฌ์ฆ์ ์บ์ ํฌ๊ธฐ ํ๋ผ๋ฏธํฐ ์์ด ์ฌ๊ท์ ์ผ๋ก ํ๋ ฌ์ ๋ถํ ํ์ฌ ์ต์ ์ ์บ์ ์ ๊ทผ ํจํด์ ๋ฌ์ฑํ๋ค:
- ์ฌ๊ท์ ๋ถํ : ํ๋ ฌ์ 2ร2 ๋ธ๋ก์ผ๋ก ์ฌ๊ท์ ์ผ๋ก ๋ถํ
- ํ๋ผ๋ฏธํฐ ๋ถํ์: ์บ์ ํฌ๊ธฐ๋ ๋ธ๋ก ํฌ๊ธฐ๋ฅผ ์ง์ ํ ํ์ ์์
- ๋์ ์ ์: ๋ฉํฐํ๋ก๊ทธ๋๋ฐ ํ๊ฒฝ์์ ์บ์ ํฌ๊ธฐ๊ฐ ๋ณํด๋ ์ต์ ์ฑ๋ฅ ์ ์ง
- ์บ์ ๋ฏธ์ค ๋ณต์ก๋: ฮ(Nยณ/(bโM)) (b: ์บ์ ๋ผ์ธ ํฌ๊ธฐ, M: ์บ์ ํฌ๊ธฐ)
Cache-Oblivious GEMM์ ์ด๋ก ์ ์ผ๋ก ์ฐ์ํ์ง๋ง, ์ค์ ๊ตฌํ์์๋ ์บ์ ๋ผ์ธ ๋จ์์ ๋ณ๋ ฌํ์ SIMD ๋ฒกํฐํ๊ฐ ์ด๋ ค์ ์ค์ฉ์ฑ์์๋ Cache-Aware tiling์ด ๋ ๋๋ฆฌ ์ฌ์ฉ๋๋ค.
๋น๊ต/๋ถ์
ํ์ผ๋ง ๊ธฐ๋ฒ ๋น๊ต
| ๊ธฐ๋ฒ | ํ์ผ ํํ | ์ฅ์ | ๋จ์ | ํ์ฉ |
|---|---|---|---|---|
| Square Tiling | ์ ์ฌ๊ฐํ (TรT) | ๊ตฌํ ์ฉ์ด, ๊ท ํ ์กํ locality | ๋ชจ๋ ์ฐจ์์์ ๋์ผํ locality | ๊ธฐ๋ณธ GEMM ์ปค๋ |
| Rectangular Tiling | ์ง์ฌ๊ฐํ (TrรTc) | ํน์ ์ฐจ์ locality ์ต์ ํ | ํ๋ผ๋ฏธํฐ ํ๋ ๋ณต์ก | ํนํ๋ ์ปค๋ |
| Triangular Tiling | ์ผ๊ฐํ | ๋์นญ ํ๋ ฌ ์ต์ ํ | ์ผ๋ฐ์ ์ฌ์ฉ ์ด๋ ค์ | ํน์ ์ฐ์ฐ |
| Register Tiling | ๋ ์ง์คํฐ ์์ค ๋ฏธ์ธ ํ์ผ | ๊ทนํ locality | ๊ตฌํ ๋ณต์ก๋ ๋์ | ๊ณ ์ฑ๋ฅ ์ปค๋ |
ํ๋์จ์ด๋ณ ํ์ผ ํฌ๊ธฐ ์์
| ํ๋์จ์ด | L1 ์บ์ | ๊ถ์ฅ ํ์ผ ํฌ๊ธฐ | ํน์ง |
|---|---|---|---|
| NVIDIA A100 | 192KB/SM | 32ร32 (FP16) | ํ ์ ์ฝ์ด์ ํธํ |
| NVIDIA H100 | 228KB/SM | 32ร32~64ร64 | Transformer ์์ง ๊ณ ๋ ค |
| AMD MI250X | 64KB/WGP | 16ร16~32ร32 | Wavefront ํฌ๊ธฐ ๊ณ ๋ ค |
| CPU (x86) | 32~64KB | 8ร8~16ร16 | SIMD ๋ ์ง์คํฐ ์ ๊ณ ๋ ค |
์ฐ์ฐ ์ฌ์ฌ์ฉ ๋ถ์
ํ์ผ ํฌ๊ธฐ T์ผ ๋ ์ฐ์ฐ ์ฌ์ฌ์ฉ๋ฅ :
์ฌ์ฌ์ฉ๋ฅ = N / T
- T = 1 (ํ์ผ๋ง ์์): ์ฌ์ฌ์ฉ๋ฅ = N (์ต์ )
- T = N (์ ์ฒด๋ฅผ ํ๋์ ํ์ผ): ์ฌ์ฌ์ฉ๋ฅ = 1 (์ด๋ก ์ ์ต์ , ๋ฉ๋ชจ๋ฆฌ ๋ถ์กฑ)
- T = โN: ์ฌ์ฌ์ฉ๋ฅ = โN (๊ท ํ์ )
๋์ ์๋ฆฌ
๊ธฐ๋ณธ GEMM ์๊ณ ๋ฆฌ์ฆ (ํ์ผ๋ง ์์)
# ๋จ์ GEMM - O(Nยณ) ์บ์ ๋ฏธ์ค
for i in range(N):
for j in range(N):
C[i][j] = 0
for k in range(N):
C[i][j] += A[i][k] * B[k][j] # ๋งค๋ฒ ๋ฉ๋ชจ๋ฆฌ์์ ๋ก๋
Square Tiling ์ ์ฉ ์๊ณ ๋ฆฌ์ฆ
# Square Tiling GEMM - O(Nยณ/Tยฒ) ์บ์ ๋ฏธ์ค
T = 32 # ํ์ผ ํฌ๊ธฐ
for i in range(0, N, T):
for j in range(0, N, T):
C_tile = C[i:i+T, j:j+T] # C ํ์ผ ๋ก๋ (TรT)
for k in range(0, N, T):
A_tile = A[i:i+T, k:k+T] # A ํ์ผ ๋ก๋ (TรT)
B_tile = B[k:k+T, j:j+T] # B ํ์ผ ๋ก๋ (TรT)
# ํ์ผ ๊ณฑ์
(์ด์ ์บ์์์ ์ฐ์ฐ)
for ii in range(T):
for jj in range(T):
for kk in range(T):
C_tile[ii][jj] += A_tile[ii][kk] * B_tile[kk][jj]
C[i:i+T, j:j+T] = C_tile # ๊ฒฐ๊ณผ ์ ์ฅ
GPU ์ปค๋ ๊ตฌํ
GPU์์๋ ์ํ(Warp) ๋จ์๋ก ํ์ผ์ ์ฒ๋ฆฌํ๋ค:
// GPU GEMM ์ปค๋ (๋จ์ํ)
__global__ void gemm_tiled(float* A, float* B, float* C, int N) {
__shared__ float As[TILE][TILE]; // ๊ณต์ ๋ฉ๋ชจ๋ฆฌ A ํ์ผ
__shared__ float Bs[TILE][TILE]; // ๊ณต์ ๋ฉ๋ชจ๋ฆฌ B ํ์ผ
int bx = blockIdx.x, by = blockIdx.y;
int tx = threadIdx.x, ty = threadIdx.y;
int row = by * TILE + ty;
int col = bx * TILE + tx;
float sum = 0.0f;
// ํ์ผ ๋ฃจํ
for (int t = 0; t < N; t += TILE) {
// ํ์ผ์ ๊ณต์ ๋ฉ๋ชจ๋ฆฌ๋ก ๋ก๋
As[ty][tx] = A[row * N + (t + tx)];
Bs[ty][tx] = B[(t + ty) * N + col];
__syncthreads();
// ํ์ผ ๊ณฑ์
for (int k = 0; k < TILE; k++) {
sum += As[ty][k] * Bs[k][tx];
}
__syncthreads();
}
C[row * N + col] = sum;
}
์ฅ๋จ์
์ฅ์
- ์บ์ ํจ์จ์ฑ ๊ทน๋ํ: ํ์ผ ํฌ๊ธฐ๋งํผ์ ๋ฐ์ดํฐ๋ฅผ ์บ์์ ์ ์ฌํ ํ ์ฌ๋ฌ ๋ฒ ์ฌ์ฌ์ฉํ์ฌ ์บ์ ๋ฏธ์ค๋ฅผ ํฌ๊ฒ ๊ฐ์์ํจ๋ค
- ๊ตฌํ ์ฉ์ด์ฑ: ๋ค๋ฅธ ํ์ผ๋ง ๊ธฐ๋ฒ(์: Register Tiling)์ ๋นํด ์๊ณ ๋ฆฌ์ฆ์ด ๋จ์ํ๊ณ ์ง๊ด์ ์ด๋ค
- ํ๋์จ์ด ๋ ๋ฆฝ์ฑ: CPU, GPU, FPGA ๋ฑ ๋ค์ํ ํ๋ซํผ์์ ๋์ผํ ์๋ฆฌ๋ก ์ ์ฉ ๊ฐ๋ฅํ๋ค
- ์ ํ์ ํ์ฅ์ฑ: ํ์ผ ํฌ๊ธฐ๋ฅผ ์กฐ์ ํ์ฌ ๋ค์ํ ์บ์ ํฌ๊ธฐ์ ๋์ํ ์ ์๋ค
- ํ ์ ์ฝ์ด ํธํ์ฑ: NVIDIA ํ ์ ์ฝ์ด์ WMMA/MMA ๋ช ๋ น๊ณผ ์์ฐ์ค๋ฝ๊ฒ ๊ฒฐํฉ๋๋ค
๋จ์
- ํ๋ผ๋ฏธํฐ ํ๋ ํ์: ํ์ผ ํฌ๊ธฐ๋ ํ๋์จ์ด์ ๋ฐ๋ผ ์ต์ ๊ฐ์ด ๋ฌ๋ผ์ง๋ฏ๋ก ์๋ ํ๋์ด ํ์ํ๋ค
- ๋น์ ํ ํ๋ ฌ ๋นํจ์จ: ํ๋ ฌ ์ฐจ์์ด ํ์ผ ํฌ๊ธฐ๋ก ๋๋์ด๋จ์ด์ง์ง ์์ผ๋ฉด ํจ๋ฉ(padding)์ด ๋ฐ์ํ์ฌ ์ฐ์ฐ ๋ญ๋น๊ฐ ์์ ์ ์๋ค
- ๋ฉ๋ชจ๋ฆฌ ์ค๋ฒํค๋: ํ์ผ์ ๊ณต์ ๋ฉ๋ชจ๋ฆฌ๋ ๋ ์ง์คํฐ์ ๋ก๋ํ๋ ๊ณผ์ ์์ ์ถ๊ฐ์ ์ธ ๋ฉ๋ชจ๋ฆฌ ์ด๋์ด ๋ฐ์ํ๋ค
- ๋๊ธฐํ ๋น์ฉ: GPU ์ปค๋์์๋
__syncthreads()๊ฐ์ ๋๊ธฐํ ๋ช ๋ น์ด ํ์ํ์ฌ ์ค๋ฒํค๋๊ฐ ๋ฐ์ํ ์ ์๋ค
๊ด๋ จ ๊ธฐ์
ํ๋ ์์ํฌ ๋ฐ ๋ผ์ด๋ธ๋ฌ๋ฆฌ
- NVIDIA cuBLAS: Square Tiling์ ๊ธฐ๋ณธ์ผ๋ก ์ฌ์ฉํ๋ GEMM ๋ผ์ด๋ธ๋ฌ๋ฆฌ
- NVIDIA CUTLASS: ํ ํ๋ฆฟ ๊ธฐ๋ฐ GEMM ์ปค๋๋ก ๋ค์ํ ํ์ผ๋ง ์ ๋ต ์ง์
- OpenBLAS: CPU ๊ธฐ๋ฐ GEMM์์ L1/L2 ์บ์ ํ์ผ๋ง ๊ตฌํ
- Intel MKL: x86 ํ๋ก์ธ์์ ์ต์ ํ๋ GEMM ์ปค๋
- OpenAI Triton: ํ์ผ๋ง์ ์๋ํํ๋ ์ปดํ์ผ๋ฌ ํ๋ ์์ํฌ
๊ด๋ จ ๊ฐ๋
- GEMV (General Matrix-Vector): ๋ฒกํฐ ์ฐจ์์ ๋ํ locality ์ต์ ํ
- Strassen ์๊ณ ๋ฆฌ์ฆ: O(N^2.81) ๊ณฑ์ ํ์, ํ์ผ๋ง๊ณผ ๊ฒฐํฉ ๊ฐ๋ฅ. 2ร2 ํ๋ ฌ ๊ณฑ์ ์ 7๋ฒ์ ๊ณฑ์ ์ผ๋ก ์ค์ด๊ณ ์ฌ๊ท์ ์ผ๋ก ์ ์ฉ
- Register Tiling: ๋ ์ง์คํฐ ์์ค์ ๋ฏธ์ธ ํ์ผ๋ง
- Loop Unrolling: ํ์ผ ๋ด๋ถ ๋ฃจํ ์ต์ ํ
- SIMD Vectorization: ํ์ผ ๋ด ์ฐ์ฐ์ ๋ฒกํฐํ
- Cannon's Algorithm: ๋ถ์ฐ ๋ฉ๋ชจ๋ฆฌ ํ๊ฒฝ์์ ํต์ ์ ํํผํ๋ 2D ํ์ผ๋ง ์๊ณ ๋ฆฌ์ฆ
- AlphaTensor: DeepMind์์ ๋ฐ๊ฒฌํ ํ๋ ฌ ๊ณฑ์ ์ต์ ํ ์๊ณ ๋ฆฌ์ฆ์ผ๋ก, Strassen๋ณด๋ค ๋ ์ ์ ๊ณฑ์ ์ผ๋ก n^2.778 ๋ณต์ก๋ ๋ฌ์ฑ
์ฐธ๊ณ ๋ฌธํ
- "Optimizing Matrix Multiplication on GPUs" - NVIDIA Developer Blog
- "Anatomy of High-Performance Matrix Multiplication" - Goto & Van De Geijn (2008)
- "CUTLASS: CUDA Templates for Linear Algebra Subroutines" - NVIDIA
- "FlashAttention: Fast and Memory-Efficient Exact Attention" - Dao et al. (2022)
- "Triton: An Intermediate Language and Compiler for Tiled Neural Network Computations" - Tillet et al. (2019)
- "Cache-Oblivious Algorithms" - Harald Prokop (MIT, 1999)
- "AlphaTensor: Discovering novel, efficient and exact matrix multiplication algorithms" - DeepMind (2022)
ํต์ฌ ์ ๋ฆฌ
- Square Tiling์ ํฐ ํ๋ ฌ์ ์์ ์ ์ฌ๊ฐํ ํ์ผ๋ก ๋ถํ ํ์ฌ ์บ์ locality๋ฅผ ๊ทน๋ํํ๋ ๊ธฐ๋ณธ GEMM ์ต์ ํ ๊ธฐ๋ฒ์ด๋ค
- ํ์ผ ํฌ๊ธฐ๋ ํ๋์จ์ด์ L1 ์บ์ ํฌ๊ธฐ์ ๋ฐ๋ผ ๊ฒฐ์ ๋๋ฉฐ, ์ผ๋ฐ์ ์ผ๋ก โ(์บ์ ํฌ๊ธฐ/3) ๊ณต์์ผ๋ก ๊ณ์ฐํ๋ค
- ํ์ผ๋ง์ ์ ์ฉํ๋ฉด ์บ์ ๋ฏธ์ค๊ฐ O(Nยณ)์์ O(Nยณ/Tยฒ)๋ก ๊ฐ์ํ์ฌ, ํ์ผ ํฌ๊ธฐ T๋งํผ์ ๋ฐ์ดํฐ ์ฌ์ฌ์ฉ์ด ๊ฐ๋ฅํ๋ค
- GPU์์๋ ๊ณต์ ๋ฉ๋ชจ๋ฆฌ์ ํ์ผ์ ๋ก๋ํ๊ณ ์ํ ๋จ์๋ก ๋ณ๋ ฌ ๊ณฑ์ ์ ์ํํ๋ ๋ฐฉ์์ผ๋ก ๊ตฌํ๋๋ค
- Square Tiling์ ๊ตฌํ์ด ์ฉ์ดํ๊ณ ํ๋์จ์ด ๋ ๋ฆฝ์ ์ด์ด์, cuBLAS, CUTLASS, Triton ๋ฑ ๋ค์ํ ๋ผ์ด๋ธ๋ฌ๋ฆฌ์ ๊ธฐ๋ณธ ์ ๋ต์ผ๋ก ์ฌ์ฉ๋๋ค