본문으로 건너뛰기

[sglang] [Diffusion] Qwen 모델의 Varlen Mask 메타데이터 호스트 측 빌드 최적화

PR 링크: sgl-project/sglang#33954 상태: Merged | 변경: +57 / -0

들어가며

안녕하세요, 기술 블로거입니다. 오늘은 sgl-project/sglang 레포지토리의 흥미로운 성능 최적화 PR(#31852에서 분리된 PR)을 분석해보고자 합니다. 이 PR은 Diffusion 모델, 특히 Qwen 모델의 denoising 과정에서 발생하는 불필요한 GPU 동기화(device sync)를 제거하여 전체적인 성능을 크게 향상시키는 것을 목표로 합니다.

Diffusion 모델은 이미지 생성 과정에서 여러 번의 denoising 단계를 거치며, 각 단계마다 어텐션 메커니즘이 중요한 역할을 합니다. SGLang과 같은 프레임워크에서는 가변 길이(Varlen) 시퀀스를 효율적으로 처리하기 위해 varlen 마스크 메타데이터를 사용합니다. 기존 Qwen 모델의 마스크 처리 로직은 이 varlen 메타데이터를 생성할 때 GPU의 nonzero() 연산을 사용했는데, 이 연산은 매 denoising 스텝마다 디바이스 동기화를 강제하여 심각한 성능 병목을 초래했습니다.

이 PR은 txt_seq_lens라는 이미 존재하는 텍스트 시퀀스 길이를 활용하여, varlen 메타데이터를 호스트(CPU) 측에서 미리 구성함으로써 GPU 왕복(round-trip)과 디바이스 동기화를 회피하는 영리한 최적화를 제안합니다. 이는 특히 반복적인 연산이 많은 Diffusion 모델에서 큰 이점을 가져다줄 것입니다.

코드 분석

이 PR은 크게 두 가지 변경 사항을 포함합니다. 첫째는 Qwen 모델의 forward 함수 내에서 varlen 마스크 메타데이터를 생성하는 로직을 변경하는 것이고, 둘째는 이 변경 사항의 정확성을 검증하는 단위 테스트를 추가하는 것입니다.

1. python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py

이 파일은 Qwen 이미지 모델의 핵심 로직을 담고 있습니다. 변경의 핵심은 block_attention_kwargs["attn_mask_meta"]를 설정하는 부분입니다. 기존에는 txt_seq_lens 정보가 있더라도 마스크 기반의 build_varlen_mask_meta 함수를 사용했을 가능성이 높습니다. 이 함수는 내부적으로 GPU nonzero() 연산을 포함하여 디바이스 동기화를 유발합니다.

Before:

                # once, so build varlen metadata replay-locally from the current
                # static mask instead of closing over stale cu_seqlens/indices.
                block_attention_kwargs["attn_mask_meta"] = DynamicVarlenMaskMeta()
            else:
                # Precompute varlen metadata once per request so every block
                # reuses the same cu_seqlens / indices instead of rebuilding.
                block_attention_kwargs["attn_mask_meta"] = build_varlen_mask_meta(mask)

After:

--- a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py
+++ b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py
@@ -39,6 +39,7 @@
     DynamicVarlenMaskMeta,
     USPAttention,
     build_varlen_mask_meta,
+    build_varlen_mask_meta_from_ranges,
 )
 from sglang.multimodal_gen.runtime.layers.elementwise import MulAdd
 from sglang.multimodal_gen.runtime.layers.fused_scale_shift_gate import (
@@ -1565,6 +1566,26 @@ def forward(
                 # once, so build varlen metadata replay-locally from the current
                 # static mask instead of closing over stale cu_seqlens/indices.
                 block_attention_kwargs["attn_mask_meta"] = DynamicVarlenMaskMeta()
+            elif (
+                txt_seq_lens is not None
+                and len(txt_seq_lens) == batch_size
+                and all(0 <= n <= encoder_hidden_states.shape[1] for n in txt_seq_lens)
+            ):
+                # txt_seq_lens already carries each row's valid text prefix
+                # (the mask is built from it), so the varlen metadata can be
+                # assembled host-side; the mask-based builder costs a GPU
+                # nonzero plus a device sync on every denoising step.
+                txt_len = encoder_hidden_states.shape[1]
+                block_attention_kwargs["attn_mask_meta"] = (
+                    build_varlen_mask_meta_from_ranges(
+                        [
+                            [(0, int(n)), (txt_len, txt_len + image_seq_len)]
+                            for n in txt_seq_lens
+                        ],
+                        max_seqlen=txt_len + image_seq_len,
+                        device=hidden_states.device,
+                    )
+                )
             else:
                 # Precompute varlen metadata once per request so every block
                 # reuses the same cu_seqlens / indices instead of rebuilding.

이 변경의 핵심은 새로운 elif 블록입니다. txt_seq_lens가 유효하게 제공될 경우, 즉 각 배치 항목의 텍스트 시퀀스 길이를 알고 있을 때, build_varlen_mask_meta_from_ranges 함수를 사용하여 attn_mask_meta를 생성합니다. 이 함수는 txt_seq_lens를 기반으로 각 시퀀스의 유효한 범위를 [(0, n), (txt_len, txt_len + image_seq_len)] 형태로 구성하여 호스트 측에서 varlen 메타데이터를 빌드합니다. 이는 GPU nonzero() 연산을 피하고, 결과적으로 매 denoising 스텝마다 발생하는 디바이스 동기화를 제거합니다.

build_varlen_mask_meta_from_rangescu_seqlens, indices, inv_indices와 같은 varlen 어텐션에 필요한 메타데이터를 CPU에서 효율적으로 계산합니다. 기존의 else 블록은 txt_seq_lens를 사용할 수 없는 경우를 위한 폴백(fallback)으로 남겨두어 유연성을 유지합니다.

2. python/sglang/multimodal_gen/test/unit/test_varlen_meta_host_build.py

성능 최적화는 항상 정확성 검증이 수반되어야 합니다. 이 PR은 호스트 측에서 빌드된 메타데이터가 기존 GPU nonzero() 기반 빌더의 결과와 동일함을 보장하는 단위 테스트를 추가합니다.

New File:

--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_varlen_meta_host_build.py
@@ -0,0 +1,36 @@
+"""Host-built varlen metadata must match the mask-based (nonzero) builder."""
+
+import unittest
+
+import torch
+
+from sglang.multimodal_gen.runtime.layers.attention.layer import (
+    build_varlen_mask_meta,
+    build_varlen_mask_meta_from_ranges,
+)
+
+
+class TestHostVarlenMetaEquivalence(unittest.TestCase):
+    def test_prefix_text_plus_full_image_matches_nonzero_builder(self):
+        txt_len, img_len = 7, 5
+        txt_seq_lens = [3, 7, 0]
+        bs = len(txt_seq_lens)
+        mask = torch.zeros(bs, txt_len + img_len, dtype=torch.bool)
+        for row, n in enumerate(txt_seq_lens):
+            mask[row, :n] = True
+            mask[row, txt_len:] = True
+
+        ref = build_varlen_mask_meta(mask)
+        host = build_varlen_mask_meta_from_ranges(
+            [[(0, n), (txt_len, txt_len + img_len)] for n in txt_seq_lens],
+            max_seqlen=txt_len + img_len,
+            device=mask.device,
+        )
+
+        for key in ("cu_seqlens", "indices", "inv_indices"):
+            torch.testing.assert_close(host[key], ref[key], rtol=0, atol=0)
+        self.assertEqual(host["max_seqlen"], ref["max_seqlen"])
+
+
+if __name__ == "__main__":
+    unittest.main()

TestHostVarlenMetaEquivalence 클래스의 test_prefix_text_plus_full_image_matches_nonzero_builder 함수는 다음을 수행합니다:

  1. 테스트 시나리오 설정: txt_len, img_len, txt_seq_lens (다양한 길이 포함)를 정의하여 실제 Diffusion 모델에서 발생할 수 있는 텍스트-이미지 시퀀스 조합을 시뮬레이션합니다.
  2. 마스크 생성: txt_seq_lens와 이미지 길이를 기반으로 torch.bool 타입의 마스크를 생성합니다. 이 마스크는 텍스트 프리픽스와 전체 이미지 영역을 포함합니다.
  3. 참조(Reference) 메타데이터 생성: build_varlen_mask_meta(mask)를 호출하여 기존의 마스크 기반(GPU nonzero() 사용) 방식으로 ref 메타데이터를 생성합니다.
  4. 호스트 빌드 메타데이터 생성: build_varlen_mask_meta_from_ranges를 호출하여 호스트 측에서 host 메타데이터를 생성합니다.
  5. 결과 비교: torch.testing.assert_close를 사용하여 hostrefcu_seqlens, indices, inv_indices가 요소별로 동일한지 확인합니다. 또한 max_seqlen도 동일한지 검증합니다.

이 테스트는 새로운 호스트 측 빌드 로직이 기존 GPU 기반 로직과 기능적으로 동등하며, 정확한 메타데이터를 생성함을 보장합니다.

왜 이게 좋은가

이 PR은 다음과 같은 중요한 이유로 좋은 최적화 및 개선 사항입니다.

1. 치명적인 성능 병목 제거

가장 큰 이점은 매 denoising 스텝마다 발생하는 GPU 디바이스 동기화를 제거한다는 점입니다. Diffusion 모델은 수십에서 수백 번의 denoising 스텝을 거치며 이미지를 점진적으로 생성합니다. 각 스텝마다 GPU가 CPU의 명령을 기다려야 하는 device sync는 전체 생성 시간을 크게 늘리는 주범이 됩니다. 이 최적화는 이 반복적인 동기화 오버헤드를 완전히 없애, Diffusion 모델의 추론 속도를 획기적으로 향상시킬 수 있습니다.

2. 효율적인 자원 활용

varlen 메타데이터를 생성하는 작업은 본질적으로 데이터 구조를 재구성하는 CPU-intensive 작업에 가깝습니다. 이를 GPU에서 nonzero()와 같은 연산으로 처리하는 것은 GPU의 강력한 병렬 처리 능력을 비효율적으로 사용하는 것입니다. 이 PR은 이 작업을 CPU로 옮김으로써 GPU는 어텐션 계산과 같은 핵심 병렬 연산에 집중하고, CPU는 메타데이터 준비와 같은 순차적인 데이터 처리에 집중하게 하여 전체 시스템의 자원 활용 효율성을 극대화합니다.

3. 기존 정보의 재활용

txt_seq_lens는 이미 각 배치 항목의 유효한 텍스트 시퀀스 길이를 담고 있는 정보입니다. 이 PR은 이 정보를 단순히 마스크를 만드는 데 사용하는 것을 넘어, varlen 메타데이터를 직접 구성하는 데 활용함으로써 정보의 가치를 높였습니다. 이는 불필요한 중간 계산(마스크 생성 후 nonzero())을 줄이고, 더 직접적이고 효율적인 경로를 택한 것입니다.

4. 견고한 단위 테스트

성능 최적화는 종종 미묘한 버그를 유발할 수 있습니다. 이 PR은 새로운 호스트 측 빌드 로직이 기존 GPU 기반 로직과 동일한 결과를 생성함을 엄격하게 검증하는 단위 테스트를 포함합니다. 이는 최적화가 정확성을 희생하지 않았음을 보장하며, 향후 코드 변경 시 회귀(regression)를 방지하는 데 기여합니다.

일반적인 교훈

  • 병목 지점 식별: 반복적인 루프 내에서 발생하는 device sync는 잠재적인 성능 병목입니다. 프로파일링을 통해 이러한 지점을 식별하는 것이 중요합니다.
  • CPU-GPU 역할 분담: GPU는 병렬 연산에, CPU는 데이터 준비 및 제어 흐름에 강점을 가집니다. 각자의 강점을 활용하여 작업을 분담하는 것이 전체 시스템 성능에 유리합니다.
  • 데이터 재활용: 이미 존재하는 정보를 최대한 활용하여 중복 계산을 피하고 효율성을 높입니다.
  • 정확성 검증: 성능 최적화 후에는 반드시 기능적 정확성을 검증하는 테스트를 추가해야 합니다.

리뷰 댓글 분석

제공된 리뷰 댓글([mickqian] /tag-and-rerun-ci)은 주로 CI/CD 파이프라인 재실행과 관련된 운영상의 요청이었습니다. 이 PR의 기술적인 내용에 대한 직접적인 피드백이나 논의는 포함되어 있지 않습니다. 다만, PR 설명에서 이 변경 사항이 #31852 PR에서 분리되어 나왔다고 언급된 점을 미루어 볼 때, 이미 내부적으로 충분한 기술적 검토와 논의가 이루어졌음을 짐작할 수 있습니다. 이 PR 자체는 명확한 성능 개선 목표와 검증 로직을 가지고 있어, 기술적인 이견보다는 구현의 정확성과 통합에 초점을 맞춘 것으로 보입니다.

결론

이 PR은 Diffusion 모델의 핵심 구성 요소인 varlen 어텐션 메타데이터 생성 과정을 최적화하여, 불필요한 GPU 동기화를 제거하고 전체 추론 성능을 크게 향상시키는 모범적인 사례입니다. 기존에 존재하는 txt_seq_lens 정보를 영리하게 활용하여 호스트 측에서 메타데이터를 빌드하는 방식은 효율적인 자원 활용과 성능 개선이라는 두 마리 토끼를 모두 잡았습니다. 이러한 종류의 미세한 최적화들이 모여 대규모 AI 모델의 실용성을 높이는 데 기여합니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글