[diffusion] optimization: reduce Qwen-Image 2.1 vae and graph warmup memory (#40481)
Co-authored-by: Mick Qian <mickqian@users.noreply.github.com>
This commit is contained in:
@@ -197,6 +197,9 @@ class QwenImage21RMS_norm(nn.Module):
|
|||||||
|
|
||||||
class QwenImage21Upsample(nn.Upsample):
|
class QwenImage21Upsample(nn.Upsample):
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
|
# Nearest interpolation copies values; no FP32 arithmetic is needed.
|
||||||
|
if self.mode == "nearest-exact" and x.dtype in (torch.float16, torch.bfloat16):
|
||||||
|
return super().forward(x)
|
||||||
return super().forward(x.float()).type_as(x)
|
return super().forward(x.float()).type_as(x)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+7
@@ -234,6 +234,13 @@ def prepare_qwen21_mu(batch, server_args):
|
|||||||
|
|
||||||
|
|
||||||
class QwenImage21DenoisingStage(DenoisingStage):
|
class QwenImage21DenoisingStage(DenoisingStage):
|
||||||
|
def _bcg_pad_prompt_kwargs(
|
||||||
|
self, call_kwargs, current_model=None, force_bucket=None
|
||||||
|
):
|
||||||
|
# Prefill runs eagerly. Later steps use exact-length prefix KV, so text
|
||||||
|
# padding only creates duplicate graphs without enabling more replay.
|
||||||
|
return call_kwargs
|
||||||
|
|
||||||
def _predict_noise(
|
def _predict_noise(
|
||||||
self,
|
self,
|
||||||
current_model,
|
current_model,
|
||||||
|
|||||||
@@ -36,6 +36,7 @@ from sglang.multimodal_gen.runtime.models.encoders.qwen3vl_vision import (
|
|||||||
from sglang.multimodal_gen.runtime.models.vaes.autoencoder_kl_qwenimage21 import (
|
from sglang.multimodal_gen.runtime.models.vaes.autoencoder_kl_qwenimage21 import (
|
||||||
AutoencoderKLQwenImage21,
|
AutoencoderKLQwenImage21,
|
||||||
QwenImage21RMS_norm,
|
QwenImage21RMS_norm,
|
||||||
|
QwenImage21Upsample,
|
||||||
_patchify,
|
_patchify,
|
||||||
_unpatchify,
|
_unpatchify,
|
||||||
)
|
)
|
||||||
@@ -50,6 +51,28 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.q
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||||
|
@pytest.mark.parametrize("layout", ["contiguous", "channels_last", "transposed"])
|
||||||
|
@pytest.mark.parametrize("device", ["cpu", "cuda"])
|
||||||
|
def test_nearest_upsample_preserves_every_finite_low_precision_value(
|
||||||
|
dtype, layout, device
|
||||||
|
):
|
||||||
|
if device == "cuda" and not torch.cuda.is_available():
|
||||||
|
pytest.skip("CUDA required")
|
||||||
|
values = torch.arange(65536, dtype=torch.int32).to(torch.int16).view(dtype)
|
||||||
|
values = values[torch.isfinite(values)].reshape(1, 2, -1, 128).to(device)
|
||||||
|
if layout == "channels_last":
|
||||||
|
values = values.contiguous(memory_format=torch.channels_last)
|
||||||
|
elif layout == "transposed":
|
||||||
|
values = values.transpose(2, 3)
|
||||||
|
upsample = QwenImage21Upsample(scale_factor=2, mode="nearest-exact")
|
||||||
|
expected = torch.nn.functional.interpolate(
|
||||||
|
values.float(), scale_factor=2, mode="nearest-exact"
|
||||||
|
).to(dtype)
|
||||||
|
actual = upsample(values)
|
||||||
|
assert torch.equal(actual.view(torch.int16), expected.view(torch.int16))
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("prompt", ["edit", ""])
|
@pytest.mark.parametrize("prompt", ["edit", ""])
|
||||||
@pytest.mark.parametrize("image_count", [0, 1, 2])
|
@pytest.mark.parametrize("image_count", [0, 1, 2])
|
||||||
def test_prompt_conditioning_uses_training_template_and_pre_norm(prompt, image_count):
|
def test_prompt_conditioning_uses_training_template_and_pre_norm(prompt, image_count):
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ from sglang.multimodal_gen.configs.models.dits.qwenimage21 import (
|
|||||||
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image21 import (
|
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image21 import (
|
||||||
QwenImage21PipelineConfig,
|
QwenImage21PipelineConfig,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.configs.sample.qwenimage21 import QwenImage21SamplingParams
|
||||||
from sglang.multimodal_gen.runtime.breakable_cuda_graph.runner import (
|
from sglang.multimodal_gen.runtime.breakable_cuda_graph.runner import (
|
||||||
DiffusionBreakableCudaGraphRunner,
|
DiffusionBreakableCudaGraphRunner,
|
||||||
)
|
)
|
||||||
@@ -33,6 +34,10 @@ from sglang.multimodal_gen.runtime.pipelines.qwen_image21 import QwenImage21Pipe
|
|||||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||||
ComposedPipelineBase,
|
ComposedPipelineBase,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.qwen_image21 import (
|
||||||
|
QwenImage21DenoisingStage,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.server_args import (
|
from sglang.multimodal_gen.runtime.server_args import (
|
||||||
ServerArgs,
|
ServerArgs,
|
||||||
set_global_server_args,
|
set_global_server_args,
|
||||||
@@ -261,14 +266,35 @@ def test_cached_prefix_matches_full_recomputation(model, edit):
|
|||||||
def test_graph_replay_uses_new_request_prefix(model, edit, sample_count):
|
def test_graph_replay_uses_new_request_prefix(model, edit, sample_count):
|
||||||
first = batched_inputs([inputs(5 + i, edit) for i in range(sample_count)])
|
first = batched_inputs([inputs(5 + i, edit) for i in range(sample_count)])
|
||||||
second = batched_inputs([inputs(9 + i, edit) for i in range(sample_count)])
|
second = batched_inputs([inputs(9 + i, edit) for i in range(sample_count)])
|
||||||
|
for kwargs in (first, second):
|
||||||
|
kwargs["encoder_hidden_states_mask"] = torch.ones(
|
||||||
|
kwargs["encoder_hidden_states"].shape[:2], device="cuda", dtype=torch.bool
|
||||||
|
)
|
||||||
|
stage = object.__new__(QwenImage21DenoisingStage)
|
||||||
runner = DiffusionBreakableCudaGraphRunner(model, torch.device("cuda"))
|
runner = DiffusionBreakableCudaGraphRunner(model, torch.device("cuda"))
|
||||||
try:
|
try:
|
||||||
with torch.no_grad(), set_forward_context(None, None):
|
with (
|
||||||
|
torch.no_grad(),
|
||||||
|
set_forward_context(
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
Req(sampling_params=QwenImage21SamplingParams(), is_warmup=True),
|
||||||
|
),
|
||||||
|
):
|
||||||
model(**first)
|
model(**first)
|
||||||
assert runner.capture(**first)
|
stage._bcg_run(runner, first, model)
|
||||||
|
assert len(runner.entries) == 1
|
||||||
|
with (
|
||||||
|
torch.no_grad(),
|
||||||
|
set_forward_context(
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
Req(sampling_params=QwenImage21SamplingParams()),
|
||||||
|
),
|
||||||
|
):
|
||||||
model(**second)
|
model(**second)
|
||||||
expected = model(**second)
|
expected = model(**second)
|
||||||
actual = runner(**second)
|
actual = stage._bcg_run(runner, second, model)
|
||||||
assert len(runner.entries) == 1
|
assert len(runner.entries) == 1
|
||||||
torch.testing.assert_close(actual, expected, atol=1e-6, rtol=1e-6)
|
torch.testing.assert_close(actual, expected, atol=1e-6, rtol=1e-6)
|
||||||
finally:
|
finally:
|
||||||
|
|||||||
Reference in New Issue
Block a user