(1/n) Triton matmul 커널 탐험 - persistent kernel은 그 자체로 빠르지 않다
요새 triton 쪽에 흥미를 가지고 있어, 관련된 글을 천천히 적어나가 보고자 한다. 사실 코드를 작성하고 테스트해보다가 글을 쓰겠다고 마음을 먹은것은 10일 전에 마음을 먹었지만, 도통 글이 적히질 않더라... 그다지 llm의 도움을 통해서 적고 싶고자하는 마음도 없어서, 공부하고 탐구하는데에만 도움을 받고 글을 쓰는건 직접 써보려고 한다. 그래서 그냥 내멋대로 쓰고 정리하는 것으로 만족하려고 한다.
주로 다루어보고 싶은 주제들은 성능 튜닝 중에 어떤 고민을 하는가인데 내가 전문가도 아니고....순수 흥미에 따라 주제를 보고 있기 때문에 규칙적으로 많은 것들을 적어볼 수 있을지는 모르겠다.
naive tiled matmul
일단은 가장 원시적인 tile 형태의 matmul kernel을 작성해보았다. naive tiled matmul 커널이다. tiled matmul은 GPU에서 matmul kernel 속도를 높이기 위해 l1 캐시와 같은 수준으로 빠른 shared memory를 활용하는 방법이다.
@triton.jit
def _matmul_kernel(
a_ptr,
b_ptr,
c_ptr,
M,
N,
K,
stride_am,
stride_ak,
stride_bk,
stride_bn,
stride_cm,
stride_cn,
NUM_STAGES: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
"""1 program이 C tile 하나를 계산한다."""
tile_id = tl.program_id(0)
num_n_tiles = tl.cdiv(N, BLOCK_N)
tile_m = tile_id // num_n_tiles
tile_n = tile_id % num_n_tiles
# TODO 1: a_ptr/b_ptr와 stride로 tile pointer와 M/N/K mask를 만든다.
offsets_m = tile_m * BLOCK_M + tl.arange(0, BLOCK_M)
offsets_n = tile_n * BLOCK_N + tl.arange(0, BLOCK_N)
a_row = a_ptr + offsets_m[:, None] * stride_am
b_col = b_ptr + offsets_n[None, :] * stride_bn
mask_m = offsets_m[:, None] < M
mask_n = offsets_n[None, :] < N
# TODO 2: K축을 순회하며 FP32 acc에 tl.dot을 누적한다.
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in tl.range(0, K, BLOCK_K, num_stages=NUM_STAGES):
offsets_k = k + tl.arange(0, BLOCK_K)
mask_a = mask_m & (offsets_k[None, :] < K)
mask_b = mask_n & (offsets_k[:, None] < K)
a_tile = tl.load(a_row + offsets_k[None, :] * stride_ak, mask=mask_a, other=0.0)
b_tile = tl.load(b_col + offsets_k[:, None] * stride_bk, mask=mask_b, other=0.0)
acc = tl.dot(a_tile, b_tile, acc)
# TODO 3: C tile을 저장한다.
tl.store(c_ptr + offsets_m[:, None] * stride_cm + offsets_n[None, :] * stride_cn, acc, mask=mask_m & mask_n)
naive persistent matmul
이번에 비교하고자 하는 대상은 naive persistent kernel 정도라고 보면 될 것 같다. persistent matmul은 GPU에서 matmul kernel 속도를 높이기 위해 SM 개수만큼만 program을 띄우고 각 program이 여러 tile을 순회하게 해서, tile 배정을 하드웨어 스케줄러 대신 kernel이 직접 제어할 수 있게 하는 방법이다. 이제 여기에 cross tile pipelining이나 Stream-K 같은 기법들을 엮어서 고도화할 수 있다.
@triton.jit
def _persistent_matmul_kernel(
a_ptr,
b_ptr,
c_ptr,
M,
N,
K,
stride_am,
stride_ak,
stride_bk,
stride_bn,
stride_cm,
stride_cn,
NUM_PROGRAMS: tl.constexpr,
NUM_STAGES: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
"""program 하나가 일정한 간격으로 여러 C tile을 계산한다."""
start_tile = tl.program_id(0)
num_m_tiles = tl.cdiv(M, BLOCK_M)
num_n_tiles = tl.cdiv(N, BLOCK_N)
num_tiles = num_m_tiles * num_n_tiles
for tile_id in tl.range(start_tile, num_tiles, NUM_PROGRAMS):
tile_m = tile_id // num_n_tiles
tile_n = tile_id % num_n_tiles
# TODO 4: 기본 kernel의 stride 기반 tile GEMM 본문을 이곳에 옮긴다.
# acc와 pointer는 tile마다 새로 초기화해야 한다.
# a_ptr/b_ptr와 stride로 tile pointer와 M/N/K mask를 만든다.
offsets_m = tile_m * BLOCK_M + tl.arange(0, BLOCK_M)
offsets_n = tile_n * BLOCK_N + tl.arange(0, BLOCK_N)
a_row = a_ptr + offsets_m[:, None] * stride_am
b_col = b_ptr + offsets_n[None, :] * stride_bn
mask_m = offsets_m[:, None] < M
mask_n = offsets_n[None, :] < N
# K축을 순회하며 FP32 acc에 tl.dot을 누적한다.
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in tl.range(0, K, BLOCK_K, num_stages=NUM_STAGES):
offsets_k = k + tl.arange(0, BLOCK_K)
mask_a = mask_m & (offsets_k[None, :] < K)
mask_b = mask_n & (offsets_k[:, None] < K)
a_tile = tl.load(a_row + offsets_k[None, :] * stride_ak, mask=mask_a, other=0.0)
b_tile = tl.load(b_col + offsets_k[:, None] * stride_bk, mask=mask_b, other=0.0)
acc = tl.dot(a_tile, b_tile, acc)
# C tile을 저장한다.
tl.store(c_ptr + offsets_m[:, None] * stride_cm + offsets_n[None, :] * stride_cn, acc, mask=mask_m & mask_n)
persistent matmul kernel의 실행시 grid 설정은 다음과 같이 한다. 여태까지 타일의 갯수에 맞추어 프로그램을 실행했던 것과는 다르게 앞서 얘기했던 대로 SM 개수에 맞추어 실행하였다.
num_tiles = triton.cdiv(M, block_m) * triton.cdiv(N, block_n)
num_sms = torch.cuda.get_device_properties(a.device).multi_processor_count
num_programs = min(num_sms, num_tiles)
grid = (num_programs,)
측정 전 기대
측정 전에 내가 기대하고 있는 것은 naive persistent matmul이 naive tiled matmul과 비슷하거나 미세하게 빠른 성능을 내는 것이었다. persistent matmul이 전체 과정에서 큰 비중은 아니지만 block scheduling과 prologue 비용을 줄여 줄 수 있다고 봤기 때문이다. 반면 전체 블록이 SM에 배정되어 계산하는 시간은 대충 ceil(num_tiles / NUM_SMS) 정도로 비슷할거라고 판단했다.
측정 결과
M=4096, N=4096, K=4096, BLOCK_M=128, BLOCK_N=64, BLOCK_K=32, NUM_STAGES=4에서의 측정 결과 (RTX 5090, torch 2.13.0+cu130, triton 3.7.1)
torch: 0.6277 ms, 218.95 TFLOPS
naive tiled: 0.6448 ms, 213.16 TFLOPS
naive persistent: 0.7042 ms, 195.17 TFLOPS
생각과는 매우 다른 결과가 나왔다. naive tiled matmul은 torch matmul과 97.4%로 근접했고 naive persistent matmul은 89.1%로 다소 떨어져 있었다. 원인을 찾기 위해 익숙하지는 않지만 ncu로 분석을 해보고자 한다.
ncu로 원인 찾기
ncu --target-processes all \
--section LaunchStats \
--section Occupancy \
--section SpeedOfLight \
--section ComputeWorkloadAnalysis \
--section MemoryWorkloadAnalysis \
--section SchedulerStats \
--section WarpStateStats \
--kernel-name 'regex:^_persistent_matmul_kernel$' \
--launch-count 1 \
uv run --frozen python trition_tutorial/persistent_matmul/benchmark.py
Section: Launch Statistics
-------------------------------- --------------- ---------------
Metric Name Metric Unit Metric Value
-------------------------------- --------------- ---------------
Block Size 128
Cluster Scheduling Policy PolicySpread
Cluster Size 0
Function Cache Configuration CachePreferNone
Grid Size 170
Preferred Cluster Size 0
Registers Per Thread register/thread 174
Shared Memory Configuration Size Kbyte 102.40
Driver Shared Memory Per Block Kbyte/block 1.02
Dynamic Shared Memory Per Block Kbyte/block 36.86
Static Shared Memory Per Block byte/block 0
# SMs SM 170
Stack Size 1024
Threads thread 21760
# TPCs 85
Enabled TPC IDs all
Uses Green Context 0
Waves Per SM 0.50
-------------------------------- --------------- ---------------
OPT If you execute __syncthreads() to synchronize the threads of a block, it is recommended to have at least two
blocks per multiprocessor (compared to the currently executed 1.0 blocks) This way, blocks that aren't
waiting for __syncthreads() can keep the hardware busy.
Section: Occupancy
------------------------------- ----------- ------------
Metric Name Metric Unit Metric Value
------------------------------- ----------- ------------
Max Active Clusters cluster 0
Max Cluster Size block 8
Overall GPU Occupancy % 0
Cluster Occupancy % 0
Block Limit Barriers block 24
Block Limit SM block 24
Block Limit Registers block 2
Block Limit Shared Mem block 2
Block Limit Warps block 12
Theoretical Active Warps per SM warp 8
Theoretical Occupancy % 16.67
Achieved Occupancy % 8.33
Achieved Active Warps Per SM warp 4.00
------------------------------- ----------- ------------
OPT Est. Speedup: 17.4%
The 2.00 theoretical warps per scheduler this kernel can issue according to its occupancy are below the
hardware maximum of 12. This kernel's theoretical occupancy (16.7%) is limited by the number of required
registers, and the required amount of shared memory.
Waves Per SM
많은 section을 켰지만, 아직 다 의미는 잘 모르고 지금 중점적으로 보려고 하는것은 LaunchStats와 Occupancy 부분이다.
우선 LaunchStats의 Waves Per SM이 0.5인데 이 값이 작다고 생각했다. persistent 커널의 정의와 내가 실행했던 그리드 정의에 따르면 SM당 한개의 커널이 돌아가기 때문에 1이되어야한다고 생각했던것 같다.
실제로 Waves Per SM은 SM당 block 슬롯을 기준으로 계산되는데 이게 0.5인 이유는 실제로 SM에 두개의 블록이 올라가기 때문인 것이다.
naive matmul 커널에서의 이 값은 6.02인데 설정상 2048의 그리드의 크기를 가졌고 SM 크기 170으로 나누면 12.047 정도 나오는데 SM당 2개의 블록이 들어갈수 있다면 추가적으로 2로 나누어 같은 값을 얻을 수 있다.
1이어야 할 것이 0.5이기 때문에 SM당 2개의 블록이 들어갈 수 있다는 의미로 추측해 볼 수 있었다.
Occupancy와 Block Limit
Occupancy = SM에 실제로 상주(resident)하는 warp 수 / SM이 담을 수 있는 최대 warp 수로 정의된다. 분모는 아키텍쳐에 따라 결정되는 값으로 RTX 5090에 대해서는 48이다. 그리고 우리가 현재 1개의 SM에 하나의 프로그램만 실행하고 triton의 기본 num_warps 값인 4로 컴파일된 결과이므로 4 / 48 = 8.33%를 얻을 수 있다. Occupancy 섹션에서 현재의 Achieved Occupancy 값 8.33과 Achieved Active Warps Per SM 값 4.00 warp로 확인할 수도 있다. Theoretical Occupancy 값은 16.67로 측정된 것도 볼 수 있다.
Block Limit은 현재 하드웨어에 따른 자원 제약이 SM당 몇 개의 block을 허용하는지를 항목별로 보여준다. 리포트에서는 Registers 2, Shared Mem 2, Warps 12, Barriers 24, SM 24로 나타났고, 이 중 최솟값인 2가 실제 상한이 된다.
Block Limit 값 직접 계산해보기
triton 커널을 launch하면 컴파일된 커널이 사용하는 register의 수나 spills, shared memory 크기, num_warps 값들을 알 수 있다.
def _print_kernel_metadata_once(
name: str,
kernel,
config: str,
cache_key: tuple[object, ...],
) -> None:
if cache_key in _PRINTED_KERNEL_METADATA:
return
_PRINTED_KERNEL_METADATA.add(cache_key)
print("-" * 80)
print(f"{name} kernel: {config}")
print(f"regs/thread : {kernel.n_regs}")
print(f"spills : {kernel.n_spills}")
print(f"shared/CTA : {kernel.metadata.shared}")
print(f"num_warps : {kernel.metadata.num_warps}")
print("-" * 80)
naive tiled와 naive persistent 모두 0 spills, 36864 shared memory / CTA, 4 num_warps 값을 얻었고 thread당 레지스터수만 각각 128과 174로 차이가 났다. shared memory 사용량의 경우, TILE_MK와 TILE_KN의 크기의 배수인데 이경우 (128 * 32 + 32 * 64) * 2 = 12288 bytes이다. 우리 커널의 shared memory는 딱 3배를 사용했으므로 num_stages 값을 4로 줬더라도 3단계 스테이징이 되었다고 보면 된다.
다시 원래 하려던 RTX 5090에서의 Block Limit 값을 직접 계산해보자. Warp, Shared Memory, Register, SM, Barriers에 제한이 있다.
SM 제한은 하드웨어적으로 최신에서는 16~32 값사이로 정해진다고 한다. 문서상에서는 32이나 ncu는 24로 보고한다.
barriers는 kernel.asm["ptx"]를 print해서 bar.sync, barrier.cta.sync 같은 명령어를 직접 세야한다는거 같은데 세어보니 1개 있는듯 하다. 아마도 제한이 24인거보니 전체 상한도 24일듯하다.
- regs / SM (32 bit): 65,536 (256KB)
- max regs / thread: 255
- max threads / SM: 1,536 (48 warps)
- shared memory + L1 / SM: 128 KB
- max shared / CTA: 99 KB
가장 쉬운 warp 제한부터 계산해보면 warp 제한의 경우 max threads / SM 값과 관련 있는데 이 값을 통해 얻을 수 있는 최대 warps는 48이다. 현재 하나의 블록당 4개의 warp를 실행하므로 12이다.
다음으로 register 제한을 따져보면 일단 max regs / thread 값인 255보다 둘다 작으므로 여기서는 문제가 없다. 그럼 regs / SM이 문제인데 하나의 블록당 사용하는 레지스터 개수는 174 * 128 = 22272개이다.
따라서 65536 / 22272 = 2.94로 register로 인한 최대 블록 갯수는 2이다.
마지막으로 shared memory 제한을 따지면, shared memory + L1 / SM 값은 128KB인데, 128KB를 shared memory와 L1이 나뉘어 써야하고 이것은 사전에 정의된 비율로 나눠질 수 있다.
CC 12.0에서 0, 8, 16, 32, 64, 100KB로 설정 가능하고 1KB는 CUDA 런타임이 자기몫으로 따로 잡아가서 CTA가 차지할 수 있는 최대 크기는 99KB라고 한다.
그래서 100KB를 (36864B + 1KB)로 나누면 2.70으로 shared memory로 인한 최대 블록 갯수는 2이다.
위와 같이 보통 아키텍쳐마다 하드웨어나 소프트웨어적으로 정해진 값들이 있기 때문에 triton 커널의 warmup을 통해 컴파일된 커널에서 코드에서 사용량을 얻어서 미리 SM별 최대 프로그램 수를 구하고 grid 값에 넣어줄 수도 있다.
grid 크기 조정
실제로 Occupancy 섹션을 확인해보면 각 하드웨어 제한으로 인한 block의 수를 확인할 수 있었다. Block Limit으로 시작하는 부분들인데 여기서의 최소 값이 SM당 block 슬롯이라고 한다.
위 결과에서 보듯이 그 값은 2이다. 따라서 programs_per_sm인자를 추가하여 이 값을 grid 생성시에 num_sms에 곱하여 num_programs를 구할 때 사용한다.
num_programs = min(num_sms * programs_per_sm, num_tiles) # programs_per_sm=2
실제로 이 후에 Waves Per SM 값은 1을 얻을 수 있었고, naive persistent: 0.6655 ms, 206.52 TFLOPS로 이전 대비 5.5% 개선되었다. 다만 여전히 naive tiled보다 3.1% 느리다.
또한, Achieved Occupancy는 15.11%로 Theoritical Occupancy에 근접해진 걸 확인할 수 있었다. 약 9% 차이가 나는 것인데 grid가 340개이고 num_tiles가 2048이니 대부분의 program은 6개를 처리하지만 일부 타일에 대해서 7개를 실행하는 프로그램이 생기므로 program간 종료 시간 불균형 때문인 것으로 추정된다.
정리
실제로 programs_per_sm을 증가시키면서 실험해보았을 때, naive tiled matmul(non-persistent 커널)에 가까워질 수록 성능 향상이 있었다. naive한 persistent 커널로 변경함으로써 하드웨어 스케줄러가 하던 타일 배정을 커널로 가져왔지만 더 나은 정책을 추가하지 않았기 때문에 추가 비용만 남은 것으로 보인다.
나머지 3.1% (2026-08-31 추가)
왜 여전히 3.1%의 차이가 나는지를 자세히 파헤쳐 보았다. 첫번째 가설은 load imbalance 때문일 수 있지 않냐는 것이었다. tiled에서도 wave간 불균형은 발생하기 때문에 이 이유는 아닐 것 같았다. SpeedOfLight 섹션에서 SM Active Cycles / Elapsed Cycles(직접 계산), Compute (SM) Throughput, Duration을 통해 확인하였다.
| 지표 | persistent | tiled |
|---|---|---|
| SM Active / Elapsed | 91.61% | 92.45% |
| Compute Throughput | 88.06% | 90.35% |
| Duration | 748.26 µs | 730.05 µs |
SM Active Cycles / Elapsed Cycles 값이 0.84% 차이로 3.1% 성능 차이를 설명하지 못하게 때문에 load imbalance 때문은 아니였다. Duration 값이 2.5%, Compute Throughput이 2.3%로 비슷한 차이를 보이고 또한 성능차인 3.1%와 값이 비슷했다. 이제 다음 가설은 persistent 커널 자체의 오버헤드이다. divsi/remsi 정수 나눗셈 2회, offsets 재계산, acc 초기화, mask 재계산이 program당 6~7회 반복되면서 issue slot을 잡아먹고, 그만큼 tensor core가 놀게 될 수 있다는 가설이다. ComputeWorkloadAnalysis 섹션에서 대략적으로 확인할 수 있다.
| 지표 | tiled | persistent |
|---|---|---|
| Executed IPC (Active) | 0.44 | 0.47 |
| Issue Slots Busy | 10.29% | 10.97% |
| SM Busy | 90.35% | 88.06% |
persistent가 명령어를 더 많이 실행한다. IPC(Instruction Per Cycle)가 6.8% 높고 issue slot도 6.6% 더 많다. persistent에서 명령어가 늘어 났는데 연산량은 같다. 늘어난 명령어는 연산에 아주 직접적인 것은 아닌 것으로 추정해볼 수 있다.
실제로 --section InstructionStats를 켜서 명령어 통계를 확인하면 Issued Instructions 값이 125,485,056(tiled) vs 136,401,632(persistent)로 8.7% 가량 persistent가 높다.
그래서 어떤 부분이 늘었는데가 정말 궁금해서 계속 진행해봤다.
ncu --metrics \
smsp__inst_executed.sum,\
smsp__inst_executed_pipe_alu.sum,\
smsp__inst_executed_pipe_fma.sum,\
smsp__inst_executed_pipe_lsu.sum,\
smsp__inst_executed_pipe_tensor_subpipe_hmma.sum,\
smsp__inst_executed_pipe_cbu.sum,\
smsp__inst_executed_pipe_uniform.sum,\
smsp__inst_executed_op_ldgsts.sum,\
smsp__inst_executed_op_shared_ld.sum,\
smsp__inst_executed_op_global_st.sum
이런식으로 발행된 op 종류별로 분류해서 어느 부분이 늘었는지 비교해보았다.
| 파이프 | tiled | persistent | 차이 |
|---|---|---|---|
| 총합 | 125,485,056 | 136,401,632 | +10,916,576 |
| tensor (HMMA) | 33,554,432 | 33,554,432 | 0 |
| ldgsts | 6,438,912 | 6,438,912 | 0 |
| shared_ld | 3,284,992 | 3,284,992 | 0 |
| global_st | 65,536 | 65,536 | 0 |
| lsu | 24,592,384 | 24,585,552 | −6,832 |
| cbu | 16,384 | 9,552 | −6,832 |
| fma | 14,286,848 | 13,001,936 | −1,284,912 |
| alu | 38,109,184 | 44,017,920 | +5,908,736 |
| uniform | 10,493,952 | 16,799,040 | +6,305,088 |
- tensor, ldgsts(async copy), shared_ld(shared -> register), global_st(global memory store) 연산은 완전히 동일했다.
- alu(정수 산술과 논리 연산, offset 계산과 mask 비교), uniform(uniform datapath. warp 안의 모든 thread가 같은 값을 다루는 scalar 연산 전용)은 늘었다.
- lsu(load/store unit), cbu(convergence barrier unit, 분기와 warp 수렴 제어), fma(fused multiply-add)는 줄었다.
결론
tile 배정을 커널 루프로 가져오면서 tile 인덱스와 offset과 mask 계산이 루프 변수에 의존하는 형태로 바뀌었다. 계산 횟수 자체는 tiled와 같지만 컴파일 결과가 달라져 alu & uniform 연산 명령어가 8.7% 늘었고, HMMA와 메모리 명령어는 그대로인 채 issue slot만 더 소모하면서 실행 시간이 2.5% 증가했다.