โšก AI Optimization

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% ์ด์ƒ์˜ ํ”ผํฌ ์„ฑ๋Šฅ์„ ๋‹ฌ์„ฑํ•  ์ˆ˜ ์žˆ๊ฒŒ ํ•˜๋Š” ํ•ต์‹ฌ ์ตœ์ ํ™” ๊ธฐ๋ฒ•์ด๋‹ค.

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

Register/Shared Memory Tiling ๋ฉ”๋ชจ๋ฆฌ ๊ณ„์ธต ๊ตฌ์กฐ

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

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();

๋”๋ธ” ๋ฒ„ํผ๋ง

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

  1. ๋ฒ„ํผ 0: ํ˜„์žฌ ํƒ€์ผ ์—ฐ์‚ฐ ์ˆ˜ํ–‰
  2. ๋ฒ„ํผ 1: ๋‹ค์Œ ํƒ€์ผ์„ ๊ธ€๋กœ๋ฒŒ ๋ฉ”๋ชจ๋ฆฌ์—์„œ ๋กœ๋”ฉ
  3. ๋‹ค์Œ ๋ฐ˜๋ณต์—์„œ ๋ฒ„ํผ ๊ต์ฒด

์ด๋ฅผ ํ†ตํ•ด ๋ฉ”๋ชจ๋ฆฌ ๋กœ๋”ฉ ์ง€์—ฐ์„ ์—ฐ์‚ฐ์œผ๋กœ ์ˆจ๊ธธ ์ˆ˜ ์žˆ๋‹ค.

๋น„๊ต/๋ถ„์„

Register vs Shared Memory Tiling ๋น„๊ต

ํƒ€์ผ๋ง ์ „๋žต ๋น„๊ต

์ „๋žต ๋ฐ์ดํ„ฐ ์žฌ์‚ฌ์šฉ ์œ„์น˜ ์Šค๋ ˆ๋“œ๋‹น ์ถœ๋ ฅ ๋ฉ”๋ชจ๋ฆฌ ์ œ์•ฝ ์„ฑ๋Šฅ ํŠน์„ฑ
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๋ธ”๋ก)

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

Register/Shared Memory Tiling ๋™์ž‘ ํ๋ฆ„

์ „์ฒด ์•Œ๊ณ ๋ฆฌ์ฆ˜ ๊ตฌ์กฐ

// 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();
}

์žฅ๋‹จ์ 

์žฅ์ 

  1. ๊ทน๋Œ€ํ™”๋œ ์—ฐ์‚ฐ ๋ฐ€๋„: TMร—TN ํฌ๊ธฐ์˜ ์™ธ์  ์—ฐ์‚ฐ์„ ํ†ตํ•ด ์Šค๋ ˆ๋“œ๋‹น FLOP์„ ๋Œ€ํญ ์ฆ๊ฐ€์‹œ์ผœ ๋ฉ”๋ชจ๋ฆฌ ๋Œ€์—ญํญ ๋ณ‘๋ชฉ์„ ์ตœ์†Œํ™”ํ•œ๋‹ค
  2. ๊ณ„์ธต์  ๋ฐ์ดํ„ฐ ์žฌ์‚ฌ์šฉ: ๊ธ€๋กœ๋ฒŒ โ†’ ๊ณต์œ  ๋ฉ”๋ชจ๋ฆฌ โ†’ ๋ ˆ์ง€์Šคํ„ฐ ์„ธ ๊ณ„์ธต์—์„œ ๊ฐ๊ฐ BMร—BN/BK, BM/BK, 1๋ฐฐ์˜ ๋ฐ์ดํ„ฐ ์žฌ์‚ฌ์šฉ์„ ๋‹ฌ์„ฑํ•œ๋‹ค
  3. ํ…์„œ ์ฝ”์–ด์™€์˜ ์ž์—ฐ์Šค๋Ÿฌ์šด ๊ฒฐํ•ฉ: NVIDIA Tensor Core์˜ WMMA/MMA ๋ช…๋ น๊ณผ ํƒ€์ผ ๊ตฌ์กฐ๊ฐ€ ํ˜ธํ™˜๋˜์–ด ํ˜ผํ•ฉ ์ •๋ฐ€๋„ ์—ฐ์‚ฐ์—์„œ ์ถ”๊ฐ€์ ์ธ ์„ฑ๋Šฅ ํ–ฅ์ƒ์„ ์ œ๊ณตํ•œ๋‹ค
  4. ์œ ์—ฐํ•œ ํŒŒ๋ผ๋ฏธํ„ฐ ์กฐ์ •: BM/BN/BK/TM/TN์„ ํ•˜๋“œ์›จ์–ด ํŠน์„ฑ์— ๋งž๊ฒŒ ์กฐ์ •ํ•˜์—ฌ ๋‹ค์–‘ํ•œ ์•„ํ‚คํ…์ฒ˜์— ์ตœ์ ํ™”ํ•  ์ˆ˜ ์žˆ๋‹ค
  5. ๋†’์€ ์‹ค์šฉ์„ฑ: cuBLAS, CUTLASS, Triton ๋“ฑ ๊ฒ€์ฆ๋œ ๊ตฌํ˜„์ฒด๊ฐ€ ์กด์žฌํ•˜์—ฌ ํ”„๋กœ๋•์…˜ ํ™˜๊ฒฝ์—์„œ ์ฆ‰์‹œ ํ™œ์šฉ ๊ฐ€๋Šฅํ•˜๋‹ค

