본문으로 건너뛰기

[transformers] Hugging Face Transformers: NoRepeatNGramLogitsProcessor 벡터화 및 성능 최적화

PR 링크: huggingface/transformers#47571 상태: Merged | 변경: +42 / -22

들어가며

NoRepeatNGramLogitsProcessor는 텍스트 생성 시 동일한 N-gram이 반복되는 것을 방지하는 중요한 컴포넌트입니다. 하지만 기존 구현은 매 디코딩 단계마다 시퀀스 내의 모든 N-gram을 딕셔너리로 재구축하고, .tolist()를 호출하여 GPU에서 CPU로 데이터를 동기화하는 병목 현상이 있었습니다. 이로 인해 생성 속도가 저하되고 torch.compile 사용 시 그래프 브레이크(graph break)가 발생하는 문제가 있었습니다.

본 PR은 이러한 비효율적인 파이썬 루프와 호스트 동기화를 제거하고, 순수 텐서 연산(Vectorization)으로 로직을 재설계하여 성능을 최적화했습니다.

코드 분석

src/transformers/generation/logits_process.py

가장 큰 변화는 _calc_banned_ngram_tokens 함수를 제거하고, __call__ 메서드 내에서 torch.unfold를 사용하여 모든 N-gram을 한 번에 처리하도록 변경한 점입니다.

Before:

def _calc_banned_ngram_tokens(ngram_size, prev_input_ids, num_hypos, cur_len):
    # ... 딕셔너리 생성 및 .tolist() 호출로 인한 호스트 동기화 발생
    generated_ngrams = _get_ngrams(ngram_size, prev_input_ids, num_hypos)
    # ...

After:

prefix = input_ids[:, cur_len + 1 - self.ngram_size :]
windows = input_ids.unfold(dimension=1, size=self.ngram_size, step=1)
matches = (windows[..., :-1] == prefix.unsqueeze(1)).all(dim=-1)

vocab_size = scores.shape[-1]
banned_mask = scores.new_zeros((scores.shape[0], vocab_size + 1), dtype=torch.bool)
banned_mask.scatter_(1, torch.where(matches, windows[..., -1], vocab_size), True)

return scores.masked_fill(banned_mask[:, :vocab_size], -float("inf"))

unfold를 통해 현재 시퀀스에서 가능한 모든 N-gram 윈도우를 생성하고, 현재 접미사(suffix)와 일치하는지 비교합니다. 일치하는 경우 해당 N-gram의 마지막 토큰을 금지 목록에 추가합니다. 여기서 vocab_size + 1 크기의 마스크를 사용하여 일치하지 않는 윈도우의 토큰이 잘못된 금지 처리를 하지 않도록 방지하는 기법이 핵심입니다.

왜 이게 좋은가

  1. 호스트 동기화 제거: .tolist()를 제거함으로써 GPU-CPU 간의 데이터 전송 오버헤드를 없앴습니다. 이는 특히 긴 시퀀스 생성 시 성능 향상에 결정적입니다.
  2. torch.compile 최적화: 기존 구현은 2개의 그래프 브레이크를 발생시켰으나, 개선 후 0개의 브레이크로 1개의 통합 그래프를 생성합니다. 이는 향후 torch.compile을 통한 추론 가속화에 유리합니다.
  3. 성능 수치: 8x2048 배치 환경에서 기존 30.90ms에서 1.42ms로 약 20배 이상의 성능 향상을 보였습니다. 특히 시퀀스가 길어질수록 파이썬 루프 기반의 기존 방식보다 압도적인 효율을 보여줍니다.

일반적 교훈

  • Vectorization: 파이썬 레벨의 루프와 딕셔너리 조작은 텐서 연산으로 대체할 수 있는지 항상 검토해야 합니다.
  • Graph Breaks: torch.compile을 사용할 때 tolist(), item() 등 호스트 동기화 연산은 그래프를 끊는 주범입니다. 이를 제거하는 것만으로도 컴파일 효율을 극대화할 수 있습니다.

리뷰 피드백 반영

리뷰어들은 이 최적화가 단순히 속도뿐만 아니라 torch.compile과의 호환성을 높였다는 점을 높게 평가했습니다. 또한, ngram_size=1인 경우나 시퀀스 길이가 짧은 경우의 엣지 케이스를 테스트하기 위해 새로운 테스트 케이스를 추가하여 안정성을 검증했습니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글