[diffusion] Opt in Qwen and Wan multi-output conditioning expansion (#31233)
This commit is contained in:
@@ -675,6 +675,10 @@ class PipelineConfig:
|
|||||||
def get_neg_prompt_embeds(self, batch):
|
def get_neg_prompt_embeds(self, batch):
|
||||||
return batch.negative_prompt_embeds
|
return batch.negative_prompt_embeds
|
||||||
|
|
||||||
|
def expand_conditioning_to_sample_batch(self, batch):
|
||||||
|
"""Used for single-request multi-output generation case."""
|
||||||
|
return batch
|
||||||
|
|
||||||
def post_denoising_loop(self, latents, batch):
|
def post_denoising_loop(self, latents, batch):
|
||||||
latents = maybe_unpad_latents(latents, batch)
|
latents = maybe_unpad_latents(latents, batch)
|
||||||
return latents
|
return latents
|
||||||
|
|||||||
@@ -22,6 +22,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
|||||||
from sglang.multimodal_gen.configs.post_training.pipeline_configs import (
|
from sglang.multimodal_gen.configs.post_training.pipeline_configs import (
|
||||||
QwenImageRolloutPipelineMixin,
|
QwenImageRolloutPipelineMixin,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.utils.condition_expansion import (
|
||||||
|
PromptToSampleBatchExpander,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.vision import resize
|
from sglang.multimodal_gen.runtime.utils.vision import resize
|
||||||
from sglang.multimodal_gen.utils import calculate_dimensions
|
from sglang.multimodal_gen.utils import calculate_dimensions
|
||||||
|
|
||||||
@@ -181,6 +184,24 @@ class QwenImagePipelineConfig(QwenImageRolloutPipelineMixin, ImagePipelineConfig
|
|||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def expand_conditioning_to_sample_batch(self, batch):
|
||||||
|
expander = PromptToSampleBatchExpander.from_batch(batch)
|
||||||
|
if expander is None:
|
||||||
|
return batch
|
||||||
|
|
||||||
|
for field_name in (
|
||||||
|
"prompt_embeds",
|
||||||
|
"negative_prompt_embeds",
|
||||||
|
"prompt_attention_mask",
|
||||||
|
"negative_attention_mask",
|
||||||
|
"prompt_embeds_mask",
|
||||||
|
"negative_prompt_embeds_mask",
|
||||||
|
"prompt_seq_lens",
|
||||||
|
"negative_prompt_seq_lens",
|
||||||
|
):
|
||||||
|
expander.expand_field(batch, field_name)
|
||||||
|
return batch
|
||||||
|
|
||||||
def tokenize_prompt(self, prompts: list[str], tokenizer, tok_kwargs) -> dict:
|
def tokenize_prompt(self, prompts: list[str], tokenizer, tok_kwargs) -> dict:
|
||||||
tok_kwargs.setdefault("truncation", True)
|
tok_kwargs.setdefault("truncation", True)
|
||||||
|
|
||||||
|
|||||||
@@ -21,6 +21,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
|||||||
from sglang.multimodal_gen.configs.pipeline_configs.model_deployment_config import (
|
from sglang.multimodal_gen.configs.pipeline_configs.model_deployment_config import (
|
||||||
ModelDeploymentConfig,
|
ModelDeploymentConfig,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.utils.condition_expansion import (
|
||||||
|
PromptToSampleBatchExpander,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
@@ -97,6 +100,20 @@ class WanT2V480PConfig(PipelineConfig):
|
|||||||
auto_dit_layerwise_offload=True,
|
auto_dit_layerwise_offload=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def expand_conditioning_to_sample_batch(self, batch):
|
||||||
|
expander = PromptToSampleBatchExpander.from_batch(batch)
|
||||||
|
if expander is None:
|
||||||
|
return batch
|
||||||
|
|
||||||
|
for field_name in (
|
||||||
|
"prompt_embeds",
|
||||||
|
"negative_prompt_embeds",
|
||||||
|
"image_embeds",
|
||||||
|
"image_latent",
|
||||||
|
):
|
||||||
|
expander.expand_field(batch, field_name)
|
||||||
|
return batch
|
||||||
|
|
||||||
def get_pos_prompt_embeds(self, batch):
|
def get_pos_prompt_embeds(self, batch):
|
||||||
return batch.prompt_embeds[0]
|
return batch.prompt_embeds[0]
|
||||||
|
|
||||||
|
|||||||
@@ -754,6 +754,8 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
Returns:
|
Returns:
|
||||||
A context object containing the invariant state for the denoising loop.
|
A context object containing the invariant state for the denoising loop.
|
||||||
"""
|
"""
|
||||||
|
batch = server_args.pipeline_config.expand_conditioning_to_sample_batch(batch)
|
||||||
|
|
||||||
assert self.transformer is not None
|
assert self.transformer is not None
|
||||||
pipeline = self.pipeline() if self.pipeline else None
|
pipeline = self.pipeline() if self.pipeline else None
|
||||||
scheduler = batch.scheduler
|
scheduler = batch.scheduler
|
||||||
|
|||||||
@@ -0,0 +1,111 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class PromptToSampleBatchExpander:
|
||||||
|
"""Expand selected conditioning from prompt order to sample order."""
|
||||||
|
|
||||||
|
prompt_batch_size: int
|
||||||
|
sample_batch_size: int
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_batch(cls, batch):
|
||||||
|
num_outputs = int(batch.num_outputs_per_prompt or 1)
|
||||||
|
if num_outputs <= 1:
|
||||||
|
return None
|
||||||
|
if isinstance(batch.prompt, list):
|
||||||
|
prompt_batch_size = len(batch.prompt)
|
||||||
|
elif batch.prompt is not None:
|
||||||
|
prompt_batch_size = 1
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
"Multi-output conditioning requires prompt text so the prompt "
|
||||||
|
"batch size is unambiguous."
|
||||||
|
)
|
||||||
|
if prompt_batch_size <= 0:
|
||||||
|
raise ValueError("Multi-output conditioning requires at least one prompt.")
|
||||||
|
return cls(prompt_batch_size, prompt_batch_size * num_outputs)
|
||||||
|
|
||||||
|
def _expand_tensor(self, value: torch.Tensor, name: str) -> torch.Tensor:
|
||||||
|
current_batch_size = value.shape[0]
|
||||||
|
if current_batch_size == self.sample_batch_size:
|
||||||
|
return value
|
||||||
|
if current_batch_size != self.prompt_batch_size:
|
||||||
|
raise ValueError(
|
||||||
|
f"{name} has batch dim {current_batch_size} (shape "
|
||||||
|
f"{tuple(value.shape)}); expected {self.prompt_batch_size} "
|
||||||
|
f"(per-prompt) or {self.sample_batch_size} (per-sample)."
|
||||||
|
)
|
||||||
|
repeats = self.sample_batch_size // self.prompt_batch_size
|
||||||
|
return value.repeat_interleave(repeats, dim=0)
|
||||||
|
|
||||||
|
def _expand_tensors(self, value, name: str):
|
||||||
|
"""Expand a tensor or each tensor in a list, preserving its container."""
|
||||||
|
if value is None:
|
||||||
|
return None
|
||||||
|
if isinstance(value, torch.Tensor):
|
||||||
|
return self._expand_tensor(value, name)
|
||||||
|
if not isinstance(value, list):
|
||||||
|
raise TypeError(f"{name} must be a tensor, list of tensors, or None.")
|
||||||
|
if any(
|
||||||
|
item is not None and not isinstance(item, torch.Tensor) for item in value
|
||||||
|
):
|
||||||
|
raise TypeError(f"{name} entries must be tensors or None.")
|
||||||
|
return [
|
||||||
|
self._expand_tensor(item, f"{name}[{index}]") if item is not None else None
|
||||||
|
for index, item in enumerate(value)
|
||||||
|
]
|
||||||
|
|
||||||
|
def _expand_sequence_lengths(
|
||||||
|
self, value: list[list[int] | None] | None, name: str
|
||||||
|
) -> list[list[int] | None] | None:
|
||||||
|
if value is None:
|
||||||
|
return None
|
||||||
|
repeats = self.sample_batch_size // self.prompt_batch_size
|
||||||
|
expanded = []
|
||||||
|
for index, sequence_lengths in enumerate(value):
|
||||||
|
if (
|
||||||
|
sequence_lengths is None
|
||||||
|
or len(sequence_lengths) == self.sample_batch_size
|
||||||
|
):
|
||||||
|
expanded.append(sequence_lengths)
|
||||||
|
elif len(sequence_lengths) == self.prompt_batch_size:
|
||||||
|
expanded.append(
|
||||||
|
[
|
||||||
|
sequence_length
|
||||||
|
for sequence_length in sequence_lengths
|
||||||
|
for _ in range(repeats)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
f"{name}[{index}] has {len(sequence_lengths)} entries; expected "
|
||||||
|
f"{self.prompt_batch_size} (per-prompt) or "
|
||||||
|
f"{self.sample_batch_size} (per-sample)."
|
||||||
|
)
|
||||||
|
return expanded
|
||||||
|
|
||||||
|
def expand_field(self, batch, field_name: str) -> None:
|
||||||
|
"""Expand one field in place, dispatching from its value type."""
|
||||||
|
value = getattr(batch, field_name)
|
||||||
|
if value is None:
|
||||||
|
return
|
||||||
|
if isinstance(value, torch.Tensor) or (
|
||||||
|
isinstance(value, list)
|
||||||
|
and all(item is None or isinstance(item, torch.Tensor) for item in value)
|
||||||
|
):
|
||||||
|
expanded = self._expand_tensors(value, field_name)
|
||||||
|
elif isinstance(value, list) and all(
|
||||||
|
item is None or isinstance(item, list) for item in value
|
||||||
|
):
|
||||||
|
expanded = self._expand_sequence_lengths(value, field_name)
|
||||||
|
else:
|
||||||
|
raise TypeError(
|
||||||
|
f"{field_name} must be a tensor, list of tensors, "
|
||||||
|
"list of sequence-length lists, or None."
|
||||||
|
)
|
||||||
|
setattr(batch, field_name, expanded)
|
||||||
Reference in New Issue
Block a user