본문으로 건너뛰기

[sglang] [MoE] SwiGLU 퓨전: Triton 커널 최적화로 메모리 대역폭 한계 돌파하기

PR 링크: sgl-project/sglang#32944 상태: Merged | 변경: +409 / -13

들어가며

LLM 추론, 특히 Mixture of Experts(MoE) 모델의 성능을 최적화할 때 가장 큰 적 중 하나는 '메모리 대역폭(Memory Bandwidth)'입니다. 기존의 MoE 레이어 실행 방식은 다음과 같은 세 단계로 이루어집니다:

  1. Up-GEMM: 입력 데이터에 가중치를 곱해 중간 버퍼(intermediate_cache1)에 저장.
  2. Activation (SwiGLU): 중간 버퍼를 다시 읽어 Silu 및 곱셈 연산을 수행 후 다른 버퍼(intermediate_cache2)에 저장.
  3. Down-GEMM: 활성화된 데이터를 다시 읽어 최종 출력 계산.

여기서 두 번째 단계인 활성화 함수는 순수하게 데이터 이동(Data Movement)이 지배적인 연산입니다. Up-GEMM이 이미 레지스터에 들고 있던 데이터를 굳이 HBM(High Bandwidth Memory)에 썼다가 다시 읽어오는 'Round-trip'이 발생하기 때문입니다. 특히 Batch Size가 작은 Decode 단계에서는 이러한 커널 런치 오버헤드와 메모리 트래픽이 전체 지연 시간(Latency)의 상당 부분을 차지합니다.

이번 PR은 가중치 레이아웃을 로드 타임에 미리 변경(Permute)하여, Up-GEMM의 에필로그(Epilogue) 단계에서 활성화 함수를 인-레지스터(In-register)로 처리하고 중간 버퍼를 제거하는 최적화를 구현했습니다.

코드 분석: 핵심 변경 사항

1. Triton 커널의 에필로그 퓨전 (fused_moe_triton_kernels.py)

가장 핵심적인 변화는 fused_moe_kernel 내부에 SwiGLU 연산을 직접 삽입한 것입니다. 기존에는 GEMM 결과인 accumulator를 그대로 저장했지만, 이제는 레지스터 내에서 gateup 값을 분리하여 연산합니다.

[Before & After: Kernel Epilogue]

# Before: 단순히 결과를 메모리에 저장
# (기존 코드의 일반적인 흐름)
# tl.store(c_ptrs, accumulator, mask=c_mask)

