[diffusion] fix: fix Qwen-Image-Layered latent shape (#28832)
This commit is contained in:
@@ -783,22 +783,31 @@ class QwenImageLayeredPipelineConfig(QwenImageEditPipelineConfig):
|
|||||||
return cond_kwargs
|
return cond_kwargs
|
||||||
|
|
||||||
def _unpad_and_unpack_latents(self, latents, batch):
|
def _unpad_and_unpack_latents(self, latents, batch):
|
||||||
vae_scale_factor = self.get_vae_scale_factor()
|
|
||||||
channels = self.dit_config.arch_config.in_channels
|
channels = self.dit_config.arch_config.in_channels
|
||||||
batch_size = latents.shape[0]
|
batch_size = latents.shape[0]
|
||||||
layers = batch.num_frames
|
|
||||||
|
|
||||||
height = 2 * (int(batch.height) // (vae_scale_factor * 2))
|
img_shapes = batch.img_shapes
|
||||||
width = 2 * (int(batch.width) // (vae_scale_factor * 2))
|
generated_shapes = img_shapes[0][:-1] if img_shapes and img_shapes[0] else []
|
||||||
|
if not generated_shapes:
|
||||||
|
raise ValueError("Qwen-Image-Layered requires generated latent shapes.")
|
||||||
|
if len({tuple(shape) for shape in generated_shapes}) != 1:
|
||||||
|
raise ValueError(
|
||||||
|
"Qwen-Image-Layered generated latent shapes must match, got "
|
||||||
|
f"{generated_shapes}."
|
||||||
|
)
|
||||||
|
layers = len(generated_shapes)
|
||||||
|
_, latent_height, latent_width = generated_shapes[0]
|
||||||
|
height = 2 * int(latent_height)
|
||||||
|
width = 2 * int(latent_width)
|
||||||
|
|
||||||
latents = maybe_unpad_latents(latents, batch)
|
latents = maybe_unpad_latents(latents, batch)
|
||||||
latents = latents.view(
|
latents = latents.view(
|
||||||
batch_size, layers + 1, height // 2, width // 2, channels // 4, 2, 2
|
batch_size, layers, height // 2, width // 2, channels // 4, 2, 2
|
||||||
)
|
)
|
||||||
latents = latents.permute(0, 1, 4, 2, 5, 3, 6)
|
latents = latents.permute(0, 1, 4, 2, 5, 3, 6)
|
||||||
|
|
||||||
latents = latents.reshape(
|
latents = latents.reshape(
|
||||||
batch_size, layers + 1, channels // (2 * 2), height, width
|
batch_size, layers, channels // (2 * 2), height, width
|
||||||
)
|
)
|
||||||
latents = latents.permute(0, 2, 1, 3, 4) # (b, c, f, h, w)
|
latents = latents.permute(0, 2, 1, 3, 4) # (b, c, f, h, w)
|
||||||
return latents, batch_size, channels, height, width
|
return latents, batch_size, channels, height, width
|
||||||
|
|||||||
@@ -112,6 +112,11 @@ class ComponentLoader(ABC):
|
|||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
|
def should_raise_customized_load_error(
|
||||||
|
self, _server_args: ServerArgs, _component_name: str
|
||||||
|
) -> bool:
|
||||||
|
return False
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _is_component_set_as_layerwise_load(
|
def _is_component_set_as_layerwise_load(
|
||||||
server_args: ServerArgs, component_name: str
|
server_args: ServerArgs, component_name: str
|
||||||
@@ -219,6 +224,12 @@ class ComponentLoader(ABC):
|
|||||||
)
|
)
|
||||||
source = "sgl-diffusion"
|
source = "sgl-diffusion"
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
if self.should_raise_customized_load_error(server_args, component_name):
|
||||||
|
traceback.print_exc()
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Failed to load customized {component_name}; native fallback "
|
||||||
|
"is disabled for this component configuration."
|
||||||
|
) from e
|
||||||
if "Unsupported model architecture" in str(e):
|
if "Unsupported model architecture" in str(e):
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Component: {component_name} doesn't have a customized version yet, using native version"
|
f"Component: {component_name} doesn't have a customized version yet, using native version"
|
||||||
|
|||||||
@@ -79,6 +79,14 @@ class TransformerLoader(ComponentLoader):
|
|||||||
]
|
]
|
||||||
expected_library = "diffusers"
|
expected_library = "diffusers"
|
||||||
|
|
||||||
|
def should_raise_customized_load_error(
|
||||||
|
self, server_args: ServerArgs, component_name: str
|
||||||
|
) -> bool:
|
||||||
|
component_server_args = _server_args_for_transformer_component(
|
||||||
|
server_args, component_name
|
||||||
|
)
|
||||||
|
return component_server_args.transformer_weights_path is not None
|
||||||
|
|
||||||
def load_customized(
|
def load_customized(
|
||||||
self, component_model_path: str, server_args: ServerArgs, component_name: str
|
self, component_model_path: str, server_args: ServerArgs, component_name: str
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -22,7 +22,10 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config i
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.utils import _list_safetensors_files
|
from sglang.multimodal_gen.runtime.loader.utils import _list_safetensors_files
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model
|
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
||||||
|
maybe_download_model,
|
||||||
|
snapshot_download,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
|
from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
|
||||||
from sglang.multimodal_gen.runtime.utils.quantization_utils import (
|
from sglang.multimodal_gen.runtime.utils.quantization_utils import (
|
||||||
@@ -317,12 +320,31 @@ def resolve_transformer_safetensors_to_load(
|
|||||||
quantized_path = server_args.transformer_weights_path
|
quantized_path = server_args.transformer_weights_path
|
||||||
|
|
||||||
if quantized_path:
|
if quantized_path:
|
||||||
quantized_path = maybe_download_model(quantized_path)
|
original_quantized_path = quantized_path
|
||||||
|
quantized_path = maybe_download_model(original_quantized_path)
|
||||||
logger.info("using quantized transformer weights from: %s", quantized_path)
|
logger.info("using quantized transformer weights from: %s", quantized_path)
|
||||||
if os.path.isfile(quantized_path) and quantized_path.endswith(".safetensors"):
|
if os.path.isfile(quantized_path) and quantized_path.endswith(".safetensors"):
|
||||||
safetensors_list = [quantized_path]
|
safetensors_list = [quantized_path]
|
||||||
else:
|
else:
|
||||||
safetensors_list = _list_safetensors_files(quantized_path)
|
safetensors_list = _list_safetensors_files(quantized_path)
|
||||||
|
if not safetensors_list and not os.path.exists(original_quantized_path):
|
||||||
|
logger.warning(
|
||||||
|
"No safetensors files found in cached transformer weights path "
|
||||||
|
"%s; refreshing snapshot for %s",
|
||||||
|
quantized_path,
|
||||||
|
original_quantized_path,
|
||||||
|
)
|
||||||
|
quantized_path = snapshot_download(
|
||||||
|
repo_id=original_quantized_path,
|
||||||
|
ignore_patterns=["*.onnx", "*.msgpack"],
|
||||||
|
allow_patterns=[
|
||||||
|
"*.json",
|
||||||
|
"*.safetensors",
|
||||||
|
"*.safetensors.index.json",
|
||||||
|
],
|
||||||
|
max_workers=8,
|
||||||
|
)
|
||||||
|
safetensors_list = _list_safetensors_files(quantized_path)
|
||||||
else:
|
else:
|
||||||
safetensors_list = _list_safetensors_files(component_model_path)
|
safetensors_list = _list_safetensors_files(component_model_path)
|
||||||
|
|
||||||
|
|||||||
+2
@@ -512,6 +512,8 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
|
|||||||
multiple_of = self.vae_scale_factor * 2
|
multiple_of = self.vae_scale_factor * 2
|
||||||
width = width // multiple_of * multiple_of
|
width = width // multiple_of * multiple_of
|
||||||
height = height // multiple_of * multiple_of
|
height = height // multiple_of * multiple_of
|
||||||
|
batch.width = width
|
||||||
|
batch.height = height
|
||||||
|
|
||||||
# if image is not None and not (isinstance(image, torch.Tensor) and image.size(1) == self.latent_channels):
|
# if image is not None and not (isinstance(image, torch.Tensor) and image.size(1) == self.latent_channels):
|
||||||
image = self.image_processor.resize(image, calculated_height, calculated_width)
|
image = self.image_processor.resize(image, calculated_height, calculated_width)
|
||||||
|
|||||||
@@ -0,0 +1,47 @@
|
|||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
|
||||||
|
QwenImageLayeredPipelineConfig,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestQwenImageLayeredPipelineConfig(unittest.TestCase):
|
||||||
|
def test_unpack_uses_layered_img_shapes_not_stale_request_size(self):
|
||||||
|
config = QwenImageLayeredPipelineConfig()
|
||||||
|
channels = config.dit_config.arch_config.in_channels
|
||||||
|
generated_layers = 2
|
||||||
|
latent_height = 40
|
||||||
|
latent_width = 40
|
||||||
|
latents = torch.empty(
|
||||||
|
1,
|
||||||
|
generated_layers * latent_height * latent_width,
|
||||||
|
channels,
|
||||||
|
)
|
||||||
|
batch = SimpleNamespace(
|
||||||
|
height=512,
|
||||||
|
width=512,
|
||||||
|
raw_latent_shape=latents.shape,
|
||||||
|
img_shapes=[
|
||||||
|
[
|
||||||
|
(1, latent_height, latent_width),
|
||||||
|
(1, latent_height, latent_width),
|
||||||
|
(1, latent_height, latent_width),
|
||||||
|
]
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
unpacked, batch_size, unpacked_channels, height, width = (
|
||||||
|
config._unpad_and_unpack_latents(latents, batch)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(batch_size, 1)
|
||||||
|
self.assertEqual(unpacked_channels, channels)
|
||||||
|
self.assertEqual((height, width), (80, 80))
|
||||||
|
self.assertEqual(unpacked.shape, (1, channels // 4, generated_layers, 80, 80))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -3,6 +3,7 @@ This unittest is introduced in #22360, preventing duplicate transformer safetens
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
|
import os
|
||||||
import sys
|
import sys
|
||||||
import tempfile
|
import tempfile
|
||||||
import types
|
import types
|
||||||
@@ -126,6 +127,34 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
|||||||
|
|
||||||
self.assertEqual(resolved, [mixed])
|
self.assertEqual(resolved, [mixed])
|
||||||
|
|
||||||
|
@patch(
|
||||||
|
"sglang.multimodal_gen.runtime.loader.transformer_load_utils.snapshot_download",
|
||||||
|
)
|
||||||
|
@patch(
|
||||||
|
"sglang.multimodal_gen.runtime.loader.transformer_load_utils.maybe_download_model",
|
||||||
|
)
|
||||||
|
def test_resolve_transformer_safetensors_to_load_refreshes_empty_cached_repo(
|
||||||
|
self, mock_download_model, mock_snapshot_download
|
||||||
|
):
|
||||||
|
with tempfile.TemporaryDirectory() as cached_dir:
|
||||||
|
repo_id = "black-forest-labs/FLUX.2-dev-NVFP4"
|
||||||
|
mixed = os.path.join(cached_dir, "flux2-dev-nvfp4-mixed.safetensors")
|
||||||
|
mock_download_model.return_value = cached_dir
|
||||||
|
|
||||||
|
def _snapshot_download(**_kwargs):
|
||||||
|
open(mixed, "a").close()
|
||||||
|
return cached_dir
|
||||||
|
|
||||||
|
mock_snapshot_download.side_effect = _snapshot_download
|
||||||
|
|
||||||
|
server_args = self._make_server_args(transformer_weights_path=repo_id)
|
||||||
|
resolved = resolve_transformer_safetensors_to_load(
|
||||||
|
server_args, "/unused/component/path"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(resolved, [mixed])
|
||||||
|
mock_snapshot_download.assert_called_once()
|
||||||
|
|
||||||
def test_filter_transformer_precision_variants_prefers_canonical_file(self):
|
def test_filter_transformer_precision_variants_prefers_canonical_file(self):
|
||||||
files = [
|
files = [
|
||||||
"/tmp/transformer/diffusion_pytorch_model.fp16.safetensors",
|
"/tmp/transformer/diffusion_pytorch_model.fp16.safetensors",
|
||||||
|
|||||||
Reference in New Issue
Block a user