[diffusion] model: GLM-Image (#16894)

Co-authored-by: jianyingzhu <53300651@qq.com>
This commit is contained in:
Yuhao Yang
2026-01-14 02:02:03 +08:00
committed by GitHub
co-authored by jianyingzhu
parent 2a7b67adff
commit a0b4ba9032
13 changed files with 1951 additions and 3 deletions
@@ -0,0 +1,39 @@
from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
@dataclass
class GlmImageArchConfig(DiTArchConfig):
patch_size: int = 2
in_channels: int = 16
out_channels: int | None = 16
num_layers: int = 30
attention_head_dim: int = 128
num_attention_heads: int = 32
condition_dim: int = 256
prior_vq_quantizer_codebook_size: int = 16384
text_embed_dim: int = 1472
time_embed_dim: int = 512
stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list)
param_names_mapping: dict = field(
default_factory=lambda: {
# LoRA mappings
r"^(transformer_blocks\.\d+\.attn\..*\.lora_[AB])\.default$": r"\1",
}
)
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.out_channels
@dataclass
class GlmImageDitConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=GlmImageArchConfig)
prefix: str = "glmimage"
@@ -0,0 +1,60 @@
from dataclasses import dataclass, field
import torch
from sglang.multimodal_gen.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class GlmImageVAEArchConfig(VAEArchConfig):
spatial_compression_ratio: int = 1
base_dim: int = 96
decoder_base_dim: int | None = None
z_dim: int = 16
dim_mult: tuple[int, ...] = (1, 2, 4, 4)
num_res_blocks: int = 2
attn_scales: tuple[float, ...] = ()
temperal_downsample: tuple[bool, ...] = (False, True, True)
dropout: float = 0.0
is_residual: bool = False
input_channels: int = 3
out_channels: int = 3
patch_size: int | None = None
scale_factor_temporal: int = 4
scale_factor_spatial: int = 8
clip_output: bool = True
scaling_factor: float | torch.Tensor = 0
latents_mean: tuple[float, ...] | None = None
latents_std: tuple[float, ...] | None = None
shift_factor: float | None = None
latent_channels: int = 16
in_channels: int = 16
@dataclass
class GlmImageVAEConfig(VAEConfig):
arch_config: GlmImageVAEArchConfig = field(default_factory=GlmImageVAEArchConfig)
use_feature_cache: bool = True
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
def get_vae_scale_factor(self):
return 2 ** len(self.arch_config.temperal_downsample)
def __post_init__(self):
self.blend_num_frames = (
self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
) * 2
def post_init(self):
self.arch_config.vae_scale_factor = 2 ** (
len(self.arch_config.temperal_downsample)
)
self.arch_config.spatial_compression_ratio = self.arch_config.vae_scale_factor
@@ -366,6 +366,9 @@ class PipelineConfig:
latents = maybe_unpad_latents(latents, batch)
return latents
def post_decoding(self, frames, server_args):
return frames
def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype):
return {}
@@ -0,0 +1,86 @@
from dataclasses import dataclass, field
import torch
from diffusers.image_processor import VaeImageProcessor
from sglang.multimodal_gen.configs.models import DiTConfig, VAEConfig
from sglang.multimodal_gen.configs.models.dits.glmimage import GlmImageDitConfig
from sglang.multimodal_gen.configs.models.vaes.glmimage import GlmImageVAEConfig
from sglang.multimodal_gen.configs.pipeline_configs.base import (
ImagePipelineConfig,
ModelTaskType,
)
@dataclass
class GlmImagePipelineConfig(ImagePipelineConfig):
"""Configuration for the GlmImage pipeline."""
vae_precision: str = "bf16"
should_use_guidance: bool = False
task_type: ModelTaskType = ModelTaskType.T2I
vae_tiling: bool = False
vae_sp: bool = False
dit_config: DiTConfig = field(default_factory=GlmImageDitConfig)
# VAE
vae_config: VAEConfig = field(default_factory=GlmImageVAEConfig)
enable_autocast: bool = False
def __post_init__(self):
self.vae_scale_factor = self.vae_config.get_vae_scale_factor()
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
def get_freqs_cis(self, batch, device, rotary_emb, dtype):
height = batch.height // self.vae_scale_factor
width = batch.width // self.vae_scale_factor
hidden_states = torch.empty(1, 1, height, width, device=device, dtype=dtype)
freqs_cis = rotary_emb(hidden_states)
return freqs_cis
def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype):
return {
"prior_token_id": batch.prior_token_id,
"prior_token_drop": batch.prior_token_drop_cond,
"crop_coords": batch.crop_coords,
"target_size": batch.target_size,
"kv_caches": batch.kv_caches,
"kv_caches_mode": "read",
"freqs_cis": self.get_freqs_cis(batch, device, rotary_emb, dtype),
}
def prepare_neg_cond_kwargs(self, batch, device, rotary_emb, dtype):
return {
"prior_token_id": batch.prior_token_id,
"prior_token_drop": batch.prior_token_drop_uncond,
"crop_coords": batch.crop_coords,
"target_size": batch.target_size,
"kv_caches": batch.kv_caches,
"kv_caches_mode": "skip",
"freqs_cis": self.get_freqs_cis(batch, device, rotary_emb, dtype),
}
def get_decode_scale_and_shift(self, device, dtype, vae):
latents_mean = (
torch.tensor(self.vae_config.latents_mean)
.view(1, self.vae_config.latent_channels, 1, 1)
.to(device, dtype)
)
latents_std = (
torch.tensor(self.vae_config.latents_std)
.view(1, self.vae_config.latent_channels, 1, 1)
.to(device, dtype)
)
return 1.0 / latents_std, latents_mean
def post_denoising_loop(self, latents, batch):
if getattr(batch, "kv_caches", None) is not None:
batch.kv_caches.clear()
return latents.bfloat16()
def post_decoding(self, frames, server_args):
return self.image_processor.postprocess(frames, output_type="latent")
@@ -0,0 +1,12 @@
from dataclasses import dataclass
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
@dataclass
class GlmImageSamplingParams(SamplingParams):
negative_prompt = ""
num_frames: int = 1
guidance_scale: float = 1.5
num_inference_steps: int = 30