โšก AI Optimization

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ยฒ)๋กœ ๊ฐ์†Œํ•œ๋‹ค.

ํ•ต์‹ฌ ๊ฐœ๋…

Square Tiling ๊ธฐ๋ณธ ๊ตฌ์กฐ

ํƒ€์ผ๋ง(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์„ ์ ์šฉํ•˜๋ฉด:

  1. ํ–‰๋ ฌ ๋ถ„ํ• : ํ–‰๋ ฌ A, B๋ฅผ Tร—T ํฌ๊ธฐ์˜ ์ •์‚ฌ๊ฐํ˜• ํƒ€์ผ๋กœ ๋ถ„ํ• 
  2. ํƒ€์ผ ๋‹จ์œ„ ์—ฐ์‚ฐ: ๊ฐ ํƒ€์ผ ์Œ(A ํƒ€์ผ, B ํƒ€์ผ)์— ๋Œ€ํ•ด ๋ถ€๋ถ„ ๊ณฑ์…ˆ ์ˆ˜ํ–‰
  3. ๋ˆ„์‚ฐ: ๋ถ€๋ถ„ ๊ณฑ์…ˆ ๊ฒฐ๊ณผ๋ฅผ C ํƒ€์ผ์— ๋ˆ„์‚ฐ
  4. ๋ฐ˜๋ณต: ๋ชจ๋“  ํƒ€์ผ ์Œ์— ๋Œ€ํ•ด ์œ„ ๊ณผ์ •์„ ๋ฐ˜๋ณต

์บ์‹œ ๋ ˆ๋ฒจ๊ณผ ํƒ€์ผ ํฌ๊ธฐ

ํƒ€์ผ ํฌ๊ธฐ๋Š” ํ•˜๋“œ์›จ์–ด์˜ ๋ฉ”๋ชจ๋ฆฌ ๊ณ„์ธต ๊ตฌ์กฐ์— ๋”ฐ๋ผ ๊ฒฐ์ •๋œ๋‹ค:

๋ฉ”๋ชจ๋ฆฌ ๊ณ„์ธต ํฌ๊ธฐ ๋Œ€ํ‘œ ์˜ˆ์‹œ ํƒ€์ผ ํฌ๊ธฐ ์˜ํ–ฅ
๋ ˆ์ง€์Šคํ„ฐ ํŒŒ์ผ ์ˆ˜ 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 (๊ท ํ˜•์ )

๋™์ž‘ ์›๋ฆฌ

Square Tiling ๋™์ž‘ ํ๋ฆ„

๊ธฐ๋ณธ 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;
}

์žฅ๋‹จ์ 

์žฅ์ 

  1. ์บ์‹œ ํšจ์œจ์„ฑ ๊ทน๋Œ€ํ™”: ํƒ€์ผ ํฌ๊ธฐ๋งŒํผ์˜ ๋ฐ์ดํ„ฐ๋ฅผ ์บ์‹œ์— ์ ์žฌํ•œ ํ›„ ์—ฌ๋Ÿฌ ๋ฒˆ ์žฌ์‚ฌ์šฉํ•˜์—ฌ ์บ์‹œ ๋ฏธ์Šค๋ฅผ ํฌ๊ฒŒ ๊ฐ์†Œ์‹œํ‚จ๋‹ค
  2. ๊ตฌํ˜„ ์šฉ์ด์„ฑ: ๋‹ค๋ฅธ ํƒ€์ผ๋ง ๊ธฐ๋ฒ•(์˜ˆ: Register Tiling)์— ๋น„ํ•ด ์•Œ๊ณ ๋ฆฌ์ฆ˜์ด ๋‹จ์ˆœํ•˜๊ณ  ์ง๊ด€์ ์ด๋‹ค
  3. ํ•˜๋“œ์›จ์–ด ๋…๋ฆฝ์„ฑ: CPU, GPU, FPGA ๋“ฑ ๋‹ค์–‘ํ•œ ํ”Œ๋žซํผ์—์„œ ๋™์ผํ•œ ์›๋ฆฌ๋กœ ์ ์šฉ ๊ฐ€๋Šฅํ•˜๋‹ค
  4. ์„ ํ˜•์  ํ™•์žฅ์„ฑ: ํƒ€์ผ ํฌ๊ธฐ๋ฅผ ์กฐ์ ˆํ•˜์—ฌ ๋‹ค์–‘ํ•œ ์บ์‹œ ํฌ๊ธฐ์— ๋Œ€์‘ํ•  ์ˆ˜ ์žˆ๋‹ค
  5. ํ…์„œ ์ฝ”์–ด ํ˜ธํ™˜์„ฑ: NVIDIA ํ…์„œ ์ฝ”์–ด์˜ WMMA/MMA ๋ช…๋ น๊ณผ ์ž์—ฐ์Šค๋Ÿฝ๊ฒŒ ๊ฒฐํ•ฉ๋œ๋‹ค

