Compare commits

...
Author SHA1 Message Date
H1yori233 ee623198c9 fix baed on review 2026-07-07 15:00:49 -07:00
k1kong 1e2f2bfbca fix fir transformer 5.10.2 2026-07-07 15:00:49 -07:00
H1yori233 5e90f421d5 fix 2026-07-07 15:00:49 -07:00
H1yori233 546ad1fb32 deterministic pipeline parity 2026-07-07 15:00:49 -07:00
H1yori233 9e261fef21 fix 2026-07-07 15:00:49 -07:00
H1yori233 ea4b2d3195 some fix 2026-07-07 15:00:49 -07:00
H1yori233 c7767910d2 precommit 2026-07-07 15:00:49 -07:00
H1yori233 cc61fc2714 lower SSIM resolution 2026-07-07 15:00:49 -07:00
H1yori233 a0749e98fd cleanup and local test 2026-07-07 15:00:49 -07:00
Shreejith SG 2d7e262b13 chore: remove unused reference image 2026-07-07 15:00:49 -07:00
Shreejith SG 65de01f1f1 Baseline: Deterministic priors, debug logs, SSIM 0.715 2026-07-07 15:00:49 -07:00
Shreejith SG e20deeb5eb fix(glm-image): resolve SSIM issues by fixing T5 glyph shapes and KV cache dimensions 2026-07-07 15:00:49 -07:00
Shreejith SG 51b9761334 Revert fix 2026-07-07 15:00:49 -07:00
Shreejith SG e99d2d9f16 Numerical alignment fix 2026-07-07 15:00:49 -07:00
Shreejith SG 8a9ec8621e enable sampling for quality 2026-07-07 15:00:49 -07:00
Shreejith SG 53440d4562 CI fix - Upgrade transformers and accelerate versions 2026-07-07 15:00:49 -07:00
Shreejith SG 3d82e5e24c Fix pre-commit errors and SSIM test ImportError 2026-07-07 15:00:49 -07:00
Shreejith SG cc2c419c4b Fix mypy error: add __init__ to GlmImageDecodingStage 2026-07-07 15:00:49 -07:00
Shreejith SG 81fb8f2c10 Add GLM-Image SSIM test and fix deterministic latent generation 2026-07-07 15:00:49 -07:00
Shreejith SG 0d80a668d2 Add L40S reference image for GLM-Image 2026-07-07 15:00:49 -07:00
Shreejith SG 9dec070d4c Fix GLM-Image seed handling and add SSIM test 2026-07-07 15:00:49 -07:00
ShreejithSG 326287dffa WIP: add GLM-Image SSIM test and diffusers reference generator 2026-07-07 15:00:49 -07:00
ShreejithSG 5c68193009 Refactor GLM-Image decoding stage to inherit from our DecodingStage 2026-07-07 15:00:49 -07:00
ShreejithSG 19ef9fa56f Refactor attn mask handling to use forward context 2026-07-07 15:00:49 -07:00
ShreejithSG 7833575da8 Fix PR comments: simplify example script and properly handle attention_mask through forward context
- Clean up basic_glm_image.py to match other example scripts (remove unnecessary comments, follow LTX2 pattern)
- Fix attention_mask handling in LocalAttention to use set_forward_context properly instead of creating SDPAMetadata inline
- SDPAMetadata import remains at top of file (no local import)
2026-07-07 15:00:49 -07:00
ShreejithSG 9ea387e872 Address more PR review comments 2026-07-07 15:00:49 -07:00
ShreejithSG 615da36296 Fix precommit error 2026-07-07 15:00:49 -07:00
ShreejithSG c54e0737d0 Address PR comments: GLM-Image Port to FastVideo- Replaced diffusers FeedForward with FastVideo MLP in DiT.- Moved SDPAMetadata import to top of attention layer for cleaner code.- Removed unnecessary __call__ override from pipeline to use standard composition.- Consolidated CFG into a single batched pass with concatenated embeddings for performance.- Fixed attention mask compatibility issues with optimized FLASH_ATTN backend.- Robustly patched flash_attn unpad utility for variable version return values. 2026-07-07 15:00:48 -07:00
30 changed files with 3366 additions and 3 deletions
+107
View File
@@ -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()
+120
View File
@@ -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()
+3 -1
View File
@@ -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
+46
View File
@@ -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
+776
View File
@@ -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."""
+6
View File
@@ -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)
+13
View File
@@ -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
View File
@@ -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'".
+3 -1
View File
@@ -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",
+74
View File
@@ -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)