[diffusion] refactor: refactor condition image resize logic (#14079)

This commit is contained in:
Mick
2025-11-28 14:06:34 +08:00
committed by GitHub
parent e12c78aab6
commit 3543a04a48
13 changed files with 268 additions and 233 deletions
@@ -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,8 +199,7 @@ 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:
self.vae.to("cpu") self.vae.to("cpu")
@@ -97,8 +97,7 @@ 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:
self.vae.to("cpu") self.vae.to("cpu")
@@ -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 num_frames = batch.num_frames
width = batch.width
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,21 +282,18 @@ class ImageVAEEncodingStage(PipelineStage):
video_condition video_condition
) )
if server_args.mode == ExecutionMode.PREPROCESS: generator = batch.generator
latent_condition = encoder_output.mean if generator is None:
else: raise ValueError("Generator must be provided")
generator = batch.generator # TODO: verify
if generator is None: sample_mode = (
raise ValueError("Generator must be provided") "argmax"
# TODO: verify if server_args.pipeline_config.task_type == ModelTaskType.I2I
sample_mode = ( else "sample"
"argmax" )
if server_args.pipeline_config.task_type == ModelTaskType.I2I latent_condition = self.retrieve_latents(
else "sample" encoder_output, generator, sample_mode=sample_mode
) )
latent_condition = self.retrieve_latents(
encoder_output, generator, sample_mode=sample_mode
)
latent_condition = self.server_args.pipeline_config.post_process_vae_encode( latent_condition = self.server_args.pipeline_config.post_process_vae_encode(
latent_condition, self.vae latent_condition, self.vae
@@ -343,97 +319,87 @@ class ImageVAEEncodingStage(PipelineStage):
batch_size = batch.batch_size batch_size = batch.batch_size
if server_args.mode == ExecutionMode.PREPROCESS: # TODO: abstract this
batch.image_latent = latent_condition if isinstance(server_args.pipeline_config, QwenImageEditPipelineConfig):
else: if (
# TODO: abstract this batch_size > latent_condition.shape[0]
if isinstance(server_args.pipeline_config, QwenImageEditPipelineConfig): and batch_size % latent_condition.shape[0] == 0
if ( ):
batch_size > latent_condition.shape[0] # expand init_latents for batch_size
and batch_size % latent_condition.shape[0] == 0 additional_image_per_prompt = batch_size // latent_condition.shape[0]
): image_latents = torch.cat(
# expand init_latents for batch_size [latent_condition] * additional_image_per_prompt, dim=0
additional_image_per_prompt = (
batch_size // latent_condition.shape[0]
)
image_latents = torch.cat(
[latent_condition] * additional_image_per_prompt, dim=0
)
elif (
batch_size > latent_condition.shape[0]
and batch_size % latent_condition.shape[0] != 0
):
raise ValueError(
f"Cannot duplicate `image` of batch size {latent_condition.shape[0]} to {batch_size} text prompts."
)
else:
image_latents = torch.cat([latent_condition], dim=0)
image_latent_height, image_latent_width = image_latents.shape[3:]
num_channels_latents = (
self.server_args.pipeline_config.dit_config.arch_config.in_channels
// 4
) )
image_latents = _pack_latents( elif (
image_latents, batch_size > latent_condition.shape[0]
batch_size, and batch_size % latent_condition.shape[0] != 0
num_channels_latents, ):
image_latent_height, raise ValueError(
image_latent_width, f"Cannot duplicate `image` of batch size {latent_condition.shape[0]} to {batch_size} text prompts."
) )
elif isinstance(server_args.pipeline_config, Flux2PipelineConfig):
# Pack each latent and concatenate
image_latents = [latent_condition]
# get image_latent_ids right after scale & shift
image_latent_ids = _prepare_image_ids(image_latents)
image_latent_ids = image_latent_ids.repeat(batch_size, 1, 1)
image_latent_ids = image_latent_ids.to(get_local_torch_device())
batch.condition_image_latent_ids = image_latent_ids
packed_latents = []
for latent in image_latents:
# latent: (1, 128, 32, 32)
packed = server_args.pipeline_config.maybe_pack_latents(
latent, None, None
) # (1, 1024, 128)
packed = packed.squeeze(0) # (1024, 128) - remove batch dim
packed_latents.append(packed)
# Concatenate all reference tokens along sequence dimension
image_latents = torch.cat(packed_latents, dim=0) # (N*1024, 128)
image_latents = image_latents.unsqueeze(0) # (1, N*1024, 128)
image_latents = image_latents.repeat(batch_size, 1, 1)
else: else:
latent_height = height // self.vae.spatial_compression_ratio image_latents = torch.cat([latent_condition], dim=0)
latent_width = width // self.vae.spatial_compression_ratio image_latent_height, image_latent_width = image_latents.shape[3:]
mask_lat_size = torch.ones( num_channels_latents = (
1, 1, num_frames, latent_height, latent_width self.server_args.pipeline_config.dit_config.arch_config.in_channels // 4
) )
mask_lat_size[:, :, list(range(1, num_frames))] = 0 image_latents = _pack_latents(
first_frame_mask = mask_lat_size[:, :, 0:1] image_latents,
first_frame_mask = torch.repeat_interleave( batch_size,
first_frame_mask, num_channels_latents,
repeats=self.vae.temporal_compression_ratio, image_latent_height,
dim=2, image_latent_width,
) )
mask_lat_size = torch.concat( elif isinstance(server_args.pipeline_config, Flux2PipelineConfig):
[first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2 # Pack each latent and concatenate
) image_latents = [latent_condition]
mask_lat_size = mask_lat_size.view( # get image_latent_ids right after scale & shift
1, image_latent_ids = _prepare_image_ids(image_latents)
-1, image_latent_ids = image_latent_ids.repeat(batch_size, 1, 1)
self.vae.temporal_compression_ratio, image_latent_ids = image_latent_ids.to(get_local_torch_device())
latent_height, batch.condition_image_latent_ids = image_latent_ids
latent_width,
)
mask_lat_size = mask_lat_size.transpose(1, 2)
mask_lat_size = mask_lat_size.to(latent_condition.device)
image_latents = torch.concat([mask_lat_size, latent_condition], dim=1)
batch.image_latent = image_latents packed_latents = []
for latent in image_latents:
# latent: (1, 128, 32, 32)
packed = server_args.pipeline_config.maybe_pack_latents(
latent, None, None
) # (1, 1024, 128)
packed = packed.squeeze(0) # (1024, 128) - remove batch dim
packed_latents.append(packed)
# Offload models if needed # Concatenate all reference tokens along sequence dimension
if hasattr(self, "maybe_free_model_hooks"): image_latents = torch.cat(packed_latents, dim=0) # (N*1024, 128)
self.maybe_free_model_hooks() image_latents = image_latents.unsqueeze(0) # (1, N*1024, 128)
image_latents = image_latents.repeat(batch_size, 1, 1)
else:
latent_height = batch.height // self.vae.spatial_compression_ratio
latent_width = batch.width // self.vae.spatial_compression_ratio
mask_lat_size = torch.ones(1, 1, num_frames, latent_height, latent_width)
mask_lat_size[:, :, list(range(1, num_frames))] = 0
first_frame_mask = mask_lat_size[:, :, 0:1]
first_frame_mask = torch.repeat_interleave(
first_frame_mask,
repeats=self.vae.temporal_compression_ratio,
dim=2,
)
mask_lat_size = torch.concat(
[first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2
)
mask_lat_size = mask_lat_size.view(
1,
-1,
self.vae.temporal_compression_ratio,
latent_height,
latent_width,
)
mask_lat_size = mask_lat_size.transpose(1, 2)
mask_lat_size = mask_lat_size.to(latent_condition.device)
image_latents = torch.concat([mask_lat_size, latent_condition], dim=1)
batch.image_latent = image_latents
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,14 +439,9 @@ 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.positive_int)
result.add_check("height", batch.height, V.list_not_empty) result.add_check("width", batch.width, V.positive_int)
result.add_check("width", batch.width, V.list_not_empty) result.add_check("num_frames", batch.num_frames, V.positive_int)
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("width", batch.width, V.positive_int)
result.add_check("num_frames", batch.num_frames, V.positive_int)
return result return result
def verify_output(self, batch: Req, server_args: ServerArgs) -> VerificationResult: def verify_output(self, batch: Req, server_args: ServerArgs) -> VerificationResult:
@@ -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":