본문으로 건너뛰기

[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)을 사용하여 원본 모델의 상태를 오염시키지 않도록 수정했습니다.

왜 이게 좋은가

  1. 성능 향상: Apple M-series 환경에서 EMA 업데이트 시간이 12.8ms에서 5.8ms로 약 2.2배 단축되었습니다. CUDA 환경에서는 커널 호출 오버헤드가 더 크기 때문에 훨씬 더 큰 폭의 성능 향상이 기대됩니다.
  2. 안정성: deepcopy를 통해 프로파일링 도구가 모델의 state_dict에 불필요한 float64 버퍼를 남기는 문제를 해결하여, 학습 도중 발생하는 타입 불일치 오류를 원천 차단했습니다.
  3. 호환성: TORCH_2_0TORCH_2_4 게이트를 통해 최신 기능을 활용하면서도, 구형 PyTorch 버전이나 MPS 환경에서의 호환성을 완벽하게 유지했습니다.

교훈: 파이썬 루프 내에서 텐서 연산을 반복하는 대신, PyTorch에서 제공하는 foreach 계열의 API를 활용하면 커널 호출 오버헤드를 최소화할 수 있습니다. 또한, 외부 라이브러리를 통한 프로파일링 시 원본 객체를 직접 수정하지 않도록 방어적인 코드를 작성하는 것이 중요합니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글