๋‹จ์ 

  1. ํŒŒ๋ผ๋ฏธํ„ฐ ํŠœ๋‹ ํ•„์š”: ํƒ€์ผ ํฌ๊ธฐ๋Š” ํ•˜๋“œ์›จ์–ด์— ๋”ฐ๋ผ ์ตœ์  ๊ฐ’์ด ๋‹ฌ๋ผ์ง€๋ฏ€๋กœ ์ˆ˜๋™ ํŠœ๋‹์ด ํ•„์š”ํ•˜๋‹ค
  2. ๋น„์ •ํ˜• ํ–‰๋ ฌ ๋น„ํšจ์œจ: ํ–‰๋ ฌ ์ฐจ์›์ด ํƒ€์ผ ํฌ๊ธฐ๋กœ ๋‚˜๋ˆ„์–ด๋–จ์–ด์ง€์ง€ ์•Š์œผ๋ฉด ํŒจ๋”ฉ(padding)์ด ๋ฐœ์ƒํ•˜์—ฌ ์—ฐ์‚ฐ ๋‚ญ๋น„๊ฐ€ ์žˆ์„ ์ˆ˜ ์žˆ๋‹ค
  3. ๋ฉ”๋ชจ๋ฆฌ ์˜ค๋ฒ„ํ—ค๋“œ: ํƒ€์ผ์„ ๊ณต์œ  ๋ฉ”๋ชจ๋ฆฌ๋‚˜ ๋ ˆ์ง€์Šคํ„ฐ์— ๋กœ๋“œํ•˜๋Š” ๊ณผ์ •์—์„œ ์ถ”๊ฐ€์ ์ธ ๋ฉ”๋ชจ๋ฆฌ ์ด๋™์ด ๋ฐœ์ƒํ•œ๋‹ค
  4. ๋™๊ธฐํ™” ๋น„์šฉ: 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)

ํ•ต์‹ฌ ์ •๋ฆฌ

  1. Square Tiling์€ ํฐ ํ–‰๋ ฌ์„ ์ž‘์€ ์ •์‚ฌ๊ฐํ˜• ํƒ€์ผ๋กœ ๋ถ„ํ• ํ•˜์—ฌ ์บ์‹œ locality๋ฅผ ๊ทน๋Œ€ํ™”ํ•˜๋Š” ๊ธฐ๋ณธ GEMM ์ตœ์ ํ™” ๊ธฐ๋ฒ•์ด๋‹ค
  2. ํƒ€์ผ ํฌ๊ธฐ๋Š” ํ•˜๋“œ์›จ์–ด์˜ L1 ์บ์‹œ ํฌ๊ธฐ์— ๋”ฐ๋ผ ๊ฒฐ์ •๋˜๋ฉฐ, ์ผ๋ฐ˜์ ์œผ๋กœ โˆš(์บ์‹œ ํฌ๊ธฐ/3) ๊ณต์‹์œผ๋กœ ๊ณ„์‚ฐํ•œ๋‹ค
  3. ํƒ€์ผ๋ง์„ ์ ์šฉํ•˜๋ฉด ์บ์‹œ ๋ฏธ์Šค๊ฐ€ O(Nยณ)์—์„œ O(Nยณ/Tยฒ)๋กœ ๊ฐ์†Œํ•˜์—ฌ, ํƒ€์ผ ํฌ๊ธฐ T๋งŒํผ์˜ ๋ฐ์ดํ„ฐ ์žฌ์‚ฌ์šฉ์ด ๊ฐ€๋Šฅํ•˜๋‹ค
  4. GPU์—์„œ๋Š” ๊ณต์œ  ๋ฉ”๋ชจ๋ฆฌ์— ํƒ€์ผ์„ ๋กœ๋“œํ•˜๊ณ  ์›Œํ”„ ๋‹จ์œ„๋กœ ๋ณ‘๋ ฌ ๊ณฑ์…ˆ์„ ์ˆ˜ํ–‰ํ•˜๋Š” ๋ฐฉ์‹์œผ๋กœ ๊ตฌํ˜„๋œ๋‹ค
  5. Square Tiling์€ ๊ตฌํ˜„์ด ์šฉ์ดํ•˜๊ณ  ํ•˜๋“œ์›จ์–ด ๋…๋ฆฝ์ ์ด์–ด์„œ, cuBLAS, CUTLASS, Triton ๋“ฑ ๋‹ค์–‘ํ•œ ๋ผ์ด๋ธŒ๋Ÿฌ๋ฆฌ์˜ ๊ธฐ๋ณธ ์ „๋žต์œผ๋กœ ์‚ฌ์šฉ๋œ๋‹ค