[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배의 전체 학습 속도 향상을 확인했습니다. - 교훈:
- Kernel Fusion: 여러 개의 작은 연산을 하나로 묶는 것은 GPU 활용도를 높이는 가장 효과적인 방법입니다.
- Batching: 데이터가 작더라도 여러 개를 모아서 배치로 처리하면 커널 실행 오버헤드를 극적으로 줄일 수 있습니다.
- Foreach API: PyTorch에서 제공하는
_foreach계열 함수는 리스트 단위 연산 시 매우 강력한 성능을 발휘하므로, 옵티마이저 구현 시 적극 고려해야 합니다.
리뷰 과정에서 Glenn Jocher는 "Net negative(코드 라인 수 감소)"를 달성한 점을 높이 평가했습니다. 이는 복잡한 로직을 통합하고 중복을 제거함으로써 성능과 유지보수성을 동시에 잡은 사례입니다.
참고 자료
- https://pytorch.org/docs/stable/generated/torch.baddbmm.html
- https://pytorch.org/docs/stable/torch.compiler_jit.html
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [vllm] vLLM Qwen3.5 GDN 최적화: `einops.rearrange`를 `torch.flatten`으로 교체하여 20배 성능 향상!
- [sglang] CUDA 그래프 호환성을 위한 LoRA 연산 최적화: 스칼라 할당 대신 슬라이스 제로화 사용
- [sglang] [SGLang] VLM 추론 성능의 비약적 향상: Cross-request ViT Batching과 Metadata 재사용 기법
- [vllm] vLLM 성능 최적화: token_to_req_indices 캐싱을 통한 6배 성능 향상
- [sglang] SGLang LTX-2.3 Diffusion 모델 최적화: Residual-Gate 연산 CUDA Fast Path 도입
PR Analysis 의 다른글
- 이전글 [openclaw] 프론트엔드 성능 최적화: 엔트리 CSS 번들 다이어트와 코드 스플리팅 전략
- 현재글 : [ultralytics] MuSGD 최적화: Batched Newton-Schulz와 Fused Kernel로 8배 성능 향상
- 다음글 [vllm] vLLM의 대규모 모델 추론을 위한 MRV2 기반 Prefill Context Parallelism(PCP) 도입
댓글