# After: FUSE_SWIGLU가 활성화된 경우
if FUSE_SWIGLU:
    # 가중치가 인터리빙되어 있어 gate/up 쌍이 인접한 컬럼에 위치함
    acc_pairs = tl.reshape(accumulator, (BLOCK_SIZE_M, BLOCK_SIZE_N // 2, 2))
    gate_b, up_b = tl.split(acc_pairs)
    gate_f = gate_b.to(tl.float32)

    # Bit-parity를 맞추기 위한 인라인 어셈블리 사용
    exp_neg = tl.inline_asm_elementwise(
        "{ mul.ftz.f32 $0, $1, 0fBFB8AA3B; ex2.approx.ftz.f32 $0, $0; }",
        "=f,f", [gate_f], dtype=tl.float32, is_pure=True, pack=1,
    )
    silu_f = tl.inline_asm_elementwise(
        "div.approx.ftz.f32 $0, $1, $2;",
        "=f,f,f", [gate_f, 1.0 + exp_neg], dtype=tl.float32, is_pure=True, pack=1,
    )
    out_act = (silu_f * up_b.to(tl.float32)).to(compute_type)
    
    # 결과값은 원래 너비(N)의 절반만 저장 (intermediate_cache1 제거)
    offs_half = pid_n * (BLOCK_SIZE_N // 2) + tl.arange(0, BLOCK_SIZE_N // 2)
    # ... (생략) ...
    tl.store(c_ptrs, out_act, mask=c_mask)

여기서 주목할 점은 tl.inline_asm_elementwise를 사용했다는 것입니다. Triton의 기본 연산자는 IEEE 표준을 따르지만, 기존 PyTorch 커널은 성능을 위해 fast-math 인트린직을 사용합니다. 비트 단위의 정확도(Bit-parity)를 맞추기 위해 직접 PTX 어셈블리를 호출하여 ex2.approxdiv.approx를 구현한 점이 인상적입니다.

2. 시퀀스 제어 및 버퍼 제거 (fused_moe.py)

커널이 합쳐졌으므로, 상위 레벨의 실행 로직에서도 변화가 필요합니다. 중간 단계의 캐시 할당을 건너뛰고 커널 호출 파라미터를 조정합니다.

[Before & After: Runner Sequence]

# Before: 중간 버퍼 할당 및 별도 활성화 커널 실행
intermediate_cache1 = torch.empty((total_tokens, N), ...)
# invoke_up_gemm(...)
# silu_and_mul(intermediate_cache1, intermediate_cache2)

# After: 퓨전 모드일 때 버퍼 할당 생략 및 플래그 전달
if fuse_swiglu_interleaved:
    # ... (검증 로직) ...
    # intermediate_cache1 할당 및 별도 활성화 런치를 건너뜀
    # GEMM1이 직접 최종 활성화 결과를 N/2 너비로 작성함

이 변경을 통해 GEMM1 -> Activation -> GEMM2로 이어지는 의존성 체인에서 중간 단계가 완전히 사라지고 GEMM1 -> GEMM2로 단축되었습니다.

3. 가중치 레이아웃 변경 (unquant.pyenviron.py)

이 최적화가 가능하려면 gate 가중치와 up 가중치가 같은 타일(Tile) 안에 들어와야 합니다. 원래 체크포인트에서는 [W1, W3] 형태로 멀리 떨어져 저장되어 있지만, 로드 타임에 이를 [gate0, up0, gate1, up1, ...] 형태로 인터리빙(Interleaving)합니다.

# environ.py에 추가된 옵션
SGLANG_OPT_FUSE_SWIGLU_INTERLEAVED = EnvBool(False)

이 옵션은 기본적으로 꺼져 있는데, 그 이유는 가중치 레이아웃을 영구적으로 바꾸기 때문에 LoRA나 EPLB(Expert Parallel Load Balancing)와 같이 가중치 구조에 의존하는 다른 기능들과 충돌할 수 있기 때문입니다.

왜 이게 좋은가?

1. 성능 향상 (TPOT 개선)

벤치마크 결과에 따르면, Kimi-Linear-48B 모델 기준 TPOT(Time Per Output Token)이 약 1.2% 개선되었습니다. 수치상으로는 작아 보일 수 있지만, 이는 모델의 가중치를 건드리지 않고 순수하게 커널 실행 구조만 개선하여 얻은 값진 결과입니다. 특히 연산량보다 메모리 대역폭이 병목인 소규모 배치 추론에서 효과가 극대화됩니다.

2. 메모리 트래픽 감소

intermediate_cache1은 전체 너비 N을 가집니다. 이를 HBM에 쓰고 다시 읽는 과정을 생략함으로써, 레이어당 수십 MB에 달하는 메모리 대역폭 낭비를 막았습니다.

3. 엄격한 정확도 보장

단순히 tl.sigmoid 등을 쓰지 않고 인라인 어셈블리를 통해 기존 커널과 Bit-identical(바이트 단위 일치)한 결과를 냈습니다. 이는 딥러닝 모델에서 미세한 오차가 누적되어 성능 저하를 일으키는 문제를 원천 차단합니다.

결론

이번 PR은 현대적인 GPU 가속기에서 연산(Compute)보다 메모리(Memory)와 지연 시간(Latency)이 얼마나 중요한지를 잘 보여줍니다. 가중치 레이아웃을 미리 변경하는 '정적 최적화'와 Triton의 '커널 퓨전'을 결합하여 실질적인 성능 이득을 이끌어냈습니다.

SGLang과 같은 고성능 서빙 엔진에서 이러한 로우 레벨 최적화는 대규모 트래픽 처리 시 큰 비용 절감으로 이어집니다. 다만, 가중치 레이아웃 변경에 따른 호환성 이슈를 해결하기 위해 옵션화(Opt-in)한 설계 판단 역시 시니어 엔지니어다운 신중함이 돋보이는 부분입니다.

References

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글