๋‹จ์ 

  1. ๊ตฌํ˜„ ๋ณต์žก๋„: ๋ฉ”๋ชจ๋ฆฌ ๊ณ„์ธต ๊ฐ„ ๋™๊ธฐํ™”, ๋ ˆ์ง€์Šคํ„ฐ ํ• ๋‹น, ํƒ€์ผ ๋กœ๋”ฉ ์ˆœ์„œ ๋“ฑ ๊ณ ๋ คํ•  ์š”์†Œ๊ฐ€ ๋งŽ์•„ ๋””๋ฒ„๊น…์ด ์–ด๋ ต๋‹ค
  2. ๋ ˆ์ง€์Šคํ„ฐ ์••๋ฐ•(Register Pressure): TMร—TN ๋ˆ„์‚ฐ๊ธฐ๊ฐ€ ๋งŽ์€ ๋ ˆ์ง€์Šคํ„ฐ๋ฅผ ์†Œ๋ชจํ•˜์—ฌ, ์ ์œ ์œจ(Occupancy)์ด ๋‚ฎ์•„์งˆ ์ˆ˜ ์žˆ๋‹ค. ์Šค๋ ˆ๋“œ ์ˆ˜ ๊ฐ์†Œ๋Š” ๋ณ‘๋ ฌ์„ฑ์„ ์ €ํ•˜์‹œํ‚จ๋‹ค
  3. ๊ณต์œ  ๋ฉ”๋ชจ๋ฆฌ ์ œ์•ฝ: SMEM ํฌ๊ธฐ๊ฐ€ ์ œํ•œ์ ์ด๋ฏ€๋กœ ํƒ€์ผ ํฌ๊ธฐ๋ฅผ ๊ณผ๋„ํ•˜๊ฒŒ ํ‚ค์šธ ์ˆ˜ ์—†์œผ๋ฉฐ, ๋ฑ…ํฌ ์ถฉ๋Œ(Bank Conflict) ํšŒํ”ผ๋ฅผ ์œ„ํ•œ ํŒจ๋”ฉ์ด ํ•„์š”ํ•˜๋‹ค
  4. ๋น„์ •ํ˜• ํ–‰๋ ฌ ์ฒ˜๋ฆฌ: ํ–‰๋ ฌ ์ฐจ์›์ด ํƒ€์ผ ํฌ๊ธฐ๋กœ ๋‚˜๋ˆ„์–ด๋–จ์–ด์ง€์ง€ ์•Š์œผ๋ฉด ํŒจ๋”ฉ์ด ๋ฐœ์ƒํ•˜์—ฌ ์—ฐ์‚ฐ ๋‚ญ๋น„์™€ ๋ณต์žก๋„ ์ฆ๊ฐ€๋ฅผ ์ดˆ๋ž˜ํ•œ๋‹ค
  5. ํ•˜๋“œ์›จ์–ด ์˜์กด์„ฑ: ์ตœ์  ํƒ€์ผ ํฌ๊ธฐ๋Š” 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)

