CODA: GEMM-Epilogue Fusion을 통한 Transformer 학습 최적화

Transformer 학습 시스템은 근본적으로 밀집 선형 대수, 특히 General Matrix Multiplication (GEMM)을 기반으로 구축됩니다. 그러나 학습 스택이 더욱 최적화됨에 따라, 엔드투엔드 실행 시간의 상당 부분이 정규화(normalization), 활성화(activations), 잔차 업데이트(residual updates), 그리고 리덕션(reductions)과 같은 "주변" 연산자에 의해 점점 더 많이 소비되고 있습니다. 이러한 메모리 대역폭 제한(memory-bound) 연산들은 산술 연산을 거의 수행하지 않으면서 거대한 중간 텐서를 글로벌 메모리로 반복적으로 이동시키며, 하드웨어의 전반적인 효율성을 제한하는 병목 현상을 생성합니다.

이를 해결하기 위해 연구자들은 이러한 계산을 "GEMM-plus-epilogue" 프로그램으로 표현하도록 설계된 GPU 커널 추상화인 CODA를 도입했습니다. Transformer 블록의 실행 방식을 재고함으로써, CODA는 불필요한 데이터 이동을 제거하고 온칩(on-chip) 활용도를 극대화하는 것을 목표로 합니다.

문제점: 메모리 대역폭 제한 병목 현상

표준 Transformer 블록에서 주요 작업은 GEMM 연산에 의해 수행됩니다. 그러나 이러한 연산 뒤에는 일반적으로 일련의 요소별(element-wise) 또는 리덕션 연산이 뒤따릅니다. 전통적인 프레임워크 수준의 실행에서는 각 연산자가 별도의 커널로 실행됩니다. 이는 GEMM의 출력이 글로벌 메모리에 기록된 후, 즉시 다음 커널(예: LayerNorm 또는 ReLU 활성화)에 의해 다시 읽혀야 함을 의미합니다.

글로벌 메모리에 쓰고 읽는 이 사이클은 비용이 많이 듭니다. 연산 능력이 증가함에 따라 산술 처리량과 메모리 대역폭 사이의 격차가 벌어지며, 이로 인해 이러한 "메모리 대역폭 제한" 연산자가 전체 학습 시간의 무시할 수 없는 부분을 차지하게 됩니다.

CODA 접근 방식: GEMM-Epilogue Fusion

CODA는 많은 Transformer 연산자가 GEMM 출력이 글로벌 메모리에 기록되기 전, 즉 GEMM 출력 타일이 여전히 칩 내부(공유 메모리 또는 레지스터)에 머물러 있는 동안 실행되도록 대수적으로 재매개변수화될 수 있다는 관찰에 기반합니다.

작동 방식

CODA는 행렬 곱셈의 고도로 최적화된 핵심인 GEMM 메인루프(mainloop)를 고정하고, 조합 가능한 소수의 에필로그(epilogue) 프리미티브를 노출합니다. 이러한 프리미티브를 통해 개발자는 다음과 같은 연산을 정의할 수 있습니다:

  • Scaling: 출력을 상수 또는 벡터로 곱함.
  • Reductions: 특정 차원에 대해 합계 또는 평균을 수행함.
  • Pairwise Transformations: 활성화 함수 또는 기타 요소별 함수를 적용함.
  • Accumulation: 결과를 잔차 연결(residual connection)에 더함.

이러한 프리미티브로 인터페이스를 제한함으로써, CODA는 전문가가 작성한 GEMM의 성능 특성을 유지하면서도, 표준 Transformer 블록의 순방향 및 역방향 패스 모두에서 어텐션(attention)을 제외한 거의 모든 계산을 처리할 수 있는 충분한 표현력을 제공합니다.

LLM 기반 코드 생성(Codegen)에 미치는 영향

이 연구의 가장 중요한 발견 중 하나는 사람이 작성한 CODA 커널과 LLM이 작성한 CODA 커널 모두 높은 성능을 달성한다는 것입니다. 이는 저수준 하드웨어 최적화에 접근하는 방식의 변화를를 시사합니다.

LLM은 레지스터 압박(register pressure) 관리나 복잡한 메모리 타일링(memory tiling)과 같은 저수준 하드웨어 최적화의 복잡함에는 종종 어려움을 겪지만, 고수준 조합(composition)에는 탁월합니다. 제한적이고 조합 가능한 API를 제공함으로써, CODA는 하드웨어의 복잡성을 효과적으로 "샌드박스화"합니다. LLM은 GEMM 메인루프를 작성할 필요가 없으며, 전문가가 작성한 에필로그 블록들을 결합하기만 하면 됩니다.

Hacker News의 커뮤니티 관찰자들이 언급한 바와 같이:

"제한적이고 조합 가능한 API를 가진 컴파일러 추상화를 설계하여 LLM이 전문가가 작성한 블록들을 쉽게 결합할 수 있게 만드는 것은 영리한 스마트한 선택입니다. 에이전트 기반 개발로 이동함에 따라 이것이 결국 코드 생성의 표준이 될 것이라고 생각합니다."

결론

CODA는 프레임워크 수준의 생산성과 하드웨어 수준의 효율성을 결합하는 실용적인 경로를를 나타냅니다. Transformer의 메모리 대역폭 제한 연산자를 GEMM의 에필로그로 취급함으로써, CODA는 글로벌 메모리 왕복 횟수를 크게 줄이고 차세대 대규모 모델의 학습을 최적화하는 방식을 혁리신합니다.

Sources