[ultralytics] PyTorch EMA 업데이트 최적화: _foreach_lerp_를 활용한 성능 개선
PR 링크: ultralytics/ultralytics#25315 상태: Merged | 변경: +10 / -4
들어가며
딥러닝 모델 학습 과정에서 Exponential Moving Average(EMA)는 모델의 가중치를 안정적으로 유지하는 데 필수적입니다. 하지만 기존의 ModelEMA.update() 구현은 모델의 모든 파라미터를 순회하며 개별적으로 연산을 수행했습니다. 이는 파라미터 수가 많은 모델에서 수천 번의 커널 호출(kernel launch)을 유발하여 학습 속도를 저하시키는 병목 현상을 초래했습니다. 본 글에서는 ultralytics 레포지토리에서 진행된 EMA 업데이트 최적화와 그 과정에서 발견된 프로파일링 버그 수정 사례를 분석합니다.
코드 분석
1. EMA 업데이트 최적화 (ultralytics/utils/torch_utils.py)
기존 코드는 파라미터마다 mul_과 add_를 반복 호출했습니다. 이를 torch._foreach_lerp_를 활용한 배치 연산으로 변경하여 커널 호출 횟수를 획기적으로 줄였습니다.
Before:
for k, v in self.ema.state_dict().items():
if v.dtype.is_floating_point:
v *= d
v += (1 - d) * msd[k].detach()
After:
ema_v, model_v = [], []
for k, v in self.ema.state_dict().items():
if v.dtype.is_floating_point:
ema_v.append(v)
model_v.append(msd[k])
if ema_v and TORCH_2_0 and (TORCH_2_4 or ema_v[0].device.type != "mps"):
torch._foreach_lerp_(ema_v, model_v, 1 - d)
else:
for v, m in zip(ema_v, model_v):
v.mul_(d).add_(m, alpha=1 - d)
이 변경은 torch._foreach_lerp_가 지원되는 환경(PyTorch 2.0+)에서 배치 처리를 수행하며, MPS(Apple Silicon) 환경의 제약을 고려하여 버전별로 분기 처리(fallback)를 구현했습니다.
2. 프로파일링 버그 수정 (ultralytics/nn/tasks.py)
기존에는 thop.profile이 모델 레이어에 직접 접근하여 float64 버퍼를 남기는 문제가 있었습니다. 이는 EMA 업데이트 시 데이터 타입 불일치 오류를 유발했습니다.
Before:
flops = thop.profile(m, inputs=[x.copy() if c else x], verbose=False)[0] / 1e9 * 2 if thop else 0
After:
flops = thop.profile(deepcopy(m), inputs=[x.copy() if c else x], verbose=False)[0] / 1e9 * 2 if thop else 0
deepcopy(m)을 사용하여 원본 모델의 상태를 오염시키지 않도록 수정했습니다.
왜 이게 좋은가
- 성능 향상: Apple M-series 환경에서 EMA 업데이트 시간이 12.8ms에서 5.8ms로 약 2.2배 단축되었습니다. CUDA 환경에서는 커널 호출 오버헤드가 더 크기 때문에 훨씬 더 큰 폭의 성능 향상이 기대됩니다.
- 안정성:
deepcopy를 통해 프로파일링 도구가 모델의state_dict에 불필요한float64버퍼를 남기는 문제를 해결하여, 학습 도중 발생하는 타입 불일치 오류를 원천 차단했습니다. - 호환성:
TORCH_2_0및TORCH_2_4게이트를 통해 최신 기능을 활용하면서도, 구형 PyTorch 버전이나 MPS 환경에서의 호환성을 완벽하게 유지했습니다.
교훈: 파이썬 루프 내에서 텐서 연산을 반복하는 대신, PyTorch에서 제공하는 foreach 계열의 API를 활용하면 커널 호출 오버헤드를 최소화할 수 있습니다. 또한, 외부 라이브러리를 통한 프로파일링 시 원본 객체를 직접 수정하지 않도록 방어적인 코드를 작성하는 것이 중요합니다.
참고 자료
- https://pytorch.org/docs/stable/generated/torch.lerp.html
- https://pytorch.org/docs/stable/generated/torch.Tensor.mul_.html
- https://pytorch.org/docs/stable/generated/torch.Tensor.add_.html
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [ultralytics] MuSGD 최적화: Batched Newton-Schulz와 Fused Kernel로 8배 성능 향상
- [vllm] vLLM Qwen3.5 GDN 최적화: `einops.rearrange`를 `torch.flatten`으로 교체하여 20배 성능 향상!
- [sglang] ERNIE-Image의 RoPE와 GELU-mul 융합 및 RoPE cos/sin 호이스팅을 통한 성능 최적화
- [transformers] Hugging Face Transformers: NoRepeatNGramLogitsProcessor 벡터화 및 성능 최적화
- [ultralytics] Ultralytics FLOPs 프로파일링 최적화: deepcopy 제거를 통한 성능 향상
PR Analysis 의 다른글
- 이전글 [sglang] SGLang ReplaySSM: GDN 추론 최적화 및 메모리 효율 개선
- 현재글 : [ultralytics] PyTorch EMA 업데이트 최적화: _foreach_lerp_를 활용한 성능 개선
- 다음글 [sglang] [SGLang] MoE Prefill의 혁신: DWDP(Distributed Weight Data Parallelism) 도입 분석
댓글