본문으로 건너뛰기

[vllm] [vLLM 분석] Speculative Decoding 성능 최적화: DSpark Markov Head 복제 전략

PR 링크: vllm-project/vllm#49731 상태: Merged | 변경: +94 / -23

들어가며

대규모 언어 모델(LLM)의 추론 속도를 높이기 위한 핵심 기술 중 하나는 Speculative Decoding(추측 디코딩)입니다. vLLM은 이를 위해 DSpark와 같은 효율적인 Draft 모델 구조를 활용합니다. 하지만 분산 환경(Tensor Parallelism, TP)에서 모델을 실행할 때, Draft 모델의 크기가 작음에도 불구하고 관성적으로 레이어를 샤딩(Sharding)하면 오히려 성능이 저하될 수 있습니다.

이번 PR([Spec Decode][Perf] Replicate DSpark Markov head across TP ranks)은 DSpark의 Markov head를 모든 TP rank에 복제(Replicate)함으로써, 매 드래프트 토큰 생성 시마다 발생하던 불필요한 통신 오버헤드(all-reduce, gather)를 제거한 최적화 사례입니다. 이 변경을 통해 Qwen3-14B 모델 기준 약 3.3% ~ 3.8%의 처리량(Throughput) 향상을 달성했습니다.

코드 분석: 무엇이 어떻게 바뀌었나?

1. VocabParallelEmbeddingParallelLMHeaddisable_tp 옵션 추가

기존의 vLLM 레이어들은 기본적으로 TP 환경에서 가중치를 나누어 갖도록 설계되었습니다. 하지만 이번 최적화를 위해 특정 레이어에서 TP를 명시적으로 끌 수 있는 기능이 필요했습니다.

Before:

# vllm/model_executor/layers/vocab_parallel_embedding.py
def __init__(...):
    # 항상 TP rank와 size를 가져옴
    tp_rank = get_tensor_model_parallel_rank()
    self.tp_size = get_tensor_model_parallel_world_size()

After:

# vllm/model_executor/layers/vocab_parallel_embedding.py
def __init__(..., disable_tp: bool = False):
    self.disable_tp = disable_tp
    if disable_tp:
        tp_rank, self.tp_size = 0, 1
    else:
        tp_rank = get_tensor_model_parallel_rank()
        self.tp_size = get_tensor_model_parallel_world_size()
    self.tp_rank = tp_rank

또한, 가중치 로딩 시의 호환성을 위해 update_param_tp_status 메서드가 추가되었습니다. 이는 리뷰어 mgoin이 언급한 weight reloading 이슈(#48025)를 해결하기 위한 장치입니다.

2. LogitsProcessor의 통신 로직 조건부 실행

레이어가 복제되어 각 GPU가 전체 가중치를 가지고 있다면, 연산 후 결과를 모으는 gather 과정이 필요 없습니다. LogitsProcessor는 이제 tp_size를 확인하여 통신 여부를 결정합니다.

Before:

# vllm/model_executor/layers/logits_processor.py
def _get_logits(self, lm_head, hidden_states, embedding_bias):
    logits = self._apply_head(lm_head, hidden_states, embedding_bias)
    # 무조건 gather 수행
    logits = self._gather_logits(logits)

After:

# vllm/model_executor/layers/logits_processor.py
def _get_logits(self, lm_head, hidden_states, embedding_bias):
    logits = self._apply_head(lm_head, hidden_states, embedding_bias)
    # TP가 활성화된 경우(tp_size > 1)에만 gather 수행
    if lm_head.tp_size > 1:
        logits = self._gather_logits(logits)

3. DSparkMarkovHead 모델 정의 변경

핵심적인 변화는 Qwen3 DSpark 모델 정의에서 일어났습니다. markov_w1(Embedding)과 markov_w2(LM Head)를 샤딩하지 않고 복제하도록 변경했습니다.

Before:

# vllm/model_executor/models/qwen3_dspark.py
self.markov_w1 = VocabParallelEmbedding(
    vocab_size, markov_rank, prefix=maybe_prefix(prefix, "markov_w1")
)
self.markov_w2 = ParallelLMHead(
    draft_vocab_size, markov_rank, prefix=maybe_prefix(prefix, "markov_w2")
)

After:

# vllm/model_executor/models/qwen3_dspark.py
# w1은 표준 nn.Embedding으로 대체 (자동 복제)
self.markov_w1 = nn.Embedding(vocab_size, markov_rank)
# w2는 disable_tp=True를 통해 복제 모드로 설정
self.markov_w2 = ParallelLMHead(
    draft_vocab_size,
    markov_rank,
    bias=False,
    prefix=maybe_prefix(prefix, "markov_w2"),
    disable_tp=True,
)

왜 이게 좋은 최적화인가?

통신 오버헤드 vs 연산 비용 (Communication vs Computation)

일반적으로 모델의 크기가 매우 크면 가중치를 여러 GPU에 나누어 담는 것이 필수적입니다. 하지만 DSpark의 Markov head처럼 파라미터 수가 적은 레이어의 경우, 연산 자체는 매우 빠르게 끝납니다. 이때 TP로 인해 발생하는 all-reduceall-gather 통신 시간은 실제 연산 시간보다 길어질 수 있습니다.

특히 Speculative Decoding은 여러 개의 드래프트 토큰을 순차적으로 생성하는 경우가 많습니다. 매 토큰마다 GPU 간 통신이 발생하면 지연 시간(Latency)이 누적되어 전체 추론 속도를 갉아먹게 됩니다.

이 PR은 "작은 레이어는 차라리 중복 연산을 하더라도 통신을 없애는 것이 빠르다"는 원칙을 적용한 것입니다. 결과적으로 TP 4 환경에서 207.0 tok/s에서 214.8 tok/s로 약 3.8%의 성능 향상을 이끌어냈습니다.

리뷰어 피드백 반영

리뷰 과정에서 중요한 논의가 있었습니다:

  1. Gemma4 호환성: 초기 구현이 Gemma4 DSpark 모델을 깨뜨릴 수 있다는 지적(mgoin)이 있었고, 이를 수정하여 범용성을 확보했습니다.
  2. Weight Reloading: disable_tp를 사용할 때 vLLM의 가중치 로딩 시스템이 올바른 rank를 참조하도록 update_param_tp_status를 추가하여 안정성을 높였습니다.
  3. Inference Mode: _freeze 메서드(requires_grad=False 설정)가 필요한지에 대한 의문(benchislett)이 제기되었고, vLLM은 추론 전용이므로 과도한 설정은 지양하되 기존 관례를 따르는 선에서 정리되었습니다.

결론

이번 최적화는 분산 시스템 설계에서 항상 고려해야 할 통신 비용의 최소화를 잘 보여주는 사례입니다. 모델의 모든 부분을 기계적으로 샤딩하기보다, 각 레이어의 특성과 실행 빈도에 따라 복제와 샤딩 중 최적의 전략을 선택하는 것이 고성능 추론 엔진의 핵심 역량임을 알 수 있습니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글