[diffusion] refactor: refactor condition image resize logic (#14079)
This commit is contained in:
@@ -9,7 +9,6 @@ from typing import Any
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.models.base import ArchConfig, ModelConfig
|
from sglang.multimodal_gen.configs.models.base import ArchConfig, ModelConfig
|
||||||
from sglang.multimodal_gen.runtime.models.vision_utils import get_default_height_width
|
|
||||||
from sglang.multimodal_gen.utils import StoreBoolean
|
from sglang.multimodal_gen.utils import StoreBoolean
|
||||||
|
|
||||||
|
|
||||||
@@ -51,13 +50,6 @@ class VAEConfig(ModelConfig):
|
|||||||
def post_init(self):
|
def post_init(self):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
# returns width, height
|
|
||||||
def calculate_dimensions(
|
|
||||||
self, image, vae_scale_factor, width, height
|
|
||||||
) -> tuple[int, int]:
|
|
||||||
height, width = get_default_height_width(image, vae_scale_factor, height, width)
|
|
||||||
return width, height
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def add_cli_args(parser: Any, prefix: str = "vae-config") -> Any:
|
def add_cli_args(parser: Any, prefix: str = "vae-config") -> Any:
|
||||||
"""Add CLI arguments for VAEConfig fields"""
|
"""Add CLI arguments for VAEConfig fields"""
|
||||||
@@ -148,6 +140,9 @@ class VAEConfig(ModelConfig):
|
|||||||
|
|
||||||
return parser
|
return parser
|
||||||
|
|
||||||
|
def get_vae_scale_factor(self):
|
||||||
|
return 2 ** (len(self.arch_config.block_out_channels) - 1)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_cli_args(cls, args: argparse.Namespace) -> "VAEConfig":
|
def from_cli_args(cls, args: argparse.Namespace) -> "VAEConfig":
|
||||||
kwargs = {}
|
kwargs = {}
|
||||||
|
|||||||
@@ -4,7 +4,6 @@
|
|||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
from sglang.multimodal_gen.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||||
from sglang.multimodal_gen.utils import calculate_dimensions
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -28,9 +27,6 @@ class QwenImageVAEArchConfig(VAEArchConfig):
|
|||||||
scale_factor_spatial: int = 8
|
scale_factor_spatial: int = 8
|
||||||
clip_output: bool = True
|
clip_output: bool = True
|
||||||
|
|
||||||
def __post_init__(self):
|
|
||||||
self.vae_scale_factor = 2 ** len(self.temperal_downsample)
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class QwenImageVAEConfig(VAEConfig):
|
class QwenImageVAEConfig(VAEConfig):
|
||||||
@@ -42,11 +38,8 @@ class QwenImageVAEConfig(VAEConfig):
|
|||||||
use_temporal_tiling: bool = False
|
use_temporal_tiling: bool = False
|
||||||
use_parallel_tiling: bool = False
|
use_parallel_tiling: bool = False
|
||||||
|
|
||||||
def calculate_dimensions(self, image, vae_scale_factor, width, height):
|
def get_vae_scale_factor(self):
|
||||||
width = image.size[0]
|
return 2 ** len(self.arch_config.temperal_downsample)
|
||||||
height = image.size[1]
|
|
||||||
width, height, _ = calculate_dimensions(1024 * 1024, width / height)
|
|
||||||
return width, height
|
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
self.blend_num_frames = (
|
self.blend_num_frames = (
|
||||||
|
|||||||
@@ -7,8 +7,8 @@ from dataclasses import asdict, dataclass, field, fields
|
|||||||
from enum import Enum, auto
|
from enum import Enum, auto
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
import PIL
|
||||||
import torch
|
import torch
|
||||||
from diffusers.image_processor import VaeImageProcessor
|
|
||||||
from einops import rearrange
|
from einops import rearrange
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.models import (
|
from sglang.multimodal_gen.configs.models import (
|
||||||
@@ -24,6 +24,7 @@ from sglang.multimodal_gen.runtime.distributed import (
|
|||||||
get_sp_world_size,
|
get_sp_world_size,
|
||||||
sequence_model_parallel_all_gather,
|
sequence_model_parallel_all_gather,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.models.vision_utils import get_default_height_width
|
||||||
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.utils import (
|
from sglang.multimodal_gen.utils import (
|
||||||
FlexibleArgumentParser,
|
FlexibleArgumentParser,
|
||||||
@@ -174,15 +175,20 @@ class PipelineConfig:
|
|||||||
# Compilation
|
# Compilation
|
||||||
# enable_torch_compile: bool = False
|
# enable_torch_compile: bool = False
|
||||||
|
|
||||||
|
# calculate the adjust size for condition image
|
||||||
|
# width: original condition image width
|
||||||
|
# height: original condition image height
|
||||||
|
def calculate_condition_image_size(self, image, width, height) -> tuple[int, int]:
|
||||||
|
vae_scale_factor = self.vae_config.arch_config.spatial_compression_ratio
|
||||||
|
height, width = get_default_height_width(image, vae_scale_factor, height, width)
|
||||||
|
return width, height
|
||||||
|
|
||||||
|
def resize_condition_image(self, image, target_width, target_height):
|
||||||
|
return image.resize((target_width, target_height), PIL.Image.Resampling.LANCZOS)
|
||||||
|
|
||||||
def slice_noise_pred(self, noise, latents):
|
def slice_noise_pred(self, noise, latents):
|
||||||
return noise
|
return noise
|
||||||
|
|
||||||
def maybe_resize_condition_image(self, width, height, image):
|
|
||||||
"""
|
|
||||||
image: input image
|
|
||||||
"""
|
|
||||||
return image, width, height
|
|
||||||
|
|
||||||
def adjust_num_frames(self, num_frames):
|
def adjust_num_frames(self, num_frames):
|
||||||
return num_frames
|
return num_frames
|
||||||
|
|
||||||
@@ -190,10 +196,6 @@ class PipelineConfig:
|
|||||||
def tokenize_prompt(self, prompt: list[str], tokenizer, tok_kwargs) -> dict:
|
def tokenize_prompt(self, prompt: list[str], tokenizer, tok_kwargs) -> dict:
|
||||||
return tokenizer(prompt, **tok_kwargs)
|
return tokenizer(prompt, **tok_kwargs)
|
||||||
|
|
||||||
# called in ImageEncodingStage, preprocess the image
|
|
||||||
def preprocess_image(self, image, image_processor: VaeImageProcessor):
|
|
||||||
return image
|
|
||||||
|
|
||||||
def prepare_latent_shape(self, batch, batch_size, num_frames):
|
def prepare_latent_shape(self, batch, batch_size, num_frames):
|
||||||
height = batch.height // self.vae_config.arch_config.spatial_compression_ratio
|
height = batch.height // self.vae_config.arch_config.spatial_compression_ratio
|
||||||
width = batch.width // self.vae_config.arch_config.spatial_compression_ratio
|
width = batch.width // self.vae_config.arch_config.spatial_compression_ratio
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ import math
|
|||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Callable, List, Optional
|
from typing import Callable, List, Optional
|
||||||
|
|
||||||
import PIL.Image
|
import PIL
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
from sglang.multimodal_gen.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||||
@@ -448,18 +448,21 @@ class Flux2PipelineConfig(FluxPipelineConfig):
|
|||||||
def get_neg_prompt_embeds(self, batch):
|
def get_neg_prompt_embeds(self, batch):
|
||||||
return batch.negative_prompt_embeds[0]
|
return batch.negative_prompt_embeds[0]
|
||||||
|
|
||||||
def maybe_resize_condition_image(self, width, height, image):
|
def calculate_condition_image_size(
|
||||||
|
self, image, width, height
|
||||||
|
) -> Optional[tuple[int, int]]:
|
||||||
target_area: int = 1024 * 1024
|
target_area: int = 1024 * 1024
|
||||||
|
|
||||||
if width is not None and height is not None:
|
if width is not None and height is not None:
|
||||||
if width * height > target_area:
|
if width * height > target_area:
|
||||||
scale = math.sqrt(target_area / (width * height))
|
scale = math.sqrt(target_area / (width * height))
|
||||||
width = int(width * scale)
|
width = int(width * scale)
|
||||||
height = int(height * scale)
|
height = int(height * scale)
|
||||||
image = image.resize((width, height), PIL.Image.Resampling.LANCZOS)
|
return width, height
|
||||||
width, height = image.size
|
|
||||||
|
|
||||||
return image, width, height
|
return None
|
||||||
|
|
||||||
|
def resize_condition_image(self, image, target_width, target_height):
|
||||||
|
return image.resize((target_width, target_height), PIL.Image.Resampling.LANCZOS)
|
||||||
|
|
||||||
def get_freqs_cis(self, prompt_embeds, width, height, device, rotary_emb, batch):
|
def get_freqs_cis(self, prompt_embeds, width, height, device, rotary_emb, batch):
|
||||||
|
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
|||||||
ModelTaskType,
|
ModelTaskType,
|
||||||
shard_rotary_emb_for_sp,
|
shard_rotary_emb_for_sp,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.models.vision_utils import resize
|
||||||
from sglang.multimodal_gen.utils import calculate_dimensions
|
from sglang.multimodal_gen.utils import calculate_dimensions
|
||||||
|
|
||||||
|
|
||||||
@@ -215,8 +216,7 @@ class QwenImageEditPipelineConfig(QwenImagePipelineConfig):
|
|||||||
assert batch_size == 1
|
assert batch_size == 1
|
||||||
height = batch.height
|
height = batch.height
|
||||||
width = batch.width
|
width = batch.width
|
||||||
image = batch.condition_image
|
image_size = batch.original_condition_image_size
|
||||||
image_size = image[0].size if isinstance(image, list) else image.size
|
|
||||||
edit_width, edit_height, _ = calculate_dimensions(
|
edit_width, edit_height, _ = calculate_dimensions(
|
||||||
1024 * 1024, image_size[0] / image_size[1]
|
1024 * 1024, image_size[0] / image_size[1]
|
||||||
)
|
)
|
||||||
@@ -262,6 +262,9 @@ class QwenImageEditPipelineConfig(QwenImagePipelineConfig):
|
|||||||
"freqs_cis": ((img_cos, img_sin), (txt_cos, txt_sin)),
|
"freqs_cis": ((img_cos, img_sin), (txt_cos, txt_sin)),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
def resize_condition_image(self, image, target_width, target_height):
|
||||||
|
return resize(image, target_height, target_width, resize_mode="default")
|
||||||
|
|
||||||
def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype):
|
def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype):
|
||||||
return self._prepare_edit_cond_kwargs(
|
return self._prepare_edit_cond_kwargs(
|
||||||
batch, batch.prompt_embeds, rotary_emb, device, dtype
|
batch, batch.prompt_embeds, rotary_emb, device, dtype
|
||||||
@@ -272,26 +275,11 @@ class QwenImageEditPipelineConfig(QwenImagePipelineConfig):
|
|||||||
batch, batch.negative_prompt_embeds, rotary_emb, device, dtype
|
batch, batch.negative_prompt_embeds, rotary_emb, device, dtype
|
||||||
)
|
)
|
||||||
|
|
||||||
def preprocess_image(self, image, image_processor):
|
def calculate_condition_image_size(self, image, width, height) -> tuple[int, int]:
|
||||||
image_size = image[0].size if isinstance(image, list) else image.size
|
|
||||||
calculated_width, calculated_height, _ = calculate_dimensions(
|
calculated_width, calculated_height, _ = calculate_dimensions(
|
||||||
1024 * 1024, image_size[0] / image_size[1]
|
1024 * 1024, width / height
|
||||||
)
|
)
|
||||||
image = image_processor.resize(image, calculated_height, calculated_width)
|
return calculated_width, calculated_height
|
||||||
return image
|
|
||||||
|
|
||||||
def maybe_resize_condition_image(self, width, height, image):
|
|
||||||
image_size = image[0].size if isinstance(image, list) else image.size
|
|
||||||
calculated_width, calculated_height, _ = calculate_dimensions(
|
|
||||||
1024 * 1024, image_size[0] / image_size[1]
|
|
||||||
)
|
|
||||||
height = height or calculated_height
|
|
||||||
width = width or calculated_width
|
|
||||||
|
|
||||||
multiple_of = self.get_vae_scale_factor() * 2
|
|
||||||
width = width // multiple_of * multiple_of
|
|
||||||
height = height // multiple_of * multiple_of
|
|
||||||
return width, height
|
|
||||||
|
|
||||||
def slice_noise_pred(self, noise, latents):
|
def slice_noise_pred(self, noise, latents):
|
||||||
# remove noise over input image
|
# remove noise over input image
|
||||||
|
|||||||
@@ -288,7 +288,7 @@ class DiffGenerator:
|
|||||||
# TODO: send batch when supported
|
# TODO: send batch when supported
|
||||||
for request_idx, req in enumerate(requests):
|
for request_idx, req in enumerate(requests):
|
||||||
logger.info(
|
logger.info(
|
||||||
"Processing prompt: %d/%d: %s",
|
"Processing prompt %d/%d: %s",
|
||||||
request_idx + 1,
|
request_idx + 1,
|
||||||
len(requests),
|
len(requests),
|
||||||
req.prompt[:100],
|
req.prompt[:100],
|
||||||
|
|||||||
@@ -55,6 +55,8 @@ class Req:
|
|||||||
image_path: str | None = None
|
image_path: str | None = None
|
||||||
# Image encoder hidden states
|
# Image encoder hidden states
|
||||||
image_embeds: list[torch.Tensor] = field(default_factory=list)
|
image_embeds: list[torch.Tensor] = field(default_factory=list)
|
||||||
|
|
||||||
|
original_condition_image_size: tuple[int, int] = None
|
||||||
condition_image: torch.Tensor | PIL.Image.Image | None = None
|
condition_image: torch.Tensor | PIL.Image.Image | None = None
|
||||||
pixel_values: torch.Tensor | PIL.Image.Image | None = None
|
pixel_values: torch.Tensor | PIL.Image.Image | None = None
|
||||||
preprocessed_image: torch.Tensor | None = None
|
preprocessed_image: torch.Tensor | None = None
|
||||||
@@ -216,14 +218,6 @@ class Req:
|
|||||||
self.guidance_scale_2 = self.guidance_scale
|
self.guidance_scale_2 = self.guidance_scale
|
||||||
|
|
||||||
def adjust_size(self, server_args: ServerArgs):
|
def adjust_size(self, server_args: ServerArgs):
|
||||||
if self.height is None or self.width is None:
|
|
||||||
_image, width, height = (
|
|
||||||
server_args.pipeline_config.maybe_resize_condition_image(
|
|
||||||
self.width, self.height, self.condition_image
|
|
||||||
)
|
|
||||||
)
|
|
||||||
self.width = width
|
|
||||||
self.height = height
|
|
||||||
if self.height is None or self.width is None:
|
if self.height is None or self.width is None:
|
||||||
self.width = 1280
|
self.width = 1280
|
||||||
self.height = 720
|
self.height = 720
|
||||||
|
|||||||
@@ -92,6 +92,9 @@ class PipelineStage(ABC):
|
|||||||
# Default implementation - no verification
|
# Default implementation - no verification
|
||||||
return VerificationResult()
|
return VerificationResult()
|
||||||
|
|
||||||
|
def maybe_free_model_hooks(self):
|
||||||
|
pass
|
||||||
|
|
||||||
# execute on all ranks by default
|
# execute on all ranks by default
|
||||||
@property
|
@property
|
||||||
def parallelism_type(self) -> StageParallelismType:
|
def parallelism_type(self) -> StageParallelismType:
|
||||||
|
|||||||
@@ -199,7 +199,6 @@ class DecodingStage(PipelineStage):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Offload models if needed
|
# Offload models if needed
|
||||||
if hasattr(self, "maybe_free_model_hooks"):
|
|
||||||
self.maybe_free_model_hooks()
|
self.maybe_free_model_hooks()
|
||||||
|
|
||||||
if server_args.vae_cpu_offload:
|
if server_args.vae_cpu_offload:
|
||||||
|
|||||||
@@ -97,7 +97,6 @@ class EncodingStage(PipelineStage):
|
|||||||
batch.latents = latents
|
batch.latents = latents
|
||||||
|
|
||||||
# Offload models if needed
|
# Offload models if needed
|
||||||
if hasattr(self, "maybe_free_model_hooks"):
|
|
||||||
self.maybe_free_model_hooks()
|
self.maybe_free_model_hooks()
|
||||||
|
|
||||||
if server_args.vae_cpu_offload:
|
if server_args.vae_cpu_offload:
|
||||||
|
|||||||
@@ -29,7 +29,6 @@ from sglang.multimodal_gen.runtime.models.vision_utils import (
|
|||||||
normalize,
|
normalize,
|
||||||
numpy_to_pt,
|
numpy_to_pt,
|
||||||
pil_to_numpy,
|
pil_to_numpy,
|
||||||
resize,
|
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
|
||||||
@@ -39,7 +38,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
|||||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
||||||
VerificationResult,
|
VerificationResult,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.server_args import ExecutionMode, ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
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.utils import PRECISION_TO_TYPE
|
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||||
|
|
||||||
@@ -88,7 +87,7 @@ class ImageEncodingStage(PipelineStage):
|
|||||||
prompt_embeds = qwen_image_postprocess_text(outputs, image_inputs, 64)
|
prompt_embeds = qwen_image_postprocess_text(outputs, image_inputs, 64)
|
||||||
return prompt_embeds
|
return prompt_embeds
|
||||||
|
|
||||||
@torch.inference_mode()
|
@torch.no_grad()
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
batch: Req,
|
batch: Req,
|
||||||
@@ -112,11 +111,6 @@ class ImageEncodingStage(PipelineStage):
|
|||||||
|
|
||||||
image = batch.condition_image
|
image = batch.condition_image
|
||||||
|
|
||||||
# preprocess via vae_image_processor
|
|
||||||
prompt_image = server_args.pipeline_config.preprocess_image(
|
|
||||||
image, self.vae_image_processor
|
|
||||||
)
|
|
||||||
|
|
||||||
if batch.prompt and (
|
if batch.prompt and (
|
||||||
isinstance(server_args.pipeline_config, QwenImageEditPipelineConfig)
|
isinstance(server_args.pipeline_config, QwenImageEditPipelineConfig)
|
||||||
or isinstance(server_args.pipeline_config, QwenImagePipelineConfig)
|
or isinstance(server_args.pipeline_config, QwenImagePipelineConfig)
|
||||||
@@ -128,7 +122,7 @@ class ImageEncodingStage(PipelineStage):
|
|||||||
image_processor_kwargs = {}
|
image_processor_kwargs = {}
|
||||||
|
|
||||||
image_inputs = self.image_processor(
|
image_inputs = self.image_processor(
|
||||||
images=prompt_image, return_tensors="pt", **image_processor_kwargs
|
images=image, return_tensors="pt", **image_processor_kwargs
|
||||||
).to(cuda_device)
|
).to(cuda_device)
|
||||||
if self.image_encoder:
|
if self.image_encoder:
|
||||||
# if an image encoder is provided
|
# if an image encoder is provided
|
||||||
@@ -151,7 +145,7 @@ class ImageEncodingStage(PipelineStage):
|
|||||||
neg_image_processor_kwargs = {}
|
neg_image_processor_kwargs = {}
|
||||||
|
|
||||||
neg_image_inputs = self.image_processor(
|
neg_image_inputs = self.image_processor(
|
||||||
images=prompt_image, return_tensors="pt", **neg_image_processor_kwargs
|
images=image, return_tensors="pt", **neg_image_processor_kwargs
|
||||||
).to(get_local_torch_device())
|
).to(get_local_torch_device())
|
||||||
|
|
||||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||||
@@ -203,7 +197,7 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
Stage for encoding pixel representations into latent space.
|
Stage for encoding pixel representations into latent space.
|
||||||
|
|
||||||
This stage handles the encoding of pixel representations into the final
|
This stage handles the encoding of pixel representations into the final
|
||||||
input format (e.g., latents).
|
input format (e.g., image_latents).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, vae: ParallelTiledVAE, **kwargs) -> None:
|
def __init__(self, vae: ParallelTiledVAE, **kwargs) -> None:
|
||||||
@@ -229,35 +223,20 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
if batch.condition_image is None:
|
if batch.condition_image is None:
|
||||||
return batch
|
return batch
|
||||||
|
|
||||||
if server_args.mode == ExecutionMode.INFERENCE:
|
|
||||||
assert batch.condition_image is not None and isinstance(
|
assert batch.condition_image is not None and isinstance(
|
||||||
batch.condition_image, PIL.Image.Image
|
batch.condition_image, PIL.Image.Image
|
||||||
)
|
)
|
||||||
assert batch.height is not None and isinstance(batch.height, int)
|
assert batch.height is not None and isinstance(batch.height, int)
|
||||||
assert batch.width is not None and isinstance(batch.width, int)
|
assert batch.width is not None and isinstance(batch.width, int)
|
||||||
assert batch.num_frames is not None and isinstance(batch.num_frames, int)
|
assert batch.num_frames is not None and isinstance(batch.num_frames, int)
|
||||||
height = batch.height
|
|
||||||
width = batch.width
|
|
||||||
num_frames = batch.num_frames
|
num_frames = batch.num_frames
|
||||||
elif server_args.mode == ExecutionMode.PREPROCESS:
|
|
||||||
assert batch.condition_image is not None and isinstance(
|
|
||||||
batch.condition_image, torch.Tensor
|
|
||||||
)
|
|
||||||
assert batch.height is not None and isinstance(batch.height, list)
|
|
||||||
assert batch.width is not None and isinstance(batch.width, list)
|
|
||||||
assert batch.num_frames is not None and isinstance(batch.num_frames, list)
|
|
||||||
num_frames = batch.num_frames[0]
|
|
||||||
height = batch.height[0]
|
|
||||||
width = batch.width[0]
|
|
||||||
|
|
||||||
self.vae = self.vae.to(get_local_torch_device())
|
self.vae = self.vae.to(get_local_torch_device())
|
||||||
|
|
||||||
image = batch.condition_image
|
image = batch.condition_image
|
||||||
image = self.preprocess(
|
image = self.preprocess(
|
||||||
image,
|
image,
|
||||||
vae_scale_factor=self.vae.spatial_compression_ratio,
|
|
||||||
height=height,
|
|
||||||
width=width,
|
|
||||||
).to(get_local_torch_device(), dtype=torch.float32)
|
).to(get_local_torch_device(), dtype=torch.float32)
|
||||||
|
|
||||||
# (B, C, H, W) -> (B, C, 1, H, W)
|
# (B, C, H, W) -> (B, C, 1, H, W)
|
||||||
@@ -303,9 +282,6 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
video_condition
|
video_condition
|
||||||
)
|
)
|
||||||
|
|
||||||
if server_args.mode == ExecutionMode.PREPROCESS:
|
|
||||||
latent_condition = encoder_output.mean
|
|
||||||
else:
|
|
||||||
generator = batch.generator
|
generator = batch.generator
|
||||||
if generator is None:
|
if generator is None:
|
||||||
raise ValueError("Generator must be provided")
|
raise ValueError("Generator must be provided")
|
||||||
@@ -343,9 +319,6 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
|
|
||||||
batch_size = batch.batch_size
|
batch_size = batch.batch_size
|
||||||
|
|
||||||
if server_args.mode == ExecutionMode.PREPROCESS:
|
|
||||||
batch.image_latent = latent_condition
|
|
||||||
else:
|
|
||||||
# TODO: abstract this
|
# TODO: abstract this
|
||||||
if isinstance(server_args.pipeline_config, QwenImageEditPipelineConfig):
|
if isinstance(server_args.pipeline_config, QwenImageEditPipelineConfig):
|
||||||
if (
|
if (
|
||||||
@@ -353,9 +326,7 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
and batch_size % latent_condition.shape[0] == 0
|
and batch_size % latent_condition.shape[0] == 0
|
||||||
):
|
):
|
||||||
# expand init_latents for batch_size
|
# expand init_latents for batch_size
|
||||||
additional_image_per_prompt = (
|
additional_image_per_prompt = batch_size // latent_condition.shape[0]
|
||||||
batch_size // latent_condition.shape[0]
|
|
||||||
)
|
|
||||||
image_latents = torch.cat(
|
image_latents = torch.cat(
|
||||||
[latent_condition] * additional_image_per_prompt, dim=0
|
[latent_condition] * additional_image_per_prompt, dim=0
|
||||||
)
|
)
|
||||||
@@ -370,8 +341,7 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
image_latents = torch.cat([latent_condition], dim=0)
|
image_latents = torch.cat([latent_condition], dim=0)
|
||||||
image_latent_height, image_latent_width = image_latents.shape[3:]
|
image_latent_height, image_latent_width = image_latents.shape[3:]
|
||||||
num_channels_latents = (
|
num_channels_latents = (
|
||||||
self.server_args.pipeline_config.dit_config.arch_config.in_channels
|
self.server_args.pipeline_config.dit_config.arch_config.in_channels // 4
|
||||||
// 4
|
|
||||||
)
|
)
|
||||||
image_latents = _pack_latents(
|
image_latents = _pack_latents(
|
||||||
image_latents,
|
image_latents,
|
||||||
@@ -403,11 +373,9 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
image_latents = image_latents.unsqueeze(0) # (1, N*1024, 128)
|
image_latents = image_latents.unsqueeze(0) # (1, N*1024, 128)
|
||||||
image_latents = image_latents.repeat(batch_size, 1, 1)
|
image_latents = image_latents.repeat(batch_size, 1, 1)
|
||||||
else:
|
else:
|
||||||
latent_height = height // self.vae.spatial_compression_ratio
|
latent_height = batch.height // self.vae.spatial_compression_ratio
|
||||||
latent_width = width // self.vae.spatial_compression_ratio
|
latent_width = batch.width // self.vae.spatial_compression_ratio
|
||||||
mask_lat_size = torch.ones(
|
mask_lat_size = torch.ones(1, 1, num_frames, latent_height, latent_width)
|
||||||
1, 1, num_frames, latent_height, latent_width
|
|
||||||
)
|
|
||||||
mask_lat_size[:, :, list(range(1, num_frames))] = 0
|
mask_lat_size[:, :, list(range(1, num_frames))] = 0
|
||||||
first_frame_mask = mask_lat_size[:, :, 0:1]
|
first_frame_mask = mask_lat_size[:, :, 0:1]
|
||||||
first_frame_mask = torch.repeat_interleave(
|
first_frame_mask = torch.repeat_interleave(
|
||||||
@@ -431,8 +399,6 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
|
|
||||||
batch.image_latent = image_latents
|
batch.image_latent = image_latents
|
||||||
|
|
||||||
# Offload models if needed
|
|
||||||
if hasattr(self, "maybe_free_model_hooks"):
|
|
||||||
self.maybe_free_model_hooks()
|
self.maybe_free_model_hooks()
|
||||||
|
|
||||||
self.vae.to("cpu")
|
self.vae.to("cpu")
|
||||||
@@ -455,19 +421,9 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
def preprocess(
|
def preprocess(
|
||||||
self,
|
self,
|
||||||
image: torch.Tensor | PIL.Image.Image,
|
image: torch.Tensor | PIL.Image.Image,
|
||||||
vae_scale_factor: int,
|
|
||||||
height: int | None = None,
|
|
||||||
width: int | None = None,
|
|
||||||
resize_mode: str = "default", # "default", "fill", "crop"
|
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
|
|
||||||
if isinstance(image, PIL.Image.Image):
|
if isinstance(image, PIL.Image.Image):
|
||||||
width, height = (
|
|
||||||
self.server_args.pipeline_config.vae_config.calculate_dimensions(
|
|
||||||
image, vae_scale_factor, width, height
|
|
||||||
)
|
|
||||||
)
|
|
||||||
image = resize(image, height, width, resize_mode=resize_mode)
|
|
||||||
image = pil_to_numpy(image) # to np
|
image = pil_to_numpy(image) # to np
|
||||||
image = numpy_to_pt(image) # to pt
|
image = numpy_to_pt(image) # to pt
|
||||||
|
|
||||||
@@ -483,11 +439,6 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
"""Verify encoding stage inputs."""
|
"""Verify encoding stage inputs."""
|
||||||
result = VerificationResult()
|
result = VerificationResult()
|
||||||
result.add_check("generator", batch.generator, V.generator_or_list_generators)
|
result.add_check("generator", batch.generator, V.generator_or_list_generators)
|
||||||
if server_args.mode == ExecutionMode.PREPROCESS:
|
|
||||||
result.add_check("height", batch.height, V.list_not_empty)
|
|
||||||
result.add_check("width", batch.width, V.list_not_empty)
|
|
||||||
result.add_check("num_frames", batch.num_frames, V.list_not_empty)
|
|
||||||
else:
|
|
||||||
result.add_check("height", batch.height, V.positive_int)
|
result.add_check("height", batch.height, V.positive_int)
|
||||||
result.add_check("width", batch.width, V.positive_int)
|
result.add_check("width", batch.width, V.positive_int)
|
||||||
result.add_check("num_frames", batch.num_frames, V.positive_int)
|
result.add_check("num_frames", batch.num_frames, V.positive_int)
|
||||||
|
|||||||
@@ -53,6 +53,97 @@ class InputValidationStage(PipelineStage):
|
|||||||
# FIXME: the generator's in latent preparation stage seems to be different from seeds
|
# FIXME: the generator's in latent preparation stage seems to be different from seeds
|
||||||
batch.generator = [torch.Generator("cpu").manual_seed(seed) for seed in seeds]
|
batch.generator = [torch.Generator("cpu").manual_seed(seed) for seed in seeds]
|
||||||
|
|
||||||
|
# def preprocess_condition_image(self, batch: Req, server_args: ServerArgs, condition_image_width,
|
||||||
|
# condition_image_height):
|
||||||
|
# """
|
||||||
|
# resize condition image
|
||||||
|
# NOTE: condition image resizing is only allowed to do in InputValidationStage
|
||||||
|
# """
|
||||||
|
# if server_args.pipeline_config.task_type == ModelTaskType.I2I:
|
||||||
|
# # calculate new condition image size
|
||||||
|
# calculated_size = (
|
||||||
|
# server_args.pipeline_config.calculate_condition_image_size(
|
||||||
|
# batch.condition_image,
|
||||||
|
# condition_image_width,
|
||||||
|
# condition_image_height,
|
||||||
|
# )
|
||||||
|
# )
|
||||||
|
#
|
||||||
|
# # resize condition image if necessary
|
||||||
|
# if calculated_size is not None:
|
||||||
|
# calculated_width, calculated_height = calculated_size
|
||||||
|
# condition_image = (
|
||||||
|
# server_args.pipeline_config.resize_condition_image(
|
||||||
|
# batch.condition_image, calculated_width, calculated_height
|
||||||
|
# )
|
||||||
|
# )
|
||||||
|
# batch.condition_image = condition_image
|
||||||
|
#
|
||||||
|
# # adjust output image size
|
||||||
|
# calculated_width, calculated_height = batch.condition_image.size
|
||||||
|
# width = calculated_width if batch.width_not_provided else batch.width
|
||||||
|
# height = (
|
||||||
|
# calculated_height if batch.height_not_provided else batch.height
|
||||||
|
# )
|
||||||
|
# multiple_of = (
|
||||||
|
# server_args.pipeline_config.vae_config.get_vae_scale_factor() * 2
|
||||||
|
# )
|
||||||
|
# width = width // multiple_of * multiple_of
|
||||||
|
# height = height // multiple_of * multiple_of
|
||||||
|
# batch.width = width
|
||||||
|
# batch.height = height
|
||||||
|
# else:
|
||||||
|
# if isinstance(server_args.pipeline_config, WanI2V480PConfig):
|
||||||
|
# # TODO: could we merge with above?
|
||||||
|
# # resize image only, Wan2.1 I2V
|
||||||
|
# max_area = 720 * 1280
|
||||||
|
# aspect_ratio = condition_image_height / condition_image_width
|
||||||
|
# mod_value = (
|
||||||
|
# server_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
|
||||||
|
# * server_args.pipeline_config.dit_config.arch_config.patch_size[1]
|
||||||
|
# )
|
||||||
|
# height = round(np.sqrt(max_area * aspect_ratio)) // mod_value * mod_value
|
||||||
|
# width = round(np.sqrt(max_area / aspect_ratio)) // mod_value * mod_value
|
||||||
|
#
|
||||||
|
# batch.condition_image = batch.condition_image.resize((width, height))
|
||||||
|
# batch.height = height
|
||||||
|
# batch.width = width
|
||||||
|
#
|
||||||
|
# if (
|
||||||
|
# server_args.pipeline_config.task_type == ModelTaskType.TI2V
|
||||||
|
# ):
|
||||||
|
# # duplicate with vae_image_processor
|
||||||
|
# # further processing for ti2v task
|
||||||
|
# img = batch.condition_image
|
||||||
|
# ih, iw = img.height, img.width
|
||||||
|
# patch_size = server_args.pipeline_config.dit_config.arch_config.patch_size
|
||||||
|
# vae_stride = (
|
||||||
|
# server_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
|
||||||
|
# )
|
||||||
|
# dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
|
||||||
|
# max_area = 704 * 1280
|
||||||
|
# ow, oh = best_output_size(iw, ih, dw, dh, max_area)
|
||||||
|
#
|
||||||
|
# scale = max(ow / iw, oh / ih)
|
||||||
|
# img = img.resize((round(iw * scale), round(ih * scale)), Image.LANCZOS)
|
||||||
|
# logger.info("resized img height: %s, img width: %s", img.height, img.width)
|
||||||
|
#
|
||||||
|
# # center-crop
|
||||||
|
# x1 = (img.width - ow) // 2
|
||||||
|
# y1 = (img.height - oh) // 2
|
||||||
|
# img = img.crop((x1, y1, x1 + ow, y1 + oh))
|
||||||
|
# assert img.width == ow and img.height == oh
|
||||||
|
#
|
||||||
|
# # to tensor
|
||||||
|
# img = TF.to_tensor(img).sub_(0.5).div_(0.5).to(self.device).unsqueeze(1)
|
||||||
|
# img = img.unsqueeze(0)
|
||||||
|
# batch.height = oh
|
||||||
|
# batch.width = ow
|
||||||
|
# # TODO: should we store in a new field: pixel values?
|
||||||
|
# height = batch.height
|
||||||
|
# width = batch.width
|
||||||
|
# batch.condition_image = resize(img, height, width)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
batch: Req,
|
batch: Req,
|
||||||
@@ -117,24 +208,44 @@ class InputValidationStage(PipelineStage):
|
|||||||
image = load_image(batch.image_path)
|
image = load_image(batch.image_path)
|
||||||
batch.condition_image = image
|
batch.condition_image = image
|
||||||
condition_image_width, condition_image_height = image.width, image.height
|
condition_image_width, condition_image_height = image.width, image.height
|
||||||
else:
|
batch.original_condition_image_size = image.size
|
||||||
condition_image_width, condition_image_height = None, None
|
|
||||||
|
|
||||||
# NOTE: resizing needs to be bring in advance
|
# self.preprocess_condition_image(batch, server_args, condition_image_width, condition_image_height)
|
||||||
|
# NOTE: condition image resizing is only allowed to do in InputValidationStage
|
||||||
if server_args.pipeline_config.task_type == ModelTaskType.I2I:
|
if server_args.pipeline_config.task_type == ModelTaskType.I2I:
|
||||||
if batch.condition_image is not None:
|
if batch.condition_image is not None:
|
||||||
resized_image, resized_width, resized_height = (
|
# calculate new condition image size
|
||||||
server_args.pipeline_config.maybe_resize_condition_image(
|
calculated_size = (
|
||||||
|
server_args.pipeline_config.calculate_condition_image_size(
|
||||||
|
batch.condition_image,
|
||||||
condition_image_width,
|
condition_image_width,
|
||||||
condition_image_height,
|
condition_image_height,
|
||||||
batch.condition_image,
|
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
batch.condition_image = resized_image
|
|
||||||
batch.width = resized_width if batch.width_not_provided else batch.width
|
# resize condition image if necessary
|
||||||
batch.height = (
|
if calculated_size is not None:
|
||||||
resized_height if batch.height_not_provided else batch.height
|
calculated_width, calculated_height = calculated_size
|
||||||
|
condition_image = (
|
||||||
|
server_args.pipeline_config.resize_condition_image(
|
||||||
|
image, calculated_width, calculated_height
|
||||||
)
|
)
|
||||||
|
)
|
||||||
|
batch.condition_image = condition_image
|
||||||
|
|
||||||
|
# adjust output image size
|
||||||
|
calculated_width, calculated_height = batch.condition_image.size
|
||||||
|
width = calculated_width if batch.width_not_provided else batch.width
|
||||||
|
height = (
|
||||||
|
calculated_height if batch.height_not_provided else batch.height
|
||||||
|
)
|
||||||
|
multiple_of = (
|
||||||
|
server_args.pipeline_config.vae_config.get_vae_scale_factor() * 2
|
||||||
|
)
|
||||||
|
width = width // multiple_of * multiple_of
|
||||||
|
height = height // multiple_of * multiple_of
|
||||||
|
batch.width = width
|
||||||
|
batch.height = height
|
||||||
elif (
|
elif (
|
||||||
server_args.pipeline_config.task_type == ModelTaskType.TI2V
|
server_args.pipeline_config.task_type == ModelTaskType.TI2V
|
||||||
) and batch.condition_image is not None:
|
) and batch.condition_image is not None:
|
||||||
|
|||||||
@@ -156,9 +156,6 @@ class ExecutionMode(str, Enum):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
INFERENCE = "inference"
|
INFERENCE = "inference"
|
||||||
PREPROCESS = "preprocess"
|
|
||||||
FINETUNING = "finetuning"
|
|
||||||
DISTILLATION = "distillation"
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_string(cls, value: str) -> "ExecutionMode":
|
def from_string(cls, value: str) -> "ExecutionMode":
|
||||||
|
|||||||
Reference in New Issue
Block a user