Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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)
Expand All @@ -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
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down
Loading