[diffusion] chore: use native qwen2.5-vl generation (#34896)
This commit is contained in:
@@ -86,6 +86,7 @@ class EncoderConfig(ModelConfig):
|
||||
@dataclass
|
||||
class TextEncoderConfig(EncoderConfig):
|
||||
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
|
||||
generation_config: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -8,6 +8,7 @@ from sglang.multimodal_gen.configs.models import DiTConfig, EncoderConfig, VAECo
|
||||
from sglang.multimodal_gen.configs.models.dits.longcat_image import (
|
||||
LongCatImageDitConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.models.encoders.qwen_image import Qwen2_5VLConfig
|
||||
from sglang.multimodal_gen.configs.models.vaes.longcat_image import (
|
||||
LongCatImageVAEConfig,
|
||||
)
|
||||
@@ -16,6 +17,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
||||
ModelTaskType,
|
||||
TextConditioningOutput,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.model_deployment_config import (
|
||||
ModelDeploymentConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -219,18 +223,6 @@ def longcat_postprocess_text(outputs, text_inputs, pipeline_config):
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatImageEncoderConfig(EncoderConfig):
|
||||
"""Encoder config for the in-stage-loaded HF Qwen2.5-VL text encoder.
|
||||
|
||||
The encoder weights are loaded by `LongCatPromptRewriteStage` (not via
|
||||
`TextEncoderLoader`), so this config only supplies the fields the standard
|
||||
`TextEncodingStage` reads — primarily `tokenizer_kwargs`.
|
||||
"""
|
||||
|
||||
tokenizer_kwargs: dict = field(default_factory=lambda: {})
|
||||
|
||||
|
||||
@dataclass
|
||||
class LongCatImagePipelineConfig(ImagePipelineConfig):
|
||||
"""Configuration for the LongCat-Image T2I pipeline."""
|
||||
@@ -246,16 +238,20 @@ class LongCatImagePipelineConfig(ImagePipelineConfig):
|
||||
dit_config: DiTConfig = field(default_factory=LongCatImageDitConfig)
|
||||
vae_config: VAEConfig = field(default_factory=LongCatImageVAEConfig)
|
||||
|
||||
# The Qwen2.5-VL text encoder (~7B) is loaded in bf16; the encoder is loaded
|
||||
# in-stage by LongCatPromptRewriteStage, not via TextEncoderLoader.
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",))
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (LongCatImageEncoderConfig(),)
|
||||
default_factory=lambda: (Qwen2_5VLConfig(),)
|
||||
)
|
||||
postprocess_text_funcs: tuple[Callable, ...] = field(
|
||||
default_factory=lambda: (longcat_postprocess_text,)
|
||||
)
|
||||
|
||||
def get_model_deployment_config(self) -> ModelDeploymentConfig:
|
||||
return ModelDeploymentConfig(
|
||||
keep_resident_min_available_gb=70,
|
||||
keep_resident_components=("text_encoder", "vae"),
|
||||
)
|
||||
|
||||
# --- LatentPreparationStage hooks ---
|
||||
|
||||
def prepare_latent_shape(self, batch, batch_size, num_frames):
|
||||
|
||||
@@ -19,6 +19,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
||||
pad_text_embeddings_with_mask,
|
||||
shard_rotary_emb_for_sp,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.model_deployment_config import (
|
||||
ModelDeploymentConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.post_training.pipeline_configs import (
|
||||
QwenImageRolloutPipelineMixin,
|
||||
)
|
||||
@@ -753,6 +756,12 @@ class QwenImageLayeredPipelineConfig(QwenImageEditPipelineConfig):
|
||||
resolution: int = 640
|
||||
vae_precision: str = "bf16"
|
||||
|
||||
def get_model_deployment_config(self) -> ModelDeploymentConfig:
|
||||
return ModelDeploymentConfig(
|
||||
keep_resident_min_available_gb=70,
|
||||
keep_resident_components=("text_encoder", "vae"),
|
||||
)
|
||||
|
||||
def postprocess_cfg_noise(
|
||||
self,
|
||||
batch,
|
||||
|
||||
@@ -52,6 +52,7 @@ from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
||||
get_config,
|
||||
get_diffusers_component_config,
|
||||
load_dict,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.precision import precision_to_dtype
|
||||
@@ -344,6 +345,9 @@ class TextEncoderLoader(ComponentLoader):
|
||||
|
||||
encoder_config = server_args.pipeline_config.text_encoder_configs[encoder_index]
|
||||
encoder_config.update_model_arch(model_config)
|
||||
encoder_config.generation_config = load_dict(
|
||||
os.path.join(component_model_path, "generation_config.json")
|
||||
)
|
||||
|
||||
if encoder_index == 0:
|
||||
for key, value in diffusers_pretrained_config.__dict__.items():
|
||||
|
||||
@@ -29,6 +29,9 @@ from sglang.multimodal_gen.runtime.layers.linear import (
|
||||
from sglang.multimodal_gen.runtime.layers.quantization import QuantizationConfig
|
||||
from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader
|
||||
from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl_vision import (
|
||||
Qwen2_5VLVisionTransformer,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
||||
from sglang.multimodal_gen.runtime.utils.common import add_prefix
|
||||
|
||||
@@ -69,8 +72,6 @@ import torch
|
||||
import torch.nn as nn
|
||||
from transformers.activations import ACT2FN
|
||||
from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
|
||||
Qwen2_5_VisionRotaryEmbedding,
|
||||
Qwen2_5_VisionTransformerPretrainedModel,
|
||||
Qwen2_5_VLCausalLMOutputWithPast,
|
||||
Qwen2_5_VLModelOutputWithPast,
|
||||
Qwen2_5_VLRotaryEmbedding,
|
||||
@@ -80,6 +81,53 @@ from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _apply_repetition_penalty(
|
||||
logits: torch.Tensor,
|
||||
input_ids: torch.LongTensor,
|
||||
penalty: float,
|
||||
) -> torch.Tensor:
|
||||
if penalty == 1.0:
|
||||
return logits
|
||||
selected_logits = torch.gather(logits, 1, input_ids)
|
||||
selected_logits = torch.where(
|
||||
selected_logits < 0,
|
||||
selected_logits * penalty,
|
||||
selected_logits / penalty,
|
||||
)
|
||||
return logits.scatter(1, input_ids, selected_logits)
|
||||
|
||||
|
||||
def _select_next_token(
|
||||
logits: torch.Tensor,
|
||||
input_ids: torch.LongTensor,
|
||||
*,
|
||||
do_sample: bool,
|
||||
temperature: float,
|
||||
top_k: int,
|
||||
top_p: float,
|
||||
repetition_penalty: float,
|
||||
) -> torch.LongTensor:
|
||||
scores = _apply_repetition_penalty(logits.float(), input_ids, repetition_penalty)
|
||||
if not do_sample or top_k == 1:
|
||||
return torch.argmax(scores, dim=-1)
|
||||
|
||||
scores = scores / temperature
|
||||
if 0 < top_k < scores.shape[-1]:
|
||||
top_k_threshold = torch.topk(scores, top_k, dim=-1).values[..., -1, None]
|
||||
scores = scores.masked_fill(scores < top_k_threshold, -torch.inf)
|
||||
if top_p < 1.0:
|
||||
sorted_scores, sorted_indices = torch.sort(scores, descending=True)
|
||||
cumulative_probs = torch.softmax(sorted_scores, dim=-1).cumsum(dim=-1)
|
||||
sorted_remove = cumulative_probs > top_p
|
||||
sorted_remove[..., 1:] = sorted_remove[..., :-1].clone()
|
||||
sorted_remove[..., 0] = False
|
||||
remove = torch.zeros_like(sorted_remove).scatter(
|
||||
1, sorted_indices, sorted_remove
|
||||
)
|
||||
scores = scores.masked_fill(remove, -torch.inf)
|
||||
return torch.multinomial(torch.softmax(scores, dim=-1), num_samples=1).squeeze(1)
|
||||
|
||||
|
||||
def _tp_world_size() -> int:
|
||||
if not model_parallel_is_initialized():
|
||||
return 1
|
||||
@@ -261,7 +309,15 @@ class Qwen2_5_VLAttention(nn.Module):
|
||||
query_states = query_states.transpose(1, 2)
|
||||
key_states = key_states.transpose(1, 2)
|
||||
value_states = value_states.transpose(1, 2)
|
||||
attn_output = self.attn(query_states, key_states, value_states)
|
||||
# Diffusion text encoding is cache-free and historically uses the native
|
||||
# causal kernel; its trailing padding is removed during postprocessing.
|
||||
# Cached generation still needs the explicit mask for padded batches.
|
||||
attn_output = self.attn(
|
||||
query_states,
|
||||
key_states,
|
||||
value_states,
|
||||
attn_mask=attention_mask if use_cache else None,
|
||||
)
|
||||
|
||||
attn_output = attn_output.reshape(bsz, q_len, -1).contiguous()
|
||||
attn_output = _linear_output(self.o_proj, attn_output)
|
||||
@@ -597,28 +653,14 @@ class Qwen2_5_VLModel(nn.Module):
|
||||
_checkpoint_conversion_mapping = {"^model": "language_model"}
|
||||
# Reference: fix gemma3 grad acc #37208
|
||||
accepts_loss_kwargs = False
|
||||
_no_split_modules = ["Qwen2_5_VLDecoderLayer", "Qwen2_5_VLVisionBlock"]
|
||||
_no_split_modules = ["Qwen2_5_VLDecoderLayer", "Qwen2_5VLVisionBlock"]
|
||||
|
||||
def __init__(self, config, enable_image_understanding: bool = False):
|
||||
super().__init__()
|
||||
self.language_model = Qwen2_5_VLTextModel(config.text_config)
|
||||
|
||||
if enable_image_understanding:
|
||||
self.visual = Qwen2_5_VisionTransformerPretrainedModel._from_config(
|
||||
config.vision_config
|
||||
)
|
||||
self.visual.to(torch.get_default_dtype())
|
||||
# keeps the vision rotary frequencies in fp32 even when weights are bf16 (as HF does)
|
||||
head_dim = (
|
||||
config.vision_config.hidden_size // config.vision_config.num_heads
|
||||
)
|
||||
rotary_dim = head_dim // 2
|
||||
inv_freq = Qwen2_5_VisionRotaryEmbedding(rotary_dim).inv_freq
|
||||
self.visual.rotary_pos_emb.register_buffer(
|
||||
"inv_freq",
|
||||
inv_freq,
|
||||
persistent=False,
|
||||
)
|
||||
self.visual = Qwen2_5VLVisionTransformer(config.vision_config)
|
||||
self.rope_deltas = None # cache rope_deltas here
|
||||
self.config = config
|
||||
# Initialize weights and apply final processing
|
||||
@@ -902,11 +944,6 @@ class Qwen2_5_VLModel(nn.Module):
|
||||
"""
|
||||
pixel_values = pixel_values.type(self.visual.dtype)
|
||||
image_embeds = self.visual(pixel_values, grid_thw=image_grid_thw)
|
||||
if not isinstance(image_embeds, torch.Tensor):
|
||||
# In transformers v5, the visual encoder returns BaseModelOutputWithPooling.
|
||||
# pooler_output contains the spatially merged embeddings (what we need),
|
||||
# while last_hidden_state contains the raw unmerged output.
|
||||
image_embeds = image_embeds.pooler_output
|
||||
split_sizes = (
|
||||
image_grid_thw.prod(-1) // self.visual.spatial_merge_size**2
|
||||
).tolist()
|
||||
@@ -1106,6 +1143,9 @@ class Qwen2_5_VLModel(nn.Module):
|
||||
|
||||
|
||||
class Qwen2_5_VLForConditionalGeneration(TextEncoder):
|
||||
layer_names = [*TextEncoder.layer_names, "model.visual.blocks"]
|
||||
_fsdp_forward_methods = ("generate",)
|
||||
|
||||
# BitandBytes specific attributes
|
||||
default_bitsandbytes_target_modules = [
|
||||
".gate_up_proj.",
|
||||
@@ -1132,6 +1172,7 @@ class Qwen2_5_VLForConditionalGeneration(TextEncoder):
|
||||
) -> None:
|
||||
super().__init__(config)
|
||||
enable_image_understanding = config.enable_image_understanding
|
||||
generation_config = config.generation_config
|
||||
config = config.arch_config
|
||||
self.model = Qwen2_5_VLModel(
|
||||
config, enable_image_understanding=enable_image_understanding
|
||||
@@ -1141,6 +1182,7 @@ class Qwen2_5_VLForConditionalGeneration(TextEncoder):
|
||||
)
|
||||
|
||||
self.enable_image_understanding = enable_image_understanding
|
||||
self.generation_config = generation_config
|
||||
|
||||
self.config = config
|
||||
|
||||
@@ -1225,6 +1267,152 @@ class Qwen2_5_VLForConditionalGeneration(TextEncoder):
|
||||
rope_deltas=outputs.rope_deltas,
|
||||
)
|
||||
|
||||
@torch.no_grad()
|
||||
def generate(
|
||||
self,
|
||||
input_ids: torch.LongTensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
*,
|
||||
max_new_tokens: int,
|
||||
pixel_values: Optional[torch.Tensor] = None,
|
||||
pixel_values_videos: Optional[torch.FloatTensor] = None,
|
||||
image_grid_thw: Optional[torch.LongTensor] = None,
|
||||
video_grid_thw: Optional[torch.LongTensor] = None,
|
||||
second_per_grid_ts: Optional[torch.Tensor] = None,
|
||||
mm_token_type_ids: Optional[torch.IntTensor] = None,
|
||||
do_sample: Optional[bool] = None,
|
||||
temperature: Optional[float] = None,
|
||||
top_k: Optional[int] = None,
|
||||
top_p: Optional[float] = None,
|
||||
repetition_penalty: Optional[float] = None,
|
||||
eos_token_id: Optional[Union[int, list[int]]] = None,
|
||||
pad_token_id: Optional[int] = None,
|
||||
) -> torch.LongTensor:
|
||||
"""Generate tokens with Qwen2.5-VL's native decoder and KV cache."""
|
||||
# Transformers 5 processors emit this field. The Qwen2.5-VL checkpoint
|
||||
# derives modality positions from image/video placeholder token IDs.
|
||||
del mm_token_type_ids
|
||||
|
||||
if input_ids.ndim != 2:
|
||||
raise ValueError("input_ids must have shape [batch_size, sequence_length]")
|
||||
if max_new_tokens < 0:
|
||||
raise ValueError("max_new_tokens must be non-negative")
|
||||
if max_new_tokens == 0:
|
||||
return input_ids
|
||||
|
||||
generation_config = self.generation_config
|
||||
do_sample = (
|
||||
generation_config.get("do_sample", False)
|
||||
if do_sample is None
|
||||
else do_sample
|
||||
)
|
||||
temperature = (
|
||||
generation_config.get("temperature", 1.0)
|
||||
if temperature is None
|
||||
else temperature
|
||||
)
|
||||
top_k = generation_config.get("top_k", 0) if top_k is None else top_k
|
||||
top_p = generation_config.get("top_p", 1.0) if top_p is None else top_p
|
||||
repetition_penalty = (
|
||||
generation_config.get("repetition_penalty", 1.0)
|
||||
if repetition_penalty is None
|
||||
else repetition_penalty
|
||||
)
|
||||
eos_token_id = (
|
||||
generation_config.get("eos_token_id", self.config.eos_token_id)
|
||||
if eos_token_id is None
|
||||
else eos_token_id
|
||||
)
|
||||
pad_token_id = (
|
||||
generation_config.get("pad_token_id", self.config.pad_token_id)
|
||||
if pad_token_id is None
|
||||
else pad_token_id
|
||||
)
|
||||
|
||||
if do_sample and temperature <= 0:
|
||||
raise ValueError("temperature must be positive when sampling")
|
||||
if top_k < 0:
|
||||
raise ValueError("top_k must be non-negative")
|
||||
if not 0 < top_p <= 1:
|
||||
raise ValueError("top_p must be in (0, 1]")
|
||||
if repetition_penalty <= 0:
|
||||
raise ValueError("repetition_penalty must be positive")
|
||||
|
||||
eos_token_ids = (
|
||||
[]
|
||||
if eos_token_id is None
|
||||
else [eos_token_id] if isinstance(eos_token_id, int) else list(eos_token_id)
|
||||
)
|
||||
if pad_token_id is None:
|
||||
raise ValueError("pad_token_id must be set for generation")
|
||||
eos_tokens = torch.tensor(
|
||||
eos_token_ids, dtype=input_ids.dtype, device=input_ids.device
|
||||
)
|
||||
if attention_mask is None:
|
||||
attention_mask = torch.ones_like(input_ids)
|
||||
|
||||
generated_ids = input_ids
|
||||
unfinished = torch.ones(
|
||||
input_ids.shape[0], dtype=torch.bool, device=input_ids.device
|
||||
)
|
||||
past_key_values = None
|
||||
model_input_ids = input_ids
|
||||
cache_position = torch.arange(input_ids.shape[1], device=input_ids.device)
|
||||
self.model.rope_deltas = None
|
||||
|
||||
for _ in range(max_new_tokens):
|
||||
outputs = self(
|
||||
input_ids=model_input_ids,
|
||||
attention_mask=attention_mask,
|
||||
past_key_values=past_key_values,
|
||||
use_cache=True,
|
||||
cache_position=cache_position,
|
||||
pixel_values=pixel_values,
|
||||
pixel_values_videos=pixel_values_videos,
|
||||
image_grid_thw=image_grid_thw,
|
||||
video_grid_thw=video_grid_thw,
|
||||
second_per_grid_ts=second_per_grid_ts,
|
||||
logits_to_keep=1,
|
||||
)
|
||||
next_tokens = _select_next_token(
|
||||
outputs.logits[:, -1, :],
|
||||
generated_ids,
|
||||
do_sample=do_sample,
|
||||
temperature=temperature,
|
||||
top_k=top_k,
|
||||
top_p=top_p,
|
||||
repetition_penalty=repetition_penalty,
|
||||
)
|
||||
next_tokens = torch.where(
|
||||
unfinished,
|
||||
next_tokens,
|
||||
torch.full_like(next_tokens, pad_token_id),
|
||||
)
|
||||
generated_ids = torch.cat([generated_ids, next_tokens[:, None]], dim=-1)
|
||||
|
||||
if eos_tokens.numel() > 0:
|
||||
reached_eos = (next_tokens[:, None] == eos_tokens[None, :]).any(dim=-1)
|
||||
unfinished = unfinished & ~reached_eos
|
||||
if not unfinished.any():
|
||||
break
|
||||
|
||||
past_key_values = outputs.past_key_values
|
||||
model_input_ids = next_tokens[:, None]
|
||||
attention_mask = torch.cat(
|
||||
[attention_mask, attention_mask.new_ones((input_ids.shape[0], 1))],
|
||||
dim=-1,
|
||||
)
|
||||
cache_position = torch.tensor(
|
||||
[generated_ids.shape[1] - 1], device=input_ids.device
|
||||
)
|
||||
pixel_values = None
|
||||
pixel_values_videos = None
|
||||
image_grid_thw = None
|
||||
video_grid_thw = None
|
||||
second_per_grid_ts = None
|
||||
|
||||
return generated_ids
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
|
||||
loaded_params: set[str] = set()
|
||||
|
||||
|
||||
@@ -0,0 +1,461 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Native Qwen2.5-VL vision encoder."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from sglang.multimodal_gen.runtime.layers.attention.selector import get_attn_backend
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _PackedSequenceMetadata:
|
||||
cu_seqlens: torch.Tensor
|
||||
cu_seqlens_host: tuple[int, ...]
|
||||
max_seqlen: int
|
||||
|
||||
@classmethod
|
||||
def from_cu_seqlens(cls, cu_seqlens: torch.Tensor) -> _PackedSequenceMetadata:
|
||||
bounds = tuple(int(value) for value in cu_seqlens.tolist())
|
||||
return cls(
|
||||
cu_seqlens=cu_seqlens,
|
||||
cu_seqlens_host=bounds,
|
||||
max_seqlen=max(
|
||||
stop - start for start, stop in zip(bounds[:-1], bounds[1:])
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class Qwen2_5VLVisionRMSNorm(nn.Module):
|
||||
def __init__(self, hidden_size: int, eps: float = 1e-6) -> None:
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(hidden_size))
|
||||
self.variance_epsilon = eps
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
input_dtype = hidden_states.dtype
|
||||
hidden_states = hidden_states.float()
|
||||
variance = hidden_states.square().mean(dim=-1, keepdim=True)
|
||||
hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
|
||||
return self.weight * hidden_states.to(input_dtype)
|
||||
|
||||
|
||||
class Qwen2_5VLVisionPatchEmbed(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: int,
|
||||
temporal_patch_size: int,
|
||||
in_channels: int,
|
||||
embed_dim: int,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.temporal_patch_size = temporal_patch_size
|
||||
self.in_channels = in_channels
|
||||
self.embed_dim = embed_dim
|
||||
kernel_size = (temporal_patch_size, patch_size, patch_size)
|
||||
self.proj = nn.Conv3d(
|
||||
in_channels,
|
||||
embed_dim,
|
||||
kernel_size=kernel_size,
|
||||
stride=kernel_size,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = hidden_states.view(
|
||||
-1,
|
||||
self.in_channels,
|
||||
self.temporal_patch_size,
|
||||
self.patch_size,
|
||||
self.patch_size,
|
||||
)
|
||||
return self.proj(hidden_states.to(self.proj.weight.dtype)).view(
|
||||
-1, self.embed_dim
|
||||
)
|
||||
|
||||
|
||||
class Qwen2_5VLVisionRotaryEmbedding(nn.Module):
|
||||
def __init__(self, dim: int, theta: float = 10000.0) -> None:
|
||||
super().__init__()
|
||||
inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim))
|
||||
self.register_buffer("inv_freq", inv_freq, persistent=False)
|
||||
|
||||
def forward(self, position_ids: torch.Tensor) -> torch.Tensor:
|
||||
return (position_ids.unsqueeze(-1) * self.inv_freq).flatten(1)
|
||||
|
||||
|
||||
def _rotate_half(hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
first, second = hidden_states.chunk(2, dim=-1)
|
||||
return torch.cat((-second, first), dim=-1)
|
||||
|
||||
|
||||
def _apply_vision_rotary_embedding(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
query_dtype = query.dtype
|
||||
key_dtype = key.dtype
|
||||
query = query.float()
|
||||
key = key.float()
|
||||
cos = cos.unsqueeze(-2).float()
|
||||
sin = sin.unsqueeze(-2).float()
|
||||
query = query * cos + _rotate_half(query) * sin
|
||||
key = key * cos + _rotate_half(key) * sin
|
||||
return query.to(query_dtype), key.to(key_dtype)
|
||||
|
||||
|
||||
class Qwen2_5VLVisionAttention(nn.Module):
|
||||
def __init__(self, config: Any, prefix: str) -> None:
|
||||
super().__init__()
|
||||
self.num_heads = config.num_heads
|
||||
self.head_dim = config.hidden_size // config.num_heads
|
||||
self.scaling = self.head_dim**-0.5
|
||||
self.prefix = prefix
|
||||
self.qkv = nn.Linear(config.hidden_size, config.hidden_size * 3, bias=True)
|
||||
self.proj = nn.Linear(config.hidden_size, config.hidden_size)
|
||||
self._attention_impl = None
|
||||
self._initialize_attention(torch.get_default_dtype())
|
||||
|
||||
def _initialize_attention(self, dtype: torch.dtype) -> None:
|
||||
backend = get_attn_backend(self.head_dim, dtype)
|
||||
if backend.supports_packed_varlen():
|
||||
self._attention_impl = backend.get_impl_cls()(
|
||||
num_heads=self.num_heads,
|
||||
head_size=self.head_dim,
|
||||
num_kv_heads=self.num_heads,
|
||||
softmax_scale=self.scaling,
|
||||
causal=False,
|
||||
prefix=self.prefix,
|
||||
)
|
||||
else:
|
||||
logger.warning_once(
|
||||
"Qwen2.5-VL vision attention uses torch SDPA because "
|
||||
f"{backend.get_enum().name.lower()} does not support packed sequences"
|
||||
)
|
||||
|
||||
def _packed_attention(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
cu_seqlens: torch.Tensor,
|
||||
cu_seqlens_host: tuple[int, ...],
|
||||
max_seqlen: int,
|
||||
) -> torch.Tensor:
|
||||
if self._attention_impl is not None:
|
||||
return self._attention_impl.forward_varlen(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
cu_seqlens=cu_seqlens,
|
||||
cu_seqlens_host=cu_seqlens_host,
|
||||
max_seqlen=max_seqlen,
|
||||
)
|
||||
|
||||
output = torch.empty_like(query)
|
||||
for start, stop in zip(cu_seqlens_host[:-1], cu_seqlens_host[1:]):
|
||||
if start == stop:
|
||||
continue
|
||||
query_segment = query[start:stop].transpose(0, 1).unsqueeze(0)
|
||||
key_segment = key[start:stop].transpose(0, 1).unsqueeze(0)
|
||||
value_segment = value[start:stop].transpose(0, 1).unsqueeze(0)
|
||||
segment = F.scaled_dot_product_attention(
|
||||
query_segment,
|
||||
key_segment,
|
||||
value_segment,
|
||||
dropout_p=0.0,
|
||||
is_causal=False,
|
||||
scale=self.scaling,
|
||||
)
|
||||
output[start:stop] = segment.squeeze(0).transpose(0, 1)
|
||||
return output
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
*,
|
||||
cu_seqlens: torch.Tensor,
|
||||
cu_seqlens_host: tuple[int, ...],
|
||||
max_seqlen: int,
|
||||
position_embeddings: tuple[torch.Tensor, torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
seq_len = hidden_states.shape[0]
|
||||
query, key, value = (
|
||||
self.qkv(hidden_states)
|
||||
.reshape(seq_len, 3, self.num_heads, self.head_dim)
|
||||
.permute(1, 0, 2, 3)
|
||||
.unbind(0)
|
||||
)
|
||||
query, key = _apply_vision_rotary_embedding(query, key, *position_embeddings)
|
||||
output = self._packed_attention(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
cu_seqlens,
|
||||
cu_seqlens_host,
|
||||
max_seqlen,
|
||||
)
|
||||
return self.proj(output.reshape(seq_len, -1).contiguous())
|
||||
|
||||
|
||||
class Qwen2_5VLVisionMLP(nn.Module):
|
||||
def __init__(self, config: Any) -> None:
|
||||
super().__init__()
|
||||
if config.hidden_act != "silu":
|
||||
raise ValueError(
|
||||
f"Unsupported Qwen2.5-VL vision activation: {config.hidden_act}"
|
||||
)
|
||||
self.gate_proj = nn.Linear(
|
||||
config.hidden_size, config.intermediate_size, bias=True
|
||||
)
|
||||
self.up_proj = nn.Linear(
|
||||
config.hidden_size, config.intermediate_size, bias=True
|
||||
)
|
||||
self.down_proj = nn.Linear(
|
||||
config.intermediate_size, config.hidden_size, bias=True
|
||||
)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
return self.down_proj(
|
||||
F.silu(self.gate_proj(hidden_states)) * self.up_proj(hidden_states)
|
||||
)
|
||||
|
||||
|
||||
class Qwen2_5VLVisionBlock(nn.Module):
|
||||
def __init__(self, config: Any, layer_idx: int) -> None:
|
||||
super().__init__()
|
||||
self.norm1 = Qwen2_5VLVisionRMSNorm(config.hidden_size)
|
||||
self.norm2 = Qwen2_5VLVisionRMSNorm(config.hidden_size)
|
||||
self.attn = Qwen2_5VLVisionAttention(
|
||||
config, prefix=f"visual.blocks.{layer_idx}.attn"
|
||||
)
|
||||
self.mlp = Qwen2_5VLVisionMLP(config)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
*,
|
||||
cu_seqlens: torch.Tensor,
|
||||
cu_seqlens_host: tuple[int, ...],
|
||||
max_seqlen: int,
|
||||
position_embeddings: tuple[torch.Tensor, torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
hidden_states = hidden_states + self.attn(
|
||||
self.norm1(hidden_states),
|
||||
cu_seqlens=cu_seqlens,
|
||||
cu_seqlens_host=cu_seqlens_host,
|
||||
max_seqlen=max_seqlen,
|
||||
position_embeddings=position_embeddings,
|
||||
)
|
||||
return hidden_states + self.mlp(self.norm2(hidden_states))
|
||||
|
||||
|
||||
class Qwen2_5VLVisionPatchMerger(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
output_dim: int,
|
||||
context_dim: int,
|
||||
spatial_merge_size: int,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.hidden_size = context_dim * spatial_merge_size**2
|
||||
self.ln_q = Qwen2_5VLVisionRMSNorm(context_dim)
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(self.hidden_size, self.hidden_size),
|
||||
nn.GELU(),
|
||||
nn.Linear(self.hidden_size, output_dim),
|
||||
)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.ln_q(hidden_states).view(-1, self.hidden_size)
|
||||
return self.mlp(hidden_states)
|
||||
|
||||
|
||||
def _vision_position_ids(
|
||||
grid_thw: torch.Tensor, spatial_merge_size: int
|
||||
) -> torch.Tensor:
|
||||
position_ids = []
|
||||
for t, h, w in grid_thw.tolist():
|
||||
h_positions = torch.arange(h, device=grid_thw.device)[:, None].expand(h, w)
|
||||
h_positions = (
|
||||
h_positions.reshape(
|
||||
h // spatial_merge_size,
|
||||
spatial_merge_size,
|
||||
w // spatial_merge_size,
|
||||
spatial_merge_size,
|
||||
)
|
||||
.transpose(1, 2)
|
||||
.flatten()
|
||||
)
|
||||
w_positions = torch.arange(w, device=grid_thw.device)[None, :].expand(h, w)
|
||||
w_positions = (
|
||||
w_positions.reshape(
|
||||
h // spatial_merge_size,
|
||||
spatial_merge_size,
|
||||
w // spatial_merge_size,
|
||||
spatial_merge_size,
|
||||
)
|
||||
.transpose(1, 2)
|
||||
.flatten()
|
||||
)
|
||||
positions = torch.stack((h_positions, w_positions), dim=-1)
|
||||
position_ids.append(positions.repeat(t, 1))
|
||||
return torch.cat(position_ids, dim=0)
|
||||
|
||||
|
||||
def _vision_window_index(
|
||||
grid_thw: torch.Tensor,
|
||||
*,
|
||||
spatial_merge_size: int,
|
||||
window_size: int,
|
||||
patch_size: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
window_indices = []
|
||||
cumulative_window_lengths = [0]
|
||||
window_index_offset = 0
|
||||
merger_window_size = window_size // spatial_merge_size // patch_size
|
||||
spatial_merge_unit = spatial_merge_size**2
|
||||
|
||||
for grid_t, grid_h, grid_w in grid_thw.tolist():
|
||||
merged_height = grid_h // spatial_merge_size
|
||||
merged_width = grid_w // spatial_merge_size
|
||||
index = torch.arange(
|
||||
grid_t * merged_height * merged_width, device=grid_thw.device
|
||||
).reshape(grid_t, merged_height, merged_width)
|
||||
pad_height = merger_window_size - merged_height % merger_window_size
|
||||
pad_width = merger_window_size - merged_width % merger_window_size
|
||||
num_windows_height = (merged_height + pad_height) // merger_window_size
|
||||
num_windows_width = (merged_width + pad_width) // merger_window_size
|
||||
index = F.pad(index, (0, pad_width, 0, pad_height), value=-100)
|
||||
index = index.reshape(
|
||||
grid_t,
|
||||
num_windows_height,
|
||||
merger_window_size,
|
||||
num_windows_width,
|
||||
merger_window_size,
|
||||
)
|
||||
index = index.permute(0, 1, 3, 2, 4).reshape(
|
||||
grid_t,
|
||||
num_windows_height * num_windows_width,
|
||||
merger_window_size,
|
||||
merger_window_size,
|
||||
)
|
||||
sequence_lengths = (index != -100).sum(dim=(2, 3)).reshape(-1)
|
||||
index = index.flatten()
|
||||
window_indices.append(index[index != -100] + window_index_offset)
|
||||
cumulative = (
|
||||
sequence_lengths.cumsum(0) * spatial_merge_unit
|
||||
+ cumulative_window_lengths[-1]
|
||||
)
|
||||
cumulative_window_lengths.extend(cumulative.tolist())
|
||||
window_index_offset += grid_t * merged_height * merged_width
|
||||
|
||||
window_index = torch.cat(window_indices)
|
||||
cu_window_seqlens = torch.tensor(
|
||||
cumulative_window_lengths, device=grid_thw.device, dtype=torch.int32
|
||||
)
|
||||
return window_index, torch.unique_consecutive(cu_window_seqlens)
|
||||
|
||||
|
||||
class Qwen2_5VLVisionTransformer(nn.Module):
|
||||
def __init__(self, config: Any) -> None:
|
||||
super().__init__()
|
||||
self.spatial_merge_size = config.spatial_merge_size
|
||||
self.spatial_merge_unit = config.spatial_merge_size**2
|
||||
self.patch_size = config.patch_size
|
||||
self.window_size = config.window_size
|
||||
self.full_attention_layers = frozenset(
|
||||
int(layer_idx) for layer_idx in config.fullatt_block_indexes
|
||||
)
|
||||
self.patch_embed = Qwen2_5VLVisionPatchEmbed(
|
||||
patch_size=config.patch_size,
|
||||
temporal_patch_size=config.temporal_patch_size,
|
||||
in_channels=config.in_channels,
|
||||
embed_dim=config.hidden_size,
|
||||
)
|
||||
head_dim = config.hidden_size // config.num_heads
|
||||
self.rotary_pos_emb = Qwen2_5VLVisionRotaryEmbedding(head_dim // 2)
|
||||
self.blocks = nn.ModuleList(
|
||||
Qwen2_5VLVisionBlock(config, layer_idx) for layer_idx in range(config.depth)
|
||||
)
|
||||
self.merger = Qwen2_5VLVisionPatchMerger(
|
||||
output_dim=config.out_hidden_size,
|
||||
context_dim=config.hidden_size,
|
||||
spatial_merge_size=config.spatial_merge_size,
|
||||
)
|
||||
|
||||
@property
|
||||
def dtype(self) -> torch.dtype:
|
||||
return self.patch_embed.proj.weight.dtype
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
return self.patch_embed.proj.weight.device
|
||||
|
||||
def forward(
|
||||
self, hidden_states: torch.Tensor, grid_thw: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
hidden_states = hidden_states.to(device=self.device, dtype=self.dtype)
|
||||
grid_thw = grid_thw.to(self.device)
|
||||
hidden_states = self.patch_embed(hidden_states)
|
||||
|
||||
position_ids = _vision_position_ids(grid_thw, self.spatial_merge_size)
|
||||
window_index, cu_window_seqlens = _vision_window_index(
|
||||
grid_thw,
|
||||
spatial_merge_size=self.spatial_merge_size,
|
||||
window_size=self.window_size,
|
||||
patch_size=self.patch_size,
|
||||
)
|
||||
|
||||
seq_len = hidden_states.shape[0]
|
||||
hidden_states = hidden_states.reshape(
|
||||
seq_len // self.spatial_merge_unit, self.spatial_merge_unit, -1
|
||||
)[window_index]
|
||||
hidden_states = hidden_states.reshape(seq_len, -1)
|
||||
|
||||
rotary = self.rotary_pos_emb(position_ids)
|
||||
rotary = rotary.reshape(
|
||||
seq_len // self.spatial_merge_unit, self.spatial_merge_unit, -1
|
||||
)[window_index]
|
||||
rotary = rotary.reshape(seq_len, -1)
|
||||
rotary = torch.cat((rotary, rotary), dim=-1)
|
||||
position_embeddings = (
|
||||
rotary.cos(),
|
||||
rotary.sin(),
|
||||
)
|
||||
|
||||
cu_seqlens = torch.repeat_interleave(
|
||||
grid_thw[:, 1] * grid_thw[:, 2], grid_thw[:, 0]
|
||||
).cumsum(dim=0, dtype=torch.int32)
|
||||
cu_seqlens = F.pad(cu_seqlens, (1, 0), value=0)
|
||||
|
||||
full_metadata = _PackedSequenceMetadata.from_cu_seqlens(cu_seqlens)
|
||||
window_metadata = _PackedSequenceMetadata.from_cu_seqlens(cu_window_seqlens)
|
||||
|
||||
for layer_idx, block in enumerate(self.blocks):
|
||||
metadata = (
|
||||
full_metadata
|
||||
if layer_idx in self.full_attention_layers
|
||||
else window_metadata
|
||||
)
|
||||
hidden_states = block(
|
||||
hidden_states,
|
||||
cu_seqlens=metadata.cu_seqlens,
|
||||
cu_seqlens_host=metadata.cu_seqlens_host,
|
||||
max_seqlen=metadata.max_seqlen,
|
||||
position_embeddings=position_embeddings,
|
||||
)
|
||||
|
||||
merged = self.merger(hidden_states)
|
||||
return merged[torch.argsort(window_index)]
|
||||
@@ -27,11 +27,8 @@ class LongCatImagePipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
|
||||
pipeline_name = "LongCatImagePipeline"
|
||||
|
||||
# The Qwen2.5-VL text encoder is loaded in-stage by LongCatPromptRewriteStage
|
||||
# (not via TextEncoderLoader), so "text_encoder" is intentionally absent;
|
||||
# the stage registers the loaded module via add_module("text_encoder", ...)
|
||||
# so the standard TextEncodingStage can fetch the same instance.
|
||||
_required_config_modules = [
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"text_processor",
|
||||
"vae",
|
||||
@@ -41,23 +38,17 @@ class LongCatImagePipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
|
||||
def create_pipeline_stages(self, server_args: ServerArgs):
|
||||
# 1. Prompt rewriting (optional) + request-level setup (generator, cfg renorm).
|
||||
# Loads the HF Qwen2.5-VL encoder and shares it with TextEncodingStage.
|
||||
rewrite_stage = LongCatPromptRewriteStage(
|
||||
tokenizer=self.get_module("tokenizer"),
|
||||
text_encoder=self.get_module("text_encoder"),
|
||||
text_processor=self.get_module("text_processor"),
|
||||
model_path=self.model_path,
|
||||
text_encoder_dtype=PRECISION_TO_TYPE[
|
||||
server_args.pipeline_config.text_encoder_precisions[0]
|
||||
],
|
||||
)
|
||||
self.add_stage(rewrite_stage)
|
||||
self.add_module("text_encoder", rewrite_stage.text_encoder)
|
||||
|
||||
# 2. Text encoding via the standard stage (tokenize_prompt +
|
||||
# postprocess_text_funcs hooks on the pipeline config). Shares the
|
||||
# encoder instance registered above; both stages declare a
|
||||
# "text_encoder" ComponentUse so the residency manager keeps it
|
||||
# resident across rewrite->encode and offloads after the last use.
|
||||
# postprocess_text_funcs hooks on the pipeline config).
|
||||
self.add_standard_text_encoding_stage()
|
||||
|
||||
# 3. Latent preparation (batch-size-aware via pipeline config hooks)
|
||||
|
||||
@@ -95,6 +95,7 @@ class QwenImageLayeredPipeline(QwenImageEditPipeline):
|
||||
pipeline_name = "QwenImageLayeredPipeline"
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder",
|
||||
"vae",
|
||||
"tokenizer",
|
||||
"processor",
|
||||
@@ -106,12 +107,11 @@ class QwenImageLayeredPipeline(QwenImageEditPipeline):
|
||||
def create_before_denoising_stage():
|
||||
return QwenImageLayeredBeforeDenoisingStage(
|
||||
vae=self.get_module("vae"),
|
||||
text_encoder=None,
|
||||
text_encoder=self.get_module("text_encoder"),
|
||||
tokenizer=self.get_module("tokenizer"),
|
||||
processor=self.get_module("processor"),
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
model_path=self.model_path,
|
||||
vae_dtype=PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision],
|
||||
text_encoder_dtype=PRECISION_TO_TYPE[
|
||||
server_args.pipeline_config.text_encoder_precisions[0]
|
||||
|
||||
@@ -322,6 +322,7 @@ class ImageEncodingStage(PipelineStage):
|
||||
pixel_values=image_inputs.pixel_values,
|
||||
image_grid_thw=image_inputs.image_grid_thw,
|
||||
output_hidden_states=True,
|
||||
use_cache=False,
|
||||
)
|
||||
if batch.do_classifier_free_guidance:
|
||||
neg_outputs = self.text_encoder(
|
||||
@@ -330,6 +331,7 @@ class ImageEncodingStage(PipelineStage):
|
||||
pixel_values=neg_image_inputs.pixel_values,
|
||||
image_grid_thw=neg_image_inputs.image_grid_thw,
|
||||
output_hidden_states=True,
|
||||
use_cache=False,
|
||||
)
|
||||
|
||||
prompt_embeds, prompt_embeds_mask, prompt_seq_lens = (
|
||||
|
||||
+7
-25
@@ -1,7 +1,7 @@
|
||||
"""Prompt-rewriting stage for LongCat-Image (T2I).
|
||||
|
||||
`LongCatPromptRewriteStage` optionally rewrites the prompt via the HuggingFace
|
||||
`Qwen2_5_VLForConditionalGeneration.generate()` and sets the CPU generator for
|
||||
`LongCatPromptRewriteStage` optionally rewrites the prompt via the native
|
||||
Qwen2.5-VL encoder and sets the CPU generator for
|
||||
seed reproducibility. Text encoding, latent preparation, RoPE and denoising are
|
||||
all handled by the standard stages + `LongCatImagePipelineConfig` hooks, so this
|
||||
is the only model-specific stage.
|
||||
@@ -181,15 +181,13 @@ def _get_prompt_language(prompt):
|
||||
class LongCatPromptRewriteStage(PipelineStage):
|
||||
"""Optional prompt rewriting + request-level setup for LongCat-Image.
|
||||
|
||||
Loads the Qwen2.5-VL text encoder (HuggingFace) in-stage and, when
|
||||
`enable_prompt_rewrite` is set, rewrites the prompt via `.generate()`
|
||||
When `enable_prompt_rewrite` is set, rewrites the prompt via `.generate()`
|
||||
(using the checkpoint's generation_config.json sampling params). Always
|
||||
sets the CPU generator for seed reproducibility. CFG-renorm params are
|
||||
read directly from sampling_params in `postprocess_cfg_noise`, not set here.
|
||||
|
||||
The same encoder instance is shared with the standard `TextEncodingStage`
|
||||
(the pipeline registers it via `add_module("text_encoder", ...)`), so
|
||||
rewrite and encode run on one set of weights. Both stages declare a
|
||||
The same encoder instance is shared with the standard `TextEncodingStage`,
|
||||
so rewrite and encode run on one set of weights. Both stages declare a
|
||||
`text_encoder` ComponentUse; the residency manager keeps the encoder
|
||||
resident across the adjacent rewrite->encode uses and offloads it only
|
||||
after the last use (when `--text-encoder-cpu-offload` is enabled).
|
||||
@@ -197,35 +195,19 @@ class LongCatPromptRewriteStage(PipelineStage):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer,
|
||||
text_encoder,
|
||||
text_processor,
|
||||
model_path: str,
|
||||
text_encoder_dtype: torch.dtype,
|
||||
):
|
||||
super().__init__()
|
||||
self.text_encoder_dtype = text_encoder_dtype
|
||||
from transformers import Qwen2_5_VLForConditionalGeneration
|
||||
|
||||
cpu_offload = self.server_args.text_encoder_cpu_offload
|
||||
init_device = torch.device("cpu") if cpu_offload else get_local_torch_device()
|
||||
self.text_encoder = (
|
||||
Qwen2_5_VLForConditionalGeneration.from_pretrained(
|
||||
model_path, subfolder="text_encoder"
|
||||
)
|
||||
.to(init_device)
|
||||
.to(dtype=self.text_encoder_dtype)
|
||||
)
|
||||
self.tokenizer = tokenizer
|
||||
self.text_encoder = text_encoder
|
||||
self.text_processor = text_processor
|
||||
|
||||
def component_uses(
|
||||
self, server_args: ServerArgs, stage_name: str | None = None
|
||||
) -> list[ComponentUse]:
|
||||
stage_name = self._component_stage_name(stage_name)
|
||||
# "text_encoder" matches is_text_encoder_component_name, so
|
||||
# text_encoder_cpu_offload routes through VanillaD2HStrategy (the encoder
|
||||
# is a plain HF nn.Module, not FSDP-sharded). memory_intensive triggers
|
||||
# torch.cuda.empty_cache() after the encoder leaves CUDA.
|
||||
return [
|
||||
ComponentUse(
|
||||
stage_name,
|
||||
|
||||
+7
-15
@@ -164,7 +164,6 @@ class QwenImageLayeredBeforeDenoisingStage(PipelineStage):
|
||||
processor,
|
||||
transformer,
|
||||
scheduler,
|
||||
model_path,
|
||||
vae_dtype: torch.dtype,
|
||||
text_encoder_dtype: torch.dtype,
|
||||
) -> None:
|
||||
@@ -172,15 +171,7 @@ class QwenImageLayeredBeforeDenoisingStage(PipelineStage):
|
||||
self.vae = vae.to(dtype=vae_dtype)
|
||||
self.vae_dtype = vae_dtype
|
||||
self.text_encoder_dtype = text_encoder_dtype
|
||||
if text_encoder is None:
|
||||
from transformers import Qwen2_5_VLForConditionalGeneration
|
||||
|
||||
text_encoder = Qwen2_5_VLForConditionalGeneration.from_pretrained(
|
||||
model_path, subfolder="text_encoder"
|
||||
)
|
||||
self.text_encoder = text_encoder.to(
|
||||
device=get_local_torch_device(), dtype=self.text_encoder_dtype
|
||||
)
|
||||
self.text_encoder = text_encoder
|
||||
self.tokenizer = tokenizer
|
||||
self.processor = processor
|
||||
self.transformer = transformer
|
||||
@@ -285,11 +276,12 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
|
||||
padding=True,
|
||||
return_tensors="pt",
|
||||
).to(device)
|
||||
encoder_hidden_states = self.text_encoder(
|
||||
input_ids=txt_tokens.input_ids,
|
||||
attention_mask=txt_tokens.attention_mask,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
encoder_hidden_states = self.text_encoder(
|
||||
input_ids=txt_tokens.input_ids,
|
||||
attention_mask=txt_tokens.attention_mask,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
hidden_states = encoder_hidden_states.hidden_states[-1]
|
||||
split_hidden_states = self._extract_masked_hidden(
|
||||
hidden_states, txt_tokens.attention_mask
|
||||
|
||||
@@ -345,9 +345,6 @@ class TestPipelineSpecificExtraModules(unittest.TestCase):
|
||||
extra_allowed_modules=extras,
|
||||
)
|
||||
self.assertEqual(extras, {"vae", "transformer"})
|
||||
self.assertNotIn(
|
||||
"text_encoder", QwenImageLayeredPipeline._required_config_modules
|
||||
)
|
||||
self.assertEqual(
|
||||
set(filtered),
|
||||
{
|
||||
@@ -356,6 +353,7 @@ class TestPipelineSpecificExtraModules(unittest.TestCase):
|
||||
"processor",
|
||||
"transformer",
|
||||
"scheduler",
|
||||
"text_encoder",
|
||||
},
|
||||
)
|
||||
|
||||
@@ -464,21 +462,16 @@ class TestQwenImageLayeredDtype(_GlobalStageArgsMixin, unittest.TestCase):
|
||||
def to(self, *args, **kwargs):
|
||||
return self
|
||||
|
||||
with patch(
|
||||
"sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.qwen_image_layered.get_local_torch_device",
|
||||
return_value=torch.device("cpu"),
|
||||
):
|
||||
stage = QwenImageLayeredBeforeDenoisingStage(
|
||||
vae=_DummyVAE(),
|
||||
text_encoder=torch.nn.Linear(1, 1),
|
||||
tokenizer=object(),
|
||||
processor=object(),
|
||||
transformer=object(),
|
||||
scheduler=object(),
|
||||
model_path="/unused",
|
||||
vae_dtype=torch.float32,
|
||||
text_encoder_dtype=torch.float16,
|
||||
)
|
||||
stage = QwenImageLayeredBeforeDenoisingStage(
|
||||
vae=_DummyVAE(),
|
||||
text_encoder=torch.nn.Linear(1, 1),
|
||||
tokenizer=object(),
|
||||
processor=object(),
|
||||
transformer=object(),
|
||||
scheduler=object(),
|
||||
vae_dtype=torch.float32,
|
||||
text_encoder_dtype=torch.float16,
|
||||
)
|
||||
|
||||
uses = stage.component_uses(SimpleNamespace(), "qwen_layered")
|
||||
self.assertEqual(
|
||||
|
||||
@@ -0,0 +1,226 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
import sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl as qwen2_5vl
|
||||
from sglang.multimodal_gen.configs.models.encoders.qwen_image import Qwen2_5VLConfig
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.longcat_image import (
|
||||
LongCatImagePipelineConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
|
||||
QwenImageLayeredPipelineConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl import (
|
||||
Qwen2_5_VLAttention,
|
||||
Qwen2_5_VLForConditionalGeneration,
|
||||
_apply_repetition_penalty,
|
||||
_select_next_token,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl_vision import (
|
||||
Qwen2_5VLVisionRotaryEmbedding,
|
||||
Qwen2_5VLVisionTransformer,
|
||||
_vision_position_ids,
|
||||
_vision_window_index,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines.longcat_image import LongCatImagePipeline
|
||||
|
||||
|
||||
class _StubQwen2_5VL(Qwen2_5_VLForConditionalGeneration):
|
||||
def __init__(self, next_tokens: list[list[int]], eos_token_id=5):
|
||||
nn.Module.__init__(self)
|
||||
self.model = nn.Module()
|
||||
self.model.rope_deltas = torch.tensor([99])
|
||||
self.config = SimpleNamespace(eos_token_id=eos_token_id, pad_token_id=0)
|
||||
self.generation_config = {
|
||||
"do_sample": True,
|
||||
"temperature": 0.1,
|
||||
"top_k": 1,
|
||||
"top_p": 0.001,
|
||||
"repetition_penalty": 1.05,
|
||||
"eos_token_id": [eos_token_id, 6],
|
||||
"pad_token_id": 0,
|
||||
}
|
||||
self.next_tokens = next_tokens
|
||||
self.calls = []
|
||||
|
||||
def forward(self, input_ids, **kwargs):
|
||||
call_index = len(self.calls)
|
||||
self.calls.append((input_ids.clone(), kwargs))
|
||||
logits = torch.full((input_ids.shape[0], 1, 8), -100.0)
|
||||
for batch_index, token_id in enumerate(self.next_tokens[call_index]):
|
||||
logits[batch_index, 0, token_id] = 100.0
|
||||
return SimpleNamespace(
|
||||
logits=logits,
|
||||
past_key_values=f"cache-{call_index}",
|
||||
)
|
||||
|
||||
|
||||
class _AttentionRecorder(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.masks = []
|
||||
|
||||
def forward(self, query, key, value, attn_mask=None):
|
||||
self.masks.append(attn_mask)
|
||||
return query
|
||||
|
||||
|
||||
def test_explicit_attention_mask_is_limited_to_cached_generation(monkeypatch):
|
||||
attention = Qwen2_5_VLAttention.__new__(Qwen2_5_VLAttention)
|
||||
nn.Module.__init__(attention)
|
||||
attention.q_proj = nn.Identity()
|
||||
attention.k_proj = nn.Identity()
|
||||
attention.v_proj = nn.Identity()
|
||||
attention.o_proj = nn.Identity()
|
||||
attention.num_heads = 1
|
||||
attention.num_key_value_heads = 1
|
||||
attention.head_dim = 4
|
||||
attention.rope_scaling = {"mrope_section": [1, 1, 0]}
|
||||
attention.attn = _AttentionRecorder()
|
||||
monkeypatch.setattr(
|
||||
qwen2_5vl,
|
||||
"apply_multimodal_rotary_pos_emb",
|
||||
lambda query, key, *_args: (query, key),
|
||||
)
|
||||
|
||||
hidden_states = torch.randn(1, 2, 4)
|
||||
explicit_mask = torch.zeros(1, 1, 2, 2)
|
||||
kwargs = {
|
||||
"hidden_states": hidden_states,
|
||||
"attention_mask": explicit_mask,
|
||||
"position_embeddings": (torch.empty(0), torch.empty(0)),
|
||||
}
|
||||
|
||||
attention(**kwargs, use_cache=False)
|
||||
attention(**kwargs, use_cache=True)
|
||||
|
||||
assert attention.attn.masks[0] is None
|
||||
assert attention.attn.masks[1] is explicit_mask
|
||||
|
||||
|
||||
def test_repetition_penalty_matches_sign_dependent_scaling():
|
||||
logits = torch.tensor([[2.0, -3.0, 4.0]])
|
||||
penalized = _apply_repetition_penalty(logits, torch.tensor([[0, 1]]), penalty=2.0)
|
||||
torch.testing.assert_close(penalized, torch.tensor([[1.0, -6.0, 4.0]]))
|
||||
|
||||
|
||||
def test_top_k_one_sampling_is_deterministic():
|
||||
token = _select_next_token(
|
||||
torch.tensor([[1.0, 3.0, 2.0]]),
|
||||
torch.tensor([[0]]),
|
||||
do_sample=True,
|
||||
temperature=0.1,
|
||||
top_k=1,
|
||||
top_p=0.001,
|
||||
repetition_penalty=1.0,
|
||||
)
|
||||
assert token.tolist() == [1]
|
||||
|
||||
|
||||
def test_generate_reuses_cache_and_only_prefills_vision_once():
|
||||
model = _StubQwen2_5VL([[4], [5]])
|
||||
pixel_values = torch.ones(1, 3)
|
||||
generated = model.generate(
|
||||
torch.tensor([[1, 2]]),
|
||||
attention_mask=torch.ones(1, 2, dtype=torch.long),
|
||||
pixel_values=pixel_values,
|
||||
image_grid_thw=torch.tensor([[1, 1, 1]]),
|
||||
mm_token_type_ids=torch.tensor([[0, 1]]),
|
||||
max_new_tokens=3,
|
||||
)
|
||||
|
||||
assert generated.tolist() == [[1, 2, 4, 5]]
|
||||
assert len(model.calls) == 2
|
||||
assert model.calls[0][1]["pixel_values"] is pixel_values
|
||||
assert model.calls[1][1]["pixel_values"] is None
|
||||
assert model.calls[0][1]["past_key_values"] is None
|
||||
assert model.calls[1][1]["past_key_values"] == "cache-0"
|
||||
assert model.calls[0][1]["cache_position"].tolist() == [0, 1]
|
||||
assert model.calls[1][1]["cache_position"].tolist() == [2]
|
||||
assert model.model.rope_deltas is None
|
||||
|
||||
|
||||
def test_generate_pads_finished_rows_until_the_batch_stops():
|
||||
model = _StubQwen2_5VL([[5, 4], [7, 6]])
|
||||
generated = model.generate(
|
||||
torch.tensor([[1], [2]]),
|
||||
max_new_tokens=3,
|
||||
)
|
||||
|
||||
assert generated.tolist() == [[1, 5, 0], [2, 4, 6]]
|
||||
|
||||
|
||||
def test_native_vision_indices_preserve_merged_token_groups():
|
||||
grid_thw = torch.tensor([[1, 4, 4], [2, 2, 4]])
|
||||
position_ids = _vision_position_ids(grid_thw, spatial_merge_size=2)
|
||||
window_index, cu_window_seqlens = _vision_window_index(
|
||||
grid_thw,
|
||||
spatial_merge_size=2,
|
||||
window_size=8,
|
||||
patch_size=2,
|
||||
)
|
||||
|
||||
assert position_ids.shape == (32, 2)
|
||||
assert sorted(window_index.tolist()) == list(range(8))
|
||||
assert cu_window_seqlens[0].item() == 0
|
||||
assert cu_window_seqlens[-1].item() == 32
|
||||
|
||||
|
||||
def test_native_vision_keeps_rotary_trigonometry_in_fp32():
|
||||
class PatchEmbed(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.proj = nn.Linear(1, 8, bias=False, dtype=torch.bfloat16)
|
||||
|
||||
def forward(self, hidden_states):
|
||||
return self.proj(hidden_states)
|
||||
|
||||
class BlockRecorder(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.position_embedding_dtypes = None
|
||||
|
||||
def forward(self, hidden_states, *, position_embeddings, **_kwargs):
|
||||
self.position_embedding_dtypes = tuple(
|
||||
embedding.dtype for embedding in position_embeddings
|
||||
)
|
||||
return hidden_states
|
||||
|
||||
class Merger(nn.Module):
|
||||
def forward(self, hidden_states):
|
||||
return hidden_states.reshape(-1, 4, hidden_states.shape[-1])[:, 0]
|
||||
|
||||
model = Qwen2_5VLVisionTransformer.__new__(Qwen2_5VLVisionTransformer)
|
||||
nn.Module.__init__(model)
|
||||
model.spatial_merge_size = 2
|
||||
model.spatial_merge_unit = 4
|
||||
model.patch_size = 2
|
||||
model.window_size = 8
|
||||
model.full_attention_layers = frozenset({0})
|
||||
model.patch_embed = PatchEmbed()
|
||||
model.rotary_pos_emb = Qwen2_5VLVisionRotaryEmbedding(2)
|
||||
block = BlockRecorder()
|
||||
model.blocks = nn.ModuleList([block])
|
||||
model.merger = Merger()
|
||||
output = model(
|
||||
torch.zeros(16, 1, dtype=torch.bfloat16),
|
||||
grid_thw=torch.tensor([[1, 4, 4]]),
|
||||
)
|
||||
|
||||
assert output.dtype == torch.bfloat16
|
||||
assert block.position_embedding_dtypes == (torch.float32, torch.float32)
|
||||
|
||||
|
||||
def test_qwen_generation_pipelines_load_the_native_component():
|
||||
longcat_config = LongCatImagePipelineConfig()
|
||||
|
||||
assert isinstance(longcat_config.text_encoder_configs[0], Qwen2_5VLConfig)
|
||||
assert "text_encoder" in LongCatImagePipeline._required_config_modules
|
||||
assert Qwen2_5_VLForConditionalGeneration._fsdp_forward_methods == ("generate",)
|
||||
assert "model.visual.blocks" in Qwen2_5_VLForConditionalGeneration.layer_names
|
||||
|
||||
for pipeline_config in (longcat_config, QwenImageLayeredPipelineConfig()):
|
||||
deployment = pipeline_config.get_model_deployment_config()
|
||||
assert deployment.keep_resident_min_available_gb == 70
|
||||
assert deployment.keep_resident_components == ("text_encoder", "vae")
|
||||
@@ -20,6 +20,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.hunyuan import FastHunyuanCo
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import (
|
||||
LingBotWorldCausalDMDConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.longcat_image import (
|
||||
LongCatImagePipelineConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
|
||||
LTX2PipelineConfig,
|
||||
LTX23PipelineConfig,
|
||||
@@ -32,6 +35,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.model_deployment_config impo
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.mova import MOVAPipelineConfig
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
|
||||
QwenImageLayeredPipelineConfig,
|
||||
QwenImagePipelineConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.sana_wm import (
|
||||
@@ -1146,6 +1150,32 @@ class TestOffloadDefaults(unittest.TestCase):
|
||||
self.assertEqual(qwen_deployment.keep_resident_components, ("vae",))
|
||||
self.assertIsNone(qwen_deployment.keep_resident_min_available_gb)
|
||||
|
||||
def test_qwen_ar_generation_residency_scales_with_available_memory(self):
|
||||
pipeline_configs = (
|
||||
QwenImageLayeredPipelineConfig(),
|
||||
LongCatImagePipelineConfig(),
|
||||
)
|
||||
|
||||
for pipeline_config in pipeline_configs:
|
||||
high_memory_args = self._from_dict_with_pipeline_config(
|
||||
pipeline_config,
|
||||
memory_gb=80,
|
||||
kwargs={"performance_mode": "auto"},
|
||||
)
|
||||
self.assertNotIn(
|
||||
"text_encoder", high_memory_args.layerwise_offload_components or []
|
||||
)
|
||||
self.assertFalse(high_memory_args.text_encoder_cpu_offload)
|
||||
|
||||
constrained_args = self._from_dict_with_pipeline_config(
|
||||
pipeline_config,
|
||||
memory_gb=60,
|
||||
kwargs={"performance_mode": "auto"},
|
||||
)
|
||||
self.assertIn(
|
||||
"text_encoder", constrained_args.layerwise_offload_components or []
|
||||
)
|
||||
|
||||
def test_auto_multi_gpu_sana_wm_prefers_fsdp_and_cfg_parallel(self):
|
||||
args = self._from_dict_with_pipeline_config(
|
||||
SanaWMPipelineConfig(),
|
||||
|
||||
Reference in New Issue
Block a user