Compare commits
28
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ee623198c9 | ||
|
|
1e2f2bfbca | ||
|
|
5e90f421d5 | ||
|
|
546ad1fb32 | ||
|
|
9e261fef21 | ||
|
|
ea4b2d3195 | ||
|
|
c7767910d2 | ||
|
|
cc61fc2714 | ||
|
|
a0749e98fd | ||
|
|
2d7e262b13 | ||
|
|
65de01f1f1 | ||
|
|
e20deeb5eb | ||
|
|
51b9761334 | ||
|
|
e99d2d9f16 | ||
|
|
8a9ec8621e | ||
|
|
53440d4562 | ||
|
|
3d82e5e24c | ||
|
|
cc2c419c4b | ||
|
|
81fb8f2c10 | ||
|
|
0d80a668d2 | ||
|
|
9dec070d4c | ||
|
|
326287dffa | ||
|
|
5c68193009 | ||
|
|
19ef9fa56f | ||
|
|
7833575da8 | ||
|
|
9ea387e872 | ||
|
|
615da36296 | ||
|
|
c54e0737d0 |
@@ -0,0 +1,107 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run GLM-Image text-to-image generation through FastVideo.
|
||||
|
||||
User story:
|
||||
"I have the HF `zai-org/GLM-Image` checkpoint and want a minimal
|
||||
text-to-image generation command, saved as a PNG."
|
||||
"""
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Run GLM-Image text-to-image generation.")
|
||||
parser.add_argument(
|
||||
"--model-path",
|
||||
default="zai-org/GLM-Image",
|
||||
help="HF id or local diffusers-format GLM-Image weights directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
default="image_output/landscape.png",
|
||||
help="Output PNG path.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
default=("A beautiful landscape photography with rolling hills, "
|
||||
"a winding river, and a vibrant sunset in the background. "
|
||||
"Warm golden light, photorealistic style."),
|
||||
help="Text prompt.",
|
||||
)
|
||||
parser.add_argument("--height", type=int, default=1024)
|
||||
parser.add_argument("--width", type=int, default=1024)
|
||||
parser.add_argument("--steps", type=int, default=50)
|
||||
parser.add_argument("--guidance-scale", type=float, default=1.5)
|
||||
parser.add_argument("--seed", type=int, default=1024)
|
||||
parser.add_argument("--num-gpus", type=int, default=1)
|
||||
parser.add_argument("--tp-size", type=int, default=None)
|
||||
parser.add_argument("--sp-size", type=int, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
|
||||
output = Path(args.output)
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
|
||||
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
|
||||
|
||||
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
|
||||
# pipeline class come from the model's registered defaults — don't override.
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
trust_remote_code=True,
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
|
||||
),
|
||||
pipeline=PipelineSelection(workload_type="t2i"),
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=1,
|
||||
fps=1,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output.parent),
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
if isinstance(result, list):
|
||||
result = result[0]
|
||||
|
||||
frames = result.frames
|
||||
if frames is not None and len(frames):
|
||||
Image.fromarray(frames[0]).save(output)
|
||||
print(f"Saved image to {output}")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,120 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run GLM-Image image-to-image (edit) generation through FastVideo.
|
||||
|
||||
User story:
|
||||
"I have the HF `zai-org/GLM-Image` checkpoint and a condition image, and
|
||||
want a minimal edit command (text + image -> edited image), saved as a PNG."
|
||||
|
||||
GLM-Image is a single unified pipeline: passing a condition image switches it
|
||||
from text-to-image to the edit path (the condition enters the DiT via a KV-cache
|
||||
write pass), so the generator config is identical to `basic_glm_image.py` — the
|
||||
`inputs.pil_image` on the request is what selects the edit mode.
|
||||
"""
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Run GLM-Image image-to-image (edit) generation.")
|
||||
parser.add_argument(
|
||||
"--model-path",
|
||||
default="zai-org/GLM-Image",
|
||||
help="HF id or local diffusers-format GLM-Image weights directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--image",
|
||||
default="assets/images/couple.jpg",
|
||||
help="Condition image to edit.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
default="image_output/edited.png",
|
||||
help="Output PNG path.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
default="Change the background to a snowy mountain landscape at golden hour.",
|
||||
help="Edit instruction.",
|
||||
)
|
||||
parser.add_argument("--height", type=int, default=1024)
|
||||
parser.add_argument("--width", type=int, default=1024)
|
||||
parser.add_argument("--steps", type=int, default=50)
|
||||
parser.add_argument("--guidance-scale", type=float, default=1.5)
|
||||
parser.add_argument("--seed", type=int, default=1024)
|
||||
parser.add_argument("--num-gpus", type=int, default=1)
|
||||
parser.add_argument("--tp-size", type=int, default=None)
|
||||
parser.add_argument("--sp-size", type=int, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
|
||||
output = Path(args.output)
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
condition = Image.open(args.image).convert("RGB")
|
||||
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
|
||||
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
|
||||
|
||||
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
|
||||
# pipeline class come from the model's registered defaults — don't override.
|
||||
# The pipeline is registered as t2i; passing inputs.pil_image below switches
|
||||
# it to the edit path.
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
trust_remote_code=True,
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
|
||||
),
|
||||
pipeline=PipelineSelection(workload_type="t2i"),
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
inputs=InputConfig(pil_image=condition),
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=1,
|
||||
fps=1,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output.parent),
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
if isinstance(result, list):
|
||||
result = result[0]
|
||||
|
||||
frames = result.frames
|
||||
if frames is not None and len(frames):
|
||||
Image.fromarray(frames[0]).save(output)
|
||||
print(f"Saved image to {output}")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -2,6 +2,7 @@ from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
|
||||
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
|
||||
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARConfig, DreamXWorldConfig
|
||||
from fastvideo.configs.models.dits.flux_2 import Flux2Config
|
||||
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
|
||||
from fastvideo.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
|
||||
@@ -16,5 +17,6 @@ from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "DreamXWorldConfig",
|
||||
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig",
|
||||
"HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config"
|
||||
"HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config",
|
||||
"GlmImageDiTConfig"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_blocks(n: str, m) -> bool:
|
||||
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageDiTArchConfig(DiTArchConfig):
|
||||
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
|
||||
|
||||
hidden_size: int = 4096
|
||||
num_attention_heads: int = 32
|
||||
attention_head_dim: int = 128
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
num_layers: int = 30
|
||||
|
||||
text_embed_dim: int = 1472
|
||||
time_embed_dim: int = 512
|
||||
condition_dim: int = 256
|
||||
|
||||
prior_vq_quantizer_codebook_size: int = 16384
|
||||
|
||||
patch_size: int = 2
|
||||
|
||||
max_height: int = 2048
|
||||
max_width: int = 2048
|
||||
|
||||
qk_norm: str = "layer_norm"
|
||||
eps: float = 1e-5
|
||||
|
||||
exclude_lora_layers: list[str] = field(
|
||||
default_factory=lambda: ["image_projector", "glyph_projector", "prior_token_embedding"])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^glyph_projector\.net\.0\.proj\.(.*)$": r"glyph_projector.fc_in.\1",
|
||||
r"^glyph_projector\.net\.2\.(.*)$": r"glyph_projector.fc_out.\1",
|
||||
r"^prior_projector\.net\.0\.proj\.(.*)$": r"prior_projector.fc_in.\1",
|
||||
r"^prior_projector\.net\.2\.(.*)$": r"prior_projector.fc_out.\1",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(.*)$": r"transformer_blocks.\1.ff.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.(.*)$": r"transformer_blocks.\1.ff.fc_out.\2",
|
||||
})
|
||||
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.num_channels_latents = self.out_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageDiTConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=GlmImageDiTArchConfig)
|
||||
prefix: str = "GlmImage"
|
||||
@@ -2,6 +2,7 @@ from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
|
||||
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
|
||||
from fastvideo.configs.models.vaes.gamecraftvae import GameCraftVAEConfig
|
||||
from fastvideo.configs.models.vaes.gen3cvae import Gen3CVAEConfig
|
||||
from fastvideo.configs.models.vaes.glm_image import GlmImageVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
|
||||
@@ -21,4 +22,5 @@ __all__ = [
|
||||
"OobleckVAEArchConfig",
|
||||
"OobleckVAEConfig",
|
||||
"Flux2VAEConfig",
|
||||
"GlmImageVAEConfig",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,95 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.autoencoder_kl import (AutoencoderKLArchConfig, AutoencoderKLVAEConfig)
|
||||
|
||||
_GLM_IMAGE_LATENTS_MEAN: tuple[float, ...] = (
|
||||
-0.2080078125,
|
||||
1.875,
|
||||
-0.470703125,
|
||||
-1.265625,
|
||||
-1.421875,
|
||||
0.77734375,
|
||||
-0.3671875,
|
||||
-0.9453125,
|
||||
0.318359375,
|
||||
0.7734375,
|
||||
-0.1884765625,
|
||||
-0.022216796875,
|
||||
-0.220703125,
|
||||
-1.59375,
|
||||
-0.81640625,
|
||||
-0.255859375,
|
||||
)
|
||||
_GLM_IMAGE_LATENTS_STD: tuple[float, ...] = (
|
||||
3.0625,
|
||||
2.203125,
|
||||
2.265625,
|
||||
4.84375,
|
||||
2.5,
|
||||
3.9375,
|
||||
2.203125,
|
||||
3.03125,
|
||||
2.1875,
|
||||
2.046875,
|
||||
2.71875,
|
||||
2.390625,
|
||||
2.390625,
|
||||
2.453125,
|
||||
2.25,
|
||||
2.15625,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageVAEArchConfig(AutoencoderKLArchConfig):
|
||||
act_fn: str = "silu"
|
||||
block_out_channels: tuple[int, ...] = (128, 512, 1024, 1024)
|
||||
down_block_types: tuple[str, ...] = (
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
)
|
||||
up_block_types: tuple[str, ...] = (
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
)
|
||||
force_upcast: bool = True
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
latent_channels: int = 16
|
||||
latents_mean: tuple[float, ...] = _GLM_IMAGE_LATENTS_MEAN
|
||||
latents_std: tuple[float, ...] = _GLM_IMAGE_LATENTS_STD
|
||||
layers_per_block: int = 3
|
||||
mid_block_add_attention: bool = False
|
||||
norm_num_groups: int = 32
|
||||
sample_size: int = 1024
|
||||
scaling_factor: float = 0.18215
|
||||
shift_factor: float | None = None
|
||||
use_quant_conv: bool = False
|
||||
use_post_quant_conv: bool = False
|
||||
|
||||
temporal_compression_ratio: int = 1
|
||||
spatial_compression_ratio: int = 8
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageVAEConfig(AutoencoderKLVAEConfig):
|
||||
arch_config: GlmImageVAEArchConfig = field(default_factory=GlmImageVAEArchConfig)
|
||||
|
||||
use_tiling: bool = True
|
||||
use_temporal_tiling: bool = False
|
||||
use_parallel_tiling: bool = False
|
||||
|
||||
tile_sample_min_height: int = 512
|
||||
tile_sample_min_width: int = 512
|
||||
tile_sample_stride_height: int = 384
|
||||
tile_sample_stride_width: int = 384
|
||||
|
||||
load_encoder: bool = True
|
||||
load_decoder: bool = True
|
||||
@@ -0,0 +1,46 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||
from fastvideo.configs.models.vaes.glm_image import GlmImageVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
def glm_image_t5_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
mask: torch.Tensor = outputs.attention_mask
|
||||
hidden_state: torch.Tensor = outputs.last_hidden_state
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
|
||||
assert torch.isnan(hidden_state).sum() == 0, "T5 hidden states contain NaN"
|
||||
|
||||
max_len = 512
|
||||
prompt_embeds = [u[:min(v, max_len)] for u, v in zip(hidden_state, seq_lens, strict=True)]
|
||||
prompt_embeds_tensor: torch.Tensor = torch.stack(
|
||||
[torch.cat([u, u.new_zeros(max_len - u.size(0), u.size(1))]) for u in prompt_embeds], dim=0)
|
||||
|
||||
return prompt_embeds_tensor
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageConfig(PipelineConfig):
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=GlmImageDiTConfig)
|
||||
dit_precision: str = "bf16"
|
||||
|
||||
vae_config: VAEConfig = field(default_factory=GlmImageVAEConfig)
|
||||
vae_precision: str = "fp32"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = False
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (T5Config(), ))
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("fp32", ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (glm_image_t5_postprocess, ))
|
||||
|
||||
flow_shift: float | None = 1.0
|
||||
embedded_cfg_scale: float = 7.5
|
||||
@@ -0,0 +1,776 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from fastvideo.layers.mlp import MLP
|
||||
|
||||
from fastvideo.attention import LocalAttention
|
||||
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
|
||||
from fastvideo.layers.layernorm import ScaleResidualLayerNormScaleShift
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.layers.rotary_embedding import _apply_rotary_emb
|
||||
from fastvideo.layers.visual_embedding import Timesteps
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class GlmImageLayerKVCache:
|
||||
|
||||
def __init__(self):
|
||||
self.k_cache = None
|
||||
self.v_cache = None
|
||||
self.mode: Optional[str] = None
|
||||
|
||||
def store(self, k: torch.Tensor, v: torch.Tensor):
|
||||
# Append along seq (dim=1).
|
||||
if self.k_cache is None:
|
||||
self.k_cache = k
|
||||
self.v_cache = v
|
||||
else:
|
||||
self.k_cache = torch.cat([self.k_cache, k], dim=1)
|
||||
self.v_cache = torch.cat([self.v_cache, v], dim=1)
|
||||
|
||||
def get(self):
|
||||
return self.k_cache, self.v_cache
|
||||
|
||||
def clear(self):
|
||||
self.k_cache = None
|
||||
self.v_cache = None
|
||||
self.mode = None
|
||||
|
||||
|
||||
class GlmImageKVCache:
|
||||
|
||||
def __init__(self, num_layers: int):
|
||||
self.num_layers = num_layers
|
||||
self.caches = [GlmImageLayerKVCache() for _ in range(num_layers)]
|
||||
|
||||
def __getitem__(self, layer_idx: int) -> GlmImageLayerKVCache:
|
||||
return self.caches[layer_idx]
|
||||
|
||||
def set_mode(self, mode: Optional[str]):
|
||||
if mode is not None and mode not in ["write", "read", "skip"]:
|
||||
raise ValueError(
|
||||
f"Invalid mode: {mode}, must be one of 'write', 'read', 'skip'"
|
||||
)
|
||||
for cache in self.caches:
|
||||
cache.mode = mode
|
||||
|
||||
def clear(self):
|
||||
for cache in self.caches:
|
||||
cache.clear()
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Timestep and Text Projection
|
||||
# =============================================================================
|
||||
class GlmImageTimestepEmbedding(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
time_embed_dim: int,
|
||||
act_fn: str = "silu",
|
||||
out_dim: int = None,
|
||||
):
|
||||
super().__init__()
|
||||
if out_dim is None:
|
||||
out_dim = time_embed_dim
|
||||
self.linear_1 = ReplicatedLinear(in_channels, time_embed_dim, bias=True)
|
||||
if act_fn == "silu":
|
||||
self.act = nn.SiLU()
|
||||
elif act_fn == "gelu":
|
||||
self.act = nn.GELU(approximate="tanh")
|
||||
else:
|
||||
self.act = nn.SiLU()
|
||||
self.linear_2 = ReplicatedLinear(time_embed_dim, out_dim, bias=True)
|
||||
|
||||
def forward(self, sample: torch.Tensor) -> torch.Tensor:
|
||||
sample, _ = self.linear_1(sample)
|
||||
sample = self.act(sample)
|
||||
sample, _ = self.linear_2(sample)
|
||||
return sample
|
||||
|
||||
|
||||
class GlmImageTextProjection(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
hidden_size: int,
|
||||
out_features: int = None,
|
||||
act_fn: str = "silu",
|
||||
):
|
||||
super().__init__()
|
||||
if out_features is None:
|
||||
out_features = hidden_size
|
||||
self.linear_1 = ReplicatedLinear(in_features, hidden_size, bias=True)
|
||||
if act_fn == "silu":
|
||||
self.act_1 = nn.SiLU()
|
||||
elif act_fn == "gelu_tanh":
|
||||
self.act_1 = nn.GELU(approximate="tanh")
|
||||
else:
|
||||
self.act_1 = nn.SiLU()
|
||||
self.linear_2 = ReplicatedLinear(hidden_size, out_features, bias=True)
|
||||
|
||||
def forward(self, caption: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states, _ = self.linear_1(caption)
|
||||
hidden_states = self.act_1(hidden_states)
|
||||
hidden_states, _ = self.linear_2(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class GlmImageCombinedTimestepSizeEmbeddings(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
condition_dim: int,
|
||||
pooled_projection_dim: int,
|
||||
timesteps_dim: int = 256,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.time_proj = Timesteps(
|
||||
num_channels=timesteps_dim, flip_sin_to_cos=True, downscale_freq_shift=0
|
||||
)
|
||||
self.condition_proj = Timesteps(
|
||||
num_channels=condition_dim, flip_sin_to_cos=True, downscale_freq_shift=0
|
||||
)
|
||||
self.timestep_embedder = GlmImageTimestepEmbedding(
|
||||
in_channels=timesteps_dim, time_embed_dim=embedding_dim
|
||||
)
|
||||
self.condition_embedder = GlmImageTextProjection(
|
||||
pooled_projection_dim, embedding_dim, act_fn="silu"
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
timestep: torch.Tensor,
|
||||
target_size: torch.Tensor,
|
||||
crop_coords: torch.Tensor,
|
||||
hidden_dtype: torch.dtype,
|
||||
) -> torch.Tensor:
|
||||
timesteps_proj = self.time_proj(timestep)
|
||||
|
||||
crop_coords_proj = self.condition_proj(crop_coords.flatten()).view(
|
||||
crop_coords.size(0), -1
|
||||
)
|
||||
target_size_proj = self.condition_proj(target_size.flatten()).view(
|
||||
target_size.size(0), -1
|
||||
)
|
||||
|
||||
condition_proj = torch.cat([crop_coords_proj, target_size_proj], dim=1)
|
||||
|
||||
timesteps_emb = self.timestep_embedder(
|
||||
timesteps_proj.to(dtype=hidden_dtype)
|
||||
)
|
||||
condition_emb = self.condition_embedder(
|
||||
condition_proj.to(dtype=hidden_dtype)
|
||||
)
|
||||
|
||||
conditioning = timesteps_emb + condition_emb
|
||||
return conditioning
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Image Projector
|
||||
# =============================================================================
|
||||
class GlmImageImageProjector(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 16,
|
||||
hidden_size: int = 2560,
|
||||
patch_size: int = 2,
|
||||
):
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.proj = nn.Linear(in_channels * patch_size**2, hidden_size)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
post_patch_height = height // self.patch_size
|
||||
post_patch_width = width // self.patch_size
|
||||
|
||||
hidden_states = hidden_states.reshape(
|
||||
batch_size,
|
||||
channel,
|
||||
post_patch_height,
|
||||
self.patch_size,
|
||||
post_patch_width,
|
||||
self.patch_size,
|
||||
)
|
||||
hidden_states = (
|
||||
hidden_states.permute(0, 2, 4, 1, 3, 5).flatten(3, 5).flatten(1, 2)
|
||||
)
|
||||
hidden_states = self.proj(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# AdaLayerNorm
|
||||
# =============================================================================
|
||||
class GlmImageAdaLayerNormZero(nn.Module):
|
||||
|
||||
def __init__(self, embedding_dim: int, dim: int) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.norm = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5)
|
||||
self.norm_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5)
|
||||
self.linear = ReplicatedLinear(embedding_dim, 12 * dim, bias=True)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, ...]:
|
||||
dtype = hidden_states.dtype
|
||||
norm_hidden_states = self.norm(hidden_states).to(dtype=dtype)
|
||||
norm_encoder_hidden_states = self.norm_context(encoder_hidden_states).to(
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
emb, _ = self.linear(temb)
|
||||
(
|
||||
shift_msa,
|
||||
c_shift_msa,
|
||||
scale_msa,
|
||||
c_scale_msa,
|
||||
gate_msa,
|
||||
c_gate_msa,
|
||||
shift_mlp,
|
||||
c_shift_mlp,
|
||||
scale_mlp,
|
||||
c_scale_mlp,
|
||||
gate_mlp,
|
||||
c_gate_mlp,
|
||||
) = emb.chunk(12, dim=1)
|
||||
|
||||
hidden_states = norm_hidden_states * (
|
||||
1 + scale_msa.unsqueeze(1)
|
||||
) + shift_msa.unsqueeze(1)
|
||||
encoder_hidden_states = norm_encoder_hidden_states * (
|
||||
1 + c_scale_msa.unsqueeze(1)
|
||||
) + c_shift_msa.unsqueeze(1)
|
||||
|
||||
return (
|
||||
hidden_states,
|
||||
gate_msa,
|
||||
shift_mlp,
|
||||
scale_mlp,
|
||||
gate_mlp,
|
||||
encoder_hidden_states,
|
||||
c_gate_msa,
|
||||
c_shift_mlp,
|
||||
c_scale_mlp,
|
||||
c_gate_mlp,
|
||||
)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Attention
|
||||
# =============================================================================
|
||||
class GlmImageAttention(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
query_dim: int,
|
||||
heads: int,
|
||||
dim_head: int,
|
||||
out_dim: int,
|
||||
bias: bool = True,
|
||||
qk_norm: str = "layer_norm",
|
||||
elementwise_affine: bool = False,
|
||||
eps: float = 1e-5,
|
||||
supported_attention_backends: set[AttentionBackendEnum] | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.heads = out_dim // dim_head if out_dim is not None else heads
|
||||
self.num_kv_heads = self.heads
|
||||
self.dim_head = dim_head
|
||||
self.inner_dim = out_dim if out_dim is not None else dim_head * heads
|
||||
self.inner_kv_dim = self.inner_dim
|
||||
self.out_dim = out_dim if out_dim is not None else query_dim
|
||||
|
||||
self.to_q = ReplicatedLinear(query_dim, self.inner_dim, bias=bias)
|
||||
self.to_k = ReplicatedLinear(query_dim, self.inner_kv_dim, bias=bias)
|
||||
self.to_v = ReplicatedLinear(query_dim, self.inner_kv_dim, bias=bias)
|
||||
|
||||
self.to_out = nn.ModuleList(
|
||||
[ReplicatedLinear(self.inner_dim, self.out_dim, bias=True)]
|
||||
)
|
||||
|
||||
if qk_norm is None:
|
||||
self.norm_q = None
|
||||
self.norm_k = None
|
||||
elif qk_norm == "layer_norm":
|
||||
self.norm_q = nn.LayerNorm(
|
||||
dim_head, eps=eps, elementwise_affine=elementwise_affine
|
||||
)
|
||||
self.norm_k = nn.LayerNorm(
|
||||
dim_head, eps=eps, elementwise_affine=elementwise_affine
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"unknown qk_norm: {qk_norm}")
|
||||
|
||||
self.attn = LocalAttention(
|
||||
num_heads=self.heads,
|
||||
head_size=dim_head,
|
||||
num_kv_heads=self.heads,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
kv_cache: Optional[GlmImageLayerKVCache] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
dtype = encoder_hidden_states.dtype
|
||||
|
||||
batch_size, text_seq_length, embed_dim = encoder_hidden_states.shape
|
||||
batch_size, image_seq_length, embed_dim = hidden_states.shape
|
||||
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
|
||||
|
||||
# 1. QKV projections
|
||||
query, _ = self.to_q(hidden_states)
|
||||
key, _ = self.to_k(hidden_states)
|
||||
value, _ = self.to_v(hidden_states)
|
||||
|
||||
query = query.unflatten(2, (self.heads, -1))
|
||||
key = key.unflatten(2, (self.heads, -1))
|
||||
value = value.unflatten(2, (self.heads, -1))
|
||||
|
||||
# 2. QK normalization
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q(query).to(dtype=dtype)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(key).to(dtype=dtype)
|
||||
|
||||
# 3. Rotational positional embeddings applied to latent stream
|
||||
if image_rotary_emb is not None:
|
||||
cos, sin = image_rotary_emb
|
||||
|
||||
query[:, text_seq_length:, :, :] = _apply_rotary_emb(
|
||||
query[:, text_seq_length:, :, :], cos, sin, is_neox_style=True
|
||||
)
|
||||
key[:, text_seq_length:, :, :] = _apply_rotary_emb(
|
||||
key[:, text_seq_length:, :, :], cos, sin, is_neox_style=True
|
||||
)
|
||||
|
||||
# 4. KV Cache handling
|
||||
if kv_cache is not None:
|
||||
if kv_cache.mode == "write":
|
||||
kv_cache.store(key, value)
|
||||
elif kv_cache.mode == "read":
|
||||
# Prepend cached condition k/v along seq (dim=1).
|
||||
k_cache, v_cache = kv_cache.get()
|
||||
key = torch.cat([k_cache, key], dim=1) if k_cache is not None else key
|
||||
value = (
|
||||
torch.cat([v_cache, value], dim=1) if v_cache is not None else value
|
||||
)
|
||||
elif kv_cache.mode == "skip":
|
||||
pass
|
||||
|
||||
hidden_states = self.attn(query, key, value)
|
||||
hidden_states = hidden_states.flatten(2, 3)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# 6. Output projection
|
||||
hidden_states, _ = self.to_out[0](hidden_states)
|
||||
|
||||
encoder_hidden_states, hidden_states = hidden_states.split(
|
||||
[text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
|
||||
)
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Transformer Block
|
||||
# =============================================================================
|
||||
class GlmImageTransformerBlock(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int = 2560,
|
||||
num_attention_heads: int = 64,
|
||||
attention_head_dim: int = 40,
|
||||
time_embed_dim: int = 512,
|
||||
supported_attention_backends: set[AttentionBackendEnum] | None = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
# 1. Attention
|
||||
self.norm1 = GlmImageAdaLayerNormZero(time_embed_dim, dim)
|
||||
|
||||
self.attn1 = GlmImageAttention(
|
||||
query_dim=dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
out_dim=dim,
|
||||
bias=True,
|
||||
qk_norm="layer_norm",
|
||||
elementwise_affine=False,
|
||||
eps=1e-5,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
prefix=f"{prefix}.attn1",
|
||||
)
|
||||
|
||||
# 2. Feedforward with fused ScaleResidualLayerNorm
|
||||
self.norm2 = ScaleResidualLayerNormScaleShift(
|
||||
dim, norm_type="layer", eps=1e-5, elementwise_affine=False
|
||||
)
|
||||
self.norm2_context = ScaleResidualLayerNormScaleShift(
|
||||
dim, norm_type="layer", eps=1e-5, elementwise_affine=False
|
||||
)
|
||||
self.ff = MLP(input_dim=dim, mlp_hidden_dim=dim * 4, output_dim=dim, act_type="gelu_pytorch_tanh")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: Optional[torch.Tensor] = None,
|
||||
image_rotary_emb: Optional[
|
||||
Union[
|
||||
Tuple[torch.Tensor, torch.Tensor],
|
||||
List[Tuple[torch.Tensor, torch.Tensor]],
|
||||
]
|
||||
] = None,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
kv_cache: Optional[GlmImageLayerKVCache] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# 1. Timestep conditioning
|
||||
(
|
||||
norm_hidden_states,
|
||||
gate_msa,
|
||||
shift_mlp,
|
||||
scale_mlp,
|
||||
gate_mlp,
|
||||
norm_encoder_hidden_states,
|
||||
c_gate_msa,
|
||||
c_shift_mlp,
|
||||
c_scale_mlp,
|
||||
c_gate_mlp,
|
||||
) = self.norm1(hidden_states, encoder_hidden_states, temb)
|
||||
|
||||
# 2. Attention
|
||||
if attention_kwargs is None:
|
||||
attention_kwargs = {}
|
||||
|
||||
attn_hidden_states, attn_encoder_hidden_states = self.attn1(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=norm_encoder_hidden_states,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
kv_cache=kv_cache,
|
||||
**attention_kwargs,
|
||||
)
|
||||
|
||||
# 3. Feedforward (fused residual + norm + scale/shift)
|
||||
norm_hidden_states, hidden_states = self.norm2(
|
||||
hidden_states,
|
||||
attn_hidden_states,
|
||||
gate_msa.unsqueeze(1),
|
||||
shift_mlp.unsqueeze(1),
|
||||
scale_mlp.unsqueeze(1),
|
||||
)
|
||||
norm_encoder_hidden_states, encoder_hidden_states = self.norm2_context(
|
||||
encoder_hidden_states,
|
||||
attn_encoder_hidden_states,
|
||||
c_gate_msa.unsqueeze(1),
|
||||
c_shift_mlp.unsqueeze(1),
|
||||
c_scale_mlp.unsqueeze(1),
|
||||
)
|
||||
|
||||
ff_output = self.ff(norm_hidden_states)
|
||||
ff_output_context = self.ff(norm_encoder_hidden_states)
|
||||
hidden_states = hidden_states + ff_output * gate_mlp.unsqueeze(1)
|
||||
encoder_hidden_states = (
|
||||
encoder_hidden_states + ff_output_context * c_gate_mlp.unsqueeze(1)
|
||||
)
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Rotary Positional Embedding
|
||||
# =============================================================================
|
||||
class GlmImageRotaryPosEmbed(nn.Module):
|
||||
|
||||
def __init__(self, dim: int, patch_size: int, theta: float = 10000.0) -> None:
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.patch_size = patch_size
|
||||
self.theta = theta
|
||||
self._cache_key: tuple | None = None
|
||||
self._cache_value: tuple[torch.Tensor, torch.Tensor] | None = None
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
_, _, raw_h, raw_w = hidden_states.shape
|
||||
height = raw_h // self.patch_size
|
||||
width = raw_w // self.patch_size
|
||||
device = hidden_states.device
|
||||
|
||||
cache_key = (height, width, device.type,
|
||||
device.index if device.index is not None else -1)
|
||||
if self._cache_key == cache_key and self._cache_value is not None:
|
||||
return self._cache_value
|
||||
|
||||
dim_h, dim_w = self.dim // 2, self.dim // 2
|
||||
h_inv_freq = 1.0 / (
|
||||
self.theta
|
||||
** (
|
||||
torch.arange(0, dim_h, 2, dtype=torch.float32, device=device)[
|
||||
: (dim_h // 2)
|
||||
]
|
||||
/ dim_h
|
||||
)
|
||||
)
|
||||
w_inv_freq = 1.0 / (
|
||||
self.theta
|
||||
** (
|
||||
torch.arange(0, dim_w, 2, dtype=torch.float32, device=device)[
|
||||
: (dim_w // 2)
|
||||
]
|
||||
/ dim_w
|
||||
)
|
||||
)
|
||||
h_seq = torch.arange(height, device=device)
|
||||
w_seq = torch.arange(width, device=device)
|
||||
freqs_h = torch.outer(h_seq, h_inv_freq).unsqueeze(1).expand(height, width, -1)
|
||||
freqs_w = torch.outer(w_seq, w_inv_freq).unsqueeze(0).expand(height, width, -1)
|
||||
|
||||
freqs = torch.cat([freqs_h, freqs_w], dim=-1).reshape(height * width, -1)
|
||||
result = (freqs.cos(), freqs.sin())
|
||||
self._cache_key = cache_key
|
||||
self._cache_value = result
|
||||
return result
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Final AdaLayerNorm
|
||||
# =============================================================================
|
||||
class GlmImageAdaLayerNormContinuous(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
conditioning_embedding_dim: int,
|
||||
elementwise_affine: bool = True,
|
||||
eps: float = 1e-5,
|
||||
bias: bool = True,
|
||||
norm_type: str = "layer_norm",
|
||||
):
|
||||
super().__init__()
|
||||
self.linear = nn.Linear(
|
||||
conditioning_embedding_dim, embedding_dim * 2, bias=bias
|
||||
)
|
||||
if norm_type == "layer_norm":
|
||||
self.norm = nn.LayerNorm(embedding_dim, eps, elementwise_affine, bias)
|
||||
elif norm_type == "rms_norm":
|
||||
self.norm = nn.RMSNorm(embedding_dim, eps, elementwise_affine)
|
||||
else:
|
||||
raise ValueError(f"unknown norm_type {norm_type}")
|
||||
|
||||
def forward(
|
||||
self, x: torch.Tensor, conditioning_embedding: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
emb = self.linear(conditioning_embedding.to(x.dtype))
|
||||
scale, shift = torch.chunk(emb, 2, dim=1)
|
||||
x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :]
|
||||
return x
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Main Model
|
||||
# =============================================================================
|
||||
class GlmImageTransformer2DModel(BaseDiT):
|
||||
|
||||
_fsdp_shard_conditions = GlmImageDiTConfig().arch_config._fsdp_shard_conditions
|
||||
_compile_conditions = GlmImageDiTConfig().arch_config._compile_conditions
|
||||
_supported_attention_backends = (
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
)
|
||||
|
||||
param_names_mapping = GlmImageDiTConfig().arch_config.param_names_mapping
|
||||
reverse_param_names_mapping = {}
|
||||
lora_param_names_mapping = {}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: GlmImageDiTConfig,
|
||||
hf_config: dict[str, Any],
|
||||
):
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
arch_config = config.arch_config
|
||||
|
||||
self.in_channels = arch_config.in_channels
|
||||
self.out_channels = arch_config.out_channels
|
||||
self.patch_size = arch_config.patch_size
|
||||
self.num_layers = arch_config.num_layers
|
||||
self.attention_head_dim = arch_config.attention_head_dim
|
||||
self.num_attention_heads = arch_config.num_attention_heads
|
||||
self.text_embed_dim = arch_config.text_embed_dim
|
||||
self.time_embed_dim = arch_config.time_embed_dim
|
||||
|
||||
# GlmImage uses 2 additional SDXL-like conditions - target_size, crop_coords
|
||||
# Each of these are sincos embeddings of shape 2 * condition_dim
|
||||
pooled_projection_dim = 2 * 2 * arch_config.condition_dim
|
||||
inner_dim = arch_config.num_attention_heads * arch_config.attention_head_dim
|
||||
|
||||
self.hidden_size = inner_dim
|
||||
self.num_channels_latents = arch_config.out_channels
|
||||
|
||||
# 1. RoPE
|
||||
self.rotary_emb = GlmImageRotaryPosEmbed(
|
||||
arch_config.attention_head_dim, arch_config.patch_size, theta=10000.0
|
||||
)
|
||||
|
||||
# 2. Patch & Text-timestep embedding
|
||||
self.image_projector = GlmImageImageProjector(
|
||||
arch_config.in_channels, inner_dim, arch_config.patch_size
|
||||
)
|
||||
self.glyph_projector = MLP(
|
||||
input_dim=arch_config.text_embed_dim,
|
||||
mlp_hidden_dim=inner_dim,
|
||||
output_dim=inner_dim,
|
||||
act_type="gelu",
|
||||
)
|
||||
self.prior_token_embedding = nn.Embedding(
|
||||
arch_config.prior_vq_quantizer_codebook_size, inner_dim
|
||||
)
|
||||
self.prior_projector = MLP(
|
||||
input_dim=inner_dim,
|
||||
mlp_hidden_dim=inner_dim,
|
||||
output_dim=inner_dim,
|
||||
act_type="silu",
|
||||
)
|
||||
|
||||
self.time_condition_embed = GlmImageCombinedTimestepSizeEmbeddings(
|
||||
embedding_dim=arch_config.time_embed_dim,
|
||||
condition_dim=arch_config.condition_dim,
|
||||
pooled_projection_dim=pooled_projection_dim,
|
||||
timesteps_dim=arch_config.time_embed_dim,
|
||||
)
|
||||
|
||||
# 3. Transformer blocks
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
GlmImageTransformerBlock(
|
||||
inner_dim,
|
||||
arch_config.num_attention_heads,
|
||||
arch_config.attention_head_dim,
|
||||
arch_config.time_embed_dim,
|
||||
supported_attention_backends=self._supported_attention_backends,
|
||||
prefix=f"transformer_blocks.{i}",
|
||||
)
|
||||
for i in range(arch_config.num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# 4. Output projection
|
||||
self.norm_out = GlmImageAdaLayerNormContinuous(
|
||||
inner_dim, arch_config.time_embed_dim, elementwise_affine=False
|
||||
)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim,
|
||||
arch_config.patch_size * arch_config.patch_size * arch_config.out_channels,
|
||||
bias=True,
|
||||
)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
prior_token_id: torch.Tensor,
|
||||
prior_token_drop: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
target_size: torch.Tensor,
|
||||
crop_coords: torch.Tensor,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
kv_caches: Optional[GlmImageKVCache] = None,
|
||||
kv_caches_mode: Optional[str] = None,
|
||||
freqs_cis: Optional[
|
||||
Union[
|
||||
Tuple[torch.Tensor, torch.Tensor],
|
||||
List[Tuple[torch.Tensor, torch.Tensor]],
|
||||
]
|
||||
] = None,
|
||||
guidance: torch.Tensor = None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
if kv_caches is not None:
|
||||
kv_caches.set_mode(kv_caches_mode)
|
||||
|
||||
batch_size, num_channels, height, width = hidden_states.shape
|
||||
|
||||
if isinstance(encoder_hidden_states, list):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
|
||||
# 1. RoPE
|
||||
image_rotary_emb = freqs_cis
|
||||
if image_rotary_emb is None:
|
||||
image_rotary_emb = self.rotary_emb(hidden_states)
|
||||
|
||||
# 2. Patch & Timestep embeddings
|
||||
p = self.patch_size
|
||||
post_patch_height = height // p
|
||||
post_patch_width = width // p
|
||||
|
||||
hidden_states = self.image_projector(hidden_states)
|
||||
encoder_hidden_states = self.glyph_projector(encoder_hidden_states)
|
||||
prior_embedding = self.prior_token_embedding(prior_token_id)
|
||||
# Zero dropped priors by multiply: boolean indexing + .any() syncs each step.
|
||||
keep = (~prior_token_drop).to(device=prior_embedding.device, dtype=prior_embedding.dtype)
|
||||
while keep.dim() < prior_embedding.dim():
|
||||
keep = keep.unsqueeze(-1)
|
||||
prior_embedding = prior_embedding * keep
|
||||
prior_hidden_states = self.prior_projector(prior_embedding)
|
||||
hidden_states = hidden_states + prior_hidden_states
|
||||
|
||||
temb = self.time_condition_embed(
|
||||
timestep, target_size, crop_coords, hidden_states.dtype
|
||||
)
|
||||
temb = F.silu(temb)
|
||||
|
||||
# 3. Transformer blocks
|
||||
for idx, block in enumerate(self.transformer_blocks):
|
||||
hidden_states, encoder_hidden_states = block(
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
temb,
|
||||
image_rotary_emb,
|
||||
attention_kwargs,
|
||||
kv_cache=kv_caches[idx] if kv_caches is not None else None,
|
||||
)
|
||||
|
||||
# 4. Output norm & projection
|
||||
hidden_states = self.norm_out(hidden_states, temb)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
# 5. Unpatchify
|
||||
hidden_states = hidden_states.reshape(
|
||||
batch_size, post_patch_height, post_patch_width, -1, p, p
|
||||
)
|
||||
output = hidden_states.permute(0, 3, 1, 4, 2, 5).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
return output.float()
|
||||
|
||||
|
||||
EntryClass = GlmImageTransformer2DModel
|
||||
@@ -0,0 +1,64 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Sole HF-import boundary for GLM-Image's AR encoder (lazy-wrapper exception E001)."""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class GlmImageARLoader(nn.Module):
|
||||
|
||||
def __init__(self, model_path: str, processor_path: str | None = None,
|
||||
*, torch_dtype: torch.dtype = torch.bfloat16,
|
||||
trust_remote_code: bool = True) -> None:
|
||||
super().__init__()
|
||||
from transformers import (AutoProcessor,
|
||||
GlmImageForConditionalGeneration)
|
||||
logger.info("Loading GLM-Image AR encoder from %s", model_path)
|
||||
self._model = GlmImageForConditionalGeneration.from_pretrained(
|
||||
model_path,
|
||||
torch_dtype=torch_dtype,
|
||||
trust_remote_code=trust_remote_code,
|
||||
)
|
||||
if processor_path is not None:
|
||||
logger.info("Loading GLM-Image processor from %s", processor_path)
|
||||
self.processor = AutoProcessor.from_pretrained(
|
||||
processor_path, trust_remote_code=trust_remote_code)
|
||||
else:
|
||||
self.processor = None
|
||||
|
||||
@torch.no_grad()
|
||||
def generate(self, *args: Any, **kwargs: Any) -> torch.Tensor:
|
||||
return self._model.generate(*args, **kwargs)
|
||||
|
||||
@torch.no_grad()
|
||||
def get_image_features(self, pixel_values: torch.Tensor,
|
||||
image_grid_thw: torch.Tensor) -> Any:
|
||||
return self._model.get_image_features(pixel_values, image_grid_thw)
|
||||
|
||||
@torch.no_grad()
|
||||
def get_image_tokens(self, image_embeds: torch.Tensor,
|
||||
image_grid_thw: torch.Tensor) -> torch.Tensor:
|
||||
return self._model.get_image_tokens(image_embeds, image_grid_thw)
|
||||
|
||||
@property
|
||||
def config(self): # type: ignore[no-untyped-def]
|
||||
return self._model.config
|
||||
|
||||
@property
|
||||
def generation_config(self): # type: ignore[no-untyped-def]
|
||||
return self._model.generation_config
|
||||
|
||||
def to(self, *args, **kwargs): # type: ignore[override]
|
||||
self._model = self._model.to(*args, **kwargs)
|
||||
return super().to(*args, **kwargs)
|
||||
|
||||
def eval(self): # type: ignore[override]
|
||||
self._model = self._model.eval()
|
||||
return super().eval()
|
||||
@@ -95,6 +95,8 @@ class ComponentLoader(ABC):
|
||||
"image_processor": (ImageProcessorLoader, "transformers"),
|
||||
"feature_extractor": (ImageProcessorLoader, "transformers"),
|
||||
"image_encoder": (ImageEncoderLoader, "transformers"),
|
||||
"vision_language_encoder": (VisionLanguageEncoderLoader, "transformers"),
|
||||
"processor": (ProcessorLoader, "transformers"),
|
||||
"upsampler": (UpsamplerLoader, "diffusers"),
|
||||
"upsampler_2": (UpsamplerLoader, "diffusers"),
|
||||
# Stable Audio's `StableAudioMultiConditioner` bundles T5 +
|
||||
@@ -523,6 +525,39 @@ class ImageEncoderLoader(TextEncoderLoader):
|
||||
)
|
||||
|
||||
|
||||
class VisionLanguageEncoderLoader(ComponentLoader):
|
||||
"""Loader for vision-language autoregressive encoders."""
|
||||
|
||||
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
|
||||
from fastvideo.distributed.parallel_state import get_local_torch_device
|
||||
from fastvideo.models.encoders.glm_image_ar_loader import (
|
||||
GlmImageARLoader)
|
||||
|
||||
logger.info("Loading vision-language encoder from %s", model_path)
|
||||
target_device = get_local_torch_device()
|
||||
loader = GlmImageARLoader(
|
||||
model_path,
|
||||
torch_dtype=torch.bfloat16,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
).to(target_device).eval()
|
||||
return loader
|
||||
|
||||
|
||||
class ProcessorLoader(ComponentLoader):
|
||||
"""Loader for HF processors that pair with vision-language encoders."""
|
||||
|
||||
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
|
||||
from transformers import AutoProcessor
|
||||
|
||||
logger.info("Loading processor from %s", model_path)
|
||||
processor = AutoProcessor.from_pretrained(
|
||||
model_path,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
)
|
||||
logger.info("Loaded processor: %s", processor.__class__.__name__)
|
||||
return processor
|
||||
|
||||
|
||||
class ImageProcessorLoader(ComponentLoader):
|
||||
"""Loader for image processor."""
|
||||
|
||||
|
||||
@@ -61,6 +61,11 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
"MatrixGame3WanModel": ("dits", "matrixgame3", "MatrixGame3WanModel"),
|
||||
}
|
||||
|
||||
# Text-to-image DiT models (2D image generation)
|
||||
_TEXT_TO_IMAGE_DIT_MODELS = {
|
||||
"GlmImageTransformer2DModel": ("dits", "glm_image", "GlmImageTransformer2DModel"),
|
||||
}
|
||||
|
||||
_TEXT_ENCODER_MODELS = {
|
||||
"CLIPTextModel": ("encoders", "clip", "CLIPTextModel"),
|
||||
"CLIPTextModelWithProjection":
|
||||
@@ -134,6 +139,7 @@ _UPSAMPLERS = {
|
||||
_LEGACY_FAST_VIDEO_MODELS = {
|
||||
**_TEXT_TO_VIDEO_DIT_MODELS,
|
||||
**_IMAGE_TO_VIDEO_DIT_MODELS,
|
||||
**_TEXT_TO_IMAGE_DIT_MODELS,
|
||||
**_TEXT_ENCODER_MODELS,
|
||||
**_IMAGE_ENCODER_MODELS,
|
||||
**_VAE_MODELS,
|
||||
|
||||
@@ -0,0 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""GLM-Image pipeline package."""
|
||||
|
||||
from fastvideo.pipelines.basic.glm_image.glm_image_pipeline import (
|
||||
GlmImagePipeline, )
|
||||
|
||||
__all__ = ["GlmImagePipeline"]
|
||||
@@ -0,0 +1,82 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler, )
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.pipelines.basic.glm_image.stages import (
|
||||
GlmImageBeforeDenoisingStage,
|
||||
GlmImageConditionEncodingStage,
|
||||
GlmImageDecodingStage,
|
||||
GlmImageDenoisingStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages import InputValidationStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class GlmImagePipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
pipeline_name = "GlmImagePipeline"
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"vae",
|
||||
"transformer",
|
||||
"scheduler",
|
||||
"vision_language_encoder",
|
||||
"processor",
|
||||
]
|
||||
|
||||
_optional_config_modules: list[str] = []
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(shift=1.0)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
self.add_stage(
|
||||
stage_name="input_validation_stage",
|
||||
stage=InputValidationStage(),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="glm_image_before_denoising_stage",
|
||||
stage=GlmImageBeforeDenoisingStage(
|
||||
vae=self.get_module("vae"),
|
||||
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"),
|
||||
vision_language_encoder=self.get_module("vision_language_encoder"),
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="glm_image_condition_encoding_stage",
|
||||
stage=GlmImageConditionEncodingStage(
|
||||
vae=self.get_module("vae"),
|
||||
transformer=self.get_module("transformer"),
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=GlmImageDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
pipeline=self,
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="decoding_stage",
|
||||
stage=GlmImageDecodingStage(
|
||||
vae=self.get_module("vae"),
|
||||
pipeline=self,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
EntryClass = GlmImagePipeline
|
||||
@@ -0,0 +1,14 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""GLM-Image pipeline stages."""
|
||||
|
||||
from fastvideo.pipelines.basic.glm_image.stages.before_denoising import (GlmImageBeforeDenoisingStage)
|
||||
from fastvideo.pipelines.basic.glm_image.stages.condition_encoding import (GlmImageConditionEncodingStage)
|
||||
from fastvideo.pipelines.basic.glm_image.stages.decoding import (GlmImageDecodingStage)
|
||||
from fastvideo.pipelines.basic.glm_image.stages.denoising import (GlmImageDenoisingStage)
|
||||
|
||||
__all__ = [
|
||||
"GlmImageBeforeDenoisingStage",
|
||||
"GlmImageConditionEncodingStage",
|
||||
"GlmImageDecodingStage",
|
||||
"GlmImageDenoisingStage",
|
||||
]
|
||||
@@ -0,0 +1,269 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from math import sqrt
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def calculate_shift(
|
||||
image_seq_len: int,
|
||||
base_seq_len: int = 256,
|
||||
base_shift: float = 0.25,
|
||||
max_shift: float = 0.75,
|
||||
) -> float:
|
||||
return (image_seq_len / base_seq_len)**0.5 * max_shift + base_shift
|
||||
|
||||
|
||||
def get_glyph_texts(prompt: str | list[str]) -> list[str] | list[list[str]]:
|
||||
if isinstance(prompt, str):
|
||||
prompts: list[str] = [prompt]
|
||||
is_batch = False
|
||||
else:
|
||||
prompts = prompt
|
||||
is_batch = True
|
||||
out: list[list[str]] = []
|
||||
for p in prompts:
|
||||
out.append(
|
||||
re.findall(r"'([^']*)'", p) + re.findall(r"“([^“”]*)”", p) + re.findall(r'"([^"]*)"', p) +
|
||||
re.findall(r"「([^「」]*)」", p))
|
||||
return out if is_batch else out[0]
|
||||
|
||||
|
||||
def compute_glyph_embeds(
|
||||
prompts: list[str],
|
||||
tokenizer,
|
||||
text_encoder,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
max_sequence_length: int = 2048,
|
||||
) -> torch.Tensor:
|
||||
all_glyph_texts = get_glyph_texts(prompts)
|
||||
all_glyph_embeds = []
|
||||
for glyph_texts in all_glyph_texts:
|
||||
if len(glyph_texts) == 0:
|
||||
glyph_texts = [""]
|
||||
input_ids = tokenizer(
|
||||
glyph_texts,
|
||||
max_length=max_sequence_length,
|
||||
truncation=True,
|
||||
).input_ids
|
||||
input_ids = [[tokenizer.pad_token_id] * ((len(input_ids) + 1) % 2) + ids for ids in input_ids]
|
||||
max_length = max(len(ids) for ids in input_ids)
|
||||
attention_mask = torch.tensor(
|
||||
[[1] * len(ids) + [0] * (max_length - len(ids)) for ids in input_ids],
|
||||
device=device,
|
||||
)
|
||||
input_ids_t = torch.tensor(
|
||||
[ids + [tokenizer.pad_token_id] * (max_length - len(ids)) for ids in input_ids],
|
||||
device=device,
|
||||
)
|
||||
outputs = text_encoder(input_ids_t, attention_mask=attention_mask)
|
||||
glyph_embeds = outputs.last_hidden_state[attention_mask.bool()].unsqueeze(0)
|
||||
all_glyph_embeds.append(glyph_embeds)
|
||||
|
||||
max_seq_len = max(emb.size(1) for emb in all_glyph_embeds)
|
||||
padded = []
|
||||
for emb in all_glyph_embeds:
|
||||
if emb.size(1) < max_seq_len:
|
||||
pad = torch.zeros(emb.size(0), max_seq_len - emb.size(1), emb.size(2), device=device, dtype=emb.dtype)
|
||||
emb = torch.cat([pad, emb], dim=1)
|
||||
padded.append(emb)
|
||||
return torch.cat(padded, dim=0).to(device=device, dtype=dtype)
|
||||
|
||||
|
||||
def _grid_dims(height: int, width: int) -> tuple[int, int, int, int]:
|
||||
th, tw = height // 32, width // 32
|
||||
ratio = th / tw
|
||||
pth = int(sqrt(ratio) * 16)
|
||||
ptw = int(sqrt(1 / ratio) * 16)
|
||||
return th, tw, pth, ptw
|
||||
|
||||
|
||||
def _upsample_d32_to_d16(tokens: torch.Tensor, th: int, tw: int) -> torch.Tensor:
|
||||
tokens = tokens.view(1, 1, th, tw).float()
|
||||
tokens = torch.nn.functional.interpolate(tokens, scale_factor=2, mode="nearest").long()
|
||||
return tokens.view(1, -1)
|
||||
|
||||
|
||||
class GlmImageBeforeDenoisingStage(PipelineStage):
|
||||
|
||||
def __init__(self,
|
||||
vae,
|
||||
text_encoder,
|
||||
tokenizer,
|
||||
processor,
|
||||
transformer,
|
||||
scheduler,
|
||||
vision_language_encoder=None) -> None:
|
||||
super().__init__()
|
||||
self.vae = vae
|
||||
self.text_encoder = text_encoder
|
||||
self.tokenizer = tokenizer
|
||||
self.processor = processor
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
if isinstance(vision_language_encoder, tuple):
|
||||
self.vision_language_encoder, self.vl_processor = (vision_language_encoder[0], vision_language_encoder[1]
|
||||
or processor)
|
||||
else:
|
||||
self.vision_language_encoder = vision_language_encoder
|
||||
self.vl_processor = processor
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
device = get_local_torch_device()
|
||||
dtype = torch.bfloat16
|
||||
th, tw, pth, ptw = _grid_dims(batch.height, batch.width)
|
||||
|
||||
if batch.seed is not None:
|
||||
torch.manual_seed(batch.seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed(batch.seed)
|
||||
|
||||
# 1-3. AR token generation. I2I prepends the condition image and uses a
|
||||
# single-scale target grid; T2I is multi-scale.
|
||||
is_t2i = batch.pil_image is None
|
||||
if self.vision_language_encoder is not None:
|
||||
content = [{"type": "text", "text": batch.prompt}]
|
||||
if not is_t2i:
|
||||
content.insert(0, {"type": "image", "image": batch.pil_image})
|
||||
messages = [{"role": "user", "content": content}]
|
||||
inputs = self.vl_processor.apply_chat_template(messages,
|
||||
tokenize=True,
|
||||
target_h=batch.height,
|
||||
target_w=batch.width,
|
||||
return_dict=True,
|
||||
return_tensors="pt").to(device)
|
||||
|
||||
if is_t2i:
|
||||
up_h, up_w = th, tw
|
||||
large_start, large_count = pth * ptw, th * tw
|
||||
max_new = large_count + (pth * ptw) + 1
|
||||
else:
|
||||
# Condition grid(s) first, target grid last.
|
||||
_, t_h, t_w = inputs["image_grid_thw"][-1].tolist()
|
||||
up_h, up_w = int(t_h), int(t_w)
|
||||
large_start, large_count = 0, up_h * up_w
|
||||
max_new = large_count + 1
|
||||
|
||||
outputs = self.vision_language_encoder.generate(**inputs, max_new_tokens=max_new, do_sample=True)
|
||||
gen_tokens = outputs[0][inputs.input_ids.shape[-1]:]
|
||||
if gen_tokens.shape[0] >= large_start + large_count:
|
||||
large_tokens = gen_tokens[large_start:large_start + large_count]
|
||||
else:
|
||||
available = gen_tokens[large_start:]
|
||||
large_tokens = torch.zeros(large_count, dtype=gen_tokens.dtype, device=gen_tokens.device)
|
||||
if available.shape[0] > 0:
|
||||
large_tokens[:min(available.shape[0], large_count)] = available[:large_count]
|
||||
logger.warning("AR generated %d tokens, expected %d. Padding with zeros.", gen_tokens.shape[0],
|
||||
large_start + large_count)
|
||||
batch.prior_token_id = _upsample_d32_to_d16(large_tokens, up_h, up_w)
|
||||
batch.prior_token_drop = torch.zeros(batch.prior_token_id.shape, dtype=torch.bool, device=device)
|
||||
|
||||
if not is_t2i:
|
||||
self._compute_source_prior_tokens(batch, inputs)
|
||||
else:
|
||||
num_prior_tokens = 4 * th * tw
|
||||
logger.warning("No vision_language_encoder provided; using random dropped priors.")
|
||||
batch.prior_token_id = torch.randint(0, 16384, (1, num_prior_tokens), device=device)
|
||||
batch.prior_token_drop = torch.ones(batch.prior_token_id.shape, dtype=torch.bool, device=device)
|
||||
|
||||
# 4. Glyph T5 encoding.
|
||||
prompts = [batch.prompt] if isinstance(batch.prompt, str) else list(batch.prompt)
|
||||
prompt_embeds = compute_glyph_embeds(prompts, self.tokenizer, self.text_encoder, device, dtype)
|
||||
|
||||
# 5. CFG-side negative encoding.
|
||||
if batch.do_classifier_free_guidance:
|
||||
neg_prompts = [batch.negative_prompt or ""] * len(prompts)
|
||||
neg_embeds = compute_glyph_embeds(neg_prompts, self.tokenizer, self.text_encoder, device, dtype)
|
||||
L_pos, L_neg = prompt_embeds.shape[1], neg_embeds.shape[1]
|
||||
max_L = max(L_pos, L_neg)
|
||||
if L_pos < max_L:
|
||||
pad = torch.zeros(prompt_embeds.shape[0],
|
||||
max_L - L_pos,
|
||||
prompt_embeds.shape[2],
|
||||
device=device,
|
||||
dtype=dtype)
|
||||
prompt_embeds = torch.cat([pad, prompt_embeds], dim=1)
|
||||
if L_neg < max_L:
|
||||
pad = torch.zeros(neg_embeds.shape[0], max_L - L_neg, neg_embeds.shape[2], device=device, dtype=dtype)
|
||||
neg_embeds = torch.cat([pad, neg_embeds], dim=1)
|
||||
# Row 0 conditional (positive), row 1 unconditional (negative).
|
||||
prompt_embeds = torch.cat([prompt_embeds, neg_embeds], dim=0)
|
||||
att_pos = torch.ones((1, max_L), device=device)
|
||||
att_neg = torch.ones((1, max_L), device=device)
|
||||
if L_pos < max_L:
|
||||
att_pos[:, :max_L - L_pos] = 0
|
||||
if L_neg < max_L:
|
||||
att_neg[:, :max_L - L_neg] = 0
|
||||
attention_mask = torch.cat([att_pos, att_neg], dim=0)
|
||||
else:
|
||||
attention_mask = torch.ones((1, prompt_embeds.shape[1]), device=device)
|
||||
|
||||
batch.prompt_embeds = [prompt_embeds]
|
||||
batch.attention_mask = attention_mask
|
||||
|
||||
# 6. Latents + dynamic flow shift.
|
||||
if batch.seed is not None:
|
||||
torch.manual_seed(batch.seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed(batch.seed)
|
||||
batch.latents = torch.randn((1, 16, 1, batch.height // 8, batch.width // 8), device=device, dtype=dtype)
|
||||
|
||||
# Integer-cast linspace timesteps with resolution-dependent shift applied to
|
||||
# sigmas only; the DiT is conditioned on the unshifted integer timesteps.
|
||||
ntt = self.scheduler.config.num_train_timesteps
|
||||
patch_size = self.transformer.patch_size
|
||||
image_seq_len = ((batch.height // 8) * (batch.width // 8)) // (patch_size**2)
|
||||
sched_timesteps = np.linspace(ntt, 1.0, batch.num_inference_steps + 1)[:-1].astype(np.int64).astype(np.float32)
|
||||
sched_sigmas = sched_timesteps / ntt
|
||||
self.scheduler.set_shift(calculate_shift(image_seq_len))
|
||||
self.scheduler.set_timesteps(batch.num_inference_steps,
|
||||
device=device,
|
||||
sigmas=sched_sigmas.tolist(),
|
||||
timesteps=sched_timesteps.tolist())
|
||||
batch.timesteps = self.scheduler.timesteps
|
||||
return batch
|
||||
|
||||
@torch.no_grad()
|
||||
def _compute_source_prior_tokens(self, batch: ForwardBatch, inputs) -> None:
|
||||
image_grid_thw = inputs["image_grid_thw"]
|
||||
num_condition_images = image_grid_thw.shape[0] - 1
|
||||
source_grids = image_grid_thw[:num_condition_images]
|
||||
image_features = self.vision_language_encoder.get_image_features(inputs["pixel_values"], source_grids)
|
||||
image_feature_parts = getattr(image_features, "pooler_output", image_features)
|
||||
embed = torch.cat(image_feature_parts, dim=0)
|
||||
src_ids_d32 = self.vision_language_encoder.get_image_tokens(embed, source_grids)
|
||||
split_sizes = source_grids.prod(dim=-1).tolist()
|
||||
upsampled = [
|
||||
_upsample_d32_to_d16(ids, int(grid[1]), int(grid[2])).squeeze(0)
|
||||
for ids, grid in zip(torch.split(src_ids_d32, split_sizes), source_grids, strict=False)
|
||||
]
|
||||
src_grids_up = source_grids.clone()
|
||||
src_grids_up[:, 1] *= 2
|
||||
src_grids_up[:, 2] *= 2
|
||||
batch.extra["glm_prior_token_image_ids"] = torch.cat(upsampled, dim=0)
|
||||
batch.extra["glm_source_image_grid_thw"] = src_grids_up
|
||||
|
||||
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("prompt", batch.prompt, V.string_not_empty)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("prior_token_id", batch.prior_token_id, V.is_tensor)
|
||||
return result
|
||||
@@ -0,0 +1,67 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.image_processor import ImageProcessor
|
||||
from fastvideo.models.dits.glm_image import GlmImageKVCache
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
|
||||
_CONDITION_MULTIPLE_OF = 16 # vae_scale_factor (8) * DiT patch_size (2)
|
||||
|
||||
|
||||
class GlmImageConditionEncodingStage(PipelineStage):
|
||||
|
||||
def __init__(self, vae, transformer) -> None:
|
||||
super().__init__()
|
||||
self.vae = vae
|
||||
self.transformer = transformer
|
||||
self.image_processor = ImageProcessor(vae_scale_factor=_CONDITION_MULTIPLE_OF)
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.pil_image is None:
|
||||
return batch
|
||||
|
||||
device = get_local_torch_device()
|
||||
dtype = torch.bfloat16
|
||||
self.vae.to(device)
|
||||
|
||||
prior_ids = batch.extra["glm_prior_token_image_ids"].to(device)
|
||||
if prior_ids.dim() == 1:
|
||||
prior_ids = prior_ids.unsqueeze(0)
|
||||
|
||||
# Latent patch count must match the source prior tokens; mismatch is fatal.
|
||||
src_grid = batch.extra["glm_source_image_grid_thw"][0]
|
||||
cond_h = int(src_grid[1]) * _CONDITION_MULTIPLE_OF
|
||||
cond_w = int(src_grid[2]) * _CONDITION_MULTIPLE_OF
|
||||
cond_img = self.image_processor.preprocess(batch.pil_image, cond_h, cond_w).to(device=device,
|
||||
dtype=torch.float32)
|
||||
latent = self.vae.encode(cond_img).latent_dist.mode()
|
||||
# NOTE: at runtime self.vae.config is a diffusers FrozenDict with flat
|
||||
# latents_mean/latents_std fields. Access the flat fields directly.
|
||||
cfg = self.vae.config
|
||||
mean = torch.tensor(cfg.latents_mean, device=device, dtype=torch.float32).view(1, -1, 1, 1)
|
||||
std = torch.tensor(cfg.latents_std, device=device, dtype=torch.float32).view(1, -1, 1, 1)
|
||||
latent = ((latent - mean) / std).to(dtype)
|
||||
|
||||
kv_caches = GlmImageKVCache(num_layers=self.transformer.num_layers)
|
||||
empty_text = batch.prompt_embeds[0][:1, :0, :].to(device=device, dtype=dtype)
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None, forward_batch=batch):
|
||||
self.transformer(
|
||||
hidden_states=latent,
|
||||
encoder_hidden_states=empty_text,
|
||||
prior_token_id=prior_ids,
|
||||
prior_token_drop=torch.zeros((prior_ids.shape[0], ), dtype=torch.bool, device=device),
|
||||
timestep=torch.zeros((1, ), device=device),
|
||||
target_size=torch.tensor([tuple(cond_img.shape[-2:])], device=device, dtype=torch.long),
|
||||
crop_coords=torch.zeros((1, 2), device=device, dtype=torch.long),
|
||||
kv_caches=kv_caches,
|
||||
kv_caches_mode="write",
|
||||
)
|
||||
batch.extra["glm_kv_caches"] = kv_caches
|
||||
return batch
|
||||
@@ -0,0 +1,35 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.stages.decoding import DecodingStage
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
|
||||
class GlmImageDecodingStage(DecodingStage):
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, latents: torch.Tensor, fastvideo_args: FastVideoArgs) -> torch.Tensor:
|
||||
self.vae.to(get_local_torch_device())
|
||||
latents = latents.to(get_local_torch_device())
|
||||
|
||||
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast = (vae_dtype != torch.float32 and not fastvideo_args.disable_autocast)
|
||||
|
||||
latents = self._denormalize_latents(latents)
|
||||
if latents.dim() == 5:
|
||||
latents = latents.squeeze(2)
|
||||
|
||||
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
if not vae_autocast:
|
||||
latents = latents.to(vae_dtype)
|
||||
decoded = self.vae.decode(latents)
|
||||
|
||||
image = decoded.sample if hasattr(decoded, "sample") else decoded
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
return image.unsqueeze(2)
|
||||
@@ -0,0 +1,192 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CFG convention (both denoise paths): row 0 conditional (positive), row 1 unconditional."""
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.attention.backends.sdpa import SDPAMetadata
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.denoising import DenoisingStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class GlmImageDenoisingStage(DenoisingStage):
|
||||
|
||||
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("timesteps", batch.timesteps, [V.is_tensor, V.min_dims(1)])
|
||||
latents = getattr(batch, "latent", getattr(batch, "latents", None))
|
||||
result.add_check("latents", latents, [V.is_tensor, V.with_dims(5)])
|
||||
result.add_check("num_inference_steps", batch.num_inference_steps, V.positive_int)
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
|
||||
return result
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
device = get_local_torch_device()
|
||||
dtype = torch.bfloat16
|
||||
guidance_scale = batch.guidance_scale
|
||||
do_cfg = guidance_scale > 1.0
|
||||
|
||||
latents = getattr(batch, "latent", getattr(batch, "latents", None))
|
||||
if latents is None:
|
||||
raise ValueError("No latents found in batch.")
|
||||
if latents.dim() == 5:
|
||||
latents = latents.squeeze(2)
|
||||
|
||||
prompt_embeds = batch.prompt_embeds[0]
|
||||
text_attention_mask = getattr(batch, "attention_mask", None)
|
||||
timesteps = batch.timesteps
|
||||
|
||||
patch_size = self.transformer.patch_size
|
||||
_, _, h, w = latents.shape
|
||||
image_seq_length = (h // patch_size) * (w // patch_size)
|
||||
text_seq_length = prompt_embeds.shape[1] if prompt_embeds.dim() >= 2 else 0
|
||||
|
||||
first_block = self.transformer.transformer_blocks[0]
|
||||
backend = getattr(first_block.attn1.attn, "backend", None)
|
||||
sdpa = backend == AttentionBackendEnum.TORCH_SDPA and text_attention_mask is not None
|
||||
|
||||
kv_caches = batch.extra.get("glm_kv_caches")
|
||||
if kv_caches is None:
|
||||
self._denoise_t2i(batch, latents, prompt_embeds, text_attention_mask, timesteps, do_cfg, guidance_scale,
|
||||
text_seq_length, image_seq_length, sdpa, device, dtype)
|
||||
else:
|
||||
self._denoise_i2i(batch, latents, prompt_embeds, text_attention_mask, timesteps, do_cfg, guidance_scale,
|
||||
text_seq_length, image_seq_length, sdpa, kv_caches, device, dtype)
|
||||
return batch
|
||||
|
||||
def _denoise_t2i(self, batch, latents, prompt_embeds, text_attention_mask, timesteps, do_cfg, guidance_scale,
|
||||
text_seq_length, image_seq_length, sdpa, device, dtype) -> None:
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
bs = 2 if do_cfg else 1
|
||||
target_size = torch.tensor([[batch.height, batch.width]], device=device, dtype=torch.long).repeat(bs, 1)
|
||||
crop_coords = torch.zeros((bs, 2), device=device, dtype=torch.long)
|
||||
|
||||
prior_token_id = batch.prior_token_id
|
||||
if do_cfg and prior_token_id.shape[0] == 1:
|
||||
prior_token_id = prior_token_id.repeat(2, 1)
|
||||
if do_cfg:
|
||||
prior_token_drop = torch.tensor([False, True], device=device)
|
||||
else:
|
||||
prior_token_drop = getattr(batch, "prior_token_drop", torch.tensor([False], device=device))
|
||||
|
||||
attention_mask_kv = None
|
||||
if sdpa:
|
||||
if (text_attention_mask.shape[0] == 1 and bs > 1):
|
||||
text_attention_mask = text_attention_mask.repeat(bs, 1)
|
||||
mix_attn_mask = torch.ones((bs, text_seq_length + image_seq_length), device=device, dtype=torch.float32)
|
||||
mix_attn_mask[:, :text_seq_length] = (text_attention_mask.float().to(device))
|
||||
attention_mask_kv = (mix_attn_mask > 0).unsqueeze(1).unsqueeze(2)
|
||||
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
latent_model_input = torch.cat([latents] * 2) if do_cfg else latents
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t).to(dtype)
|
||||
t_expand = t.expand(latent_model_input.shape[0]) - 1
|
||||
|
||||
attn_metadata = (SDPAMetadata(current_timestep=i, attn_mask=attention_mask_kv)
|
||||
if attention_mask_kv is not None else None)
|
||||
|
||||
with torch.no_grad(), set_forward_context(current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch):
|
||||
noise_pred = self.transformer(
|
||||
latent_model_input,
|
||||
prompt_embeds,
|
||||
prior_token_id,
|
||||
prior_token_drop,
|
||||
t_expand,
|
||||
target_size,
|
||||
crop_coords,
|
||||
)
|
||||
|
||||
if do_cfg:
|
||||
noise_pred_cond, noise_pred_uncond = noise_pred.float().chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_cond - noise_pred_uncond)
|
||||
|
||||
guidance_rescale = getattr(batch, "guidance_rescale", 0.0)
|
||||
if guidance_rescale > 0.0:
|
||||
dims = list(range(1, noise_pred_cond.ndim))
|
||||
std_text = noise_pred_cond.std(dim=dims, keepdim=True)
|
||||
std_cfg = noise_pred.std(dim=dims, keepdim=True)
|
||||
rescaled = noise_pred * (std_text / std_cfg)
|
||||
noise_pred = (guidance_rescale * rescaled + (1 - guidance_rescale) * noise_pred)
|
||||
|
||||
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
|
||||
progress_bar.update()
|
||||
|
||||
batch.latents = latents.unsqueeze(2)
|
||||
|
||||
def _denoise_i2i(self, batch, latents, prompt_embeds, text_attention_mask, timesteps, do_cfg, guidance_scale,
|
||||
text_seq_length, image_seq_length, sdpa, kv_caches, device, dtype) -> None:
|
||||
"""Two separate transformer calls (cond reads the cache, uncond skips it):
|
||||
the cache mode is one global flag with batch-1 k/v, so a 2-row CFG call
|
||||
cannot express both."""
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
target_size = torch.tensor([[batch.height, batch.width]], device=device, dtype=torch.long)
|
||||
crop_coords = torch.zeros((1, 2), device=device, dtype=torch.long)
|
||||
prior_token_id = batch.prior_token_id[:1]
|
||||
drop_keep = torch.zeros((1, ), dtype=torch.bool, device=device)
|
||||
drop_all = torch.ones((1, ), dtype=torch.bool, device=device)
|
||||
cache_len = kv_caches[0].k_cache.shape[1] if kv_caches[0].k_cache is not None else 0
|
||||
|
||||
def _mask(row: int, with_cache: bool):
|
||||
if not sdpa:
|
||||
return None
|
||||
prefix = cache_len if with_cache else 0
|
||||
m = torch.ones((1, prefix + text_seq_length + image_seq_length), device=device, dtype=torch.float32)
|
||||
m[:, prefix:prefix + text_seq_length] = text_attention_mask[row:row + 1].float().to(device)
|
||||
return (m > 0).unsqueeze(1).unsqueeze(2)
|
||||
|
||||
cond_mask = _mask(0, with_cache=True)
|
||||
uncond_mask = _mask(1, with_cache=False) if do_cfg else None
|
||||
pos_embeds = prompt_embeds[:1]
|
||||
neg_embeds = prompt_embeds[1:2] if do_cfg else None
|
||||
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
latent_model_input = self.scheduler.scale_model_input(latents, t).to(dtype)
|
||||
t_expand = t.expand(1) - 1
|
||||
|
||||
cond_meta = (SDPAMetadata(current_timestep=i, attn_mask=cond_mask) if cond_mask is not None else None)
|
||||
with torch.no_grad(), set_forward_context(current_timestep=i,
|
||||
attn_metadata=cond_meta,
|
||||
forward_batch=batch):
|
||||
noise_pred = self.transformer(latent_model_input,
|
||||
pos_embeds,
|
||||
prior_token_id,
|
||||
drop_keep,
|
||||
t_expand,
|
||||
target_size,
|
||||
crop_coords,
|
||||
kv_caches=kv_caches,
|
||||
kv_caches_mode="read")
|
||||
|
||||
if do_cfg:
|
||||
uncond_meta = (SDPAMetadata(current_timestep=i, attn_mask=uncond_mask)
|
||||
if uncond_mask is not None else None)
|
||||
with torch.no_grad(), set_forward_context(current_timestep=i,
|
||||
attn_metadata=uncond_meta,
|
||||
forward_batch=batch):
|
||||
noise_pred_uncond = self.transformer(latent_model_input,
|
||||
neg_embeds,
|
||||
prior_token_id,
|
||||
drop_all,
|
||||
t_expand,
|
||||
target_size,
|
||||
crop_coords,
|
||||
kv_caches=kv_caches,
|
||||
kv_caches_mode="skip")
|
||||
noise_pred = noise_pred_uncond.float() + guidance_scale * (noise_pred.float() -
|
||||
noise_pred_uncond.float())
|
||||
|
||||
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
|
||||
progress_bar.update()
|
||||
|
||||
kv_caches.clear()
|
||||
batch.latents = latents.unsqueeze(2)
|
||||
@@ -58,6 +58,7 @@ from fastvideo.configs.pipelines.wan import (
|
||||
WanT2V480PConfig,
|
||||
WanT2V720PConfig,
|
||||
)
|
||||
from fastvideo.configs.pipelines.glm_image import GlmImageConfig
|
||||
from fastvideo.configs.pipelines.sd35 import SD35Config
|
||||
from fastvideo.configs.pipelines.stable_audio import (StableAudioOpenSmallConfig, StableAudioT2AConfig)
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
@@ -1068,6 +1069,18 @@ def _register_configs() -> None:
|
||||
default_preset="sd35_medium",
|
||||
)
|
||||
|
||||
# GLM-Image
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=GlmImageConfig,
|
||||
hf_model_paths=[
|
||||
"zai-org/GLM-Image",
|
||||
],
|
||||
model_detectors=[lambda path: "glmimage" in path.lower() or "glm-image" in path.lower()],
|
||||
workload_types=(WorkloadType.T2I, ),
|
||||
model_family="glm_image",
|
||||
)
|
||||
|
||||
|
||||
# --- Part 3: Main Resolver ---
|
||||
|
||||
|
||||
@@ -0,0 +1,141 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""SSIM-based regression test for GLM-Image generation."""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.tests.ssim.inference_similarity_utils import (
|
||||
run_text_to_video_similarity_test,
|
||||
)
|
||||
from fastvideo.tests.ssim.reference_utils import (
|
||||
get_cuda_device_name,
|
||||
resolve_device_reference_folder,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
REQUIRED_GPUS = 1
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
LOCAL_WEIGHTS_DIR = Path(
|
||||
os.getenv("GLM_IMAGE_LOCAL_WEIGHTS_DIR",
|
||||
REPO_ROOT / "official_weights" / "glm_image"))
|
||||
GLM_IMAGE_MODEL_PATH = os.getenv("GLM_IMAGE_MODEL_DIR", str(LOCAL_WEIGHTS_DIR))
|
||||
|
||||
device_reference_folder = resolve_device_reference_folder(
|
||||
(
|
||||
("A40", "A40"),
|
||||
("L40S", "L40S"),
|
||||
("H100", "H100"),
|
||||
("H200", "H200"),
|
||||
("B200", "B200"),
|
||||
),
|
||||
device_name=get_cuda_device_name(),
|
||||
fallback_device_prefix="L40S",
|
||||
logger=logger,
|
||||
)
|
||||
|
||||
MODEL_ID = "zai-org__GLM-Image"
|
||||
|
||||
TEST_PROMPTS = [
|
||||
"A beautiful landscape photography with rolling hills, "
|
||||
"a winding river, and a vibrant sunset in the background. "
|
||||
"Warm golden light, photorealistic style.",
|
||||
]
|
||||
|
||||
GLM_IMAGE_PARAMS = {
|
||||
"num_gpus": 1,
|
||||
"model_path": GLM_IMAGE_MODEL_PATH,
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"height": 256,
|
||||
"width": 256,
|
||||
"num_frames": 1,
|
||||
"fps": 1,
|
||||
"num_inference_steps": 4,
|
||||
"guidance_scale": 1.5,
|
||||
"seed": 0,
|
||||
"neg_prompt": "",
|
||||
}
|
||||
|
||||
GLM_IMAGE_FULL_QUALITY_PARAMS = {
|
||||
"num_gpus": 1,
|
||||
"model_path": GLM_IMAGE_MODEL_PATH,
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"height": 1024,
|
||||
"width": 1024,
|
||||
"num_frames": 1,
|
||||
"fps": 1,
|
||||
"num_inference_steps": 50,
|
||||
"guidance_scale": 1.5,
|
||||
"seed": 0,
|
||||
"neg_prompt": "",
|
||||
}
|
||||
|
||||
GLM_IMAGE_MODEL_TO_PARAMS = {
|
||||
MODEL_ID: GLM_IMAGE_PARAMS,
|
||||
}
|
||||
GLM_IMAGE_FULL_QUALITY_MODEL_TO_PARAMS = {
|
||||
MODEL_ID: GLM_IMAGE_FULL_QUALITY_PARAMS,
|
||||
}
|
||||
|
||||
|
||||
def _has_weights() -> bool:
|
||||
required = ["transformer", "vae", "text_encoder",
|
||||
"vision_language_encoder", "processor", "tokenizer",
|
||||
"scheduler"]
|
||||
return all((LOCAL_WEIGHTS_DIR / r).exists() for r in required)
|
||||
|
||||
|
||||
def _upstream_glm_image_available() -> bool:
|
||||
try:
|
||||
import transformers
|
||||
except ImportError:
|
||||
return False
|
||||
return hasattr(transformers, "GlmImageForConditionalGeneration")
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="GLM-Image SSIM test requires CUDA",
|
||||
)
|
||||
@pytest.mark.skipif(
|
||||
not _has_weights(),
|
||||
reason=f"GLM-Image full weights not found at {LOCAL_WEIGHTS_DIR}.",
|
||||
)
|
||||
@pytest.mark.skipif(
|
||||
not _upstream_glm_image_available(),
|
||||
reason="GLM-Image needs transformers>=5.0.0rc0 (ships the AR encoder).",
|
||||
)
|
||||
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("attention_backend_name", ["TORCH_SDPA"])
|
||||
@pytest.mark.parametrize("model_id", list(GLM_IMAGE_MODEL_TO_PARAMS.keys()))
|
||||
def test_glm_image_similarity(
|
||||
prompt: str,
|
||||
attention_backend_name: str,
|
||||
model_id: str,
|
||||
) -> None:
|
||||
run_text_to_video_similarity_test(
|
||||
logger=logger,
|
||||
script_dir=os.path.dirname(os.path.abspath(__file__)),
|
||||
device_reference_folder=device_reference_folder,
|
||||
prompt=prompt,
|
||||
attention_backend_name=attention_backend_name,
|
||||
model_id=model_id,
|
||||
default_params_map=GLM_IMAGE_MODEL_TO_PARAMS,
|
||||
full_quality_params_map=GLM_IMAGE_FULL_QUALITY_MODEL_TO_PARAMS,
|
||||
min_acceptable_ssim=0.98,
|
||||
init_kwargs_override={
|
||||
"trust_remote_code": True,
|
||||
"use_fsdp_inference": False,
|
||||
},
|
||||
generation_kwargs_override={
|
||||
"save_video": True,
|
||||
},
|
||||
)
|
||||
+3
-1
@@ -21,7 +21,9 @@ dependencies = [
|
||||
"requests>=2.32.2",
|
||||
|
||||
# Machine Learning & Transformers
|
||||
"transformers>=4.57.3",
|
||||
# GLM-Image's AR encoder (GlmImageForConditionalGeneration) first ships in
|
||||
# transformers 5.0.0; floor bumped from >=4.57.3 to >=5.0.0 (stable, not rc).
|
||||
"transformers>=5.0.0",
|
||||
# <0.23: tokenizers 0.23 renamed RobertaProcessing's binding args, so
|
||||
# transformers' CLIP-style tokenizer loading dies with
|
||||
# "RobertaProcessing.__new__() got an unexpected keyword argument 'cls'".
|
||||
|
||||
@@ -21,7 +21,9 @@ dependencies = [
|
||||
"requests>=2.32.2",
|
||||
|
||||
# Machine Learning & Transformers
|
||||
"transformers>=4.57.3",
|
||||
# GLM-Image's AR encoder (GlmImageForConditionalGeneration) first ships in
|
||||
# transformers 5.0.0; floor bumped from >=4.57.3 to >=5.0.0 (stable, not rc).
|
||||
"transformers>=5.0.0",
|
||||
"tokenizers>=0.20.1",
|
||||
"sentencepiece>=0.2.0",
|
||||
"timm>=1.0.11",
|
||||
|
||||
@@ -0,0 +1,74 @@
|
||||
# GLM-Image Local Tests
|
||||
|
||||
Local-only parity and smoke tests for the `glm_image` port: FastVideo vs the
|
||||
official HF `transformers` + `diffusers` GLM-Image implementation. These are
|
||||
**not run in CI** — they need the full weights, CUDA, and (for the reference
|
||||
side) the diffusers GLM-Image classes; each test skips with an actionable
|
||||
message when a prerequisite is missing.
|
||||
|
||||
## What you need
|
||||
|
||||
| | |
|
||||
|---|---|
|
||||
| Weights | `zai-org/GLM-Image` → `official_weights/glm_image` (~30 GB: AR encoder 4 shards, transformer 3 shards, vae + text_encoder) |
|
||||
| Runtime deps | `transformers>=5.0.0` (first release with the AR encoder `GlmImageForConditionalGeneration`) and `diffusers>=0.38.0` — both committed in `pyproject.toml` |
|
||||
| Test-only deps | the diffusers `GlmImageTransformer2DModel` / `GlmImagePipeline` reference classes need `diffusers>=0.37.0.dev0` (used **only** by the parity tests) |
|
||||
|
||||
Download weights (from the repo root; cache via `HF_HOME` if `/` is small):
|
||||
|
||||
```bash
|
||||
python ".agents/skills/add-model-01-prep/scripts/download_hf_weights.py" \
|
||||
"zai-org/GLM-Image" "official_weights/glm_image"
|
||||
```
|
||||
|
||||
Optionally install the diffusers reference classes to run the parity side:
|
||||
|
||||
```bash
|
||||
uv pip install -U "git+https://github.com/huggingface/diffusers.git@ff3b86b4"
|
||||
```
|
||||
|
||||
Parity was checked against transformers `a65bf6c9` + diffusers `ff3b86b4`.
|
||||
|
||||
## Parity results (GB200)
|
||||
|
||||
| Component | Test | Result |
|
||||
|---|---|---|
|
||||
| `transformer` (DiT) | `test_transformer_parity.py` | fp32 cosine `1.000000` (100% within 5e-3); bf16 cosine `0.999969` (element-wise rounding over 30 layers vs a different SDPA kernel — not a bug; RoPE proven exactly equivalent). fp32 + bf16 + strict-load |
|
||||
| `vae` (`AutoencoderKL`) | `test_vae_parity.py` | decode + encode + config, 3/3 |
|
||||
| `text_encoder` (T5) + **ByT5** tokenizer | `test_t5_parity.py` | tokenizer + glyph extract + embeds, 3/3 |
|
||||
| `vision_language_encoder` (GLM-4-9B AR) | `test_ar_parity.py` | surface + to/eval + greedy determinism, 3/3 |
|
||||
| `pipeline` (T2I) | `test_pipeline_parity.py` | injected-input determinism: latent cosine `0.999997`, decoded-image MAE `0.159/255`. validity + deterministic parity, 2/2 |
|
||||
| `pipeline` (I2I / edit) | `test_edit_pipeline_parity.py` | validity / wiring gate (condition image enters via a KV-cache write pass, not noise-the-latent). diffusers distribution parity deferred (stochastic AR) |
|
||||
|
||||
`processor` (`GlmImageProcessor`) and `scheduler` (`FlowMatchEulerDiscreteScheduler`,
|
||||
shared flow-matching) are exercised via the AR and pipeline tests.
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
pytest tests/local_tests/glm_image/test_transformer_parity.py -v -s
|
||||
pytest tests/local_tests/glm_image/test_vae_parity.py -v -s
|
||||
pytest tests/local_tests/glm_image/test_t5_parity.py -v -s
|
||||
pytest tests/local_tests/glm_image/test_ar_parity.py -v -s
|
||||
pytest tests/local_tests/glm_image/test_pipeline_parity.py -v -s
|
||||
pytest tests/local_tests/glm_image/test_edit_pipeline_parity.py -v -s
|
||||
```
|
||||
|
||||
## Design notes (load-bearing — don't regress)
|
||||
|
||||
- **Pipeline parity is by injection.** Full-pipeline pixel parity is ill-posed
|
||||
(stochastic AR + independent latent RNG), so the pipeline tests inject matched
|
||||
`(prompt_embeds, prior_token_ids, latents)` into both the diffusers pipeline
|
||||
and the real `GlmImageDenoisingStage`, then compare denoised latents + decoded
|
||||
image — not raw end-to-end pixels.
|
||||
- **HF imports are confined** to the lazy-wrapper loader
|
||||
`fastvideo/models/encoders/glm_image_ar_loader.py`, never the production
|
||||
loaders (production-boundary rule).
|
||||
- **Tokenizer is ByT5**, matching the diffusers reference (not plain T5);
|
||||
`compute_glyph_embeds` mirrors diffusers `_get_glyph_embeds`.
|
||||
- **Model-specific stages** live under
|
||||
`fastvideo/pipelines/basic/glm_image/stages/` per the `add-model` Files Map.
|
||||
- **Param mapping** is handled in-place at load via `param_names_mapping`
|
||||
(`fastvideo/configs/models/dits/glm_image.py`) + the native VAE/encoder
|
||||
configs — no conversion script; the loader raises on any unmatched param and
|
||||
`test_transformer_strict_load.py` asserts completeness with `strict=True`.
|
||||
@@ -0,0 +1,109 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29519")
|
||||
os.environ.setdefault("DISABLE_SP", "1")
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
FAMILY = "glm_image"
|
||||
LOCAL_WEIGHTS_DIR = Path(
|
||||
os.getenv("GLM_IMAGE_LOCAL_WEIGHTS_DIR",
|
||||
REPO_ROOT / "official_weights" / FAMILY))
|
||||
AR_DIR = LOCAL_WEIGHTS_DIR / "vision_language_encoder"
|
||||
PROCESSOR_DIR = LOCAL_WEIGHTS_DIR / "processor"
|
||||
|
||||
|
||||
def _has_weights() -> bool:
|
||||
return AR_DIR.exists() and any(AR_DIR.glob("*.safetensors"))
|
||||
|
||||
|
||||
def _has_glm_image_transformers() -> bool:
|
||||
import importlib.util
|
||||
spec = importlib.util.find_spec("transformers")
|
||||
if spec is None:
|
||||
return False
|
||||
import transformers
|
||||
return hasattr(transformers, "GlmImageForConditionalGeneration")
|
||||
|
||||
|
||||
pytestmark = [
|
||||
pytest.mark.skipif(
|
||||
not _has_weights(),
|
||||
reason=f"GLM-Image AR encoder weights not found at {AR_DIR}.",
|
||||
),
|
||||
pytest.mark.skipif(
|
||||
not _has_glm_image_transformers(),
|
||||
reason=("transformers in this env lacks "
|
||||
"`GlmImageForConditionalGeneration` (main pins 4.57.3; needs "
|
||||
"5.0.0rc0+). Bump transformers locally to exercise the AR "
|
||||
"lazy-wrapper."),
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
SAMPLE_PROMPT = (
|
||||
"A beautiful landscape photography with rolling hills, a winding river, "
|
||||
"and a vibrant sunset in the background. Photorealistic style."
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def device() -> torch.device:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA required for GLM-Image AR encoder parity.")
|
||||
return torch.device("cuda")
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def fastvideo_ar(device):
|
||||
pytest.importorskip("fastvideo")
|
||||
try:
|
||||
from fastvideo.models.encoders.glm_image_ar_loader import (
|
||||
GlmImageARLoader)
|
||||
except ImportError as e:
|
||||
pytest.skip(f"FastVideo lazy-wrapper not yet at target path: {e}")
|
||||
loader = GlmImageARLoader(str(AR_DIR), str(PROCESSOR_DIR),
|
||||
torch_dtype=torch.bfloat16)
|
||||
loader.to(device).eval()
|
||||
return loader
|
||||
|
||||
|
||||
def test_ar_wrapper_exposes_expected_surface(fastvideo_ar):
|
||||
assert callable(getattr(fastvideo_ar, "generate", None))
|
||||
assert getattr(fastvideo_ar, "processor", None) is not None
|
||||
assert getattr(fastvideo_ar, "config", None) is not None
|
||||
assert getattr(fastvideo_ar, "generation_config", None) is not None
|
||||
from transformers import GlmImageForConditionalGeneration
|
||||
assert isinstance(fastvideo_ar._model, GlmImageForConditionalGeneration)
|
||||
|
||||
|
||||
def test_ar_to_and_eval_propagate(fastvideo_ar, device):
|
||||
assert next(fastvideo_ar._model.parameters()).device.type == device.type
|
||||
assert not fastvideo_ar._model.training
|
||||
|
||||
|
||||
def test_ar_generate_is_deterministic_under_do_sample_false(fastvideo_ar,
|
||||
device):
|
||||
processor = fastvideo_ar.processor
|
||||
messages = [{
|
||||
"role": "user",
|
||||
"content": [{"type": "text", "text": SAMPLE_PROMPT}],
|
||||
}]
|
||||
inputs = processor.apply_chat_template(messages, tokenize=True,
|
||||
target_h=1024, target_w=1024,
|
||||
return_dict=True,
|
||||
return_tensors="pt").to(device)
|
||||
with torch.no_grad():
|
||||
out_a = fastvideo_ar.generate(**inputs, max_new_tokens=32,
|
||||
do_sample=False)
|
||||
out_b = fastvideo_ar.generate(**inputs, max_new_tokens=32,
|
||||
do_sample=False)
|
||||
assert out_a.shape == out_b.shape
|
||||
assert torch.equal(out_a, out_b)
|
||||
@@ -0,0 +1,126 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""End-to-end scaffold for the GLM-Image image-to-image (edit) pipeline.
|
||||
|
||||
Drives FastVideo `VideoGenerator` in edit mode (a condition image is passed) and
|
||||
checks that the unified pipeline routes through the condition-encoding stage +
|
||||
KV-cache denoising path and emits a well-formed image. Like the T2I parity test,
|
||||
pixel-exact parity vs diffusers is NOT asserted: the AR prior is sampled
|
||||
(`do_sample=True`) and the diffusion latents draw from independent RNG streams,
|
||||
so the codebases produce different valid samples of the same edit. We gate on
|
||||
validity (and, when diffusers is present, a comparable brightness regime).
|
||||
|
||||
Skips cleanly until weights are available; GPU-heavy.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29520")
|
||||
os.environ.setdefault("DISABLE_SP", "1")
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
FAMILY = "glm_image"
|
||||
LOCAL_WEIGHTS_DIR = Path(
|
||||
os.getenv("GLM_IMAGE_LOCAL_WEIGHTS_DIR",
|
||||
REPO_ROOT / "official_weights" / FAMILY))
|
||||
CONDITION_IMAGE = REPO_ROOT / "assets" / "images" / "couple.jpg"
|
||||
|
||||
|
||||
def _has_weights() -> bool:
|
||||
required = ["transformer", "vae", "text_encoder",
|
||||
"vision_language_encoder", "processor", "tokenizer",
|
||||
"scheduler"]
|
||||
return all((LOCAL_WEIGHTS_DIR / r).exists() for r in required)
|
||||
|
||||
|
||||
def _upstream_glm_image_available() -> bool:
|
||||
try:
|
||||
import transformers
|
||||
except ImportError:
|
||||
return False
|
||||
return hasattr(transformers, "GlmImageForConditionalGeneration")
|
||||
|
||||
|
||||
pytestmark = [
|
||||
pytest.mark.skipif(
|
||||
not _has_weights(),
|
||||
reason=f"GLM-Image full weights not found at {LOCAL_WEIGHTS_DIR}.",
|
||||
),
|
||||
pytest.mark.skipif(
|
||||
not _upstream_glm_image_available(),
|
||||
reason=("Edit pipeline needs transformers>=5.0.0rc0 "
|
||||
"(GlmImageForConditionalGeneration); main pin predates it. "
|
||||
"Bump locally to run."),
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
EDIT_PROMPT = "Make the scene a snowy winter landscape."
|
||||
SEED = 0
|
||||
HEIGHT = 512
|
||||
WIDTH = 512
|
||||
STEPS = 8 # low step count keeps the test under 5 min on a single GPU
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def device():
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA required for GLM-Image edit pipeline.")
|
||||
return torch.device("cuda")
|
||||
|
||||
|
||||
def _to_uint8_hwc(a) -> np.ndarray:
|
||||
"""Coerce a frame buffer (np or torch, [0,1] float or uint8, with optional
|
||||
leading batch/time dims, HWC or CHW) to a single (H, W, 3) uint8 image."""
|
||||
if torch.is_tensor(a):
|
||||
a = a.float().cpu().numpy()
|
||||
a = np.asarray(a)
|
||||
while a.ndim > 3:
|
||||
a = a[0]
|
||||
if a.ndim == 3 and a.shape[0] in (1, 3) and a.shape[-1] not in (1, 3):
|
||||
a = np.transpose(a, (1, 2, 0))
|
||||
if a.dtype != np.uint8:
|
||||
scale = 255.0 if float(a.max()) <= 1.5 else 1.0
|
||||
a = np.clip(a * scale, 0, 255).astype(np.uint8)
|
||||
return a
|
||||
|
||||
|
||||
def _fastvideo_edit_image(device) -> np.ndarray:
|
||||
pytest.importorskip("fastvideo")
|
||||
try:
|
||||
from fastvideo import VideoGenerator
|
||||
except ImportError as e:
|
||||
pytest.skip(f"FastVideo VideoGenerator unavailable: {e}")
|
||||
condition = Image.open(CONDITION_IMAGE).convert("RGB")
|
||||
gen = VideoGenerator.from_pretrained(str(LOCAL_WEIGHTS_DIR), num_gpus=1,
|
||||
trust_remote_code=True)
|
||||
result = gen.generate_video(prompt=EDIT_PROMPT,
|
||||
pil_image=condition,
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
height=HEIGHT, width=WIDTH,
|
||||
num_inference_steps=STEPS,
|
||||
guidance_scale=1.5,
|
||||
seed=SEED)
|
||||
gen.shutdown()
|
||||
return _to_uint8_hwc(result["frames"][0])
|
||||
|
||||
|
||||
def test_edit_pipeline_produces_valid_image(device):
|
||||
"""Wiring gate for the edit path: passing a condition image routes through
|
||||
the condition-encoding (KV-cache write) stage and the read/skip denoising
|
||||
path, and the pipeline emits a well-formed, non-degenerate image of the
|
||||
requested size. Component numerical correctness is covered by the (green)
|
||||
DiT / VAE / T5 / AR component-parity tests."""
|
||||
fv = _fastvideo_edit_image(device)
|
||||
assert fv.shape == (HEIGHT, WIDTH, 3), f"unexpected shape {fv.shape}"
|
||||
assert fv.dtype == np.uint8 and np.isfinite(fv).all()
|
||||
assert fv.std() > 10.0, f"image is near-constant (std={fv.std():.2f})"
|
||||
@@ -0,0 +1,331 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29519")
|
||||
os.environ.setdefault("DISABLE_SP", "1")
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
FAMILY = "glm_image"
|
||||
LOCAL_WEIGHTS_DIR = Path(
|
||||
os.getenv("GLM_IMAGE_LOCAL_WEIGHTS_DIR",
|
||||
REPO_ROOT / "official_weights" / FAMILY))
|
||||
TRANSFORMER_DIR = LOCAL_WEIGHTS_DIR / "transformer"
|
||||
|
||||
|
||||
def _has_weights() -> bool:
|
||||
required = ["transformer", "vae", "text_encoder",
|
||||
"vision_language_encoder", "processor", "tokenizer",
|
||||
"scheduler"]
|
||||
return all((LOCAL_WEIGHTS_DIR / r).exists() for r in required)
|
||||
|
||||
|
||||
def _upstream_glm_image_available() -> bool:
|
||||
try:
|
||||
import transformers
|
||||
import diffusers
|
||||
except ImportError:
|
||||
return False
|
||||
return (hasattr(transformers, "GlmImageForConditionalGeneration")
|
||||
and hasattr(diffusers, "GlmImagePipeline"))
|
||||
|
||||
|
||||
pytestmark = [
|
||||
pytest.mark.skipif(
|
||||
not _has_weights(),
|
||||
reason=f"GLM-Image full weights not found at {LOCAL_WEIGHTS_DIR}.",
|
||||
),
|
||||
pytest.mark.skipif(
|
||||
not _upstream_glm_image_available(),
|
||||
reason=("Pipeline parity needs transformers>=5.0.0rc0 and "
|
||||
"diffusers>=0.37.0.dev0; main pins predate both. Bump locally "
|
||||
"to run."),
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
SAMPLE_PROMPT = (
|
||||
"A landscape photo with rolling green hills under a clear blue sky.")
|
||||
SEED = 0
|
||||
HEIGHT = 512
|
||||
WIDTH = 512
|
||||
STEPS = 8
|
||||
|
||||
# bf16 + a different SDPA kernel across 30 DiT layers x STEPS steps leaves a
|
||||
# small residual; a real wiring/schedule bug is far larger than these bounds.
|
||||
LATENT_COSINE_MIN = 0.995
|
||||
IMAGE_MAE_MAX = 5.0 # /255
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def device():
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA required for GLM-Image pipeline parity.")
|
||||
return torch.device("cuda")
|
||||
|
||||
|
||||
def _to_uint8_hwc(a) -> np.ndarray:
|
||||
if torch.is_tensor(a):
|
||||
a = a.detach().float().cpu().numpy()
|
||||
a = np.asarray(a)
|
||||
while a.ndim > 3:
|
||||
a = a[0]
|
||||
if a.ndim == 3 and a.shape[0] in (1, 3) and a.shape[-1] not in (1, 3):
|
||||
a = np.transpose(a, (1, 2, 0))
|
||||
if a.dtype != np.uint8:
|
||||
scale = 255.0 if float(a.max()) <= 1.5 else 1.0
|
||||
a = np.clip(a * scale, 0, 255).astype(np.uint8)
|
||||
return a
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# FastVideo native DiT loader (mirrors test_transformer_parity.py).
|
||||
# --------------------------------------------------------------------------- #
|
||||
def _load_state_dict(dir_path: Path) -> dict[str, torch.Tensor]:
|
||||
safetensors = pytest.importorskip("safetensors.torch")
|
||||
sd: dict[str, torch.Tensor] = {}
|
||||
for shard in sorted(dir_path.glob("*.safetensors")):
|
||||
sd.update(safetensors.load_file(str(shard)))
|
||||
return sd
|
||||
|
||||
|
||||
def _apply_param_mapping(sd, mapping):
|
||||
import re
|
||||
out = {}
|
||||
for k, v in sd.items():
|
||||
new_k = k
|
||||
for pat, repl in mapping.items():
|
||||
if re.match(pat, k):
|
||||
new_k = re.sub(pat, repl, k)
|
||||
break
|
||||
out[new_k] = v
|
||||
return out
|
||||
|
||||
|
||||
def _ensure_distributed():
|
||||
"""The denoising stage calls get_local_torch_device(), which needs
|
||||
FastVideo's world/TP groups. Initialize a single-rank (world_size=1) group
|
||||
in-process (mirrors tests/local_tests/sd35/test_sd35_component_parity.py)."""
|
||||
import torch.distributed as dist
|
||||
from fastvideo.distributed.parallel_state import (
|
||||
get_tp_group, init_distributed_environment, initialize_model_parallel)
|
||||
try:
|
||||
get_tp_group()
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
if not dist.is_initialized():
|
||||
os.environ.setdefault("RANK", "0")
|
||||
os.environ.setdefault("WORLD_SIZE", "1")
|
||||
os.environ.setdefault("LOCAL_RANK", "0")
|
||||
torch.cuda.set_device(int(os.environ["LOCAL_RANK"]))
|
||||
store_path = f"/tmp/fastvideo_glm_pg_{os.getpid()}.store"
|
||||
dist.init_process_group(backend="nccl",
|
||||
init_method=f"file://{store_path}",
|
||||
rank=0, world_size=1)
|
||||
init_distributed_environment(world_size=1, rank=0, local_rank=0,
|
||||
distributed_init_method="env://")
|
||||
try:
|
||||
get_tp_group()
|
||||
except Exception:
|
||||
initialize_model_parallel(tensor_model_parallel_size=1,
|
||||
sequence_model_parallel_size=1,
|
||||
data_parallel_size=1)
|
||||
|
||||
|
||||
def _load_fastvideo_transformer(device, dtype):
|
||||
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
|
||||
from fastvideo.models.dits.glm_image import GlmImageTransformer2DModel
|
||||
cfg = GlmImageDiTConfig()
|
||||
model = GlmImageTransformer2DModel(
|
||||
cfg, {"_class_name": "GlmImageTransformer2DModel"})
|
||||
sd = _load_state_dict(TRANSFORMER_DIR)
|
||||
sd = _apply_param_mapping(sd, cfg.arch_config.param_names_mapping)
|
||||
missing, unexpected = model.load_state_dict(sd, strict=False)
|
||||
assert not missing, f"FastVideo DiT missing keys: {missing[:10]}"
|
||||
assert not unexpected, f"FastVideo DiT unexpected keys: {unexpected[:10]}"
|
||||
return model.to(device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def _fastvideo_denoise_latents(device, dtype, *, prompt_embeds, prior_token_ids,
|
||||
init_latents):
|
||||
"""Drive the real GlmImageDenoisingStage with injected, matched inputs and
|
||||
return the denoised latents (1, 16, H/8, W/8)."""
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.pipelines.basic.glm_image.stages.before_denoising import (
|
||||
calculate_shift)
|
||||
from fastvideo.pipelines.basic.glm_image.stages.denoising import (
|
||||
GlmImageDenoisingStage)
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
_ensure_distributed()
|
||||
transformer = _load_fastvideo_transformer(device, dtype)
|
||||
scheduler = FlowMatchEulerDiscreteScheduler(shift=1.0)
|
||||
stage = GlmImageDenoisingStage(transformer=transformer, scheduler=scheduler)
|
||||
|
||||
# Reproduce the production timestep/sigma schedule (before_denoising.py):
|
||||
# integer-cast linspace timesteps, resolution-dependent shift on sigmas.
|
||||
ntt = scheduler.config.num_train_timesteps
|
||||
patch = transformer.patch_size
|
||||
image_seq_len = ((HEIGHT // 8) * (WIDTH // 8)) // (patch**2)
|
||||
sched_t = np.linspace(ntt, 1.0, STEPS + 1)[:-1].astype(
|
||||
np.int64).astype(np.float32)
|
||||
scheduler.set_shift(calculate_shift(image_seq_len))
|
||||
scheduler.set_timesteps(STEPS, device=device,
|
||||
sigmas=(sched_t / ntt).tolist(),
|
||||
timesteps=sched_t.tolist())
|
||||
|
||||
batch = ForwardBatch(data_type="image")
|
||||
batch.prompt_embeds = [prompt_embeds.to(device, dtype)]
|
||||
batch.attention_mask = None # match diffusers (no text padding mask)
|
||||
batch.prior_token_id = prior_token_ids.to(device)
|
||||
batch.prior_token_drop = torch.zeros_like(prior_token_ids,
|
||||
dtype=torch.bool, device=device)
|
||||
batch.latents = init_latents.clone().unsqueeze(2).to(device, dtype)
|
||||
batch.timesteps = scheduler.timesteps
|
||||
batch.height, batch.width = HEIGHT, WIDTH
|
||||
batch.num_inference_steps = STEPS
|
||||
batch.guidance_scale = 1.5
|
||||
batch.do_classifier_free_guidance = True
|
||||
batch.seed = SEED
|
||||
batch.extra = {} # no glm_kv_caches -> T2I denoise path
|
||||
|
||||
out = stage.forward(batch, fastvideo_args=None)
|
||||
latents = out.latents
|
||||
if latents.dim() == 5:
|
||||
latents = latents.squeeze(2)
|
||||
latents = latents.detach().float().cpu()
|
||||
del transformer, stage
|
||||
torch.cuda.empty_cache()
|
||||
import torch.distributed as dist
|
||||
if dist.is_initialized():
|
||||
dist.destroy_process_group()
|
||||
return latents
|
||||
|
||||
|
||||
def _fastvideo_image(device) -> np.ndarray:
|
||||
pytest.importorskip("fastvideo")
|
||||
try:
|
||||
from fastvideo import VideoGenerator
|
||||
except ImportError as e:
|
||||
pytest.skip(f"FastVideo VideoGenerator unavailable: {e}")
|
||||
gen = VideoGenerator.from_pretrained(str(LOCAL_WEIGHTS_DIR), num_gpus=1,
|
||||
trust_remote_code=True)
|
||||
result = gen.generate_video(prompt=SAMPLE_PROMPT,
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
height=HEIGHT, width=WIDTH,
|
||||
num_inference_steps=STEPS,
|
||||
guidance_scale=1.5,
|
||||
seed=SEED)
|
||||
gen.shutdown()
|
||||
return _to_uint8_hwc(result["frames"][0])
|
||||
|
||||
|
||||
def test_fastvideo_pipeline_produces_valid_image(device):
|
||||
fv = _fastvideo_image(device)
|
||||
assert fv.shape == (HEIGHT, WIDTH, 3), f"unexpected shape {fv.shape}"
|
||||
assert fv.dtype == np.uint8 and np.isfinite(fv).all()
|
||||
assert fv.std() > 10.0, f"image is near-constant (std={fv.std():.2f})"
|
||||
|
||||
|
||||
def test_pipeline_denoise_parity_deterministic(device):
|
||||
"""Inject identical (prompt_embeds, prior_token_ids, init_latents) into the
|
||||
diffusers pipeline and FastVideo's GlmImageDenoisingStage, then compare the
|
||||
denoised latents and the decoded images. With the stochastic AR and the
|
||||
independent latent RNG removed, the two implementations must agree to within
|
||||
bf16/SDPA-kernel noise."""
|
||||
# Inference only: the T5 glyph encoder leaves a grad graph on the embeds,
|
||||
# which later breaks image_processor.postprocess()'s .numpy() on the decoded
|
||||
# image. Disabling grad for the whole test also trims memory.
|
||||
torch.set_grad_enabled(False)
|
||||
diffusers = pytest.importorskip("diffusers")
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
dtype = torch.bfloat16
|
||||
pipe = diffusers.GlmImagePipeline.from_pretrained(
|
||||
str(LOCAL_WEIGHTS_DIR), torch_dtype=dtype).to(device)
|
||||
|
||||
gen = torch.Generator(device=device).manual_seed(SEED)
|
||||
|
||||
# --- shared, deterministic inputs (computed once, fed to both) ---------- #
|
||||
pos, neg = pipe.encode_prompt(SAMPLE_PROMPT,
|
||||
do_classifier_free_guidance=True,
|
||||
device=device, dtype=dtype)
|
||||
assert pos.shape[1] == neg.shape[1], (
|
||||
f"quote-free prompt should give equal pos/neg glyph lengths, got "
|
||||
f"{pos.shape[1]} vs {neg.shape[1]}")
|
||||
prior_token_ids, _, _ = pipe.generate_prior_tokens(
|
||||
SAMPLE_PROMPT, HEIGHT, WIDTH, image=None, device=device, generator=gen)
|
||||
latent_ch = pipe.transformer.config.in_channels
|
||||
init_latents = randn_tensor((1, latent_ch, HEIGHT // 8, WIDTH // 8),
|
||||
generator=gen, device=device, dtype=dtype)
|
||||
|
||||
# --- official denoised latents ----------------------------------------- #
|
||||
# prompt=None: diffusers check_inputs rejects passing prompt + prompt_embeds
|
||||
# together; the embeds (and prior tokens) fully specify the run.
|
||||
official = pipe(prompt=None,
|
||||
prompt_embeds=pos, negative_prompt_embeds=neg,
|
||||
prior_token_ids=prior_token_ids,
|
||||
latents=init_latents.clone(),
|
||||
height=HEIGHT, width=WIDTH,
|
||||
num_inference_steps=STEPS, guidance_scale=1.5,
|
||||
output_type="latent").images.float().cpu()
|
||||
|
||||
# FastVideo packs CFG as a single [pos; neg] 2-row tensor (no padding here
|
||||
# because L_pos == L_neg, asserted above).
|
||||
packed = torch.cat([pos, neg], dim=0)
|
||||
|
||||
# --- FastVideo denoised latents (real GlmImageDenoisingStage) ---------- #
|
||||
fv = _fastvideo_denoise_latents(device, dtype,
|
||||
prompt_embeds=packed,
|
||||
prior_token_ids=prior_token_ids,
|
||||
init_latents=init_latents)
|
||||
|
||||
assert official.shape == fv.shape, (
|
||||
f"latent shape mismatch: {official.shape} vs {fv.shape}")
|
||||
|
||||
# --- latent-level agreement -------------------------------------------- #
|
||||
cos = torch.nn.functional.cosine_similarity(
|
||||
official.flatten(), fv.flatten(), dim=0).item()
|
||||
lat_mae = (official - fv).abs().mean().item()
|
||||
lat_scale = official.abs().mean().item()
|
||||
print(f"[glm-image denoise parity] latent cosine={cos:.6f} "
|
||||
f"MAE={lat_mae:.5f} (|latent| mean={lat_scale:.5f}, "
|
||||
f"rel={lat_mae / max(lat_scale, 1e-6):.4f})")
|
||||
|
||||
# --- decoded-image agreement (same VAE both sides isolates denoise) ----- #
|
||||
def _decode(lat):
|
||||
lat = lat.to(device, dtype)
|
||||
mean = torch.tensor(pipe.vae.config.latents_mean).view(
|
||||
1, -1, 1, 1).to(device, dtype)
|
||||
std = torch.tensor(pipe.vae.config.latents_std).view(
|
||||
1, -1, 1, 1).to(device, dtype)
|
||||
img = pipe.vae.decode((lat * std + mean), return_dict=False)[0]
|
||||
return _to_uint8_hwc(pipe.image_processor.postprocess(
|
||||
img, output_type="np")[0])
|
||||
|
||||
off_img, fv_img = _decode(official), _decode(fv)
|
||||
img_mae = np.abs(off_img.astype(np.float32)
|
||||
- fv_img.astype(np.float32)).mean()
|
||||
print(f"[glm-image denoise parity] decoded image MAE={img_mae:.3f}/255 "
|
||||
f"(diffusers mean={off_img.mean():.1f}, fv mean={fv_img.mean():.1f})")
|
||||
|
||||
del pipe
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
assert cos > LATENT_COSINE_MIN, (
|
||||
f"denoised-latent cosine {cos:.6f} < {LATENT_COSINE_MIN}: the FastVideo "
|
||||
"and official denoise diverge well beyond kernel noise (wiring bug).")
|
||||
assert img_mae < IMAGE_MAE_MAX, (
|
||||
f"decoded-image MAE {img_mae:.3f}/255 >= {IMAGE_MAE_MAX}: outputs are "
|
||||
"not equivalent under matched inputs (wiring/schedule bug).")
|
||||
@@ -0,0 +1,180 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Parity scaffold for GLM-Image glyph T5 encoder.
|
||||
|
||||
The text_encoder is a small byte-level T5 (d_model=1472, gated-gelu, ByT5
|
||||
tokenizer, vocab_size=384). This test exercises the glyph-only encoding path
|
||||
that the diffusers pipeline implements as `_get_glyph_embeds` and compares
|
||||
against the FastVideo glyph-encoding stage logic. Until the FastVideo stage is
|
||||
refactored to mirror diffusers (per-prompt loop, even-length pad-prefix,
|
||||
attention-mask flattening, left-padded batch), the parity test asserts the
|
||||
high-impact pieces individually so failures point to the exact path that
|
||||
diverged.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29519")
|
||||
os.environ.setdefault("DISABLE_SP", "1")
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
FAMILY = "glm_image"
|
||||
LOCAL_WEIGHTS_DIR = Path(
|
||||
os.getenv("GLM_IMAGE_LOCAL_WEIGHTS_DIR",
|
||||
REPO_ROOT / "official_weights" / FAMILY))
|
||||
TEXT_ENCODER_DIR = LOCAL_WEIGHTS_DIR / "text_encoder"
|
||||
TOKENIZER_DIR = LOCAL_WEIGHTS_DIR / "tokenizer"
|
||||
|
||||
|
||||
def _has_weights() -> bool:
|
||||
return (TEXT_ENCODER_DIR / "model.safetensors").exists() and TOKENIZER_DIR.exists()
|
||||
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not _has_weights(),
|
||||
reason=f"GLM-Image text_encoder/tokenizer not found under {LOCAL_WEIGHTS_DIR}.",
|
||||
)
|
||||
|
||||
|
||||
SAMPLE_PROMPT = (
|
||||
'A photo of a coffee shop with a wooden sign reading "Daily Grind" '
|
||||
"and a chalk menu next to it"
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def device() -> torch.device:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA required for GLM-Image glyph encoder parity.")
|
||||
return torch.device("cuda")
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def official_tokenizer():
|
||||
transformers = pytest.importorskip("transformers")
|
||||
return transformers.ByT5Tokenizer.from_pretrained(str(TOKENIZER_DIR))
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def official_text_encoder(device):
|
||||
transformers = pytest.importorskip("transformers")
|
||||
return transformers.T5EncoderModel.from_pretrained(
|
||||
str(TEXT_ENCODER_DIR), torch_dtype=torch.float32).to(device).eval()
|
||||
|
||||
|
||||
def test_tokenizer_is_byt5(official_tokenizer):
|
||||
"""Glyph tokenizer must be ByT5 (byte-level), not generic T5."""
|
||||
assert type(official_tokenizer).__name__ == "ByT5Tokenizer"
|
||||
assert official_tokenizer.vocab_size <= 400, (
|
||||
"ByT5 vocab is byte-level (~256+special)")
|
||||
|
||||
|
||||
def test_glyph_text_extraction():
|
||||
"""Glyph extraction must pull single+double+CJK quoted spans."""
|
||||
pytest.importorskip("fastvideo")
|
||||
try:
|
||||
from fastvideo.pipelines.basic.glm_image.stages.before_denoising import ( # noqa
|
||||
get_glyph_texts)
|
||||
except ImportError as e:
|
||||
pytest.skip(f"FastVideo glyph extractor not yet at target path: {e}")
|
||||
out = get_glyph_texts(SAMPLE_PROMPT)
|
||||
assert "Daily Grind" in out
|
||||
|
||||
|
||||
def _diffusers_reference_get_glyph_embeds(prompts, tokenizer, text_encoder,
|
||||
device, dtype,
|
||||
max_sequence_length=2048):
|
||||
"""Inlined copy of diffusers `_get_glyph_embeds` + `get_glyph_texts`.
|
||||
|
||||
Source: `reference/diffusers/src/diffusers/pipelines/glm_image/`
|
||||
`pipeline_glm_image.py::GlmImagePipeline._get_glyph_embeds` at commit
|
||||
ff3b86b4755b46a7b5656dfcf84d25bd25ad4740 (diffusers main, 0.37.0.dev0).
|
||||
Inlined here because the installed diffusers package (`__version__`
|
||||
typically <0.37.0) does not yet ship `GlmImagePipeline`. This function is
|
||||
the reference oracle for FastVideo's `compute_glyph_embeds`.
|
||||
"""
|
||||
import re
|
||||
|
||||
def get_glyph_texts(prompt):
|
||||
if isinstance(prompt, str):
|
||||
prompt = [prompt]
|
||||
out = []
|
||||
for p in prompt:
|
||||
out.append(
|
||||
re.findall(r"'([^']*)'", p)
|
||||
+ re.findall(r"“([^“”]*)”", p)
|
||||
+ re.findall(r'"([^"]*)"', p)
|
||||
+ re.findall(r"「([^「」]*)」", p))
|
||||
return out
|
||||
|
||||
all_glyph_texts = get_glyph_texts(prompts)
|
||||
all_glyph_embeds = []
|
||||
for glyph_texts in all_glyph_texts:
|
||||
if len(glyph_texts) == 0:
|
||||
glyph_texts = [""]
|
||||
input_ids = tokenizer(
|
||||
glyph_texts,
|
||||
max_length=max_sequence_length,
|
||||
truncation=True,
|
||||
).input_ids
|
||||
input_ids = [
|
||||
[tokenizer.pad_token_id] * ((len(input_ids) + 1) % 2) + ids
|
||||
for ids in input_ids
|
||||
]
|
||||
max_length = max(len(ids) for ids in input_ids)
|
||||
attention_mask = torch.tensor(
|
||||
[[1] * len(ids) + [0] * (max_length - len(ids))
|
||||
for ids in input_ids],
|
||||
device=device,
|
||||
)
|
||||
input_ids_t = torch.tensor(
|
||||
[
|
||||
ids + [tokenizer.pad_token_id] * (max_length - len(ids))
|
||||
for ids in input_ids
|
||||
],
|
||||
device=device,
|
||||
)
|
||||
outputs = text_encoder(input_ids_t, attention_mask=attention_mask)
|
||||
glyph_embeds = outputs.last_hidden_state[attention_mask.bool()].unsqueeze(0)
|
||||
all_glyph_embeds.append(glyph_embeds)
|
||||
|
||||
max_seq_len = max(emb.size(1) for emb in all_glyph_embeds)
|
||||
padded = []
|
||||
for emb in all_glyph_embeds:
|
||||
if emb.size(1) < max_seq_len:
|
||||
pad = torch.zeros(emb.size(0), max_seq_len - emb.size(1),
|
||||
emb.size(2), device=device, dtype=emb.dtype)
|
||||
emb = torch.cat([pad, emb], dim=1)
|
||||
padded.append(emb)
|
||||
return torch.cat(padded, dim=0).to(device=device, dtype=dtype)
|
||||
|
||||
|
||||
def test_get_glyph_embeds_matches_diffusers(official_tokenizer,
|
||||
official_text_encoder, device):
|
||||
"""FastVideo `compute_glyph_embeds` must equal the diffusers reference."""
|
||||
pytest.importorskip("fastvideo")
|
||||
try:
|
||||
from fastvideo.pipelines.basic.glm_image.stages.before_denoising import ( # noqa
|
||||
compute_glyph_embeds)
|
||||
except ImportError as e:
|
||||
pytest.skip(f"FastVideo glyph encoder not yet at target path: {e}")
|
||||
|
||||
with torch.no_grad():
|
||||
official_embeds = _diffusers_reference_get_glyph_embeds(
|
||||
[SAMPLE_PROMPT], official_tokenizer, official_text_encoder, device,
|
||||
torch.float32)
|
||||
fv_embeds = compute_glyph_embeds([SAMPLE_PROMPT],
|
||||
tokenizer=official_tokenizer,
|
||||
text_encoder=official_text_encoder,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
max_sequence_length=2048)
|
||||
assert fv_embeds.shape == official_embeds.shape, (
|
||||
f"shape mismatch: {fv_embeds.shape} vs {official_embeds.shape}")
|
||||
torch.testing.assert_close(fv_embeds, official_embeds, atol=1e-4,
|
||||
rtol=1e-4)
|
||||
@@ -0,0 +1,205 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Component parity scaffold for GLM-Image DiT (GlmImageTransformer2DModel).
|
||||
|
||||
Compares the FastVideo-native port at `fastvideo.models.dits.glm_image` against
|
||||
the diffusers reference `GlmImageTransformer2DModel` from
|
||||
`diffusers.models.transformers.transformer_glm_image`.
|
||||
|
||||
Skips cleanly until both the FastVideo class and the local weights exist; the
|
||||
test is real (loads weights, forwards a fixed input, compares output tensors).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29519")
|
||||
os.environ.setdefault("DISABLE_SP", "1")
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
FAMILY = "glm_image"
|
||||
LOCAL_WEIGHTS_DIR = Path(
|
||||
os.getenv("GLM_IMAGE_LOCAL_WEIGHTS_DIR",
|
||||
REPO_ROOT / "official_weights" / FAMILY))
|
||||
TRANSFORMER_DIR = LOCAL_WEIGHTS_DIR / "transformer"
|
||||
|
||||
def _has_weights() -> bool:
|
||||
return TRANSFORMER_DIR.exists() and any(
|
||||
TRANSFORMER_DIR.glob("*.safetensors"))
|
||||
|
||||
|
||||
def _has_glm_image_diffusers() -> bool:
|
||||
"""`GlmImageTransformer2DModel` lands in diffusers>=0.37.0.dev0.
|
||||
|
||||
FastVideo main pins `diffusers>=0.33.1`; the installed env is 0.36.0. The
|
||||
numerical parity test needs the diffusers class as an oracle — until the
|
||||
installed diffusers is upgraded, this test skips. The structural surface
|
||||
is still verified by `test_glm_image_transformer_strict_load.py`.
|
||||
"""
|
||||
import importlib.util
|
||||
spec = importlib.util.find_spec("diffusers")
|
||||
if spec is None:
|
||||
return False
|
||||
import diffusers
|
||||
return hasattr(diffusers, "GlmImageTransformer2DModel")
|
||||
|
||||
|
||||
pytestmark = [
|
||||
pytest.mark.skipif(
|
||||
not _has_weights(),
|
||||
reason=(
|
||||
f"GLM-Image transformer weights not found at {TRANSFORMER_DIR}. "
|
||||
"Download via "
|
||||
"`python .agents/skills/add-model-01-prep/scripts/download_hf_weights.py "
|
||||
"zai-org/GLM-Image official_weights/glm_image`."),
|
||||
),
|
||||
pytest.mark.skipif(
|
||||
not _has_glm_image_diffusers(),
|
||||
reason=("diffusers in this env lacks `GlmImageTransformer2DModel` "
|
||||
"(main allows >=0.33.1; installed 0.36.0; class lands in "
|
||||
"0.37.0.dev0). Bump diffusers locally to run numerical parity. "
|
||||
"Structural surface is covered by "
|
||||
"`test_glm_image_transformer_strict_load.py`."),
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def device() -> torch.device:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA required for GLM-Image transformer parity.")
|
||||
return torch.device("cuda")
|
||||
|
||||
|
||||
def _load_official(device, dtype):
|
||||
diffusers = pytest.importorskip("diffusers")
|
||||
cls = diffusers.GlmImageTransformer2DModel
|
||||
return cls.from_pretrained(str(TRANSFORMER_DIR),
|
||||
torch_dtype=dtype).to(device).eval()
|
||||
|
||||
|
||||
def _load_fastvideo(device, dtype):
|
||||
pytest.importorskip("fastvideo")
|
||||
try:
|
||||
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
|
||||
from fastvideo.models.dits.glm_image import GlmImageTransformer2DModel
|
||||
except ImportError as e:
|
||||
pytest.skip(f"FastVideo GLM-Image DiT not yet ported: {e}")
|
||||
cfg = GlmImageDiTConfig()
|
||||
model = GlmImageTransformer2DModel(cfg,
|
||||
{"_class_name":
|
||||
"GlmImageTransformer2DModel"})
|
||||
sd = _load_state_dict(TRANSFORMER_DIR)
|
||||
sd = _apply_param_mapping(sd, cfg.arch_config.param_names_mapping)
|
||||
missing, unexpected = model.load_state_dict(sd, strict=False)
|
||||
assert not missing, f"FastVideo DiT missing keys: {missing[:10]}"
|
||||
assert not unexpected, f"FastVideo DiT unexpected keys: {unexpected[:10]}"
|
||||
return model.to(device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def _forward_pair(device, dtype):
|
||||
"""Run both DiTs on identical synthetic inputs; return (fv_out, ref_out)."""
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
ref = _load_official(device, dtype)
|
||||
fv = _load_fastvideo(device, dtype)
|
||||
inputs = _make_inputs(device, dtype=dtype)
|
||||
with torch.no_grad():
|
||||
ref_out = ref(**inputs, return_dict=False)[0]
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None,
|
||||
forward_batch=None):
|
||||
fv_out = fv(**inputs)
|
||||
if isinstance(fv_out, tuple):
|
||||
fv_out = fv_out[0]
|
||||
assert ref_out.shape == fv_out.shape, (
|
||||
f"shape mismatch: {ref_out.shape} vs {fv_out.shape}")
|
||||
# Free the two 7B models before the next dtype runs (single-GPU budget).
|
||||
del ref, fv
|
||||
torch.cuda.empty_cache()
|
||||
return fv_out.float(), ref_out.float()
|
||||
|
||||
|
||||
def _load_state_dict(dir_path: Path) -> dict[str, torch.Tensor]:
|
||||
safetensors = pytest.importorskip("safetensors.torch")
|
||||
sd: dict[str, torch.Tensor] = {}
|
||||
for shard in sorted(dir_path.glob("*.safetensors")):
|
||||
sd.update(safetensors.load_file(str(shard)))
|
||||
return sd
|
||||
|
||||
|
||||
def _apply_param_mapping(sd, mapping):
|
||||
import re
|
||||
out = {}
|
||||
for k, v in sd.items():
|
||||
new_k = k
|
||||
for pat, repl in mapping.items():
|
||||
if re.match(pat, k):
|
||||
new_k = re.sub(pat, repl, k)
|
||||
break
|
||||
out[new_k] = v
|
||||
return out
|
||||
|
||||
|
||||
def _make_inputs(device, batch_size=1, height=64, width=64,
|
||||
text_seq_len=32, dtype=torch.bfloat16) -> dict[str, torch.Tensor]:
|
||||
"""Synthetic inputs matching the diffusers `GlmImageTransformer2DModel.forward`
|
||||
contract. Latent shape is `(B, 16, H, W)`; post-patch grid is `(H/2, W/2)`
|
||||
so `prior_token_id` has `(H/2)*(W/2)` entries per batch element."""
|
||||
torch.manual_seed(0)
|
||||
patch_size = 2
|
||||
num_patches = (height // patch_size) * (width // patch_size)
|
||||
return {
|
||||
"hidden_states": torch.randn(batch_size, 16, height, width,
|
||||
device=device, dtype=dtype),
|
||||
"encoder_hidden_states": torch.randn(batch_size, text_seq_len, 1472,
|
||||
device=device, dtype=dtype),
|
||||
"prior_token_id": torch.randint(0, 16384, (batch_size, num_patches),
|
||||
device=device, dtype=torch.long),
|
||||
"prior_token_drop": torch.zeros(batch_size, device=device,
|
||||
dtype=torch.bool),
|
||||
"timestep": torch.tensor([500] * batch_size, device=device,
|
||||
dtype=torch.long),
|
||||
"target_size": torch.tensor([[height * 8, width * 8]] * batch_size,
|
||||
device=device, dtype=torch.long),
|
||||
"crop_coords": torch.zeros(batch_size, 2, device=device,
|
||||
dtype=torch.long),
|
||||
}
|
||||
|
||||
|
||||
# fp32 is the structural-correctness gate: with identical math the two
|
||||
# implementations agree to ~1e-3 (residual is SDPA reduction-order noise).
|
||||
ATOL_FP32 = 1e-2
|
||||
RTOL_FP32 = 1e-2
|
||||
# bf16 is the inference dtype. Rounding accumulates across 30 layers and the
|
||||
# attention kernels differ, so element-wise atol is meaningless; cosine
|
||||
# similarity is the right statistical bar for "same function, bf16 noise".
|
||||
COS_MIN_BF16 = 0.999
|
||||
|
||||
|
||||
def test_transformer_forward_matches_diffusers_fp32(device):
|
||||
"""Structural gate: in fp32 the FastVideo DiT is numerically faithful to
|
||||
the diffusers reference (cosine ~1.0, element-wise close to ~1e-3)."""
|
||||
fv_out, ref_out = _forward_pair(device, torch.float32)
|
||||
cos = torch.nn.functional.cosine_similarity(
|
||||
fv_out.flatten(), ref_out.flatten(), dim=0)
|
||||
assert cos > 0.9999, f"fp32 cosine similarity too low: {cos:.6f}"
|
||||
torch.testing.assert_close(fv_out, ref_out,
|
||||
atol=ATOL_FP32, rtol=RTOL_FP32)
|
||||
|
||||
|
||||
def test_transformer_forward_matches_diffusers_bf16(device):
|
||||
"""Inference-dtype sanity: bf16 output tracks the diffusers reference up to
|
||||
precision accumulation (high cosine similarity). Element-wise atol is not a
|
||||
meaningful bar for a 30-layer bf16 transformer vs a different attention
|
||||
kernel (~28% of elements exceed atol=5e-3 purely from rounding) — see
|
||||
PORT_STATUS I020."""
|
||||
fv_out, ref_out = _forward_pair(device, torch.bfloat16)
|
||||
cos = torch.nn.functional.cosine_similarity(
|
||||
fv_out.flatten(), ref_out.flatten(), dim=0)
|
||||
assert cos > COS_MIN_BF16, (
|
||||
f"bf16 cosine similarity {cos:.6f} below {COS_MIN_BF16}; "
|
||||
"indicates a structural divergence, not just precision.")
|
||||
@@ -0,0 +1,94 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Strict-load smoke for GLM-Image DiT.
|
||||
|
||||
Loads all three transformer shards from `official_weights/glm_image/transformer/`,
|
||||
applies the FastVideo `param_names_mapping`, and calls
|
||||
`GlmImageTransformer2DModel.load_state_dict(..., strict=True)`. Asserts the
|
||||
config-driven model instantiates and every checkpoint key has a matching
|
||||
FastVideo parameter with the same shape and dtype.
|
||||
|
||||
This is weight-load verification only — it does not run a forward pass. Numerical
|
||||
parity against the diffusers reference lives in
|
||||
`test_glm_image_transformer_parity.py` and requires `diffusers >= 0.37.0.dev0`.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29519")
|
||||
os.environ.setdefault("DISABLE_SP", "1")
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
FAMILY = "glm_image"
|
||||
LOCAL_WEIGHTS_DIR = Path(
|
||||
os.getenv("GLM_IMAGE_LOCAL_WEIGHTS_DIR",
|
||||
REPO_ROOT / "official_weights" / FAMILY))
|
||||
TRANSFORMER_DIR = LOCAL_WEIGHTS_DIR / "transformer"
|
||||
|
||||
|
||||
def _has_shards() -> bool:
|
||||
return TRANSFORMER_DIR.exists() and any(
|
||||
TRANSFORMER_DIR.glob("*.safetensors"))
|
||||
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not _has_shards(),
|
||||
reason=f"GLM-Image transformer shards not at {TRANSFORMER_DIR}.",
|
||||
)
|
||||
|
||||
|
||||
def test_dit_strict_load_against_hf_checkpoint():
|
||||
safetensors = pytest.importorskip("safetensors.torch")
|
||||
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
|
||||
from fastvideo.models.dits.glm_image import GlmImageTransformer2DModel
|
||||
|
||||
cfg = GlmImageDiTConfig()
|
||||
mapping = cfg.arch_config.param_names_mapping
|
||||
|
||||
# 1. Build the FastVideo model on CPU/meta-free path.
|
||||
torch.set_default_device("cpu")
|
||||
hf_config = {"_class_name": "GlmImageTransformer2DModel"}
|
||||
model = GlmImageTransformer2DModel(cfg, hf_config)
|
||||
|
||||
# 2. Load every shard and apply param_names_mapping.
|
||||
raw_sd: dict[str, torch.Tensor] = {}
|
||||
for shard in sorted(TRANSFORMER_DIR.glob("*.safetensors")):
|
||||
raw_sd.update(safetensors.load_file(str(shard)))
|
||||
|
||||
def rename(k: str) -> str:
|
||||
for pat, repl in mapping.items():
|
||||
if re.match(pat, k):
|
||||
return re.sub(pat, repl, k)
|
||||
return k
|
||||
|
||||
renamed_sd = {rename(k): v for k, v in raw_sd.items()}
|
||||
|
||||
# 3. Compare key sets (sanity, redundant with strict=True but gives a
|
||||
# clearer error message if it fails).
|
||||
fv_keys = set(model.state_dict().keys())
|
||||
ckpt_keys = set(renamed_sd.keys())
|
||||
missing = fv_keys - ckpt_keys
|
||||
unexpected = ckpt_keys - fv_keys
|
||||
assert not missing, f"FastVideo DiT missing {len(missing)} keys, e.g. {sorted(missing)[:5]}"
|
||||
assert not unexpected, f"checkpoint has {len(unexpected)} unexpected keys, e.g. {sorted(unexpected)[:5]}"
|
||||
|
||||
# 4. Shape and dtype sanity — strict_load with shape mismatch raises early.
|
||||
fv_sd = model.state_dict()
|
||||
shape_mismatches = []
|
||||
for k, v in renamed_sd.items():
|
||||
if v.shape != fv_sd[k].shape:
|
||||
shape_mismatches.append((k, tuple(fv_sd[k].shape), tuple(v.shape)))
|
||||
assert not shape_mismatches, (
|
||||
f"shape mismatches: {shape_mismatches[:5]}")
|
||||
|
||||
# 5. Strict load.
|
||||
incompatible = model.load_state_dict(renamed_sd, strict=True)
|
||||
assert not incompatible.missing_keys
|
||||
assert not incompatible.unexpected_keys
|
||||
@@ -0,0 +1,106 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Component parity scaffold for GLM-Image VAE (AutoencoderKL).
|
||||
|
||||
Compares the FastVideo-native AutoencoderKL port against diffusers'
|
||||
`AutoencoderKL` loaded from `zai-org/GLM-Image/vae`. The published checkpoint
|
||||
uses `block_out_channels=[128, 512, 1024, 1024]`, `latent_channels=16`, and
|
||||
per-channel `latents_mean[16]` + `latents_std[16]` normalization (not a scalar
|
||||
`scaling_factor`).
|
||||
|
||||
Skips cleanly until both the FastVideo class and the VAE shard exist.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29519")
|
||||
os.environ.setdefault("DISABLE_SP", "1")
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
FAMILY = "glm_image"
|
||||
LOCAL_WEIGHTS_DIR = Path(
|
||||
os.getenv("GLM_IMAGE_LOCAL_WEIGHTS_DIR",
|
||||
REPO_ROOT / "official_weights" / FAMILY))
|
||||
VAE_DIR = LOCAL_WEIGHTS_DIR / "vae"
|
||||
|
||||
|
||||
def _has_weights() -> bool:
|
||||
return (VAE_DIR / "diffusion_pytorch_model.safetensors").exists()
|
||||
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not _has_weights(),
|
||||
reason=f"GLM-Image VAE weights not found at {VAE_DIR}.",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def device() -> torch.device:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA required for GLM-Image VAE parity.")
|
||||
return torch.device("cuda")
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def vae_config():
|
||||
with open(VAE_DIR / "config.json") as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def official_vae(device):
|
||||
diffusers = pytest.importorskip("diffusers")
|
||||
vae = diffusers.AutoencoderKL.from_pretrained(
|
||||
str(VAE_DIR), torch_dtype=torch.float32).to(device).eval()
|
||||
return vae
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def fastvideo_vae(device):
|
||||
pytest.importorskip("fastvideo")
|
||||
try:
|
||||
from fastvideo.configs.models.vaes.glm_image import GlmImageVAEConfig
|
||||
from fastvideo.models.vaes.autoencoder_kl import AutoencoderKL
|
||||
except ImportError as e:
|
||||
pytest.skip(f"FastVideo native AutoencoderKL not yet ported: {e}")
|
||||
cfg = GlmImageVAEConfig()
|
||||
vae = AutoencoderKL(cfg)
|
||||
safetensors = pytest.importorskip("safetensors.torch")
|
||||
sd = safetensors.load_file(
|
||||
str(VAE_DIR / "diffusion_pytorch_model.safetensors"))
|
||||
missing, unexpected = vae.load_state_dict(sd, strict=False)
|
||||
assert not missing, f"FastVideo VAE missing keys: {missing[:10]}"
|
||||
assert not unexpected, f"FastVideo VAE unexpected keys: {unexpected[:10]}"
|
||||
return vae.to(device, dtype=torch.float32).eval()
|
||||
|
||||
|
||||
def test_vae_config_matches_real_checkpoint(vae_config):
|
||||
assert vae_config["block_out_channels"] == [128, 512, 1024, 1024]
|
||||
assert vae_config["latent_channels"] == 16
|
||||
assert "latents_mean" in vae_config and len(
|
||||
vae_config["latents_mean"]) == 16
|
||||
assert "latents_std" in vae_config and len(vae_config["latents_std"]) == 16
|
||||
|
||||
|
||||
def test_vae_decode_matches_diffusers(official_vae, fastvideo_vae, device):
|
||||
torch.manual_seed(0)
|
||||
latents = torch.randn(1, 16, 32, 32, device=device, dtype=torch.float32)
|
||||
with torch.no_grad():
|
||||
official_out = official_vae.decode(latents, return_dict=False)[0]
|
||||
fv_out = fastvideo_vae.decode(latents, return_dict=False)[0]
|
||||
torch.testing.assert_close(fv_out, official_out, atol=1e-3, rtol=1e-3)
|
||||
|
||||
|
||||
def test_vae_encode_matches_diffusers(official_vae, fastvideo_vae, device):
|
||||
torch.manual_seed(0)
|
||||
pixels = torch.randn(1, 3, 256, 256, device=device, dtype=torch.float32)
|
||||
with torch.no_grad():
|
||||
official_z = official_vae.encode(pixels).latent_dist.mode()
|
||||
fv_z = fastvideo_vae.encode(pixels).latent_dist.mode()
|
||||
torch.testing.assert_close(fv_z, official_z, atol=1e-3, rtol=1e-3)
|
||||
Reference in New Issue
Block a user