본문으로 건너뛰기
PyTorch조회 1

TorchInductor를 위한 새로운 CuteDSL 백엔드: Blackwell GPU에서 GEMM 성능 극대화

TorchInductor에 NVIDIA CuteDSL 백엔드가 추가되어 Blackwell GPU 환경의 행렬 곱셈 연산 속도와 컴파일 효율이 대폭 향상됐다.

섹션별 상세

01
기존 CUTLASS C++ 백엔드는 커널 변형마다 nvcc 컴파일이 필요하여 오토튜닝 과정에서 병목 현상이 발생했다. CuteDSL은 파이썬-MLIR 컴파일러를 사용하여 다른 백엔드와 대등한 수준으로 컴파일 속도를 개선하고 하드웨어 계층 구조를 직접 제어할 수 있게 한다. 이를 통해 런타임에 수백 개의 커널 후보를 신속하게 평가하고 최적의 설정을 선택할 수 있는 환경이 구축됐다.
다양한 컴파일러 백엔드 간의 단일 커널 컴파일 시간 비교 표
ChartCuteDSL의 컴파일 시간은 0.29초로, nvcc를 사용하는 CUTLASS 4.x(14.41초) 대비 약 50배 빠름을 보여준다. 이는 런타임 오토튜닝 시 CuteDSL이 훨씬 효율적임을 입증하는 핵심 근거이다.
02
Triton은 메모리 대역폭 제한 작업에 강점이 있지만, 연산 집약적인 GEMM에서는 하드웨어 특화 기능 활용이 제한적이다. CuteDSL 백엔드는 Blackwell(B200)의 텐서 코어 파이프라인, 공유 메모리 스테이징, 분산 공유 메모리 등을 정교하게 제어하는 데 집중한다. 특히 LLM 추론의 디코드 단계와 같이 특정 형상의 행렬 연산에서 하드웨어 활용도를 극대화하는 전략을 취한다.
03
TorchInductor는 cutlass_api를 통해 호환 가능한 모든 커널 구성을 조회하고 nvMatmulHeuristics 분석 모델로 후보군을 압축한다. 수백 개의 후보 중 하드웨어 처리량 예측 점수가 높은 상위 5개 내외의 커널만 실제 벤치마킹 대상으로 선정하여 오토튜닝 시간을 최적화한다. 최종 선택된 커널은 메모리에 캐싱되어 이후 동일한 연산 요청 시 컴파일 오버헤드 없이 즉시 실행된다.
04
NVIDIA B200 GPU 기반 테스트에서 BF16 디코드 형상 연산 속도가 기존 대비 최대 1.73배 향상되는 결과가 확인됐다. MXFP8 및 NVFP4와 같은 최신 데이터 형식에서도 중간 크기 행렬 연산에서 최대 1.78배의 성능 이득을 얻었다. 이는 Blackwell 하드웨어의 잠재력을 소프트웨어 수준에서 효과적으로 끌어내고 있음을 증명하는 수치이다.
B200 GPU에서 BF16 행렬 곱셈 커널의 처리량(TFLOPS) 비교 차트
Chart다양한 행렬 크기(M, N, K)에 대해 Inductor ATen, Triton, NVGEMM(CuteDSL)의 성능을 비교한다. 특히 작은 M 값(디코드 단계)에서 NVGEMM이 기존 백엔드 대비 최대 1.73배 높은 성능을 보임을 시각화한다.
05
vLLM V1 엔진을 활용한 엔드투엔드 추론 테스트에서 Llama 3.3 70B 모델의 지연 시간이 배치 사이즈 16 기준 6.5% 감소했다. Llama 3.1 8B 모델 또한 배치 사이즈 전반에서 2~4%의 일관된 성능 향상을 보였다. 동적 형상을 사용하는 추론 환경에서도 autotune_batch_hint를 통해 런타임 형상에 최적화된 커널을 선택함으로써 실질적인 서비스 성능 개선이 가능하다.
python
import torch
import torch._inductor.config as config

# NVGEMM 백엔드 활성화
config.max_autotune_gemm_backends = "ATEN,TRITON,NVGEMM"

A = torch.randn(128, 4096, device="cuda", dtype=torch.bfloat16)
B = torch.randn(4096, 4096, device="cuda", dtype=torch.bfloat16)

@torch.compile(mode="max-autotune-no-cudagraphs")
def f(a, b):
    return a @ b

out = f(A, B)  # 첫 호출 시 오토튜닝 트리거

TorchInductor에서 CuteDSL(NVGEMM) 백엔드를 활성화하고 행렬 곱셈을 실행하는 예시 코드

vLLM 엔진을 사용한 BF16 데이터 형식의 모델별 추론 속도 향상률 차트
ChartLlama 3.1, Qwen3, Llama 3.3 모델에 대해 배치 사이즈별 성능 향상을 보여준다. Llama 3.3 70B 모델이 배치 사이즈 16에서 6.5%로 가장 높은 성능 향상을 기록했음을 확인할 수 있다.

용어 해설

일반 행렬 곱셈(GEMM)
General Matrix Multiply의 약자로, 딥러닝 모델 연산의 대부분을 차지하는 핵심 연산이다. 하드웨어 가속기의 성능을 최대한 끌어내기 위해 타일링, 워프 스케줄링 등 정교한 최적화가 필수적이다.
오토튜닝(Autotuning)
특정 하드웨어와 데이터 형상에 최적화된 커널 설정을 런타임에 자동으로 탐색하고 선택하는 기법이다. 타일 크기나 워프 구성 등 다양한 후보군을 벤치마킹하여 가장 빠른 결과물을 도출한다.
에필로그 퓨전(Epilogue Fusion)
행렬 곱셈 연산 직후에 이어지는 활성화 함수나 덧셈 연산을 별도의 커널 호출 없이 하나의 커널 안에서 처리하는 최적화 기법이다. 메모리 대역폭 낭비를 줄여 전체 추론 속도를 향상시킨다.
MXFP8
NVIDIA Blackwell 아키텍처에서 도입된 8비트 부동소수점 데이터 형식이다. 기존 FP8보다 정밀도와 효율성을 개선하여 대규모 언어 모델의 추론 및 학습 속도를 획기적으로 높인다.
적시 컴파일(JIT Compilation)
프로그램 실행 시점에 코드를 기계어로 컴파일하는 방식이다. TorchInductor는 JIT 방식을 통해 실제 입력 데이터의 크기와 타입을 확인하고 그에 최적화된 GPU 커널을 생성한다.

코드 예제

bash
pip install nvidia-cutlass-dsl==4.3.5
pip install nvidia-matmul-heuristics
# Clone and install cutlass_api from the cutlass_api branch
git clone --branch cutlass_api https://github.com/NVIDIA/cutlass.git
cd cutlass/python/cutlass_api
pip install -e ".[torch]"

CuteDSL 백엔드 사용을 위한 필수 라이브러리 및 API 설치 명령어

기술

  • TorchInductor
  • CuteDSL
  • Triton
  • CUTLASS
  • vLLM
  • B200

활용 사례

  • LLM 추론 가속화
  • MXFP8/NVFP4 기반 저정밀도 추론 최적화
  • 실시간 GPU 커널 오토튜닝
AI 분석 전체 내용 보기

AI 요약 · 북마크 · 개인 피드 설정 — 무료

출처 · 인용 안내

원문 발행 2026. 04. 07.수집 2026. 04. 07.출처 타입 RSS

인용 시 "요약 출처: AI Trends (aitrends.kr)"를 표기하고, 사실 확인은 원문 보기 기준으로 진행해 주세요. 자세한 기준은 운영 정책을 참고해 주세요.