본문으로 건너뛰기

[ultralytics] MuSGD 최적화: Batched Newton-Schulz와 Fused Kernel로 8배 성능 향상

PR 링크: ultralytics/ultralytics#25288 상태: Merged | 변경: +101 / -97

들어가며

딥러닝 모델 학습에서 옵티마이저의 연산 속도는 전체 학습 처리량(throughput)에 큰 영향을 미칩니다. 특히 Muon과 같은 고성능 옵티마이저는 Newton-Schulz 반복법을 통해 가중치를 직교화(orthogonalization)하는데, 기존 구현 방식은 모델의 각 파라미터마다 개별적으로 연산을 수행했습니다. 이는 수천 개의 작은 CUDA 커널을 실행하게 만들어, 연산 자체의 복잡도보다 커널 실행 오버헤드(launch bound)가 병목이 되는 문제를 야기했습니다. 본 PR은 이 문제를 해결하기 위해 연산을 배치화하고 PyTorch의 _foreach API를 활용하여 성능을 8배 이상 개선했습니다.

코드 분석

1. ultralytics/optim/muon.py: Batched Newton-Schulz

기존에는 단일 텐서에 대해서만 Newton-Schulz를 수행했으나, 이제는 3D 배치를 지원하도록 변경되었습니다.

Before:

for a, b, c in [...]:
    A = X @ X.T
    B = b * A + c * A @ A
    X = a * X + B @ X

After:

for _ in range(5):
    A = X @ X.transpose(-2, -1)
    B = torch.baddbmm(A, A, A, beta=b, alpha=c)
    X = torch.baddbmm(X, B, X, beta=a)

torch.baddbmm을 사용하여 여러 행렬 연산을 하나의 커널로 융합(fuse)했습니다. 또한, 행렬 크기가 다른 경우 제로 패딩(zero-padding)을 통해 동일한 배치 그룹으로 묶어 처리함으로써 커널 실행 횟수를 획기적으로 줄였습니다.

2. ultralytics/optim/muon.py: Fused Foreach Ops

모멘텀과 SGD 연산 시 파라미터별 루프를 제거하고 torch._foreach_* 계열 함수를 사용했습니다.

Before:

for p in group["params"]:
    state["momentum_buffer"].mul_(group["momentum"]).add_(grad)

After:

torch._foreach_mul_(momentums, beta)
torch._foreach_add_(momentums, grads, alpha=1 - beta)

_foreach 함수들은 리스트 형태의 텐서를 입력받아 내부적으로 최적화된 단일 커널에서 연산을 수행하므로, 파이썬 루프 오버헤드와 커널 실행 횟수를 동시에 최소화합니다.

왜 이게 좋은가

이번 최적화의 핵심은 'Launch Bound'에서 'Compute Bound'로의 전환입니다.

  • 성능 수치: RTX PRO 6000 환경에서 yolo26n 모델 기준, 기존 62.2ms/step이던 MuSGD 연산 시간이 7.7ms/step으로 약 8배 단축되었습니다. 실제 학습 환경(COCO 데이터셋)에서도 1.33배의 전체 학습 속도 향상을 확인했습니다.
  • 교훈:
    1. Kernel Fusion: 여러 개의 작은 연산을 하나로 묶는 것은 GPU 활용도를 높이는 가장 효과적인 방법입니다.
    2. Batching: 데이터가 작더라도 여러 개를 모아서 배치로 처리하면 커널 실행 오버헤드를 극적으로 줄일 수 있습니다.
    3. Foreach API: PyTorch에서 제공하는 _foreach 계열 함수는 리스트 단위 연산 시 매우 강력한 성능을 발휘하므로, 옵티마이저 구현 시 적극 고려해야 합니다.

리뷰 과정에서 Glenn Jocher는 "Net negative(코드 라인 수 감소)"를 달성한 점을 높이 평가했습니다. 이는 복잡한 로직을 통합하고 중복을 제거함으로써 성능과 유지보수성을 동시에 잡은 사례입니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글