본문으로 건너뛰기

[triton] Triton Blackwell 커널의 MXFP8 활성화 스케일 최적화: Zero Padding 전략

PR 링크: triton-lang/triton#11030 상태: Merged | 변경: +62 / -0

들어가며

최신 NVIDIA Blackwell 아키텍처에서 MXFP8 연산을 수행할 때, 타일 크기가 블록 크기의 배수가 아닌 경우(non-clean multiple) 성능 저하 문제가 발생합니다. 기존에는 소비자(consumer) 측에서 마스킹(masking)을 통해 불필요한 데이터를 처리했는데, 이 방식은 컴파일러가 스케일 값을 SMEM(Shared Memory)에서 TMEM(Tensor Memory)으로 유지하는 것을 방해하여 성능을 크게 저하시켰습니다. 본 PR은 이 문제를 해결하기 위해 생산자(producer) 측에서 패딩 영역에 0을 미리 채워 넣는 방식을 도입했습니다.

코드 분석

1. _matmul.py_p_matmul.py 수정

핵심 변경 사항은 tl.where를 사용하여 패딩 영역을 명시적으로 0으로 초기화하는 것입니다. Blackwell의 스케일 레이아웃은 블록이 인터리브(interleaved)되어 있으므로, 패딩 또한 할당 내부에 인터리브됩니다. 이를 통해 소비자 측에서 복잡한 마스킹 로직을 제거할 수 있게 되었습니다.

Before:

mask_n_scale = offs_y_n_scale < N_MX_BLOCK
scale_store_mask = mask_m[:, None] if OUT_N_TILE_ALIGNED else mask_m[:, None] & mask_n_scale[None, :]
tl.store(YActualScalePtrs, out_scale, mask=scale_store_mask)

After:

mask_n_scale = offs_y_n_scale < N_MX_BLOCK
if Y_MX_SCALE_LAYOUT == "BLACKWELL_ACT_SCALE" and not OUT_N_TILE_ALIGNED:
    out_scale = tl.where(mask_n_scale[None, :], out_scale, 0)
    mask_n_scale = offs_y_n_scale < tl.cdiv(N_MX_BLOCK, 4) * 4
scale_store_mask = mask_m[:, None] if OUT_N_TILE_ALIGNED else mask_m[:, None] & mask_n_scale[None, :]
tl.store(YActualScalePtrs, out_scale, mask=scale_store_mask)

2. 테스트 코드 추가

test_mxfp8_act_scale_store_zeroes_partial_group 테스트를 통해 패딩 영역이 실제로 0으로 채워지는지 검증합니다. 특히 BlackwellActMXScaleLayout을 사용하는 경우, 데이터가 정렬되지 않은 상황에서도 올바르게 0이 채워지는지 확인합니다.

왜 이게 좋은가

이 최적화의 핵심은 '소비자 측의 복잡도를 생산자 측으로 전이'시킨 점입니다.

  1. 컴파일러 최적화 효율 증대: 소비자 측에서 마스킹을 수행하면 컴파일러는 메모리 접근 패턴을 최적화하기 어렵습니다. 생산자 측에서 미리 0을 채워 넣음으로써, 소비자 커널은 마스킹 없이 연속적인 메모리 접근을 수행할 수 있어 SMEM/TMEM 활용도가 극대화됩니다.
  2. Blackwell 레이아웃 최적화: Blackwell의 인터리브된 스케일 레이아웃 특성을 활용하여, 패딩 영역을 0으로 채우는 것이 하드웨어 가속기 입장에서 훨씬 효율적인 데이터 정렬을 보장합니다.

일반적인 교훈으로, GPU 커널 최적화 시 '데이터를 읽을 때 마스킹하여 걸러내는 비용'보다 '데이터를 쓸 때 미리 정제해두는 비용'이 훨씬 저렴할 수 있음을 보여줍니다. 특히 고성능 연산 유닛(Tensor Core 등)으로 데이터를 넘기기 전, 메모리 레이아웃을 하드웨어 친화적으로 만드는 것이 성능의 핵심입니다.

참고 자료

⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.

댓글

관련 포스트

PR Analysis 의 다른글