diff --git a/tests/diffusion/models/qwen_image/test_qwen_image_pipeline_device.py b/tests/diffusion/models/qwen_image/test_qwen_image_pipeline_device.py index 7f0ce51c8dc..c65fb1731db 100644 --- a/tests/diffusion/models/qwen_image/test_qwen_image_pipeline_device.py +++ b/tests/diffusion/models/qwen_image/test_qwen_image_pipeline_device.py @@ -64,7 +64,7 @@ def temperal_downsample(self): return [True, True, True] -def _run_pipeline_init(monkeypatch, *, enable_cpu_offload, loader_device): +def _run_pipeline_init(monkeypatch, *, enable_cpu_offload, loader_device, diffusion_offload_config=None): import vllm_omni.diffusion.models.qwen_image.pipeline_qwen_image as pipe_mod from vllm_omni.diffusion.models.qwen_image.pipeline_qwen_image import QwenImagePipeline @@ -110,6 +110,7 @@ def _prefetch(factory, *args, **kwargs): tf_model_config={}, quantization_config=None, enable_cpu_offload=enable_cpu_offload, + diffusion_offload_config=diffusion_offload_config, enable_diffusion_pipeline_profiler=False, ) return QwenImagePipeline(od_config=od_config) @@ -129,3 +130,19 @@ def test_pipeline_dit_follows_loader_cpu_without_override(monkeypatch): assert pipe.transformer.probe_device_type == "cpu" assert pipe.text_encoder.placed_device.type == "cuda" assert pipe.vae.placed_device.type == "cuda" + + +@pytest.mark.parametrize( + ("mode", "loader_device", "expected_device"), + [("module", "cuda", "cpu"), ("layer", "cpu", "cuda")], +) +def test_pipeline_compact_offload_places_encoder_vae(monkeypatch, mode, loader_device, expected_device): + pipe = _run_pipeline_init( + monkeypatch, + enable_cpu_offload=False, + loader_device=loader_device, + diffusion_offload_config={"mode": mode, "components": ["dit"]}, + ) + assert pipe.text_encoder.placed_device.type == expected_device + assert pipe.vae.placed_device.type == expected_device + assert pipe.transformer.probe_device_type == loader_device diff --git a/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image.py b/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image.py index 5e262ff1e01..6d7a2a7d359 100644 --- a/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image.py +++ b/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image.py @@ -37,6 +37,7 @@ QwenImageTransformer2DModel, ) from vllm_omni.diffusion.models.qwen_image.rope_utils import txt_seq_lens_from_embeds +from vllm_omni.diffusion.offloader.config import OffloadStrategy, resolve_offload_strategy from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin from vllm_omni.diffusion.utils.prompt_utils import ( validate_prompt_sequence_lengths, @@ -341,7 +342,7 @@ def __init__( # do not share VRAM with DiT construction (#7555). The DiT follows the # loader's default-device context (CUDA for online / AutoRound INT under # offload, CPU for layerwise / unquantized HSDP defer). - cpu_offload = bool(getattr(self.od_config, "enable_cpu_offload", False)) + cpu_offload = resolve_offload_strategy(self.od_config) is OffloadStrategy.MODEL_LEVEL enc_vae_device = torch.device("cpu") if cpu_offload else self.device self.text_encoder = self.text_encoder.to(enc_vae_device) self.vae = from_pretrained_with_prefetch(