[diffusion] model: GLM-Image (#16894)
Co-authored-by: jianyingzhu <53300651@qq.com>
This commit is contained in:
co-authored by
jianyingzhu
parent
2a7b67adff
commit
a0b4ba9032
@@ -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
|
||||
Reference in New Issue
Block a user