본문으로 건너뛰기

[vllm] vLLM, Pixtral 모델의 멀티모달 인코더 어텐션 최적화: Packed Sequence Metadata 도입

PR 링크: vllm-project/vllm#52185 상태: Merged | 변경: +157 / -71

들어가며

최근 대규모 언어 모델(LLM)은 텍스트뿐만 아니라 이미지와 같은 다양한 모달리티를 이해하는 멀티모달(Multimodal) 능력으로 발전하고 있습니다. Pixtral 모델 역시 이러한 멀티모달 LLM의 한 예로, 이미지 패치를 이해하고 처리하는 인코더를 포함합니다. 하지만 기존 vLLM 구현에서 Pixtral 모델은 여러 이미지의 패치 시퀀스를 단순히 연결하여 처리했습니다. 특히 xFormers 라이브러리가 사용 불가능할 경우, 이는 메모리 사용량과 지연 시간을 증가시키는 주요 원인이 되었습니다. 이 PR은 이러한 문제를 해결하기 위해 Pixtral 모델의 멀티모달 인코더 어텐션 로직을 vLLM의 MMEncoderAttention으로 통합하고, 'Packed Sequence Metadata' 개념을 도입하여 효율성을 극대화합니다.

코드 분석

이번 변경의 핵심은 Pixtral 모델의 멀티모달 인코더 어텐션 처리 방식을 vLLM의 범용적인 MMEncoderAttention 레이어로 위임하고, 이를 통해 'Packed Sequence Metadata'를 활용하는 것입니다. 이전에는 각 이미지의 패치 시퀀스를 개별적으로 처리하는 대신, 모든 이미지의 패치 시퀀스를 하나의 긴 시퀀스로 간주하고 이를 위한 복잡한 마스크를 생성했습니다. 이 방식은 이미지 수가 많아질수록 비효율적이었습니다.

1. vllm/model_executor/models/pixtral.py 변경사항

가장 큰 변화는 PixtralVisionTransformerBlockPixtralVisionTransformer 클래스에서 나타납니다. 기존에는 xformers.opstorch.nn.functional.scaled_dot_product_attention을 직접 사용하여 어텐션을 계산하고, xformers.ops.fmha.attn_bias.BlockDiagonalMask 또는 transformers 라이브러리의 generate_block_attention_mask를 통해 이미지 간의 경계를 나타내는 마스크를 생성했습니다.

이전 코드 (예시):

--- a/vllm/model_executor/models/pixtral.py
+++ b/vllm/model_executor/models/pixtral.py
@@ -766,15 +766,14 @@
 
     def forward(
         self,
-        x: torch.Tensor,
-        mask: torch.Tensor,
+        x: torch.Tensor,
         freqs_cis: torch.Tensor,
+        cu_seqlens: torch.Tensor,
+        max_seqlen: torch.Tensor,
+        sequence_lengths: torch.Tensor | None,
     ) -> torch.Tensor:
         batch, patches, _ = x.shape
 
-        q, k, v = self.qkv_proj(x).chunk(3, dim=-1)
-
         q, k, v = q.reshape(batch, patches, self.n_heads, self.head_dim), \
                   k.reshape(batch, patches, self.n_heads, self.head_dim), \
                   v.reshape(batch, patches, self.n_heads, self.head_dim)
@@ -782,15 +781,14 @@
         v = v.reshape(batch, patches, self.n_heads, self.head_dim)
 
         q, k = apply_rotary_emb_vit(q, k, freqs_cis=freqs_cis)
-
-        if USE_XFORMERS_OPS:
-            out = xops.memory_efficient_attention(q, k, v, attn_bias=mask)
-        else:
-            q = q.transpose(1, 2)
-            k = k.transpose(1, 2)
-            v = v.transpose(1, 2)
-            out = nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
-            out = out.transpose(1, 2)
+        out = self.attn(
+            q,
+            k,
+            v,
+            cu_seqlens=cu_seqlens,
+            max_seqlen=max_seqlen,
+            sequence_lengths=sequence_lengths,
+        )
 
         out = out.reshape(batch, patches, self.n_heads * self.head_dim)
         out, _ = self.o_proj(out)

새로운 코드:

