최근 'Show HN'을 통해 JAX XLA 기반의 새로운 아키텍처 'jax-softmax-bypass'가 공개되며 대규모 언어모델(LLM)의 추론(inference) 병목 현상을 해결할 혁신적인 접근법을 제시했습니다. 이 프로젝트는 트랜스포머(Transformer) 모델의 핵심 연산인 소프트맥스(Softmax)에서 발생하는 $e^x$ (초월 지수 함수) 계산의 비효율성을 우회하는 것을 목표로 합니다. 이는 기존 소프트맥스 연산이 전역적인 행 단위 합산(global row-wise reduction)을 요구하여 메모리 버스(memory bus)를 점유하고 고대역폭 메모리(HBM) 병목을 유발하는 문제를 해결하기 위함입니다.
'jax-softmax-bypass'는 활성화 궤적을 단일 패스 2차 테일러 다항식(Taylor polynomial) FMA(Fused Multiply-Add) 대수 평면으로 인수분해하여, 4개의 브랜치리스(branchless) 통합 가속 엔진을 배포합니다. 이를 통해 수치적 발산(numerical divergence)을 물리적 경계 조건 내에 엄격히 제한합니다. 특히 메타(Meta)의 LLaMA(LLaMA-2, LLaMA-3, LLaMA-3.1) 및 구글(Google)의 Gemma(Gemma, Gemma-2) 모델 제품군에 최적화되어 있으며, SPMD(Single Program, Multiple Data) 텐서 병렬 파티션 제약을 통해 대규모 가속기 클러스터 간 통신 지연을 억제하고, SiLU(Swish) 활성화 병목을 우회하는 등 특정 모델의 구조적 특성을 활용합니다. 이 아키텍처는 하드웨어 수정 없이 순수하게 수학적 장치를 변경하여 컴퓨팅 밀도와 하드웨어 처리량 효율성을 극대화합니다.
이 기술은 LLM 추론 속도와 효율성을 획기적으로 개선할 잠재력을 가지고 있습니다. 기존 소프트맥스 연산으로 인한 HBM 병목과 유휴 사이클 낭비 문제를 해결함으로써, 특히 장문 맥락(long-context) 워크로드에서 컴퓨팅 밀도와 처리량을 높일 수 있습니다. 이는 대규모 AI 모델을 더 빠르고 저렴하게 운영할 수 있게 하여, AI 서비스 제공업체와 연구자들에게 큰 이점을 제공할 것입니다. 다만, MoE(Mixture-of-Experts) 아키텍처(예: Mixtral, DeepSeek)는 동적 토큰 라우팅(dynamic token routing)으로 인해 정적 컴파일러 추적을 방해하므로, 이 프레임워크의 주요 가속 엔진 범위에는 포함되지 않습니다.