ํ•ต์‹ฌ ์ •๋ฆฌ

  1. Register/Shared Memory Tiling์€ ๊ธ€๋กœ๋ฒŒ ๋ฉ”๋ชจ๋ฆฌ โ†’ ๊ณต์œ  ๋ฉ”๋ชจ๋ฆฌ โ†’ ๋ ˆ์ง€์Šคํ„ฐ ์„ธ ๊ณ„์ธต์—์„œ ํƒ€์ผ์„ ์ฒด๊ณ„์ ์œผ๋กœ ๊ด€๋ฆฌํ•˜๋Š” GPU GEMM ์ตœ์ ํ™” ๊ธฐ๋ฒ•์ด๋‹ค
  2. ๊ฐ ์Šค๋ ˆ๋“œ๊ฐ€ TMร—TN ํฌ๊ธฐ์˜ ์ถœ๋ ฅ ์„œ๋ธŒ ๋งคํŠธ๋ฆญ์Šค๋ฅผ ์™ธ์ (Outer Product) ๋ฐฉ์‹์œผ๋กœ ๊ณ„์‚ฐํ•˜์—ฌ, ์Šค๋ ˆ๋“œ๋‹น ์—ฐ์‚ฐ ๋ฐ€๋„๋ฅผ ๊ทน๋Œ€ํ™”ํ•œ๋‹ค
  3. ๊ณต์œ  ๋ฉ”๋ชจ๋ฆฌ ํƒ€์ผ์€ ๋ธ”๋ก ๋‚ด ๋ชจ๋“  ์Šค๋ ˆ๋“œ๊ฐ€ ํ˜‘๋ ฅ์ ์œผ๋กœ ๋กœ๋”ฉํ•˜์—ฌ ๋ฐ์ดํ„ฐ ์žฌ์‚ฌ์šฉ์„ ๋‹ฌ์„ฑํ•˜๊ณ , ๋ ˆ์ง€์Šคํ„ฐ ํƒ€์ผ์€ ์Šค๋ ˆ๋“œ ๋‚ด๋ถ€์—์„œ ์—ฐ์‚ฐ์„ ์ˆ˜ํ–‰ํ•œ๋‹ค
  4. ๋”๋ธ” ๋ฒ„ํผ๋ง์„ ํ†ตํ•ด ํƒ€์ผ ๋กœ๋”ฉ๊ณผ ์—ฐ์‚ฐ์„ ๊ฒน์ณ ์‹คํ–‰ํ•˜์—ฌ ๋ฉ”๋ชจ๋ฆฌ ์ง€์—ฐ ์‹œ๊ฐ„์„ ์ˆจ๊ธธ ์ˆ˜ ์žˆ๋‹ค
  5. cuBLAS, CUTLASS, Triton ๋“ฑ ๋ชจ๋“  ๊ณ ์„ฑ๋Šฅ GEMM ๋ผ์ด๋ธŒ๋Ÿฌ๋ฆฌ๊ฐ€ ์ด ์ „๋žต์„ ๊ธฐ๋ณธ์œผ๋กœ ์ฑ„ํƒํ•˜๋ฉฐ, NVIDIA ํ…์„œ ์ฝ”์–ด์™€ ๊ฒฐํ•ฉํ•˜์—ฌ ํ”ผํฌ ์„ฑ๋Šฅ์˜ 90% ์ด์ƒ์„ ๋‹ฌ์„ฑํ•  ์ˆ˜ ์žˆ๋‹ค