본문으로 건너뛰기

[triton] Triton JIT의 메모리 누수 방지: Reference Cycle 제거를 통한 성능 최적화

PR 링크: triton-lang/triton#11093 상태: Merged | 변경: +14 / -13

들어가며

Python에서 함수 내부의 클로저(Closure)는 편리하지만, 의도치 않은 메모리 누수를 유발하는 주범이 되기도 합니다. 특히 재귀적으로 호출되는 클로저가 외부 스코프의 객체를 참조할 경우, Python의 가비지 컬렉터(GC)가 이를 즉시 회수하지 못하는 'Reference Cycle'이 발생할 수 있습니다. 최근 Triton 레포지토리에서는 triton.runtime.jit.compute_cache_key 함수 내에서 발생하던 이러한 문제를 해결하기 위해 재귀적 클로저를 전역 함수로 분리하는 리팩토링이 진행되었습니다. 이 글에서는 해당 PR이 왜 필요한지, 그리고 어떻게 개선되었는지 살펴봅니다.

코드 분석

python/triton/runtime/jit.py 리팩토링

기존 코드에서는 compute_cache_key 함수 내부에서 replace_callables라는 중첩 함수를 정의하여 재귀적으로 객체를 순회했습니다. 이 방식은 compute_cache_key가 호출될 때마다 새로운 클로저 객체가 생성되며, 이 과정에서 불필요한 참조 사이클이 형성될 위험이 있었습니다.

Before: 클로저를 사용한 재귀 호출

def compute_cache_key(kernel_key_cache, specialization, options):
    # ... (생략)
    def replace_callables(obj):
        if isinstance(obj, list):
            return [replace_callables(arg) for arg in obj]
        # ... (중략)
        elif isinstance(obj, JITCallable):
            return obj.cache_key
        return obj

    cache_key = str(replace_callables(specialization)) + str(options)
    # ... (생략)

After: 전역 함수로 분리

def _replace_jit_callables(obj):
    if isinstance(obj, list):
        return [_replace_jit_callables(arg) for arg in obj]
    elif is_namedtuple(obj):
        results = [_replace_jit_callables(arg) for arg in obj]
        return obj.__class__(*results)
    # ... (중략)
    elif isinstance(obj, JITCallable):
        return obj.cache_key
    return obj

def compute_cache_key(kernel_key_cache, specialization, options):
    # ... (생략)
    cache_key = str(_replace_jit_callables(specialization)) + str(options)
    # ... (생략)

핵심 변경 사항은 replace_callables를 모듈 레벨의 _replace_jit_callables 함수로 추출한 것입니다. 이제 더 이상 compute_cache_key 호출 시마다 클로저 객체가 생성되지 않으며, Python의 메모리 관리 효율성이 향상되었습니다.

왜 이게 좋은가

  1. Reference Cycle 방지: 클로저가 외부 스코프를 캡처하면 순환 참조가 발생하기 쉽습니다. 이를 전역 함수로 분리함으로써 GC의 부담을 줄이고 메모리 누수 가능성을 원천 차단했습니다.
  2. 성능 향상: 매번 함수 객체를 새로 생성하는 오버헤드를 제거했습니다. 이는 반복적으로 호출되는 JIT 컴파일 경로에서 미세하지만 유의미한 성능 이득을 가져옵니다.
  3. 코드 가독성 및 테스트 용이성: 로직이 분리됨에 따라 _replace_jit_callables만 별도로 유닛 테스트를 수행하기가 훨씬 수월해졌습니다.

이러한 패턴은 PyTorch 등 대규모 Python 프로젝트에서도 자주 사용되는 최적화 기법입니다. 특히 고성능이 요구되는 라이브러리에서는 클로저 사용 시 항상 순환 참조 가능성을 염두에 두어야 합니다.

결론

이번 Triton의 PR은 단순한 리팩토링처럼 보이지만, Python의 메모리 모델과 가비지 컬렉션의 동작 방식을 깊이 이해하고 적용한 사례입니다. 성능 최적화는 거창한 알고리즘 변경뿐만 아니라, 이처럼 언어의 특성을 활용한 세심한 코드 구조 개선에서 시작된다는 점을 다시 한번 상기시켜 줍니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글