PixtralVisionTransformerBlockPixtralVisionTransformerforward 메서드는 이제 MMEncoderAttention 레이어를 직접 호출합니다. 이 레이어는 cu_seqlens, max_seqlen, sequence_lengths와 같은 'Packed Sequence Metadata'를 인자로 받습니다. 이는 여러 시퀀스를 하나의 텐서로 묶어 처리할 때 각 시퀀스의 시작점과 길이를 명시적으로 알려주는 정보입니다. 이를 통해 MMEncoderAttention은 내부적으로 각 시퀀스를 독립적으로 처리하면서도 효율적인 연산을 수행할 수 있습니다.

기존의 USE_XFORMERS_OPS 관련 조건부 로직과 xformers 또는 scaled_dot_product_attention 직접 호출 부분이 제거되고, self.attn = MMEncoderAttention(...)으로 대체되었습니다.

또한, PixtralModel.forward 메서드에서는 이미지 패치 임베딩을 처리한 후, _make_packed_sequence_metadata 함수를 호출하여 cu_seqlens, max_seqlen, sequence_lengths를 계산합니다. 이 정보는 self.transformer (즉, PixtralVisionTransformer)의 forward 메서드로 전달됩니다.

--- a/vllm/model_executor/models/pixtral.py
+++ b/vllm/model_executor/models/pixtral.py
@@ -971,20 +971,21 @@
         positions = position_meshgrid(patch_embeds_list).to(self.device)
         freqs_cis = self.freqs_cis[positions[:, 0], positions[:, 1]]
 
-        # pass through Transformer with a block diagonal mask delimiting images
-        if USE_XFORMERS_OPS:
-            mask = xops.fmha.attn_bias.BlockDiagonalMask.from_seqlens(
-                [p.shape[-2] * p.shape[-1] for p in patch_embeds_list],
-            )
-        else:
-            from transformers.models.pixtral.modeling_pixtral import (
-                generate_block_attention_mask,
-            )
-
-            mask = generate_block_attention_mask(
-                [p.shape[-2] * p.shape[-1] for p in patch_embeds_list], patch_embeds
-            )
-        out = self.transformer(patch_embeds, mask=mask, freqs_cis=freqs_cis)
+        attention = self.transformer.layers[0].attention.attn
+        cu_seqlens, max_seqlen, sequence_lengths = _make_packed_sequence_metadata(
+            embed_sizes,
+            attention.attn_backend,
+            self.args.hidden_size,
+            1 if is_vit_use_data_parallel() else get_tensor_model_parallel_world_size(),
+            patch_embeds.device,
+        )
+        out = self.transformer(
+            patch_embeds,
+            freqs_cis=freqs_cis,
+            cu_seqlens=cu_seqlens,
+            max_seqlen=max_seqlen,
+            sequence_lengths=sequence_lengths,
+        )
 
         # squeeze dim 0 and split into separate tensors for each image
         return torch.split(out.squeeze(0), embed_sizes)

2. tests/models/multimodal/generation/test_pixtral.py 변경사항

새로운 테스트 함수 test_packed_sequence_metadata가 추가되었습니다. 이 테스트는 _make_packed_sequence_metadata 함수가 다양한 어텐션 백엔드(FLASH_ATTN, FLASHINFER, TORCH_SDPA)에 대해 올바르게 cu_seqlens, max_seqlen, sequence_lengths를 생성하는지 검증합니다. 이는 새로운 어텐션 처리 방식의 정확성을 보장하는 중요한 단계입니다.

--- a/tests/models/multimodal/generation/test_pixtral.py
+++ b/tests/models/multimodal/generation/test_pixtral.py
@@ -15,7 +16,9 @@
 from vllm import SamplingParams, TextPrompt, TokensPrompt
 from vllm.inputs import MultiModalDataBuiltins
 from vllm.logprobs import Logprob, SampleLogprobs
+from vllm.model_executor.models.pixtral import _make_packed_sequence_metadata
 from vllm.platforms import current_platform
+from vllm.v1.attention.backends.registry import AttentionBackendEnum
 
 from ....utils import VLLM_PATH, large_gpu_test
 from ...utils import check_logprobs_close
@@ -123,6 +126,37 @@ def _create_engine_inputs_hf(urls: list[str]) -> TextPrompt:
 OutputsLogprobs = list[tuple[list[int], str, SampleLogprobs | None]]
 
 
+@pytest.mark.parametrize(
+    

## 참고 자료
- https://pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html
- https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/layers/attention.py#L100
- https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/models/pixtral.py#L69
- https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/models/pixtral.py#L860
- https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/models/pixtral.py#L1244

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

댓글

관련 포스트

PR Analysis 의 다른글