currybab's blog

(4/n) Triton matmul 커널 탐험 - TMA, Warp Specialization, Epilogue subtiling

Tensor Memory Accelerator

Hopper 아키텍처부터 도입된 하드웨어로 global memory와 shared memory 사이에서 다차원 타일을 비동기로 옮기는데 사용한다. 스레드가 개별 load/store를 발행하는 대신 descriptor 기반으로 명령 하나로 타일 전체를 비동기 전송하므로, 주소 계산과 경계 처리를 전용 하드웨어가 대신해 스레드·레지스터 부담이 줄고 연산과 전송을 겹치기 쉽다. 사용하려면 가장 안쪽 차원이 연속이고, 나머지 차원의 stride는 16바이트 배수, 시작 주소는 16바이트 정렬이어야 한다.

Triton에서는 host에서 텐서의 base, shape, stride, block_shape을 담은 TensorDescriptor를 만들고, 기존에 포인터 텐서를 넘기던 자리에 이 descriptor를 커널 인자로 넘긴다. 커널 안에서는 포인터 연산 대신 타일의 시작 좌표만 지정해 desc.load([m, k]), desc.store([m, n], value)로 접근한다. 범위 밖 처리는 하드웨어가 해 주므로 mask를 직접 만들 필요가 없다(load는 0으로 채움).

# @host
a_desc = TensorDescriptor(
    base=a, shape=[M, K], strides=list(a.stride()),
    block_shape=[TMA_BLOCK_M, TMA_BLOCK_K],
)

# @device
a_tile = a_desc.load([start_m, start_k])
c_desc.store([store_m, store_n], acc.to(tl.float16))

Warp Specialization

같은 thread block 안에서 warp별로 역할을 나누는 기법이다. CUDA 기준으로는 데이터를 로드하는 producer warp와 MMA를 수행하는 consumer warp로 분리하고, shared memory의 다단계 버퍼와 mbarrier로 동기화한다.

기존 방식 vs Warp specialization

기존 (모든 warp 동일 역할) Warp specialization
루프 구조 각 warp가 로드 → 대기 → MMA를 순서대로 반복 producer는 로드만, consumer는 MMA만 반복
로드/연산 겹침 software pipelining(num_stages)으로 겹치지만, 같은 warp가 둘 다 발행하므로 명령어 스케줄링이 얽힘 역할이 분리되어 서로 기다리지 않고 진행
레지스터 모든 warp가 같은 양 (accumulator 때문에 많이 필요) 역할별 배분 가능 (producer 적게, consumer 많이)
병목 로드 발행이 MMA 발행을 방해하거나, 그 반대 발행 경쟁이 줄어듦

Triton에서는 루프를 tl.range(..., warp_specialize=True)로 바꾸면 컴파일러가 load 파트와 MMA 파트로 자동 분할한다.

Epilogue subtiling

epilogue에서 accumulator를 BM × BN 전체가 아니라 BM × BN/2 같은 조각으로 나눠 변환하고 store하는 기법이다.

기존에는 출력 타일 전체 크기의 shared memory가 epilogue에 필요해서, 그만큼 mainloop의 stage 수를 줄여야 했다. Subtiling은 조각 단위로 같은 버퍼를 재사용하므로 epilogue shared memory가 절반으로 줄고, 남는 공간으로 stage를 더 늘려 메모리 지연을 더 잘 숨길 수 있다. 레지스터 압력도 줄고, 한 조각의 store와 다음 조각의 변환을 겹칠 수 있다.

triton에서는 다음과 같이 구현한다.

if EPILOGUE_SUBTILE:
    #EPI 1: acc를 N 방향의 왼쪽/오른쪽 절반으로 나눈다.
    acc_left, acc_right = acc.reshape(BLOCK_M, 2, BLOCK_N // 2).permute(0, 2, 1).split()
    # EPI 2: 두 조각을 FP16으로 바꿔 c_desc.store로 저장한다.
    c_desc.store([store_m, store_n], acc_left.to(tl.float16))
    c_desc.store([store_m, store_n + BLOCK_N // 2], acc_right.to(tl.float16))
else:
    c_desc.store([off_m, off_n], acc.to(tl.float16))

성능 측정 (B200 @modal)

그동안 내가 만들어온 triton kernel 성능

구현·설정 상태 F/W/E BM/BN/BK warps stages / P 시간 TFLOPS
torch . . . . 110.5 μs 1,244
Raw tiled · 기본값 —/0/0 128/64/32 4→4 3* / 전체 237.6 μs 579
Raw persistent · 기본값 0/0/0 128/64/32 4→4 4 / 2 265.4 μs 518
Swizzle tiled · 기본값 —/0/0 128/64/32 4→4 3* / 전체 239.8 μs 573
Swizzle persistent · 기본값 0/0/0 128/64/32 4→4 4 / 2 274.1 μs 501
Cross-tile · 기본값 1/0/0 128/64/32 4→4 4 / 2 271.2 μs 507
TMA tiled · 기존 튜닝값 —/0/0 128/256/64 4→4 4 / 전체 124.6 μs 1,103
TMA persistent · 재튜닝값 1/0/0 128/256/64 4→4 3 / 2 144.6 μs 950
WS off · 재튜닝값 1/0/0 128/256/64 4→4 3 / 2 144.9 μs 948
WS on · 재튜닝값 1/1/0 128/128/128 8→12 3 / 1 132.6 μs 1,037
Epilogue · WS off 1/0/1 128/256/64 4→4 4 / 1 122.1 μs 1,126
Epilogue · WS on 1/1/1 128/256/64 4→8 4 / 1 122.0 μs 1,127

공식 커널 성능

공식 구현 F/W/E warps P 시간 TFLOPS
torch . . . 111.1 μs 1,237
Raw tiled —/0/0 4→4 전체 132.3 μs 1,039
Raw persistent 1/0/0 4→4 1 138.5 μs 992
TMA tiled —/0/0 8→8 전체 126.2 μs 1,089
TMA tiled + WS —/1/0 4→8 전체 125.7 μs 1,093
TMA persistent + subtiling 1/0/1 8→8 1 123.0 μs 1,118
TMA persistent + WS + subtiling 1/1/1 4→8 1 120.2 μs 1,144

그래도 동일한 성능이라고 볼수 있을정도 까지는 되지 않았나? 싶어서 만족했다.

블록 크기를 더 늘리면,

커널 기존 128×256×64 큰 타일의 좋은 설정 시간 처리량 향상
TMA tiled 124.28 μs 256×256×64, stage 3 117.94 μs +5.4%
TMA persistent 141.24 μs 큰 타일은 자원 한도 초과 — —
WS 130.75 μs 256×256×32, stage 3 128.19 μs +2.0%
WS + subtiling 121.18 μs 256×256×32, stage 4 117.26 μs +3.3%

성능이 더 좋아짐을 알수 있었다. 아마 TMA box는 차원당 최대 256이기 때문에 더 크게 하기는 어려울 것이다.

triton matmul 커널 탐험을 우선 마무리하면서 다음 방향에 대한 고민....

triton matrix multiplication을 공부하면서 느끼고 있는것중에 가장 신기한 건 compiler가 너무 많은 부분을 해줘서 생각보다는 쉽다는 것이였다. flatten, warp_specialize 같은건 정말 공짜로 구현할 수 있는 너무 사기스럽고... 올해 초에 cuda 할떄는 이런 부분이 어렵게 느껴져서 결국 직접하지는 못했었는데 말이다. 그래서 지금은 컴파일러가 마법을 부리는 부분이 좀더 관심이 가고 알아가고 싶다.

#blog #gpu #matmul #tma #triton