diff --git a/docs/source/en/_toctree.yml b/docs/source/en/_toctree.yml index 525734a2e4cb..ee2c2b53e292 100644 --- a/docs/source/en/_toctree.yml +++ b/docs/source/en/_toctree.yml @@ -379,6 +379,8 @@ title: Krea2Transformer2DModel - local: api/models/latte_transformer3d title: LatteTransformer3DModel + - local: api/models/llada_image_transformer2d + title: LLaDA-Image - local: api/models/longcat_image_transformer2d title: LongCatImageTransformer2DModel - local: api/models/ltx2_video_transformer3d @@ -617,6 +619,8 @@ title: Latent Diffusion - local: api/pipelines/ledits_pp title: LEDITS++ + - local: api/pipelines/llada_image + title: LLaDA-Image - local: api/pipelines/longcat_image title: LongCat-Image - local: api/pipelines/lumina2 diff --git a/docs/source/en/api/models/llada_image_transformer2d.md b/docs/source/en/api/models/llada_image_transformer2d.md new file mode 100644 index 000000000000..ecf445967764 --- /dev/null +++ b/docs/source/en/api/models/llada_image_transformer2d.md @@ -0,0 +1,23 @@ +# LLaDA-Image + +The LLaDA-Image model family combines a denoising transformer with a QueryFormer, text projection model, and SigVQ +image tokenizer. Together, these components support text-to-image generation, VQ-conditioned generation, and +instruction-guided image editing. + +The original code and checkpoints are available in the [LLaDA-Image repository](https://github.com/inclusionAI/LLaDA-Image). + +## LLaDAImageTransformer2DModel + +[[autodoc]] LLaDAImageTransformer2DModel + +## LLaDAImageQueryFormerModel + +[[autodoc]] LLaDAImageQueryFormerModel + +## LLaDAImageTextProjectionModel + +[[autodoc]] LLaDAImageTextProjectionModel + +## LLaDAImageSigVQModel + +[[autodoc]] LLaDAImageSigVQModel diff --git a/docs/source/en/api/pipelines/llada_image.md b/docs/source/en/api/pipelines/llada_image.md new file mode 100644 index 000000000000..65f44aab7d62 --- /dev/null +++ b/docs/source/en/api/pipelines/llada_image.md @@ -0,0 +1,57 @@ +# LLaDA-Image + +[LLaDA-Image](https://huggingface.co/inclusionAI/LLaDA-Image) is a unified image generation and editing model. The +same pipeline supports text-to-image generation, VQ-conditioned generation, and instruction-guided editing with a +reference image. The Base checkpoint is designed for 50 sampling steps, while +[LLaDA-Image-Turbo](https://huggingface.co/inclusionAI/LLaDA-Image-Turbo) is distilled for 4 steps. + +The checkpoint includes a custom LLaDA2 text encoder, so pass `trust_remote_code=True` when loading it. + +```python +import torch + +from diffusers import LLaDAImagePipeline + +pipe = LLaDAImagePipeline.from_pretrained( + "inclusionAI/LLaDA-Image", + dtype=torch.bfloat16, + trust_remote_code=True, +) +pipe.enable_model_cpu_offload() + +image = pipe( + prompt="A cinematic photograph of a red fox standing in fresh snow", + height=1024, + width=1024, + num_inference_steps=50, + guidance_scale=5.0, + generator=torch.Generator("cuda").manual_seed(42), +).images[0] +``` + +For image editing, pass a reference image and select the editing mode. + +```python +from diffusers.utils import load_image + +reference_image = load_image("https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/cat.png") +image = pipe( + prompt="Turn it into a watercolor painting", + image=reference_image, + generation_mode="editing", + height=1024, + width=1024, + num_inference_steps=50, + guidance_scale=5.0, +).images[0] +``` + +## LLaDAImagePipeline + +[[autodoc]] LLaDAImagePipeline + - all + - call + +## LLaDAImagePipelineOutput + +[[autodoc]] pipelines.LLaDAImagePipelineOutput diff --git a/src/diffusers/__init__.py b/src/diffusers/__init__.py index 2825e9888c98..b7a27e4777b2 100644 --- a/src/diffusers/__init__.py +++ b/src/diffusers/__init__.py @@ -309,6 +309,10 @@ "Kandinsky5Transformer3DModel", "Krea2Transformer2DModel", "LatteTransformer3DModel", + "LLaDAImageQueryFormerModel", + "LLaDAImageSigVQModel", + "LLaDAImageTextProjectionModel", + "LLaDAImageTransformer2DModel", "LongCatAudioDiTTransformer", "LongCatAudioDiTVae", "LongCatImageTransformer2DModel", @@ -727,6 +731,8 @@ "LEditsPPPipelineStableDiffusionXL", "LLaDA2Pipeline", "LLaDA2PipelineOutput", + "LLaDAImagePipeline", + "LLaDAImagePipelineOutput", "LongCatAudioDiTPipeline", "LongCatImageEditPipeline", "LongCatImagePipeline", @@ -1191,6 +1197,10 @@ Kandinsky5Transformer3DModel, Krea2Transformer2DModel, LatteTransformer3DModel, + LLaDAImageQueryFormerModel, + LLaDAImageSigVQModel, + LLaDAImageTextProjectionModel, + LLaDAImageTransformer2DModel, LongCatAudioDiTTransformer, LongCatAudioDiTVae, LongCatImageTransformer2DModel, @@ -1584,6 +1594,8 @@ LEditsPPPipelineStableDiffusionXL, LLaDA2Pipeline, LLaDA2PipelineOutput, + LLaDAImagePipeline, + LLaDAImagePipelineOutput, LongCatAudioDiTPipeline, LongCatImageEditPipeline, LongCatImagePipeline, diff --git a/src/diffusers/models/__init__.py b/src/diffusers/models/__init__.py index 1a396b312441..189c27c9171f 100755 --- a/src/diffusers/models/__init__.py +++ b/src/diffusers/models/__init__.py @@ -131,6 +131,12 @@ _import_structure["transformers.transformer_joyimage_edit_plus"] = ["JoyImageEditPlusTransformer3DModel"] _import_structure["transformers.transformer_kandinsky"] = ["Kandinsky5Transformer3DModel"] _import_structure["transformers.transformer_krea2"] = ["Krea2Transformer2DModel"] + _import_structure["transformers.transformer_llada_image"] = [ + "LLaDAImageQueryFormerModel", + "LLaDAImageSigVQModel", + "LLaDAImageTextProjectionModel", + "LLaDAImageTransformer2DModel", + ] _import_structure["transformers.transformer_longcat_audio_dit"] = ["LongCatAudioDiTTransformer"] _import_structure["transformers.transformer_longcat_image"] = ["LongCatImageTransformer2DModel"] _import_structure["transformers.transformer_ltx"] = ["LTXVideoTransformer3DModel"] @@ -273,6 +279,10 @@ Kandinsky5Transformer3DModel, Krea2Transformer2DModel, LatteTransformer3DModel, + LLaDAImageQueryFormerModel, + LLaDAImageSigVQModel, + LLaDAImageTextProjectionModel, + LLaDAImageTransformer2DModel, LongCatAudioDiTTransformer, LongCatImageTransformer2DModel, LTX2VideoTransformer3DModel, diff --git a/src/diffusers/models/transformers/__init__.py b/src/diffusers/models/transformers/__init__.py index ffb0cbc0318b..11ac52798025 100755 --- a/src/diffusers/models/transformers/__init__.py +++ b/src/diffusers/models/transformers/__init__.py @@ -46,6 +46,12 @@ from .transformer_joyimage_edit_plus import JoyImageEditPlusTransformer3DModel from .transformer_kandinsky import Kandinsky5Transformer3DModel from .transformer_krea2 import Krea2Transformer2DModel + from .transformer_llada_image import ( + LLaDAImageQueryFormerModel, + LLaDAImageSigVQModel, + LLaDAImageTextProjectionModel, + LLaDAImageTransformer2DModel, + ) from .transformer_longcat_audio_dit import LongCatAudioDiTTransformer from .transformer_longcat_image import LongCatImageTransformer2DModel from .transformer_ltx import LTXVideoTransformer3DModel diff --git a/src/diffusers/models/transformers/transformer_llada_image.py b/src/diffusers/models/transformers/transformer_llada_image.py new file mode 100644 index 000000000000..d7db5b004751 --- /dev/null +++ b/src/diffusers/models/transformers/transformer_llada_image.py @@ -0,0 +1,1878 @@ +# Copyright 2026 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import math +from dataclasses import dataclass + +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch.nn.utils.rnn import pad_sequence + +from ...configuration_utils import ConfigMixin, register_to_config +from ...utils import BaseOutput +from ...utils.torch_utils import maybe_allow_in_graph +from ..attention import AttentionMixin, AttentionModuleMixin, FeedForward +from ..attention_dispatch import dispatch_attention_fn +from ..modeling_outputs import Transformer2DModelOutput +from ..modeling_utils import ModelMixin +from ..normalization import RMSNorm + + +ADALN_EMBED_DIM = 256 +SEQUENCE_MULTIPLE = 32 + + +@dataclass +class _LLaDAImageSequence: + features: list[torch.Tensor] + position_ids: list[torch.Tensor] + padding_masks: list[torch.Tensor] + noise_masks: list[list[int]] | None = None + + +class LLaDAImageTimestepEmbedder(nn.Module): + def __init__(self, output_dim: int, hidden_dim: int = 1024, frequency_embedding_dim: int = 256): + super().__init__() + self.mlp = nn.Sequential( + nn.Linear(frequency_embedding_dim, hidden_dim, bias=True), + nn.SiLU(), + nn.Linear(hidden_dim, output_dim, bias=True), + ) + self.frequency_embedding_dim = frequency_embedding_dim + + def forward(self, timestep: torch.Tensor, hidden_dtype: torch.dtype) -> torch.Tensor: + half_dim = self.frequency_embedding_dim // 2 + frequencies = torch.exp( + -math.log(10000) * torch.arange(half_dim, dtype=torch.float32, device=timestep.device) / half_dim + ) + arguments = timestep[:, None].float() * frequencies[None] + embedding = torch.cat([torch.cos(arguments), torch.sin(arguments)], dim=-1) + if self.frequency_embedding_dim % 2: + embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) + return self.mlp(embedding.to(dtype=hidden_dtype)) + + +class LLaDAImageRopeEmbedder(nn.Module): + def __init__(self, theta: float, axes_dims: tuple[int, ...], axes_lens: tuple[int, ...]): + super().__init__() + self.theta = theta + self.axes_dims = axes_dims + self.axes_lens = axes_lens + self.freqs_cis = None + + def _create_frequencies(self, device: torch.device) -> list[torch.Tensor]: + frequencies = [] + for axis_dim, axis_len in zip(self.axes_dims, self.axes_lens): + inverse_frequencies = 1.0 / ( + self.theta ** (torch.arange(0, axis_dim, 2, dtype=torch.float32, device=device) / axis_dim) + ) + positions = torch.arange(axis_len, dtype=torch.float32, device=device) + angles = torch.outer(positions, inverse_frequencies) + frequencies.append(torch.complex(torch.cos(angles), torch.sin(angles))) + return frequencies + + def forward(self, position_ids: torch.Tensor) -> torch.Tensor: + if torch.compiler.is_compiling(): + freqs_cis = self._create_frequencies(position_ids.device) + else: + if self.freqs_cis is None or self.freqs_cis[0].device != position_ids.device: + self.freqs_cis = self._create_frequencies(position_ids.device) + freqs_cis = self.freqs_cis + + frequencies = [] + for axis, axis_frequencies in enumerate(freqs_cis): + frequencies.append(axis_frequencies[position_ids[:, axis]]) + return torch.cat(frequencies, dim=-1) + + +class LLaDAImageAttnProcessor: + _attention_backend = None + _parallel_config = None + + def __call__( + self, + attn: "LLaDAImageAttention", + hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None = None, + freqs_cis: torch.Tensor | None = None, + ) -> torch.Tensor: + query = attn.to_q(hidden_states).unflatten(-1, (attn.heads, attn.head_dim)) + key = attn.to_k(hidden_states).unflatten(-1, (attn.heads, attn.head_dim)) + value = attn.to_v(hidden_states).unflatten(-1, (attn.heads, attn.head_dim)) + + if attn.norm_q is not None: + query = attn.norm_q(query) + key = attn.norm_k(key) + + if freqs_cis is not None: + with torch.autocast(device_type=hidden_states.device.type, enabled=False): + query_complex = torch.view_as_complex(query.float().reshape(*query.shape[:-1], -1, 2)) + key_complex = torch.view_as_complex(key.float().reshape(*key.shape[:-1], -1, 2)) + frequencies = freqs_cis.unsqueeze(2) + query = torch.view_as_real(query_complex * frequencies).flatten(3).to(dtype=query.dtype) + key = torch.view_as_real(key_complex * frequencies).flatten(3).to(dtype=key.dtype) + + if attention_mask is not None and attention_mask.ndim == 2: + attention_mask = attention_mask[:, None, None, :] + + hidden_states = dispatch_attention_fn( + query, + key, + value, + attn_mask=attention_mask, + dropout_p=0.0, + is_causal=False, + backend=self._attention_backend, + parallel_config=self._parallel_config, + ) + hidden_states = hidden_states.flatten(2, 3) + return attn.to_out[0](hidden_states) + + +class LLaDAImageAttention(nn.Module, AttentionModuleMixin): + _default_processor_cls = LLaDAImageAttnProcessor + _available_processors = [LLaDAImageAttnProcessor] + _supports_qkv_fusion = False + + def __init__(self, dim: int, num_heads: int, norm_eps: float, qk_norm: bool): + super().__init__() + self.heads = num_heads + self.head_dim = dim // num_heads + self.to_q = nn.Linear(dim, dim, bias=False) + self.to_k = nn.Linear(dim, dim, bias=False) + self.to_v = nn.Linear(dim, dim, bias=False) + self.norm_q = RMSNorm(self.head_dim, eps=norm_eps, elementwise_affine=False) if qk_norm else None + self.norm_k = RMSNorm(self.head_dim, eps=norm_eps, elementwise_affine=False) if qk_norm else None + self.to_out = nn.ModuleList([nn.Linear(dim, dim, bias=False), nn.Dropout(0.0)]) + self.set_processor(self._default_processor_cls()) + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None, + freqs_cis: torch.Tensor, + ) -> torch.Tensor: + return self.processor(self, hidden_states, attention_mask, freqs_cis) + + +class LLaDAImageFeedForward(nn.Module): + def __init__(self, dim: int): + super().__init__() + hidden_dim = int(dim / 3 * 8) + self.w1 = nn.Linear(dim, hidden_dim, bias=False) + self.w2 = nn.Linear(hidden_dim, dim, bias=False) + self.w3 = nn.Linear(dim, hidden_dim, bias=False) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + return self.w2(F.silu(self.w1(hidden_states)) * self.w3(hidden_states)) + + +def _select_per_token( + noisy_value: torch.Tensor, + clean_value: torch.Tensor, + noise_mask: torch.Tensor, + sequence_length: int, +) -> torch.Tensor: + noise_mask = noise_mask.unsqueeze(-1) + return torch.where( + noise_mask == 1, + noisy_value.unsqueeze(1).expand(-1, sequence_length, -1), + clean_value.unsqueeze(1).expand(-1, sequence_length, -1), + ) + + +@maybe_allow_in_graph +class LLaDAImageTransformerBlock(nn.Module): + def __init__(self, dim: int, num_heads: int, norm_eps: float, qk_norm: bool, modulation: bool): + super().__init__() + self.modulation = modulation + self.attention = LLaDAImageAttention(dim, num_heads, norm_eps, qk_norm) + self.feed_forward = LLaDAImageFeedForward(dim) + self.attention_norm1 = RMSNorm(dim, eps=norm_eps, elementwise_affine=False) + self.ffn_norm1 = RMSNorm(dim, eps=norm_eps, elementwise_affine=False) + self.attention_norm2 = RMSNorm(dim, eps=norm_eps, elementwise_affine=False) + self.ffn_norm2 = RMSNorm(dim, eps=norm_eps, elementwise_affine=False) + if modulation: + self.adaLN_modulation = nn.Sequential(nn.Linear(min(dim, ADALN_EMBED_DIM), 4 * dim, bias=True)) + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None, + freqs_cis: torch.Tensor, + adaln_input: torch.Tensor | None = None, + noise_mask: torch.Tensor | None = None, + adaln_noisy: torch.Tensor | None = None, + adaln_clean: torch.Tensor | None = None, + ) -> torch.Tensor: + if self.modulation: + sequence_length = hidden_states.shape[1] + if noise_mask is None: + scale_msa, gate_msa, scale_mlp, gate_mlp = ( + self.adaLN_modulation(adaln_input).unsqueeze(1).chunk(4, dim=2) + ) + gate_msa = gate_msa.tanh() + gate_mlp = gate_mlp.tanh() + scale_msa = 1.0 + scale_msa + scale_mlp = 1.0 + scale_mlp + else: + noisy_modulation = self.adaLN_modulation(adaln_noisy) + clean_modulation = self.adaLN_modulation(adaln_clean) + noisy_scale_msa, noisy_gate_msa, noisy_scale_mlp, noisy_gate_mlp = noisy_modulation.chunk(4, dim=1) + clean_scale_msa, clean_gate_msa, clean_scale_mlp, clean_gate_mlp = clean_modulation.chunk(4, dim=1) + scale_msa = _select_per_token( + 1.0 + noisy_scale_msa, 1.0 + clean_scale_msa, noise_mask, sequence_length + ) + scale_mlp = _select_per_token( + 1.0 + noisy_scale_mlp, 1.0 + clean_scale_mlp, noise_mask, sequence_length + ) + gate_msa = _select_per_token(noisy_gate_msa.tanh(), clean_gate_msa.tanh(), noise_mask, sequence_length) + gate_mlp = _select_per_token(noisy_gate_mlp.tanh(), clean_gate_mlp.tanh(), noise_mask, sequence_length) + + attention_output = self.attention( + self.attention_norm1(hidden_states) * scale_msa, + attention_mask, + freqs_cis, + ) + hidden_states = hidden_states + gate_msa * self.attention_norm2(attention_output) + hidden_states = hidden_states + gate_mlp * self.ffn_norm2( + self.feed_forward(self.ffn_norm1(hidden_states) * scale_mlp) + ) + else: + attention_output = self.attention( + self.attention_norm1(hidden_states), + attention_mask, + freqs_cis, + ) + hidden_states = hidden_states + self.attention_norm2(attention_output) + hidden_states = hidden_states + self.ffn_norm2(self.feed_forward(self.ffn_norm1(hidden_states))) + return hidden_states + + +class LLaDAImageFinalLayer(nn.Module): + def __init__(self, dim: int, out_channels: int): + super().__init__() + self.norm_final = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) + self.linear = nn.Linear(dim, out_channels, bias=True) + self.adaLN_modulation = nn.Sequential( + nn.SiLU(), + nn.Linear(min(dim, ADALN_EMBED_DIM), dim, bias=True), + ) + + def forward( + self, + hidden_states: torch.Tensor, + adaln_input: torch.Tensor | None = None, + noise_mask: torch.Tensor | None = None, + adaln_noisy: torch.Tensor | None = None, + adaln_clean: torch.Tensor | None = None, + ) -> torch.Tensor: + if noise_mask is None: + scale = 1.0 + self.adaLN_modulation(adaln_input) + scale = scale.unsqueeze(1) + else: + sequence_length = hidden_states.shape[1] + noisy_scale = 1.0 + self.adaLN_modulation(adaln_noisy) + clean_scale = 1.0 + self.adaLN_modulation(adaln_clean) + scale = _select_per_token(noisy_scale, clean_scale, noise_mask, sequence_length) + hidden_states = self.norm_final(hidden_states) * scale + return self.linear(hidden_states) + + +class LLaDAImageTransformer2DModel(ModelMixin, ConfigMixin, AttentionMixin): + r""" + The denoising transformer used by LLaDAImage for text-to-image generation and single-image editing. + + This component consumes caption features that have already passed through LLaDAImage's QueryFormer, connector, and + projector. For editing, it additionally consumes GLM/SigVQ features and source-image latents. + + Args: + all_patch_size (`tuple[int, ...]`, defaults to `(1,)`): + Supported spatial patch sizes. + all_f_patch_size (`tuple[int, ...]`, defaults to `(1,)`): + Supported temporal patch sizes paired with `all_patch_size`. + in_channels (`int`, defaults to `128`): + Number of channels in the patchified Flux2 VAE latents. + dim (`int`, defaults to `3840`): + Transformer hidden dimension. + n_layers (`int`, defaults to `30`): + Number of main transformer blocks. + n_refiner_layers (`int`, defaults to `2`): + Number of noise, caption, and SigVQ refiner blocks. + n_heads (`int`, defaults to `30`): + Number of attention heads. + norm_eps (`float`, defaults to `1e-5`): + Epsilon used by RMS normalization layers. + qk_norm (`bool`, defaults to `True`): + Whether to apply RMS normalization to query and key tensors. + cap_feat_dim (`int`, defaults to `2560`): + Dimension of projected QueryFormer caption features. + semantic_feat_dim (`int`, defaults to `4096`): + Dimension of GLM/SigVQ semantic features. + rope_theta (`float`, defaults to `256.0`): + RoPE frequency base. + t_scale (`float`, defaults to `1000.0`): + Scale applied to diffusion timesteps. + axes_dims (`tuple[int, ...]`, defaults to `(32, 48, 48)`): + RoPE dimensions for sequence, height, and width axes. + axes_lens (`tuple[int, ...]`, defaults to `(32768, 1024, 1024)`): + Maximum RoPE positions for sequence, height, and width axes. + """ + + _supports_gradient_checkpointing = True + _no_split_modules = ["LLaDAImageTransformerBlock"] + _repeated_blocks = ["LLaDAImageTransformerBlock"] + _skip_layerwise_casting_patterns = [ + "t_embedder", + "cap_embedder", + "semantic_embedder", + "sigvq_embedder", + ] + + @register_to_config + def __init__( + self, + all_patch_size: tuple[int, ...] = (1,), + all_f_patch_size: tuple[int, ...] = (1,), + in_channels: int = 128, + dim: int = 3840, + n_layers: int = 30, + n_refiner_layers: int = 2, + n_heads: int = 30, + norm_eps: float = 1e-5, + qk_norm: bool = True, + cap_feat_dim: int = 2560, + semantic_feat_dim: int = 4096, + rope_theta: float = 256.0, + t_scale: float = 1000.0, + axes_dims: tuple[int, ...] = (32, 48, 48), + axes_lens: tuple[int, ...] = (32768, 1024, 1024), + ): + super().__init__() + if len(all_patch_size) != len(all_f_patch_size): + raise ValueError("`all_patch_size` and `all_f_patch_size` must have the same length.") + if dim % n_heads != 0: + raise ValueError(f"`dim` ({dim}) must be divisible by `n_heads` ({n_heads}).") + if dim // n_heads != sum(axes_dims): + raise ValueError("The attention head dimension must equal the sum of `axes_dims`.") + + self.in_channels = in_channels + self.out_channels = in_channels + self.all_patch_size = all_patch_size + self.all_f_patch_size = all_f_patch_size + self.t_scale = t_scale + self.gradient_checkpointing = False + + self.all_x_embedder = nn.ModuleDict() + self.all_final_layer = nn.ModuleDict() + for patch_size, f_patch_size in zip(all_patch_size, all_f_patch_size): + patch_key = f"{patch_size}-{f_patch_size}" + patch_dim = f_patch_size * patch_size * patch_size * in_channels + self.all_x_embedder[patch_key] = nn.Linear(patch_dim, dim, bias=True) + self.all_final_layer[patch_key] = LLaDAImageFinalLayer(dim, patch_dim) + + self.noise_refiner = nn.ModuleList( + [ + LLaDAImageTransformerBlock(dim, n_heads, norm_eps, qk_norm, modulation=True) + for _ in range(n_refiner_layers) + ] + ) + self.context_refiner = nn.ModuleList( + [ + LLaDAImageTransformerBlock(dim, n_heads, norm_eps, qk_norm, modulation=False) + for _ in range(n_refiner_layers) + ] + ) + self.sigvq_refiner = nn.ModuleList( + [ + LLaDAImageTransformerBlock(dim, n_heads, norm_eps, qk_norm, modulation=False) + for _ in range(n_refiner_layers) + ] + ) + self.layers = nn.ModuleList( + [LLaDAImageTransformerBlock(dim, n_heads, norm_eps, qk_norm, modulation=True) for _ in range(n_layers)] + ) + + self.t_embedder = LLaDAImageTimestepEmbedder(min(dim, ADALN_EMBED_DIM)) + self.cap_embedder = nn.Sequential( + RMSNorm(cap_feat_dim, eps=norm_eps, elementwise_affine=False), + nn.Linear(cap_feat_dim, dim, bias=True), + ) + self.semantic_embedder = nn.Sequential( + RMSNorm(semantic_feat_dim, eps=norm_eps, elementwise_affine=False), + nn.Linear(semantic_feat_dim, dim, bias=True), + ) + self.sigvq_embedder = nn.Sequential( + RMSNorm(semantic_feat_dim, eps=norm_eps, elementwise_affine=False), + nn.Linear(semantic_feat_dim, dim, bias=True), + ) + + nn.init.normal_(self.semantic_embedder[1].weight, mean=0.0, std=0.02) + nn.init.zeros_(self.semantic_embedder[1].bias) + nn.init.normal_(self.sigvq_embedder[1].weight, mean=0.0, std=0.02) + nn.init.zeros_(self.sigvq_embedder[1].bias) + + self.x_pad_token = nn.Parameter(torch.zeros(1, dim)) + self.cap_pad_token = nn.Parameter(torch.zeros(1, dim)) + self.sigvq_pad_token = nn.Parameter(torch.zeros(1, dim)) + nn.init.normal_(self.sigvq_pad_token, mean=0.0, std=0.02) + + self.rope_embedder = LLaDAImageRopeEmbedder(rope_theta, axes_dims, axes_lens) + + @staticmethod + def _create_coordinate_grid( + size: tuple[int, int, int], + start: tuple[int, int, int], + device: torch.device, + ) -> torch.Tensor: + axes = [ + torch.arange(start_value, start_value + span, dtype=torch.int32, device=device) + for start_value, span in zip(start, size) + ] + return torch.stack(torch.meshgrid(axes, indexing="ij"), dim=-1) + + def _patchify_image( + self, + image: torch.Tensor, + patch_size: int, + f_patch_size: int, + ) -> tuple[torch.Tensor, tuple[int, int, int], tuple[int, int, int]]: + channels, frames, height, width = image.shape + frame_tokens = frames // f_patch_size + height_tokens = height // patch_size + width_tokens = width // patch_size + image = image.view( + channels, + frame_tokens, + f_patch_size, + height_tokens, + patch_size, + width_tokens, + patch_size, + ) + image = image.permute(1, 3, 5, 2, 4, 6, 0).reshape( + frame_tokens * height_tokens * width_tokens, + f_patch_size * patch_size * patch_size * channels, + ) + return image, (frames, height, width), (frame_tokens, height_tokens, width_tokens) + + def _pad_with_ids( + self, + features: torch.Tensor, + position_grid_size: tuple[int, int, int], + position_start: tuple[int, int, int], + noise_value: int | None = None, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, list[int] | None]: + original_length = len(features) + padding_length = (-original_length) % SEQUENCE_MULTIPLE + padded_length = original_length + padding_length + device = features.device + + position_ids = self._create_coordinate_grid( + position_grid_size, + position_start, + device, + ).flatten(0, 2) + if padding_length > 0: + padding_position_ids = ( + self._create_coordinate_grid( + (1, 1, 1), + (0, 0, 0), + device, + ) + .flatten(0, 2) + .repeat(padding_length, 1) + ) + position_ids = torch.cat([position_ids, padding_position_ids], dim=0) + features = torch.cat([features, features[-1:].repeat(padding_length, 1)], dim=0) + padding_mask = torch.cat( + [ + torch.zeros(original_length, dtype=torch.bool, device=device), + torch.ones(padding_length, dtype=torch.bool, device=device), + ] + ) + else: + padding_mask = torch.zeros(original_length, dtype=torch.bool, device=device) + + noise_mask = [noise_value] * padded_length if noise_value is not None else None + return features, position_ids, padding_mask, padded_length, noise_mask + + @staticmethod + def _batch_sequences( + features: list[torch.Tensor], + frequencies: list[torch.Tensor], + inner_padding_masks: list[torch.Tensor], + pad_token: torch.Tensor, + noise_masks: list[list[int]] | None = None, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None, list[int], torch.Tensor | None]: + sequence_lengths = [len(item) for item in features] + max_sequence_length = max(sequence_lengths) + features = torch.cat(features, dim=0) + inner_padding_mask = torch.cat(inner_padding_masks).unsqueeze(-1) + features = torch.where( + inner_padding_mask.to(device=features.device), + pad_token.to(device=features.device, dtype=features.dtype), + features, + ) + features = list(features.split(sequence_lengths, dim=0)) + + features = pad_sequence(features, batch_first=True, padding_value=0.0) + frequencies = pad_sequence(frequencies, batch_first=True, padding_value=0.0)[:, : features.shape[1]] + + attention_mask = None + if not all(length == max_sequence_length for length in sequence_lengths): + attention_mask = torch.zeros( + (len(sequence_lengths), max_sequence_length), + dtype=torch.bool, + device=features.device, + ) + for batch_index, sequence_length in enumerate(sequence_lengths): + attention_mask[batch_index, :sequence_length] = True + + noise_mask = None + if noise_masks is not None: + noise_mask = pad_sequence( + [torch.tensor(mask, dtype=torch.long, device=features.device) for mask in noise_masks], + batch_first=True, + padding_value=0, + )[:, : features.shape[1]] + + return features, frequencies, attention_mask, sequence_lengths, noise_mask + + def _unpatchify( + self, + hidden_states: list[torch.Tensor], + sizes: list[tuple[int, int, int] | list[tuple[int, int, int]]], + patch_size: int, + f_patch_size: int, + image_offsets: list[tuple[int, int]] | None = None, + ) -> list[torch.Tensor]: + outputs = [] + for batch_index, batch_hidden_states in enumerate(hidden_states): + if image_offsets is None: + batch_sizes = [sizes[batch_index]] + image_hidden_states = batch_hidden_states + else: + batch_sizes = sizes[batch_index] + start, end = image_offsets[batch_index] + image_hidden_states = batch_hidden_states[start:end] + + current_offset = 0 + output = None + for frames, height, width in batch_sizes: + original_length = (frames // f_patch_size) * (height // patch_size) * (width // patch_size) + padding_length = (-original_length) % SEQUENCE_MULTIPLE + output = ( + image_hidden_states[current_offset : current_offset + original_length] + .view( + frames // f_patch_size, + height // patch_size, + width // patch_size, + f_patch_size, + patch_size, + patch_size, + self.out_channels, + ) + .permute(6, 0, 3, 1, 4, 2, 5) + .reshape(self.out_channels, frames, height, width) + ) + current_offset += original_length + padding_length + outputs.append(output) + return outputs + + def _prepare_t2i_sequences( + self, + x: list[torch.Tensor], + cap_feats: list[torch.Tensor] | None, + glm_features: list[torch.Tensor] | None, + patch_size: int, + f_patch_size: int, + ) -> tuple[ + _LLaDAImageSequence, + _LLaDAImageSequence | None, + _LLaDAImageSequence | None, + list[tuple[int, int, int]], + ]: + image_sequence = _LLaDAImageSequence([], [], []) + cap_sequence = _LLaDAImageSequence([], [], []) if cap_feats is not None else None + glm_sequence = _LLaDAImageSequence([], [], []) if glm_features is not None else None + image_sizes = [] + + for batch_index, latent in enumerate(x): + position_cursor = 1 + if cap_sequence is not None: + padded_features, position_ids, padding_mask, sequence_length, _ = self._pad_with_ids( + cap_feats[batch_index], + (len(cap_feats[batch_index]), 1, 1), + (position_cursor, 0, 0), + ) + cap_sequence.features.append(padded_features) + cap_sequence.position_ids.append(position_ids) + cap_sequence.padding_masks.append(padding_mask) + position_cursor += sequence_length + + if glm_sequence is not None: + padded_features, position_ids, padding_mask, sequence_length, _ = self._pad_with_ids( + glm_features[batch_index], + (len(glm_features[batch_index]), 1, 1), + (position_cursor, 0, 0), + ) + glm_sequence.features.append(padded_features) + glm_sequence.position_ids.append(position_ids) + glm_sequence.padding_masks.append(padding_mask) + position_cursor += sequence_length + + patches, image_size, token_grid_size = self._patchify_image(latent, patch_size, f_patch_size) + padded_features, position_ids, padding_mask, _, _ = self._pad_with_ids( + patches, + token_grid_size, + (position_cursor, 0, 0), + ) + image_sequence.features.append(padded_features) + image_sequence.position_ids.append(position_ids) + image_sequence.padding_masks.append(padding_mask) + image_sizes.append(image_size) + + return image_sequence, cap_sequence, glm_sequence, image_sizes + + def _prepare_editing_sequences( + self, + x: list[torch.Tensor], + cap_feats: list[torch.Tensor], + glm_cap_feats: list[torch.Tensor], + source_latents: list[torch.Tensor], + patch_size: int, + f_patch_size: int, + ) -> tuple[ + _LLaDAImageSequence, + _LLaDAImageSequence, + _LLaDAImageSequence, + list[list[tuple[int, int, int]]], + list[tuple[int, int]], + ]: + image_sequence = _LLaDAImageSequence([], [], [], []) + cap_sequence = _LLaDAImageSequence([], [], [], []) + sigvq_sequence = _LLaDAImageSequence([], [], [], []) + image_sizes = [] + image_offsets = [] + + for batch_index, latent in enumerate(x): + cap_end_positions = [] + position_cursor = 1 + batch_cap_features = [] + batch_cap_positions = [] + batch_cap_padding = [] + batch_cap_noise = [] + for noise_value in (0, 1): + padded_features, position_ids, padding_mask, _, noise_mask = self._pad_with_ids( + cap_feats[batch_index], + (len(cap_feats[batch_index]), 1, 1), + (position_cursor, 0, 0), + noise_value, + ) + batch_cap_features.append(padded_features) + batch_cap_positions.append(position_ids) + batch_cap_padding.append(padding_mask) + batch_cap_noise.extend(noise_mask) + position_cursor += len(cap_feats[batch_index]) + cap_end_positions.append(position_cursor) + position_cursor += 2 + + batch_image_features = [] + batch_image_sizes = [] + batch_image_positions = [] + batch_image_padding = [] + batch_image_noise = [] + for image, position_start, noise_value in zip( + (source_latents[batch_index], latent), + cap_end_positions, + (0, 1), + ): + patches, image_size, token_grid_size = self._patchify_image(image, patch_size, f_patch_size) + padded_features, position_ids, padding_mask, _, noise_mask = self._pad_with_ids( + patches, + token_grid_size, + (position_start, 0, 0), + noise_value, + ) + batch_image_features.append(padded_features) + batch_image_sizes.append(image_size) + batch_image_positions.append(position_ids) + batch_image_padding.append(padding_mask) + batch_image_noise.extend(noise_mask) + + batch_cap_features = torch.cat(batch_cap_features, dim=0) + batch_image_features = torch.cat(batch_image_features, dim=0) + cap_sequence.features.append(batch_cap_features) + cap_sequence.position_ids.append(torch.cat(batch_cap_positions, dim=0)) + cap_sequence.padding_masks.append(torch.cat(batch_cap_padding, dim=0)) + cap_sequence.noise_masks.append(batch_cap_noise) + image_sequence.features.append(batch_image_features) + image_sequence.position_ids.append(torch.cat(batch_image_positions, dim=0)) + image_sequence.padding_masks.append(torch.cat(batch_image_padding, dim=0)) + image_sequence.noise_masks.append(batch_image_noise) + image_sizes.append(batch_image_sizes) + image_offsets.append( + ( + len(batch_cap_features), + len(batch_cap_features) + len(batch_image_features), + ) + ) + + padded_features, position_ids, padding_mask, _, noise_mask = self._pad_with_ids( + glm_cap_feats[batch_index], + (len(glm_cap_feats[batch_index]), 1, 1), + (len(batch_cap_features) + len(batch_image_features) + 1, 0, 0), + 0, + ) + sigvq_sequence.features.append(padded_features) + sigvq_sequence.position_ids.append(position_ids) + sigvq_sequence.padding_masks.append(padding_mask) + sigvq_sequence.noise_masks.append(noise_mask) + + return image_sequence, cap_sequence, sigvq_sequence, image_sizes, image_offsets + + @staticmethod + def _merge_padded_sequences( + feature_groups: tuple[torch.Tensor, ...], + frequency_groups: tuple[torch.Tensor, ...], + length_groups: tuple[list[int], ...], + noise_mask_groups: tuple[torch.Tensor, ...] | None = None, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None, list[int], torch.Tensor | None]: + batch_size = feature_groups[0].shape[0] + merged_features = [] + merged_frequencies = [] + merged_noise_masks = [] if noise_mask_groups is not None else None + + for batch_index in range(batch_size): + device = feature_groups[0].device + merged_features.append( + torch.cat( + [ + features[batch_index, : lengths[batch_index]].to(device) + for features, lengths in zip(feature_groups, length_groups) + ], + dim=0, + ) + ) + merged_frequencies.append( + torch.cat( + [ + frequencies[batch_index, : lengths[batch_index]].to(device) + for frequencies, lengths in zip(frequency_groups, length_groups) + ], + dim=0, + ) + ) + if merged_noise_masks is not None: + merged_noise_masks.append( + torch.cat( + [ + noise_masks[batch_index, : lengths[batch_index]].to(device) + for noise_masks, lengths in zip(noise_mask_groups, length_groups) + ], + dim=0, + ) + ) + + merged_lengths = [len(features) for features in merged_features] + merged_features = pad_sequence(merged_features, batch_first=True, padding_value=0.0) + merged_frequencies = pad_sequence(merged_frequencies, batch_first=True, padding_value=0.0) + + attention_mask = None + max_length = max(merged_lengths) + if not all(length == max_length for length in merged_lengths): + attention_mask = torch.zeros( + (batch_size, max_length), + dtype=torch.bool, + device=merged_features.device, + ) + for batch_index, sequence_length in enumerate(merged_lengths): + attention_mask[batch_index, :sequence_length] = True + + noise_mask = None + if merged_noise_masks is not None: + noise_mask = pad_sequence(merged_noise_masks, batch_first=True, padding_value=0)[ + :, : merged_features.shape[1] + ] + + return merged_features, merged_frequencies, attention_mask, merged_lengths, noise_mask + + def forward( + self, + x: list[torch.Tensor], + t: torch.Tensor, + cap_feats: list[torch.Tensor] | None, + glm_cap_feats: list[torch.Tensor] | None = None, + source_latents: list[torch.Tensor] | None = None, + patch_size: int = 1, + f_patch_size: int = 1, + return_dict: bool = True, + ) -> Transformer2DModelOutput | tuple[list[torch.Tensor]]: + r""" + Args: + x (`list[torch.Tensor]`): + Target latents. Each tensor has shape `(channels, frames, height, width)`. + t (`torch.Tensor`): + Denoising timestep for each batch item. + cap_feats (`list[torch.Tensor]`, *optional*): + Projected QueryFormer features, each with shape `(sequence_length, cap_feat_dim)`. + glm_cap_feats (`list[torch.Tensor]`, *optional*): + GLM/SigVQ features, each with shape `(sequence_length, semantic_feat_dim)`. + source_latents (`list[torch.Tensor]`, *optional*): + Source-image latents for editing. When provided, `cap_feats` and `glm_cap_feats` are required. + patch_size (`int`, defaults to `1`): + Spatial patch size. + f_patch_size (`int`, defaults to `1`): + Temporal patch size. + return_dict (`bool`, defaults to `True`): + Whether to return a [`~models.modeling_outputs.Transformer2DModelOutput`]. + + Returns: + [`~models.modeling_outputs.Transformer2DModelOutput`] or `tuple`: + The denoised target latents. + """ + patch_key = f"{patch_size}-{f_patch_size}" + if patch_key not in self.all_x_embedder: + raise ValueError(f"Unsupported patch sizes: patch_size={patch_size}, f_patch_size={f_patch_size}.") + if source_latents is None and cap_feats is None and glm_cap_feats is None: + raise ValueError("Text-to-image inference requires `cap_feats` or `glm_cap_feats`.") + if source_latents is not None and (cap_feats is None or glm_cap_feats is None): + raise ValueError("Editing requires `cap_feats`, `glm_cap_feats`, and `source_latents`.") + + batch_size = len(x) + is_editing = source_latents is not None + adaln_input = None + noisy_embedding = None + clean_embedding = None + image_offsets = None + + if is_editing: + if t.shape[0] == 1: + t = t.repeat(batch_size) + dual_timestep = torch.cat([t, torch.zeros_like(t)], dim=0) + dual_embedding = self.t_embedder(dual_timestep.abs() * self.t_scale, x[0].dtype) + noisy_embedding = dual_embedding[:batch_size] + clean_embedding = dual_embedding[batch_size:] + image_sequence, cap_sequence, sigvq_sequence, image_sizes, image_offsets = self._prepare_editing_sequences( + x, + cap_feats, + glm_cap_feats, + source_latents, + patch_size, + f_patch_size, + ) + else: + adaln_input = self.t_embedder(t * self.t_scale, x[0].dtype) + glm_features = ( + [self.semantic_embedder(batch_features) for batch_features in glm_cap_feats] + if glm_cap_feats is not None + else None + ) + image_sequence, cap_sequence, glm_sequence, image_sizes = self._prepare_t2i_sequences( + x, + cap_feats, + glm_features, + patch_size, + f_patch_size, + ) + + image_lengths = [len(features) for features in image_sequence.features] + image_features = self.all_x_embedder[patch_key](torch.cat(image_sequence.features, dim=0)) + image_frequencies = list( + self.rope_embedder(torch.cat(image_sequence.position_ids, dim=0)).split( + [len(position_ids) for position_ids in image_sequence.position_ids], + dim=0, + ) + ) + image_features, image_frequencies, image_attention_mask, image_lengths, image_noise_mask = ( + self._batch_sequences( + list(image_features.split(image_lengths, dim=0)), + image_frequencies, + image_sequence.padding_masks, + self.x_pad_token, + image_sequence.noise_masks, + ) + ) + + for layer in self.noise_refiner: + if torch.is_grad_enabled() and self.gradient_checkpointing: + if is_editing: + image_features = self._gradient_checkpointing_func( + layer, + image_features, + image_attention_mask, + image_frequencies, + None, + image_noise_mask, + noisy_embedding, + clean_embedding, + ) + else: + image_features = self._gradient_checkpointing_func( + layer, + image_features, + image_attention_mask, + image_frequencies, + adaln_input, + ) + elif is_editing: + image_features = layer( + image_features, + image_attention_mask, + image_frequencies, + noise_mask=image_noise_mask, + adaln_noisy=noisy_embedding, + adaln_clean=clean_embedding, + ) + else: + image_features = layer( + image_features, + image_attention_mask, + image_frequencies, + adaln_input, + ) + + if is_editing: + cap_lengths = [len(features) for features in cap_sequence.features] + cap_features = self.cap_embedder(torch.cat(cap_sequence.features, dim=0)) + cap_frequencies = list( + self.rope_embedder(torch.cat(cap_sequence.position_ids, dim=0)).split( + [len(position_ids) for position_ids in cap_sequence.position_ids], + dim=0, + ) + ) + cap_features, cap_frequencies, cap_attention_mask, cap_lengths, cap_noise_mask = self._batch_sequences( + list(cap_features.split(cap_lengths, dim=0)), + cap_frequencies, + cap_sequence.padding_masks, + self.cap_pad_token, + cap_sequence.noise_masks, + ) + + for layer in self.context_refiner: + if torch.is_grad_enabled() and self.gradient_checkpointing: + cap_features = self._gradient_checkpointing_func( + layer, + cap_features, + cap_attention_mask, + cap_frequencies, + ) + else: + cap_features = layer( + cap_features, + cap_attention_mask, + cap_frequencies, + ) + + sigvq_lengths = [len(features) for features in sigvq_sequence.features] + sigvq_features = self.sigvq_embedder(torch.cat(sigvq_sequence.features, dim=0)) + sigvq_frequencies = list( + self.rope_embedder(torch.cat(sigvq_sequence.position_ids, dim=0)).split( + [len(position_ids) for position_ids in sigvq_sequence.position_ids], + dim=0, + ) + ) + ( + sigvq_features, + sigvq_frequencies, + sigvq_attention_mask, + sigvq_lengths, + sigvq_noise_mask, + ) = self._batch_sequences( + list(sigvq_features.split(sigvq_lengths, dim=0)), + sigvq_frequencies, + sigvq_sequence.padding_masks, + self.sigvq_pad_token, + sigvq_sequence.noise_masks, + ) + + for layer in self.sigvq_refiner: + if torch.is_grad_enabled() and self.gradient_checkpointing: + sigvq_features = self._gradient_checkpointing_func( + layer, + sigvq_features, + sigvq_attention_mask, + sigvq_frequencies, + ) + else: + sigvq_features = layer( + sigvq_features, + sigvq_attention_mask, + sigvq_frequencies, + ) + + ( + unified_features, + unified_frequencies, + unified_attention_mask, + _, + unified_noise_mask, + ) = self._merge_padded_sequences( + (cap_features, image_features, sigvq_features), + (cap_frequencies, image_frequencies, sigvq_frequencies), + (cap_lengths, image_lengths, sigvq_lengths), + (cap_noise_mask, image_noise_mask, sigvq_noise_mask), + ) + else: + condition_feature_groups = [] + condition_frequency_groups = [] + condition_length_groups = [] + + if cap_sequence is not None: + cap_lengths = [len(features) for features in cap_sequence.features] + cap_features = self.cap_embedder(torch.cat(cap_sequence.features, dim=0)) + cap_padding_mask = torch.cat(cap_sequence.padding_masks).unsqueeze(-1).to(cap_features.device) + cap_features = torch.where( + cap_padding_mask, + self.cap_pad_token.to(device=cap_features.device, dtype=cap_features.dtype), + cap_features, + ) + cap_features = pad_sequence( + list(cap_features.split(cap_lengths, dim=0)), + batch_first=True, + padding_value=0.0, + ) + cap_frequencies = list( + self.rope_embedder(torch.cat(cap_sequence.position_ids, dim=0)).split( + [len(position_ids) for position_ids in cap_sequence.position_ids], + dim=0, + ) + ) + cap_frequencies = pad_sequence(cap_frequencies, batch_first=True, padding_value=0.0) + condition_feature_groups.append(cap_features) + condition_frequency_groups.append(cap_frequencies) + condition_length_groups.append(cap_lengths) + + if glm_sequence is not None: + glm_lengths = [len(features) for features in glm_sequence.features] + glm_features = torch.cat(glm_sequence.features, dim=0) + glm_padding_mask = torch.cat(glm_sequence.padding_masks).unsqueeze(-1).to(glm_features.device) + glm_features = torch.where( + glm_padding_mask, + self.cap_pad_token.to(device=glm_features.device, dtype=glm_features.dtype), + glm_features, + ) + glm_features = pad_sequence( + list(glm_features.split(glm_lengths, dim=0)), + batch_first=True, + padding_value=0.0, + ) + glm_frequencies = list( + self.rope_embedder(torch.cat(glm_sequence.position_ids, dim=0)).split( + [len(position_ids) for position_ids in glm_sequence.position_ids], + dim=0, + ) + ) + glm_frequencies = pad_sequence(glm_frequencies, batch_first=True, padding_value=0.0) + condition_feature_groups.append(glm_features) + condition_frequency_groups.append(glm_frequencies) + condition_length_groups.append(glm_lengths) + + condition_features, condition_frequencies, condition_attention_mask, condition_lengths, _ = ( + self._merge_padded_sequences( + tuple(condition_feature_groups), + tuple(condition_frequency_groups), + tuple(condition_length_groups), + ) + ) + + for layer in self.context_refiner: + if torch.is_grad_enabled() and self.gradient_checkpointing: + condition_features = self._gradient_checkpointing_func( + layer, + condition_features, + condition_attention_mask, + condition_frequencies, + ) + else: + condition_features = layer( + condition_features, + condition_attention_mask, + condition_frequencies, + ) + + unified_features, unified_frequencies, unified_attention_mask, _, unified_noise_mask = ( + self._merge_padded_sequences( + (image_features, condition_features), + (image_frequencies, condition_frequencies), + (image_lengths, condition_lengths), + ) + ) + + for layer in self.layers: + if torch.is_grad_enabled() and self.gradient_checkpointing: + if is_editing: + unified_features = self._gradient_checkpointing_func( + layer, + unified_features, + unified_attention_mask, + unified_frequencies, + None, + unified_noise_mask, + noisy_embedding, + clean_embedding, + ) + else: + unified_features = self._gradient_checkpointing_func( + layer, + unified_features, + unified_attention_mask, + unified_frequencies, + adaln_input, + ) + elif is_editing: + unified_features = layer( + unified_features, + unified_attention_mask, + unified_frequencies, + noise_mask=unified_noise_mask, + adaln_noisy=noisy_embedding, + adaln_clean=clean_embedding, + ) + else: + unified_features = layer( + unified_features, + unified_attention_mask, + unified_frequencies, + adaln_input, + ) + + if is_editing: + unified_features = self.all_final_layer[patch_key]( + unified_features, + noise_mask=unified_noise_mask, + adaln_noisy=noisy_embedding, + adaln_clean=clean_embedding, + ) + else: + unified_features = self.all_final_layer[patch_key]( + unified_features, + adaln_input=adaln_input, + ) + + output = self._unpatchify( + list(unified_features.unbind(dim=0)), + image_sizes, + patch_size, + f_patch_size, + image_offsets, + ) + if not return_dict: + return (output,) + return Transformer2DModelOutput(sample=output) + + +@dataclass +class LLaDAImageQueryFormerOutput(BaseOutput): + query_embeds: torch.Tensor + + +class LLaDAImageQueryAttnProcessor: + _attention_backend = None + _parallel_config = None + + def __call__( + self, + attn: "LLaDAImageQueryAttention", + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None = None, + ) -> torch.Tensor: + query = F.linear( + hidden_states, + attn.in_proj_weight[: attn.inner_dim], + attn.in_proj_bias[: attn.inner_dim], + ) + key = F.linear( + encoder_hidden_states, + attn.in_proj_weight[attn.inner_dim : 2 * attn.inner_dim], + attn.in_proj_bias[attn.inner_dim : 2 * attn.inner_dim], + ) + value = F.linear( + encoder_hidden_states, + attn.in_proj_weight[2 * attn.inner_dim :], + attn.in_proj_bias[2 * attn.inner_dim :], + ) + + query = query.unflatten(-1, (attn.heads, attn.head_dim)) + key = key.unflatten(-1, (attn.heads, attn.head_dim)) + value = value.unflatten(-1, (attn.heads, attn.head_dim)) + + if attention_mask is not None: + attention_mask = attention_mask[:, None, None, :] + + hidden_states = dispatch_attention_fn( + query, + key, + value, + attn_mask=attention_mask, + dropout_p=attn.dropout if attn.training else 0.0, + is_causal=False, + backend=self._attention_backend, + parallel_config=self._parallel_config, + ) + hidden_states = hidden_states.flatten(2, 3) + return attn.out_proj(hidden_states) + + +class LLaDAImageQueryAttention(nn.Module, AttentionModuleMixin): + _default_processor_cls = LLaDAImageQueryAttnProcessor + _available_processors = [LLaDAImageQueryAttnProcessor] + _supports_qkv_fusion = False + + def __init__(self, hidden_size: int, num_heads: int, dropout: float): + super().__init__() + self.inner_dim = hidden_size + self.heads = num_heads + self.head_dim = hidden_size // num_heads + self.dropout = dropout + + self.in_proj_weight = nn.Parameter(torch.zeros(3 * hidden_size, hidden_size)) + self.in_proj_bias = nn.Parameter(torch.zeros(3 * hidden_size)) + self.out_proj = nn.Linear(hidden_size, hidden_size, bias=True) + self.set_processor(self._default_processor_cls()) + + nn.init.xavier_uniform_(self.in_proj_weight) + nn.init.zeros_(self.in_proj_bias) + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None = None, + ) -> torch.Tensor: + return self.processor(self, hidden_states, encoder_hidden_states, attention_mask) + + +@maybe_allow_in_graph +class LLaDAImageQueryFormerBlock(nn.Module): + def __init__( + self, + hidden_size: int, + num_heads: int, + intermediate_size: int, + dropout: float, + norm_eps: float, + ): + super().__init__() + self.norm_q = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=norm_eps) + self.norm_k = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=norm_eps) + self.cross_attn = LLaDAImageQueryAttention(hidden_size, num_heads, dropout) + self.dropout = nn.Dropout(dropout) + self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=norm_eps) + self.mlp = nn.Module() + self.mlp.fc1 = nn.Linear(hidden_size, intermediate_size, bias=True) + self.mlp.fc2 = nn.Linear(intermediate_size, hidden_size, bias=True) + + def forward( + self, + query_embeds: torch.Tensor, + encoder_hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None = None, + ) -> torch.Tensor: + query_embeds = self.norm_q(query_embeds) + encoder_hidden_states = self.norm_k(encoder_hidden_states) + attention_output = self.cross_attn(query_embeds, encoder_hidden_states, attention_mask) + query_embeds = query_embeds + self.dropout(attention_output) + query_embeds = self.norm1(query_embeds) + mlp_output = self.mlp.fc2(F.gelu(self.mlp.fc1(query_embeds), approximate="tanh")) + return query_embeds + self.dropout(mlp_output) + + +class LLaDAImageQueryFormerModel(ModelMixin, ConfigMixin, AttentionMixin): + r""" + QueryFormer used by LLaDA-Image to derive learnable image-generation queries from LLaDA token embeddings. + + This model is independent from the LLaDA text encoder. It returns refined query embeddings; the pipeline appends + them to the text embeddings and invokes the text encoder backbone. + + Args: + num_queries (`int`, defaults to `256`): + Number of learnable query tokens. + hidden_size (`int`, defaults to `2048`): + Query and LLaDA token embedding dimension. + num_hidden_layers (`int`, defaults to `1`): + Number of QueryFormer blocks. + num_attention_heads (`int`, defaults to `16`): + Number of cross-attention heads. + intermediate_size (`int`, defaults to `8192`): + Hidden dimension of the QueryFormer MLP. + dropout (`float`, defaults to `0.0`): + Dropout probability. + norm_eps (`float`, defaults to `1e-6`): + Epsilon used by parameter-free layer normalization. + """ + + _supports_gradient_checkpointing = True + _no_split_modules = ["LLaDAImageQueryFormerBlock"] + _repeated_blocks = ["LLaDAImageQueryFormerBlock"] + _skip_layerwise_casting_patterns = ["norm"] + + @register_to_config + def __init__( + self, + num_queries: int = 256, + hidden_size: int = 2048, + num_hidden_layers: int = 1, + num_attention_heads: int = 16, + intermediate_size: int = 8192, + dropout: float = 0.0, + norm_eps: float = 1e-6, + ): + super().__init__() + if hidden_size % num_attention_heads != 0: + raise ValueError( + f"`hidden_size` ({hidden_size}) must be divisible by `num_attention_heads` ({num_attention_heads})." + ) + + self.meta_queries = nn.Parameter(torch.zeros(num_queries, hidden_size)) + nn.init.normal_(self.meta_queries, std=1 / math.sqrt(hidden_size)) + self.query_blocks = nn.ModuleList( + [ + LLaDAImageQueryFormerBlock( + hidden_size, + num_attention_heads, + intermediate_size, + dropout, + norm_eps, + ) + for _ in range(num_hidden_layers) + ] + ) + self.gradient_checkpointing = False + + def forward( + self, + inputs_embeds: torch.Tensor, + attention_mask: torch.Tensor, + return_dict: bool = True, + ) -> LLaDAImageQueryFormerOutput | tuple[torch.Tensor]: + r""" + Args: + inputs_embeds (`torch.Tensor` of shape `(batch_size, sequence_length, hidden_size)`): + LLaDA input token embeddings. + attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`): + Mask whose nonzero entries identify valid text tokens. + return_dict (`bool`, defaults to `True`): + Whether to return [`LLaDAImageQueryFormerOutput`] instead of a tuple. + + Returns: + [`LLaDAImageQueryFormerOutput`] or `tuple`: + The refined query embeddings. + """ + batch_size = inputs_embeds.shape[0] + query_embeds = self.meta_queries.unsqueeze(0).expand(batch_size, -1, -1) + attention_mask = attention_mask.bool() + + for query_block in self.query_blocks: + if torch.is_grad_enabled() and self.gradient_checkpointing: + query_embeds = self._gradient_checkpointing_func( + query_block, + query_embeds, + inputs_embeds, + attention_mask, + ) + else: + query_embeds = query_block(query_embeds, inputs_embeds, attention_mask) + + if not return_dict: + return (query_embeds,) + return LLaDAImageQueryFormerOutput(query_embeds=query_embeds) + + +@dataclass +class LLaDAImageTextProjectionOutput(BaseOutput): + hidden_states: torch.Tensor + + +class LLaDAImageTextProjectionAttnProcessor: + _attention_backend = None + _parallel_config = None + + def __call__( + self, + attn: "LLaDAImageTextProjectionAttention", + hidden_states: torch.Tensor, + ) -> torch.Tensor: + query = attn.q_proj(hidden_states).unflatten(-1, (attn.heads, attn.head_dim)) + key = attn.k_proj(hidden_states).unflatten(-1, (attn.heads, attn.head_dim)) + value = attn.v_proj(hidden_states).unflatten(-1, (attn.heads, attn.head_dim)) + + query = attn.q_norm(query) + key = attn.k_norm(key) + + hidden_states = dispatch_attention_fn( + query, + key, + value, + attn_mask=None, + dropout_p=attn.dropout if attn.training else 0.0, + is_causal=False, + backend=self._attention_backend, + parallel_config=self._parallel_config, + ) + hidden_states = hidden_states.flatten(2, 3) + return attn.out_proj(hidden_states) + + +class LLaDAImageTextProjectionAttention(nn.Module, AttentionModuleMixin): + _default_processor_cls = LLaDAImageTextProjectionAttnProcessor + _available_processors = [LLaDAImageTextProjectionAttnProcessor] + _supports_qkv_fusion = False + + def __init__(self, hidden_size: int, num_attention_heads: int, attention_dropout: float, norm_eps: float): + super().__init__() + self.heads = num_attention_heads + self.head_dim = hidden_size // num_attention_heads + self.dropout = attention_dropout + + self.k_proj = nn.Linear(hidden_size, hidden_size, bias=True) + self.v_proj = nn.Linear(hidden_size, hidden_size, bias=True) + self.q_proj = nn.Linear(hidden_size, hidden_size, bias=True) + self.out_proj = nn.Linear(hidden_size, hidden_size, bias=True) + self.q_norm = RMSNorm(self.head_dim, eps=norm_eps, elementwise_affine=False) + self.k_norm = RMSNorm(self.head_dim, eps=norm_eps, elementwise_affine=False) + self.set_processor(self._default_processor_cls()) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + return self.processor(self, hidden_states) + + +class LLaDAImageTextProjectionMLP(nn.Module): + def __init__(self, hidden_size: int, intermediate_size: int): + super().__init__() + self.fc1 = nn.Linear(hidden_size, intermediate_size, bias=True) + self.fc2 = nn.Linear(intermediate_size, hidden_size, bias=True) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + hidden_states = self.fc1(hidden_states) + hidden_states = F.gelu(hidden_states, approximate="tanh") + return self.fc2(hidden_states) + + +@maybe_allow_in_graph +class LLaDAImageTextProjectionBlock(nn.Module): + def __init__( + self, + hidden_size: int, + intermediate_size: int, + num_attention_heads: int, + attention_dropout: float, + norm_eps: float, + ): + super().__init__() + self.self_attn = LLaDAImageTextProjectionAttention( + hidden_size, + num_attention_heads, + attention_dropout, + norm_eps, + ) + self.layer_norm1 = RMSNorm(hidden_size, eps=norm_eps, elementwise_affine=False) + self.mlp = LLaDAImageTextProjectionMLP(hidden_size, intermediate_size) + self.layer_norm2 = RMSNorm(hidden_size, eps=norm_eps, elementwise_affine=False) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + hidden_states = hidden_states + self.self_attn(self.layer_norm1(hidden_states)) + hidden_states = hidden_states + self.mlp(self.layer_norm2(hidden_states)) + return hidden_states + + +class LLaDAImageTextProjectionModel(ModelMixin, ConfigMixin, AttentionMixin): + r""" + Connector and output projection used to map LLaDA hidden states to the LLaDA-Image denoiser context dimension. + + Args: + hidden_size (`int`, defaults to `2048`): + Input and connector hidden dimension. + intermediate_size (`int`, defaults to `8960`): + Connector MLP hidden dimension. + num_hidden_layers (`int`, defaults to `6`): + Number of connector layers. + num_attention_heads (`int`, defaults to `32`): + Number of connector self-attention heads. + projection_dim (`int`, defaults to `2560`): + Output dimension expected by the denoising transformer. + attention_dropout (`float`, defaults to `0.0`): + Attention dropout probability. + norm_eps (`float`, defaults to `1e-6`): + Epsilon used by parameter-free RMS normalization. + """ + + _supports_gradient_checkpointing = True + _no_split_modules = ["LLaDAImageTextProjectionBlock"] + _repeated_blocks = ["LLaDAImageTextProjectionBlock"] + _skip_layerwise_casting_patterns = ["layer_norm", "q_norm", "k_norm"] + + @register_to_config + def __init__( + self, + hidden_size: int = 2048, + intermediate_size: int = 8960, + num_hidden_layers: int = 6, + num_attention_heads: int = 32, + projection_dim: int = 2560, + attention_dropout: float = 0.0, + norm_eps: float = 1e-6, + ): + super().__init__() + if hidden_size % num_attention_heads != 0: + raise ValueError( + f"`hidden_size` ({hidden_size}) must be divisible by `num_attention_heads` ({num_attention_heads})." + ) + + self.layers = nn.ModuleList( + [ + LLaDAImageTextProjectionBlock( + hidden_size, + intermediate_size, + num_attention_heads, + attention_dropout, + norm_eps, + ) + for _ in range(num_hidden_layers) + ] + ) + self.projector = nn.Linear(hidden_size, projection_dim, bias=True) + self.gradient_checkpointing = False + + def forward( + self, + hidden_states: torch.Tensor, + return_dict: bool = True, + ) -> LLaDAImageTextProjectionOutput | tuple[torch.Tensor]: + r""" + Args: + hidden_states (`torch.Tensor` of shape `(batch_size, sequence_length, hidden_size)`): + Hidden states produced by the LLaDA text backbone. + return_dict (`bool`, defaults to `True`): + Whether to return [`LLaDAImageTextProjectionOutput`] instead of a tuple. + + Returns: + [`LLaDAImageTextProjectionOutput`] or `tuple`: + Hidden states projected to the denoiser caption dimension. + """ + for layer in self.layers: + if torch.is_grad_enabled() and self.gradient_checkpointing: + hidden_states = self._gradient_checkpointing_func(layer, hidden_states) + else: + hidden_states = layer(hidden_states) + + hidden_states = self.projector(hidden_states) + if not return_dict: + return (hidden_states,) + return LLaDAImageTextProjectionOutput(hidden_states=hidden_states) + + +@dataclass +class LLaDAImageSigVQOutput(BaseOutput): + semantic_features: torch.Tensor + token_ids: torch.Tensor + + +class LLaDAImageSigVQAttnProcessor: + _attention_backend = None + _parallel_config = None + + def __call__(self, attn: "LLaDAImageSigVQAttention", hidden_states: torch.Tensor) -> torch.Tensor: + query, key, value = attn.qkv(hidden_states).chunk(3, dim=-1) + query = query.unflatten(-1, (attn.heads, attn.head_dim)) + key = key.unflatten(-1, (attn.heads, attn.head_dim)) + value = value.unflatten(-1, (attn.heads, attn.head_dim)) + + hidden_states = dispatch_attention_fn( + query, + key, + value, + attn_mask=None, + dropout_p=attn.dropout if attn.training else 0.0, + is_causal=False, + backend=self._attention_backend, + parallel_config=self._parallel_config, + ) + hidden_states = hidden_states.flatten(2, 3) + return attn.proj(hidden_states) + + +class LLaDAImageSigVQAttention(nn.Module, AttentionModuleMixin): + _default_processor_cls = LLaDAImageSigVQAttnProcessor + _available_processors = [LLaDAImageSigVQAttnProcessor] + _supports_qkv_fusion = False + + def __init__( + self, + hidden_size: int, + num_attention_heads: int, + attention_bias: bool, + attention_dropout: float, + ): + super().__init__() + self.heads = num_attention_heads + self.head_dim = hidden_size // num_attention_heads + self.dropout = attention_dropout + self.qkv = nn.Linear(hidden_size, 3 * hidden_size, bias=attention_bias) + self.proj = nn.Linear(hidden_size, hidden_size, bias=attention_bias) + self.set_processor(self._default_processor_cls()) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + return self.processor(self, hidden_states) + + +class LLaDAImageSigVQMLP(nn.Module): + def __init__(self, hidden_size: int, intermediate_size: int): + super().__init__() + self.fc1 = nn.Linear(hidden_size, intermediate_size, bias=True) + self.fc2 = nn.Linear(intermediate_size, hidden_size, bias=True) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + return self.fc2(F.gelu(self.fc1(hidden_states))) + + +@maybe_allow_in_graph +class LLaDAImageSigVQVisionBlock(nn.Module): + def __init__( + self, + hidden_size: int, + intermediate_size: int, + num_attention_heads: int, + attention_bias: bool, + attention_dropout: float, + norm_eps: float, + ): + super().__init__() + self.norm1 = nn.LayerNorm(hidden_size, eps=norm_eps) + self.norm2 = nn.LayerNorm(hidden_size, eps=norm_eps) + self.attn = LLaDAImageSigVQAttention( + hidden_size, + num_attention_heads, + attention_bias, + attention_dropout, + ) + self.mlp = LLaDAImageSigVQMLP(hidden_size, intermediate_size) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + hidden_states = hidden_states + self.attn(self.norm1(hidden_states)) + hidden_states = hidden_states + self.mlp(self.norm2(hidden_states)) + return hidden_states + + +class LLaDAImageSigVQPatchEmbed(nn.Module): + def __init__(self, in_channels: int, hidden_size: int, patch_size: int): + super().__init__() + self.in_channels = in_channels + self.patch_size = patch_size + self.proj = nn.Conv2d(in_channels, hidden_size, kernel_size=patch_size, stride=patch_size) + + def forward(self, pixel_values: torch.Tensor) -> torch.Tensor: + batch_size, channels, height, width = pixel_values.shape + grid_height = height // self.patch_size + grid_width = width // self.patch_size + patches = pixel_values.reshape( + batch_size, + channels, + grid_height, + self.patch_size, + grid_width, + self.patch_size, + ) + patches = patches.permute(0, 2, 4, 1, 3, 5).reshape( + batch_size * grid_height * grid_width, + channels, + self.patch_size, + self.patch_size, + ) + hidden_states = self.proj(patches).flatten(1) + return hidden_states.reshape(batch_size, grid_height * grid_width, -1) + + +class LLaDAImageSigVQEmbeddings(nn.Module): + def __init__(self, image_size: int, patch_size: int, hidden_size: int): + super().__init__() + num_positions = (image_size // patch_size) ** 2 + self.position_embedding = nn.Embedding(num_positions, hidden_size) + + def forward(self, hidden_states: torch.Tensor, grid_height: int, grid_width: int) -> torch.Tensor: + batch_size = hidden_states.shape[0] + position_embedding = self.position_embedding.weight + hidden_size = position_embedding.shape[1] + original_size = int(position_embedding.shape[0] ** 0.5) + position_embedding = position_embedding.reshape(original_size, original_size, hidden_size) + position_embedding = position_embedding.permute(2, 0, 1).unsqueeze(0).float() + + height_coordinates = torch.arange(grid_height, device=hidden_states.device, dtype=torch.float32) + width_coordinates = torch.arange(grid_width, device=hidden_states.device, dtype=torch.float32) + height_coordinates, width_coordinates = torch.meshgrid( + height_coordinates, + width_coordinates, + indexing="ij", + ) + normalized_width = ((width_coordinates.flatten() + 0.5) / grid_width) * 2 - 1 + normalized_height = ((height_coordinates.flatten() + 0.5) / grid_height) * 2 - 1 + grid = torch.stack((normalized_width, normalized_height), dim=-1) + grid = grid.reshape(1, grid_height * grid_width, 1, 2).expand(batch_size, -1, -1, -1) + + position_embedding = F.grid_sample( + position_embedding.expand(batch_size, -1, -1, -1), + grid, + mode="bilinear", + align_corners=False, + padding_mode="border", + ) + position_embedding = position_embedding.squeeze(-1).transpose(1, 2).to(hidden_states.dtype) + return hidden_states + position_embedding + + +class LLaDAImageSigVQQuantizer(nn.Module): + def __init__(self, num_embeddings: int, embedding_dim: int): + super().__init__() + self.embedding = nn.Embedding(num_embeddings, embedding_dim) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + hidden_states = hidden_states.permute(0, 2, 3, 1).contiguous() + hidden_states = F.normalize(hidden_states.reshape(-1, hidden_states.shape[-1]), p=2, dim=-1) + embedding = F.normalize(self.embedding.weight, p=2, dim=-1) + distances = ( + torch.sum(hidden_states**2, dim=1, keepdim=True) + + torch.sum(embedding**2, dim=1) + - 2 * torch.matmul(hidden_states, embedding.t()) + ) + return torch.argmin(distances, dim=1) + + +class LLaDAImageSigVQModel(ModelMixin, ConfigMixin, AttentionMixin): + r""" + Minimal GLM SigVQ image encoder used by LLaDA-Image editing. + + The model contains only the GLM vision encoder, VQ quantizer, and prior token projection used during inference. + Input images must already be RGB tensors normalized to `[-1, 1]`, have one common size, and be divisible by + `patch_size`. + """ + + _supports_gradient_checkpointing = True + _no_split_modules = ["LLaDAImageSigVQVisionBlock"] + _repeated_blocks = ["LLaDAImageSigVQVisionBlock"] + _skip_layerwise_casting_patterns = ["patch_embed", "position_embedding", "norm", "quantize"] + + @register_to_config + def __init__( + self, + image_size: int = 2048, + patch_size: int = 16, + in_channels: int = 3, + hidden_size: int = 1536, + intermediate_size: int = 6144, + num_hidden_layers: int = 40, + num_attention_heads: int = 16, + attention_bias: bool = True, + attention_dropout: float = 0.0, + norm_eps: float = 1e-6, + codebook_size: int = 16384, + codebook_embed_dim: int = 2048, + semantic_embed_dim: int = 4096, + ): + super().__init__() + if hidden_size % num_attention_heads != 0: + raise ValueError( + f"`hidden_size` ({hidden_size}) must be divisible by `num_attention_heads` ({num_attention_heads})." + ) + + self.visual = nn.Module() + self.visual.patch_embed = LLaDAImageSigVQPatchEmbed(in_channels, hidden_size, patch_size) + self.visual.embeddings = LLaDAImageSigVQEmbeddings(image_size, patch_size, hidden_size) + self.visual.blocks = nn.ModuleList( + [ + LLaDAImageSigVQVisionBlock( + hidden_size, + intermediate_size, + num_attention_heads, + attention_bias, + attention_dropout, + norm_eps, + ) + for _ in range(num_hidden_layers) + ] + ) + + self.vqmodel = nn.Module() + self.vqmodel.quant_conv = nn.Conv2d(hidden_size, codebook_embed_dim, kernel_size=1) + self.vqmodel.quantize = LLaDAImageSigVQQuantizer(codebook_size, codebook_embed_dim) + + self.prior_token_embedding = nn.Embedding(codebook_size, semantic_embed_dim) + self.prior_projector = FeedForward( + semantic_embed_dim, + semantic_embed_dim, + inner_dim=semantic_embed_dim, + activation_fn="linear-silu", + ) + self.gradient_checkpointing = False + + def forward( + self, + pixel_values: torch.Tensor | None = None, + token_ids: torch.Tensor | None = None, + return_dict: bool = True, + ) -> LLaDAImageSigVQOutput | tuple[torch.Tensor, torch.Tensor]: + r""" + Args: + pixel_values (`torch.Tensor` of shape `(batch_size, 3, height, width)`, *optional*): + RGB images normalized to `[-1, 1]`. Mutually exclusive with `token_ids`. + token_ids (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*): + Precomputed VQ codebook IDs. Mutually exclusive with `pixel_values`. + return_dict (`bool`, defaults to `True`): + Whether to return [`LLaDAImageSigVQOutput`] instead of a tuple. + + Returns: + [`LLaDAImageSigVQOutput`] or `tuple`: + The projected semantic features and their discrete token IDs. + """ + if (pixel_values is None) == (token_ids is None): + raise ValueError("Provide exactly one of `pixel_values` or `token_ids`.") + + if pixel_values is not None: + if pixel_values.ndim != 4: + raise ValueError(f"`pixel_values` must have 4 dimensions, got shape {tuple(pixel_values.shape)}.") + height, width = pixel_values.shape[-2:] + if height % self.config.patch_size != 0 or width % self.config.patch_size != 0: + raise ValueError( + f"Image height and width must be divisible by {self.config.patch_size}, got {height}x{width}." + ) + + grid_height = height // self.config.patch_size + grid_width = width // self.config.patch_size + hidden_states = self.visual.patch_embed(pixel_values) + hidden_states = self.visual.embeddings(hidden_states, grid_height, grid_width) + + for block in self.visual.blocks: + if torch.is_grad_enabled() and self.gradient_checkpointing: + hidden_states = self._gradient_checkpointing_func(block, hidden_states) + else: + hidden_states = block(hidden_states) + + hidden_states = hidden_states.transpose(1, 2).reshape( + pixel_values.shape[0], + self.config.hidden_size, + grid_height, + grid_width, + ) + hidden_states = self.vqmodel.quant_conv(hidden_states) + token_ids = self.vqmodel.quantize(hidden_states).reshape(pixel_values.shape[0], -1) + elif token_ids.ndim != 2: + raise ValueError(f"`token_ids` must have 2 dimensions, got shape {tuple(token_ids.shape)}.") + + semantic_features = self.prior_projector(self.prior_token_embedding(token_ids)) + + if not return_dict: + return semantic_features, token_ids + return LLaDAImageSigVQOutput(semantic_features=semantic_features, token_ids=token_ids) diff --git a/src/diffusers/pipelines/__init__.py b/src/diffusers/pipelines/__init__.py index 32f193a03080..063b997ad3d3 100644 --- a/src/diffusers/pipelines/__init__.py +++ b/src/diffusers/pipelines/__init__.py @@ -330,6 +330,7 @@ ) _import_structure["latte"] = ["LattePipeline"] _import_structure["llada2"] = ["LLaDA2Pipeline", "LLaDA2PipelineOutput"] + _import_structure["llada_image"] = ["LLaDAImagePipeline", "LLaDAImagePipelineOutput"] _import_structure["ltx"] = [ "LTXPipeline", "LTXImageToVideoPipeline", @@ -795,6 +796,7 @@ LEditsPPPipelineStableDiffusionXL, ) from .llada2 import LLaDA2Pipeline, LLaDA2PipelineOutput + from .llada_image import LLaDAImagePipeline, LLaDAImagePipelineOutput from .longcat_audio_dit import LongCatAudioDiTPipeline from .longcat_image import LongCatImageEditPipeline, LongCatImagePipeline from .ltx import ( diff --git a/src/diffusers/pipelines/llada_image/__init__.py b/src/diffusers/pipelines/llada_image/__init__.py new file mode 100644 index 000000000000..af462b184843 --- /dev/null +++ b/src/diffusers/pipelines/llada_image/__init__.py @@ -0,0 +1,47 @@ +from typing import TYPE_CHECKING + +from ...utils import ( + DIFFUSERS_SLOW_IMPORT, + OptionalDependencyNotAvailable, + _LazyModule, + get_objects_from_module, + is_torch_available, + is_transformers_available, +) + + +_dummy_objects = {} +_import_structure = {} + +try: + if not (is_transformers_available() and is_torch_available()): + raise OptionalDependencyNotAvailable() +except OptionalDependencyNotAvailable: + from ...utils import dummy_torch_and_transformers_objects # noqa F403 + + _dummy_objects.update(get_objects_from_module(dummy_torch_and_transformers_objects)) +else: + _import_structure["pipeline_llada_image"] = ["LLaDAImagePipeline"] + _import_structure["pipeline_output"] = ["LLaDAImagePipelineOutput"] + +if TYPE_CHECKING or DIFFUSERS_SLOW_IMPORT: + try: + if not (is_transformers_available() and is_torch_available()): + raise OptionalDependencyNotAvailable() + except OptionalDependencyNotAvailable: + from ...utils.dummy_torch_and_transformers_objects import * # noqa F403 + else: + from .pipeline_llada_image import LLaDAImagePipeline + from .pipeline_output import LLaDAImagePipelineOutput +else: + import sys + + sys.modules[__name__] = _LazyModule( + __name__, + globals()["__file__"], + _import_structure, + module_spec=__spec__, + ) + + for name, value in _dummy_objects.items(): + setattr(sys.modules[__name__], name, value) diff --git a/src/diffusers/pipelines/llada_image/pipeline_llada_image.py b/src/diffusers/pipelines/llada_image/pipeline_llada_image.py new file mode 100644 index 000000000000..3908df498d3d --- /dev/null +++ b/src/diffusers/pipelines/llada_image/pipeline_llada_image.py @@ -0,0 +1,634 @@ +# Copyright 2026 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from collections.abc import Callable + +import torch +import torch.nn.functional as F +from transformers import PreTrainedModel, PreTrainedTokenizerBase +from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS + +from ...image_processor import PipelineImageInput, VaeImageProcessor +from ...models import ( + AutoencoderKLFlux2, + LLaDAImageQueryFormerModel, + LLaDAImageSigVQModel, + LLaDAImageTextProjectionModel, + LLaDAImageTransformer2DModel, +) +from ...schedulers import FlowMatchEulerDiscreteScheduler +from ...utils import is_transformers_version, logging +from ...utils.torch_utils import randn_tensor +from ..pipeline_utils import DiffusionPipeline +from .pipeline_output import LLaDAImagePipelineOutput + + +logger = logging.get_logger(__name__) + + +def _default_rope_parameters(config, device=None, seq_len=None, layer_type=None): + del seq_len, layer_type + device = device if device is not None else torch.device("cpu") + head_dim = getattr(config, "head_dim", None) or config.hidden_size // config.num_attention_heads + dim = int(head_dim * getattr(config, "partial_rotary_factor", 1.0)) + inv_freq = 1.0 / (config.rope_theta ** (torch.arange(0, dim, 2, dtype=torch.int64, device=device).float() / dim)) + return inv_freq, 1.0 + + +# The official LLaDA2 remote model uses this Transformers 4 compatibility entry. Transformers 5 no longer +# registers it, so make it available before DiffusionPipeline loads the custom text encoder. +if "default" not in ROPE_INIT_FUNCTIONS: + ROPE_INIT_FUNCTIONS["default"] = _default_rope_parameters + + +class LLaDAImagePipeline(DiffusionPipeline): + r""" + Pipeline for LLaDA-Image text-to-image generation, VQ-conditioned generation, and single-image editing. + + Args: + scheduler ([`FlowMatchEulerDiscreteScheduler`]): + Flow-matching scheduler used for denoising. + vae ([`AutoencoderKLFlux2`]): + Flux2 VAE used to encode reference images and decode generated latents. + text_encoder (`transformers.PreTrainedModel`): + LLaDA2 conditional-generation model. It must expose `get_input_embeddings()` and its language backbone as + `model`. + tokenizer (`transformers.PreTrainedTokenizerBase`): + Tokenizer paired with the LLaDA2 text encoder. + queryformer ([`LLaDAImageQueryFormerModel`]): + QueryFormer that refines the learnable generation queries. + text_projection ([`LLaDAImageTextProjectionModel`]): + Connector and projector that map LLaDA2 hidden states to denoiser caption features. + sigvq ([`LLaDAImageSigVQModel`]): + GLM SigVQ component that embeds MLLM-generated VQ tokens and encodes editing reference images. + transformer ([`LLaDAImageTransformer2DModel`]): + Denoising transformer. + """ + + model_cpu_offload_seq = "text_encoder->queryformer->text_projection->sigvq->transformer->vae" + _exclude_from_cpu_offload = ["text_encoder"] + _callback_tensor_inputs = ["latents"] + + def __init__( + self, + scheduler: FlowMatchEulerDiscreteScheduler, + vae: AutoencoderKLFlux2, + text_encoder: PreTrainedModel, + tokenizer: PreTrainedTokenizerBase, + queryformer: LLaDAImageQueryFormerModel, + text_projection: LLaDAImageTextProjectionModel, + sigvq: LLaDAImageSigVQModel, + transformer: LLaDAImageTransformer2DModel, + ): + super().__init__() + self.register_modules( + scheduler=scheduler, + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + queryformer=queryformer, + text_projection=text_projection, + sigvq=sigvq, + transformer=transformer, + ) + + self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1) if self.vae is not None else 8 + self.latent_scale_factor = self.vae_scale_factor * 2 + self.image_processor = VaeImageProcessor(vae_scale_factor=self.latent_scale_factor) + + text_encoder_model = getattr(self.text_encoder, "model", None) + language_model = getattr(text_encoder_model, "language_model", None) + if is_transformers_version(">=", "5.0.0") and language_model is not None: + # Transformers 5 constructs sharded models on the meta device. The remote model's RoPE buffer is + # non-persistent and is therefore not restored from the checkpoint, so materialize it after loading. + rotary_emb = language_model.rotary_emb + rotary_device = self.text_encoder.get_input_embeddings().weight.device + default_rope_init = ROPE_INIT_FUNCTIONS["default"] + inv_freq, attention_scaling = default_rope_init(rotary_emb.config, device=rotary_device) + rotary_emb.register_buffer("inv_freq", inv_freq, persistent=False) + rotary_emb.original_inv_freq = inv_freq + rotary_emb.attention_scaling = attention_scaling + + @property + def guidance_scale(self) -> float: + return self._guidance_scale + + @property + def num_timesteps(self) -> int: + return self._num_timesteps + + @staticmethod + def _patchify_latents(latents: torch.Tensor) -> torch.Tensor: + batch_size, channels, height, width = latents.shape + latents = latents.reshape(batch_size, channels, height // 2, 2, width // 2, 2) + latents = latents.permute(0, 1, 3, 5, 2, 4) + return latents.reshape(batch_size, channels * 4, height // 2, width // 2) + + @staticmethod + def _unpatchify_latents(latents: torch.Tensor) -> torch.Tensor: + batch_size, channels, height, width = latents.shape + latents = latents.reshape(batch_size, channels // 4, 2, 2, height, width) + latents = latents.permute(0, 1, 4, 2, 5, 3) + return latents.reshape(batch_size, channels // 4, height * 2, width * 2) + + def _encode_text( + self, + prompts: list[str], + max_sequence_length: int, + ) -> tuple[torch.Tensor, torch.Tensor]: + formatted_prompts = [ + "HUMAN Generate an image.\nASSISTANT\n" + if prompt is None + else f"HUMAN Generate an image: {prompt.strip()}\nASSISTANT\n" + for prompt in prompts + ] + text_inputs = self.tokenizer( + formatted_prompts, + add_special_tokens=True, + padding=True, + truncation=True, + max_length=max_sequence_length, + return_tensors="pt", + ) + text_encoder_device = self.text_encoder.get_input_embeddings().weight.device + input_ids = text_inputs.input_ids.to(text_encoder_device) + attention_mask = text_inputs.attention_mask.to(text_encoder_device).bool() + inputs_embeds = self.text_encoder.get_input_embeddings()(input_ids) + + query_embeds = self.queryformer( + inputs_embeds.to(device=self._execution_device, dtype=self.queryformer.dtype), + attention_mask.to(self._execution_device), + ).query_embeds.to(device=text_encoder_device, dtype=inputs_embeds.dtype) + text_length = inputs_embeds.shape[1] + inputs_embeds = torch.cat([inputs_embeds, query_embeds], dim=1) + attention_mask = torch.cat( + [attention_mask, attention_mask.new_ones(attention_mask.shape[0], query_embeds.shape[1])], + dim=1, + ) + position_ids = attention_mask.long().cumsum(dim=1) - 1 + position_ids.masked_fill_(position_ids < 0, 0) + + mask_value = torch.finfo(inputs_embeds.dtype).min + backbone_attention_mask = attention_mask[:, None, None, :].expand(-1, 1, attention_mask.shape[1], -1) + backbone_attention_mask = torch.where( + backbone_attention_mask, + torch.zeros((), dtype=inputs_embeds.dtype, device=text_encoder_device), + torch.full((), mask_value, dtype=inputs_embeds.dtype, device=text_encoder_device), + ) + backbone_attention_mask[:, :, :text_length, text_length:] = mask_value + + hidden_states = self.text_encoder.model( + inputs_embeds=inputs_embeds, + attention_mask=backbone_attention_mask, + position_ids=position_ids, + return_dict=True, + ).last_hidden_state + prompt_embeds = self.text_projection( + hidden_states.to(device=self._execution_device, dtype=self.text_projection.dtype) + ).hidden_states + return prompt_embeds, attention_mask.to(prompt_embeds.device) + + def encode_prompt( + self, + prompt: str | list[str] | None, + negative_prompt: str | list[str] | None = None, + do_classifier_free_guidance: bool = True, + num_images_per_prompt: int = 1, + prompt_embeds: torch.Tensor | None = None, + prompt_attention_mask: torch.Tensor | None = None, + negative_prompt_embeds: torch.Tensor | None = None, + negative_prompt_attention_mask: torch.Tensor | None = None, + max_sequence_length: int = 2048, + device: torch.device | None = None, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None, torch.Tensor | None]: + device = device or self._execution_device + + if prompt_embeds is None: + prompt = [prompt] if isinstance(prompt, str) else prompt + prompt_embeds, prompt_attention_mask = self._encode_text(prompt, max_sequence_length) + else: + prompt_embeds = prompt_embeds.to(device) + prompt_attention_mask = prompt_attention_mask.to(device).bool() + + batch_size = prompt_embeds.shape[0] + if do_classifier_free_guidance and negative_prompt_embeds is None: + if negative_prompt is None: + negative_prompt = [None] * batch_size + elif isinstance(negative_prompt, str): + negative_prompt = [negative_prompt] * batch_size + negative_prompt_embeds, negative_prompt_attention_mask = self._encode_text( + negative_prompt, max_sequence_length + ) + elif do_classifier_free_guidance: + negative_prompt_embeds = negative_prompt_embeds.to(device) + negative_prompt_attention_mask = negative_prompt_attention_mask.to(device).bool() + + prompt_embeds = prompt_embeds.repeat_interleave(num_images_per_prompt, dim=0) + prompt_attention_mask = prompt_attention_mask.repeat_interleave(num_images_per_prompt, dim=0) + if do_classifier_free_guidance: + negative_prompt_embeds = negative_prompt_embeds.repeat_interleave(num_images_per_prompt, dim=0) + negative_prompt_attention_mask = negative_prompt_attention_mask.repeat_interleave( + num_images_per_prompt, dim=0 + ) + + return prompt_embeds, prompt_attention_mask, negative_prompt_embeds, negative_prompt_attention_mask + + def generate_vq_tokens( + self, + prompt: str | list[str], + height: int, + width: int, + ) -> torch.Tensor: + text_encoder_device = self.text_encoder.get_input_embeddings().weight.device + execution_device = self._execution_device + restore_text_encoder_device = text_encoder_device != execution_device + if restore_text_encoder_device: + self.text_encoder.to(execution_device) + + prompts = [prompt] if isinstance(prompt, str) else prompt + image_token_offset = 157184 + frontend_scale = max(max(height, width) / 512, 1.0) + frontend_height = int(height / frontend_scale) + frontend_width = int(width / frontend_scale) + vq_height = frontend_height // 16 + vq_width = frontend_width // 16 + image_token_count = vq_height * vq_width + system_prompt = "You are a text-to-image generation assistant." + generated_tokens = [] + + try: + for prompt in prompts: + text_prompt = f"SYSTEM {system_prompt} HUMAN{prompt}ASSISTANT" + text_ids = self.tokenizer(text_prompt).input_ids + image_info_ids = self.tokenizer( + f"<|image|><|reserved_token_{vq_height}|><|reserved_token_{vq_width}|><|/image|>" + ).input_ids + input_ids = text_ids + image_info_ids[:-1] + + uncond_prompt = ( + f"SYSTEM {system_prompt} HUMANASSISTANT" + ) + uncond_ids = self.tokenizer(uncond_prompt).input_ids + image_info_ids[:-1] + output_ids = self.text_encoder.generate_bd_image_logic( + data={ + "input_ids": torch.tensor( + input_ids, device=self.text_encoder.get_input_embeddings().weight.device + ).unsqueeze(0), + "uncond_ids": uncond_ids, + }, + block_length=32, + steps=8, + gen_length=image_token_count, + cfg_scale=2.0, + ) + token_ids = output_ids[0, len(input_ids) : len(input_ids) + image_token_count] - image_token_offset + if len(token_ids) != image_token_count: + raise ValueError(f"The MLLM generated {len(token_ids)} VQ tokens, expected {image_token_count}.") + if torch.any((token_ids < 0) | (token_ids >= self.sigvq.config.codebook_size)): + raise ValueError("The MLLM generated token IDs outside the SigVQ codebook.") + generated_tokens.append(token_ids) + finally: + if restore_text_encoder_device: + self.text_encoder.to(text_encoder_device) + + return torch.stack(generated_tokens) + + def check_inputs( + self, + prompt: str | list[str] | None, + image: PipelineImageInput | None, + generation_mode: str, + height: int, + width: int, + num_images_per_prompt: int, + prompt_embeds: torch.Tensor | None, + prompt_attention_mask: torch.Tensor | None, + negative_prompt_embeds: torch.Tensor | None, + negative_prompt_attention_mask: torch.Tensor | None, + callback_on_step_end_tensor_inputs: list[str], + num_inference_steps: int, + ) -> None: + if generation_mode not in {"text", "vq", "editing"}: + raise ValueError("`generation_mode` must be one of 'text', 'vq', or 'editing'.") + if generation_mode in {"text", "vq"} and image is not None: + raise ValueError(f"`image` must be omitted when `generation_mode='{generation_mode}'`.") + if generation_mode == "vq" and prompt is None: + raise ValueError("`prompt` is required when `generation_mode='vq'`.") + if generation_mode == "editing" and image is None: + raise ValueError("`image` is required when `generation_mode='editing'`.") + if generation_mode == "vq" and (height % 16 != 0 or width % 16 != 0): + raise ValueError("`height` and `width` must be divisible by 16 in VQ mode.") + + required_multiple = self.latent_scale_factor * (2 if generation_mode == "editing" else 1) + if height <= 0 or width <= 0 or height % required_multiple != 0 or width % required_multiple != 0: + raise ValueError(f"`height` and `width` must be divisible by {required_multiple}.") + if num_inference_steps < 1: + raise ValueError("`num_inference_steps` must be at least 1.") + if prompt is None and prompt_embeds is None: + raise ValueError("Provide either `prompt` or `prompt_embeds`.") + if prompt is not None and prompt_embeds is not None: + raise ValueError("Provide only one of `prompt` or `prompt_embeds`.") + if prompt_embeds is not None and prompt_attention_mask is None: + raise ValueError("`prompt_attention_mask` is required with `prompt_embeds`.") + if negative_prompt_embeds is not None and negative_prompt_attention_mask is None: + raise ValueError("`negative_prompt_attention_mask` is required with `negative_prompt_embeds`.") + if num_images_per_prompt < 1: + raise ValueError("`num_images_per_prompt` must be at least 1.") + if not all(name in self._callback_tensor_inputs for name in callback_on_step_end_tensor_inputs): + raise ValueError( + f"`callback_on_step_end_tensor_inputs` must be chosen from {self._callback_tensor_inputs}." + ) + + def _encode_source_image( + self, + image: PipelineImageInput, + height: int, + width: int, + batch_size: int, + num_images_per_prompt: int, + ) -> tuple[torch.Tensor, torch.Tensor]: + image = self.image_processor.preprocess(image, height=height, width=width) + if image.shape[0] == 1 and batch_size > 1: + image = image.repeat(batch_size, 1, 1, 1) + if image.shape[0] != batch_size: + raise ValueError(f"The image batch size must be 1 or {batch_size}, but is {image.shape[0]}.") + image = image.repeat_interleave(num_images_per_prompt, dim=0) + + sigvq_pixel_values = F.interpolate( + image.float(), + size=(height // 2, width // 2), + mode="bilinear", + align_corners=False, + ) + semantic_features = self.sigvq( + sigvq_pixel_values.to(device=self._execution_device, dtype=self.sigvq.dtype) + ).semantic_features + + source_latents = self.vae.encode( + image.to(device=self._execution_device, dtype=self.vae.dtype) + ).latent_dist.mode() + source_latents = self._patchify_latents(source_latents) + latent_mean = self.vae.bn.running_mean.view(1, -1, 1, 1).to(source_latents) + latent_std = torch.sqrt(self.vae.bn.running_var.view(1, -1, 1, 1) + self.vae.config.batch_norm_eps).to( + source_latents + ) + source_latents = (source_latents - latent_mean) / latent_std + return source_latents, semantic_features + + @torch.no_grad() + def __call__( + self, + prompt: str | list[str] | None = None, + image: PipelineImageInput | None = None, + generation_mode: str = "text", + negative_prompt: str | list[str] | None = None, + height: int = 1024, + width: int = 1024, + num_inference_steps: int = 20, + guidance_scale: float = 4.5, + num_images_per_prompt: int = 1, + generator: torch.Generator | list[torch.Generator] | None = None, + latents: torch.Tensor | None = None, + prompt_embeds: torch.Tensor | None = None, + prompt_attention_mask: torch.Tensor | None = None, + negative_prompt_embeds: torch.Tensor | None = None, + negative_prompt_attention_mask: torch.Tensor | None = None, + max_sequence_length: int = 2048, + output_type: str = "pil", + return_dict: bool = True, + callback_on_step_end: Callable[["LLaDAImagePipeline", int, torch.Tensor, dict], dict] | None = None, + callback_on_step_end_tensor_inputs: list[str] = ["latents"], + ) -> LLaDAImagePipelineOutput | tuple: + r""" + Generate images using text-only, VQ-conditioned, or editing inference. + + The timestep schedule is selected by the scheduler configuration. `use_uniform_sigmas=True` uses a uniform + pre-shift grid; otherwise the source Kumaraswamy schedule is used. + + Args: + prompt (`str` or `list[str]`, *optional*): + Text prompts that describe the generated image or requested edit. + image (`PipelineImageInput`, *optional*): + Reference image or image batch. Required in `"editing"` mode and rejected in other modes. + generation_mode (`str`, defaults to `"text"`): + Inference path. `"text"` uses only the text prompt. `"vq"` uses the MLLM to generate VQ tokens from the + prompt at a maximum frontend resolution of 512 before diffusion. `"editing"` uses both reference-image + SigVQ features and source-image latents. + negative_prompt (`str` or `list[str]`, *optional*): + Text excluded from generation. The checkpoint's empty CFG prompt is used by default. + height (`int`, defaults to `1024`): + Output image height. + width (`int`, defaults to `1024`): + Output image width. + num_inference_steps (`int`, defaults to `20`): + Number of flow-matching denoising steps. + guidance_scale (`float`, defaults to `4.5`): + Classifier-free guidance scale. Guidance is disabled at values up to `1.0`. + num_images_per_prompt (`int`, defaults to `1`): + Number of images generated per prompt. + generator (`torch.Generator` or `list[torch.Generator]`, *optional*): + Random generator or generator batch used to create the initial latents. + latents (`torch.Tensor`, *optional*): + Pre-generated patchified Flux2 latents. + prompt_embeds (`torch.Tensor`, *optional*): + Precomputed, projected positive prompt embeddings. + prompt_attention_mask (`torch.Tensor`, *optional*): + Valid-token mask for `prompt_embeds`. + negative_prompt_embeds (`torch.Tensor`, *optional*): + Precomputed, projected negative prompt embeddings. + negative_prompt_attention_mask (`torch.Tensor`, *optional*): + Valid-token mask for `negative_prompt_embeds`. + max_sequence_length (`int`, defaults to `2048`): + Maximum text sequence length before the QueryFormer tokens are appended. + output_type (`str`, defaults to `"pil"`): + Output format. Choose `"pil"`, `"np"`, `"pt"`, or `"latent"`. + return_dict (`bool`, defaults to `True`): + Whether to return [`LLaDAImagePipelineOutput`] instead of a tuple. + callback_on_step_end (`Callable`, *optional*): + Function called after each denoising step. + callback_on_step_end_tensor_inputs (`list[str]`, defaults to `["latents"]`): + Tensor names forwarded to `callback_on_step_end`. + + Returns: + [`LLaDAImagePipelineOutput`] or `tuple`: + Generated images or final patchified latents. + """ + self.check_inputs( + prompt, + image, + generation_mode, + height, + width, + num_images_per_prompt, + prompt_embeds, + prompt_attention_mask, + negative_prompt_embeds, + negative_prompt_attention_mask, + callback_on_step_end_tensor_inputs, + num_inference_steps, + ) + + if prompt_embeds is not None: + batch_size = prompt_embeds.shape[0] + elif isinstance(prompt, str): + batch_size = 1 + else: + batch_size = len(prompt) + device = self._execution_device + self._guidance_scale = guidance_scale + do_classifier_free_guidance = guidance_scale > 1.0 + + prompt_embeds, prompt_attention_mask, negative_prompt_embeds, negative_prompt_attention_mask = ( + self.encode_prompt( + prompt, + negative_prompt, + do_classifier_free_guidance, + num_images_per_prompt, + prompt_embeds, + prompt_attention_mask, + negative_prompt_embeds, + negative_prompt_attention_mask, + max_sequence_length, + device, + ) + ) + effective_batch_size = batch_size * num_images_per_prompt + + source_latents = None + semantic_features = None + if generation_mode == "vq": + vq_token_ids = self.generate_vq_tokens(prompt, height, width) + vq_token_ids = vq_token_ids.repeat_interleave(num_images_per_prompt, dim=0) + semantic_features = self.sigvq(token_ids=vq_token_ids.to(self._execution_device)).semantic_features + elif generation_mode == "editing": + source_latents, semantic_features = self._encode_source_image( + image, + height, + width, + batch_size, + num_images_per_prompt, + ) + + latent_shape = ( + effective_batch_size, + self.transformer.config.in_channels, + height // self.latent_scale_factor, + width // self.latent_scale_factor, + ) + if latents is None: + latents = randn_tensor(latent_shape, generator=generator, device=device, dtype=torch.float32) + latents = latents.to(self.transformer.dtype).float() + else: + if latents.shape != latent_shape: + raise ValueError(f"Expected `latents` to have shape {latent_shape}, got {tuple(latents.shape)}.") + latents = latents.to(device=device, dtype=torch.float32) + + if self.scheduler.config.get("use_uniform_sigmas", False): + # diffusers 0.39.0 does not natively support this scheduler option. Supplying the pre-shift grid + # explicitly preserves the behavior of the patched scheduler used by LLaDA-Image-SGLang. + sigmas = torch.linspace(1.0, 0.0, num_inference_steps + 1, dtype=torch.float32)[:-1].tolist() + self.scheduler.set_timesteps(sigmas=sigmas, device=device) + else: + schedule_steps = num_inference_steps + 1 + schedule = torch.linspace(0.001, 1.0, schedule_steps, dtype=torch.float64)[:-1] + schedule = (1 - (1 - schedule**1.17) ** 0.8) ** 1.1 + sigmas = (1 - schedule).tolist() + self.scheduler.set_timesteps(sigmas=sigmas, device=device) + timesteps = self.scheduler.timesteps + self._num_timesteps = len(timesteps) + + cond_cap_feats = [ + embeds[mask].to(device=device, dtype=self.transformer.dtype) + for embeds, mask in zip(prompt_embeds, prompt_attention_mask.bool()) + ] + if do_classifier_free_guidance: + uncond_cap_feats = [ + embeds[mask].to(device=device, dtype=self.transformer.dtype) + for embeds, mask in zip(negative_prompt_embeds, negative_prompt_attention_mask.bool()) + ] + cap_feats = cond_cap_feats + uncond_cap_feats + else: + cap_feats = cond_cap_feats + + glm_cap_feats = None + source_latent_list = None + if semantic_features is not None: + cond_glm_cap_feats = [ + features.to(device=device, dtype=self.transformer.dtype) for features in semantic_features + ] + if source_latents is not None: + source_latent_list = [ + latent.unsqueeze(1).to(device=device, dtype=self.transformer.dtype) for latent in source_latents + ] + if do_classifier_free_guidance: + empty_glm = semantic_features.new_zeros((0, semantic_features.shape[-1])).to( + device=device, dtype=self.transformer.dtype + ) + glm_cap_feats = cond_glm_cap_feats + [empty_glm] * effective_batch_size + if source_latent_list is not None: + source_latent_list = source_latent_list + source_latent_list + else: + glm_cap_feats = cond_glm_cap_feats + + with self.progress_bar(total=num_inference_steps) as progress_bar: + for step_index, timestep in enumerate(timesteps): + latent_model_input = torch.cat([latents, latents], dim=0) if do_classifier_free_guidance else latents + latent_list = [latent.unsqueeze(1).to(self.transformer.dtype) for latent in latent_model_input] + model_timestep = (timestep / self.scheduler.config.num_train_timesteps).expand( + latent_model_input.shape[0] + ) + + model_output = self.transformer( + x=latent_list, + t=model_timestep.to(self.transformer.dtype), + cap_feats=cap_feats, + glm_cap_feats=glm_cap_feats, + source_latents=source_latent_list, + ).sample + model_output = -torch.stack(model_output, dim=0).squeeze(2).float() + + if do_classifier_free_guidance: + conditional_output, unconditional_output = model_output.chunk(2) + model_output = unconditional_output + self.guidance_scale * ( + conditional_output - unconditional_output + ) + + latents = self.scheduler.step(model_output, timestep, latents, return_dict=False)[0] + + if callback_on_step_end is not None: + callback_kwargs = {} + for name in callback_on_step_end_tensor_inputs: + callback_kwargs[name] = locals()[name] + callback_outputs = callback_on_step_end(self, step_index, timestep, callback_kwargs) + latents = callback_outputs.pop("latents", latents) + + progress_bar.update() + + if output_type == "latent": + images = latents + else: + latents = latents.to(device=device, dtype=self.vae.dtype) + latent_mean = self.vae.bn.running_mean.view(1, -1, 1, 1).to(latents) + latent_std = torch.sqrt(self.vae.bn.running_var.view(1, -1, 1, 1) + self.vae.config.batch_norm_eps).to( + latents + ) + latents = latents * latent_std + latent_mean + latents = self._unpatchify_latents(latents) + images = self.vae.decode(latents, return_dict=False)[0] + images = self.image_processor.postprocess(images, output_type=output_type) + + self.maybe_free_model_hooks() + if not return_dict: + return (images,) + return LLaDAImagePipelineOutput(images=images) diff --git a/src/diffusers/pipelines/llada_image/pipeline_output.py b/src/diffusers/pipelines/llada_image/pipeline_output.py new file mode 100644 index 000000000000..f79efaa726a8 --- /dev/null +++ b/src/diffusers/pipelines/llada_image/pipeline_output.py @@ -0,0 +1,34 @@ +# Copyright 2026 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from dataclasses import dataclass + +import numpy as np +import PIL.Image +import torch + +from diffusers.utils import BaseOutput + + +@dataclass +class LLaDAImagePipelineOutput(BaseOutput): + """ + Output class for the LLaDA-Image pipeline. + + Args: + images (`list[PIL.Image.Image]`, `np.ndarray`, or `torch.Tensor`): + Generated images. The format is controlled by the pipeline's `output_type` argument. + """ + + images: list[PIL.Image.Image] | np.ndarray | torch.Tensor diff --git a/src/diffusers/pipelines/pipeline_utils.py b/src/diffusers/pipelines/pipeline_utils.py index 82e6c4c2aff4..441e9cc1e6d8 100644 --- a/src/diffusers/pipelines/pipeline_utils.py +++ b/src/diffusers/pipelines/pipeline_utils.py @@ -1735,8 +1735,8 @@ def download(cls, pretrained_model_name, **kwargs) -> str | os.PathLike: # allow all patterns from non-model folders # this enables downloading schedulers, tokenizers, ... allow_patterns += [f"{k}/*" for k in folder_names if k not in model_folder_names] - # add custom component files - allow_patterns += [f"{k}/{f}.py" for k, f in custom_components.items()] + # Add custom component modules and their local Python dependencies. + allow_patterns += [f"{folder_name}/*.py" for folder_name in custom_components] # add custom pipeline file allow_patterns += [f"{custom_pipeline}.py"] if f"{custom_pipeline}.py" in filenames else [] # also allow downloading config.json files with the model diff --git a/src/diffusers/utils/dummy_pt_objects.py b/src/diffusers/utils/dummy_pt_objects.py index 3434c6416cce..012dc7873374 100644 --- a/src/diffusers/utils/dummy_pt_objects.py +++ b/src/diffusers/utils/dummy_pt_objects.py @@ -1669,6 +1669,66 @@ def from_pretrained(cls, *args, **kwargs): requires_backends(cls, ["torch"]) +class LLaDAImageQueryFormerModel(metaclass=DummyObject): + _backends = ["torch"] + + def __init__(self, *args, **kwargs): + requires_backends(self, ["torch"]) + + @classmethod + def from_config(cls, *args, **kwargs): + requires_backends(cls, ["torch"]) + + @classmethod + def from_pretrained(cls, *args, **kwargs): + requires_backends(cls, ["torch"]) + + +class LLaDAImageSigVQModel(metaclass=DummyObject): + _backends = ["torch"] + + def __init__(self, *args, **kwargs): + requires_backends(self, ["torch"]) + + @classmethod + def from_config(cls, *args, **kwargs): + requires_backends(cls, ["torch"]) + + @classmethod + def from_pretrained(cls, *args, **kwargs): + requires_backends(cls, ["torch"]) + + +class LLaDAImageTextProjectionModel(metaclass=DummyObject): + _backends = ["torch"] + + def __init__(self, *args, **kwargs): + requires_backends(self, ["torch"]) + + @classmethod + def from_config(cls, *args, **kwargs): + requires_backends(cls, ["torch"]) + + @classmethod + def from_pretrained(cls, *args, **kwargs): + requires_backends(cls, ["torch"]) + + +class LLaDAImageTransformer2DModel(metaclass=DummyObject): + _backends = ["torch"] + + def __init__(self, *args, **kwargs): + requires_backends(self, ["torch"]) + + @classmethod + def from_config(cls, *args, **kwargs): + requires_backends(cls, ["torch"]) + + @classmethod + def from_pretrained(cls, *args, **kwargs): + requires_backends(cls, ["torch"]) + + class LongCatAudioDiTTransformer(metaclass=DummyObject): _backends = ["torch"] diff --git a/src/diffusers/utils/dummy_torch_and_transformers_objects.py b/src/diffusers/utils/dummy_torch_and_transformers_objects.py index ed724e7de751..60654020c140 100644 --- a/src/diffusers/utils/dummy_torch_and_transformers_objects.py +++ b/src/diffusers/utils/dummy_torch_and_transformers_objects.py @@ -3092,6 +3092,36 @@ def from_pretrained(cls, *args, **kwargs): requires_backends(cls, ["torch", "transformers"]) +class LLaDAImagePipeline(metaclass=DummyObject): + _backends = ["torch", "transformers"] + + def __init__(self, *args, **kwargs): + requires_backends(self, ["torch", "transformers"]) + + @classmethod + def from_config(cls, *args, **kwargs): + requires_backends(cls, ["torch", "transformers"]) + + @classmethod + def from_pretrained(cls, *args, **kwargs): + requires_backends(cls, ["torch", "transformers"]) + + +class LLaDAImagePipelineOutput(metaclass=DummyObject): + _backends = ["torch", "transformers"] + + def __init__(self, *args, **kwargs): + requires_backends(self, ["torch", "transformers"]) + + @classmethod + def from_config(cls, *args, **kwargs): + requires_backends(cls, ["torch", "transformers"]) + + @classmethod + def from_pretrained(cls, *args, **kwargs): + requires_backends(cls, ["torch", "transformers"]) + + class LongCatAudioDiTPipeline(metaclass=DummyObject): _backends = ["torch", "transformers"] diff --git a/tests/models/transformers/test_models_transformer_llada_image.py b/tests/models/transformers/test_models_transformer_llada_image.py new file mode 100644 index 000000000000..7f8a781520b4 --- /dev/null +++ b/tests/models/transformers/test_models_transformer_llada_image.py @@ -0,0 +1,239 @@ +# coding=utf-8 +# Copyright 2026 HuggingFace Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import pytest +import torch + +from diffusers import ( + LLaDAImageQueryFormerModel, + LLaDAImageSigVQModel, + LLaDAImageTextProjectionModel, + LLaDAImageTransformer2DModel, +) +from diffusers.utils.torch_utils import randn_tensor + +from ...testing_utils import assert_tensors_close, enable_full_determinism, torch_device +from ..testing_utils import ( + AttentionTesterMixin, + BaseModelTesterConfig, + MemoryTesterMixin, + ModelTesterMixin, + TorchCompileTesterMixin, + TrainingTesterMixin, +) + + +enable_full_determinism() + + +def _flatten_list_output(output: list[torch.Tensor]) -> torch.Tensor: + return torch.cat([sample.flatten() for sample in output]) + + +class LLaDAImageTransformerTesterConfig(BaseModelTesterConfig): + @property + def model_class(self): + return LLaDAImageTransformer2DModel + + @property + def pretrained_model_name_or_path(self): + return "inclusionAI/LLaDA-Image" + + @property + def pretrained_model_kwargs(self): + return {"subfolder": "transformer"} + + @property + def generator(self): + return torch.Generator("cpu").manual_seed(0) + + @property + def main_input_name(self) -> str: + return "x" + + @property + def model_split_percents(self) -> list[float]: + return [0.9, 0.9, 0.9] + + def get_init_dict(self) -> dict[str, int | list[int]]: + # __init__ parameters: + # all_patch_size: tuple[int, Ellipsis] = + # all_f_patch_size: tuple[int, Ellipsis] = + # in_channels: int = 128 + # dim: int = 3840 + # n_layers: int = 30 + # n_refiner_layers: int = 2 + # n_heads: int = 30 + # norm_eps: float = 1e-05 + # qk_norm: bool = True + # cap_feat_dim: int = 2560 + # semantic_feat_dim: int = 4096 + # rope_theta: float = 256.0 + # t_scale: float = 1000.0 + # axes_dims: tuple[int, Ellipsis] = + # axes_lens: tuple[int, Ellipsis] = + return { + "in_channels": 8, + "dim": 32, + "n_layers": 1, + "n_refiner_layers": 1, + "n_heads": 2, + "cap_feat_dim": 24, + "semantic_feat_dim": 20, + "axes_dims": (4, 6, 6), + "axes_lens": (2048, 32, 32), + } + + def get_dummy_inputs(self) -> dict[str, torch.Tensor]: + # forward() parameters: + # x: list[torch.Tensor] + # t: torch.Tensor + # cap_feats: list[torch.Tensor] | None + # glm_cap_feats: list[torch.Tensor] | None + # source_latents: list[torch.Tensor] | None + # patch_size: int = 1 + # f_patch_size: int = 1 + # return_dict: bool = True + return self.get_inputs_with_shapes(4, 4) + + def get_inputs_with_shapes(self, height: int, width: int) -> dict[str, torch.Tensor]: + return { + "x": [ + randn_tensor((8, 1, height, width), generator=self.generator, device=torch_device) for _ in range(2) + ], + "t": torch.tensor([0.8, 0.8], device=torch_device), + "cap_feats": [ + randn_tensor((length, 24), generator=self.generator, device=torch_device) for length in (5, 7) + ], + } + + @property + def input_shape(self) -> tuple[int, ...]: + return (2, 8, 1, 4, 4) + + @property + def output_shape(self) -> tuple[int, ...]: + return (8, 1, 4, 4) + + +class TestLLaDAImageTransformerModel(LLaDAImageTransformerTesterConfig, ModelTesterMixin): + @torch.no_grad() + def test_determinism(self, atol=1e-5, rtol=0): + model = self.model_class(**self.get_init_dict()).to(torch_device).eval() + inputs = self.get_dummy_inputs() + first = _flatten_list_output(model(**inputs, return_dict=False)[0]) + second = _flatten_list_output(model(**inputs, return_dict=False)[0]) + mask = ~(torch.isnan(first) | torch.isnan(second)) + assert_tensors_close(first[mask], second[mask], atol=atol, rtol=rtol) + + @pytest.mark.skip("The model returns a list so it can preserve per-sample spatial shapes.") + def test_outputs_equivalence(self, atol=1e-5, rtol=0): + pass + + +class TestLLaDAImageTransformerMemory(LLaDAImageTransformerTesterConfig, MemoryTesterMixin): + @pytest.mark.skip("The shared training test does not support list-valued main inputs.") + def test_layerwise_casting_training(self): + pass + + +class TestLLaDAImageTransformerTorchCompile(LLaDAImageTransformerTesterConfig, TorchCompileTesterMixin): + @property + def different_shapes_for_compilation(self): + return [(4, 4), (4, 8), (8, 8)] + + def get_dummy_inputs(self, height: int = 4, width: int = 4) -> dict[str, torch.Tensor]: + return self.get_inputs_with_shapes(height, width) + + def test_torch_compile_repeated_blocks(self): + # The same block class is used by two refiners and the denoising stack with different attention processors. + super().test_torch_compile_repeated_blocks(recompile_limit=3) + + @pytest.mark.skip("AOTInductor package loading does not support this model's list-valued inputs.") + def test_compile_works_with_aot(self, tmp_path): + pass + + +class TestLLaDAImageTransformerTraining(LLaDAImageTransformerTesterConfig, TrainingTesterMixin): + def test_gradient_checkpointing_is_applied(self): + super().test_gradient_checkpointing_is_applied(expected_set={"LLaDAImageTransformer2DModel"}) + + @pytest.mark.skip("The shared training test does not support list-valued model outputs.") + def test_training(self): + pass + + @pytest.mark.skip("The shared training test does not support list-valued model outputs.") + def test_training_with_ema(self): + pass + + @pytest.mark.skip("The shared mixed-precision test does not support list-valued model outputs.") + def test_mixed_precision_training(self): + pass + + @pytest.mark.skip("The shared gradient comparison does not support list-valued model outputs.") + def test_gradient_checkpointing_equivalence(self, loss_tolerance=1e-5, param_grad_tol=5e-5, skip=None): + pass + + +class TestLLaDAImageTransformerAttention(LLaDAImageTransformerTesterConfig, AttentionTesterMixin): + pass + + +class TestLLaDAImageAuxiliaryModels: + def test_queryformer(self): + torch.manual_seed(0) + model = LLaDAImageQueryFormerModel( + num_queries=4, + hidden_size=16, + num_hidden_layers=1, + num_attention_heads=4, + intermediate_size=32, + ).to(torch_device) + hidden_states = torch.randn(2, 5, 16, device=torch_device) + attention_mask = torch.tensor([[1, 1, 1, 1, 1], [1, 1, 1, 0, 0]], device=torch_device) + output = model(hidden_states, attention_mask).query_embeds + assert output.shape == (2, 4, 16) + + def test_text_projection(self): + torch.manual_seed(0) + model = LLaDAImageTextProjectionModel( + hidden_size=16, + intermediate_size=32, + num_hidden_layers=1, + num_attention_heads=4, + projection_dim=24, + ).to(torch_device) + output = model(torch.randn(2, 7, 16, device=torch_device)).hidden_states + assert output.shape == (2, 7, 24) + + def test_sigvq_image_and_tokens(self): + torch.manual_seed(0) + model = LLaDAImageSigVQModel( + image_size=32, + patch_size=8, + hidden_size=16, + intermediate_size=32, + num_hidden_layers=1, + num_attention_heads=4, + codebook_size=32, + codebook_embed_dim=8, + semantic_embed_dim=20, + ).to(torch_device) + image_output = model(pixel_values=torch.randn(2, 3, 32, 32, device=torch_device)) + token_output = model(token_ids=image_output.token_ids) + + assert image_output.semantic_features.shape == (2, 16, 20) + assert image_output.token_ids.shape == (2, 16) + torch.testing.assert_close(token_output.semantic_features, image_output.semantic_features) diff --git a/tests/pipelines/llada_image/__init__.py b/tests/pipelines/llada_image/__init__.py new file mode 100644 index 000000000000..8b137891791f --- /dev/null +++ b/tests/pipelines/llada_image/__init__.py @@ -0,0 +1 @@ + diff --git a/tests/pipelines/llada_image/test_pipeline_llada_image.py b/tests/pipelines/llada_image/test_pipeline_llada_image.py new file mode 100644 index 000000000000..4a45ef8ef2be --- /dev/null +++ b/tests/pipelines/llada_image/test_pipeline_llada_image.py @@ -0,0 +1,201 @@ +# Copyright 2026 HuggingFace Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import types + +import pytest +import torch +from transformers import AutoTokenizer, LlamaConfig, LlamaForCausalLM + +from diffusers import ( + AutoencoderKLFlux2, + FlowMatchEulerDiscreteScheduler, + LLaDAImagePipeline, + LLaDAImageQueryFormerModel, + LLaDAImageSigVQModel, + LLaDAImageTextProjectionModel, + LLaDAImageTransformer2DModel, +) +from diffusers.utils import is_transformers_version + +from ...testing_utils import torch_device +from ..testing_utils import BasePipelineTesterConfig, MemoryTesterMixin, PipelineTesterMixin + + +class LLaDAImagePipelineTesterConfig(BasePipelineTesterConfig): + pipeline_class = LLaDAImagePipeline + required_input_params_in_call_signature = frozenset( + ["prompt", "image", "generation_mode", "height", "width", "guidance_scale", "prompt_embeds"] + ) + batch_input_params = frozenset(["prompt"]) + text_stack_component_names = ["text_encoder", "tokenizer", "queryformer", "text_projection"] + group_offloading_leaf_level_exclude_modules = ["text_encoder"] + group_offloading_block_level_exclude_modules = ["vae", "text_encoder"] + output_shape = (3, 8, 8) + + def get_dummy_components(self): + torch.manual_seed(0) + transformer = LLaDAImageTransformer2DModel( + in_channels=4, + dim=32, + n_layers=1, + n_refiner_layers=1, + n_heads=2, + cap_feat_dim=24, + semantic_feat_dim=20, + axes_dims=(4, 6, 6), + axes_lens=(2048, 32, 32), + ) + torch.manual_seed(0) + queryformer = LLaDAImageQueryFormerModel( + num_queries=4, + hidden_size=16, + num_hidden_layers=1, + num_attention_heads=4, + intermediate_size=32, + ) + torch.manual_seed(0) + text_projection = LLaDAImageTextProjectionModel( + hidden_size=16, + intermediate_size=32, + num_hidden_layers=1, + num_attention_heads=4, + projection_dim=24, + ) + torch.manual_seed(0) + sigvq = LLaDAImageSigVQModel( + image_size=8, + patch_size=2, + hidden_size=16, + intermediate_size=32, + num_hidden_layers=1, + num_attention_heads=4, + codebook_size=32, + codebook_embed_dim=8, + semantic_embed_dim=20, + ) + torch.manual_seed(0) + vae = AutoencoderKLFlux2( + sample_size=8, + in_channels=3, + out_channels=3, + down_block_types=("DownEncoderBlock2D",), + up_block_types=("UpDecoderBlock2D",), + block_out_channels=(4,), + layers_per_block=1, + latent_channels=1, + norm_num_groups=1, + use_quant_conv=False, + use_post_quant_conv=False, + ) + torch.manual_seed(0) + text_encoder = LlamaForCausalLM( + LlamaConfig( + vocab_size=32000, + hidden_size=16, + intermediate_size=32, + num_hidden_layers=1, + num_attention_heads=4, + num_key_value_heads=2, + max_position_embeddings=128, + pad_token_id=0, + bos_token_id=1, + eos_token_id=2, + ) + ) + tokenizer = AutoTokenizer.from_pretrained("hf-internal-testing/tiny-random-LlamaForCausalLM") + + return { + "scheduler": FlowMatchEulerDiscreteScheduler(stochastic_sampling=False), + "vae": vae, + "text_encoder": text_encoder, + "tokenizer": tokenizer, + "queryformer": queryformer, + "text_projection": text_projection, + "sigvq": sigvq, + "transformer": transformer, + } + + def get_dummy_inputs(self): + return { + "prompt": "a tiny cat", + "height": 8, + "width": 8, + "num_inference_steps": 2, + "guidance_scale": 1.0, + "generator": self.get_generator(0), + "max_sequence_length": 12, + "output_type": "pt", + } + + +class TestLLaDAImagePipeline(LLaDAImagePipelineTesterConfig, PipelineTesterMixin): + def test_transformers_5_default_rope_is_materialized(self): + if not is_transformers_version(">=", "5.0.0"): + pytest.skip("This regression only affects Transformers 5 and later.") + + components = self.get_dummy_components() + rotary_emb = torch.nn.Module() + rotary_emb.config = types.SimpleNamespace( + head_dim=8, + hidden_size=16, + num_attention_heads=2, + partial_rotary_factor=1.0, + rope_theta=10000.0, + ) + rotary_emb.rope_type = "default" + rotary_emb.register_buffer("inv_freq", torch.zeros(4), persistent=False) + language_model = torch.nn.Module() + language_model.rotary_emb = rotary_emb + components["text_encoder"].model.language_model = language_model + + self.pipeline_class(**components) + + expected = torch.tensor([1.0, 0.1, 0.01, 0.001]) + torch.testing.assert_close(rotary_emb.inv_freq, expected) + torch.testing.assert_close(rotary_emb.original_inv_freq, expected) + + def test_vq_conditioned(self): + pipe = self.get_pipeline().to(torch_device) + + def generate_bd_image_logic(text_encoder, data, block_length, steps, gen_length, cfg_scale): + input_ids = data["input_ids"] + image_tokens = torch.arange(gen_length, device=input_ids.device).unsqueeze(0) + 157184 + return torch.cat([input_ids, image_tokens], dim=1) + + pipe.text_encoder.generate_bd_image_logic = types.MethodType(generate_bd_image_logic, pipe.text_encoder) + inputs = self.get_dummy_inputs() + inputs.update(generation_mode="vq", height=16, width=16, output_type="latent") + output = pipe(**inputs).images + assert output.shape == (1, 4, 8, 8) + + def test_image_editing(self): + pipe = self.get_pipeline().to(torch_device) + inputs = self.get_dummy_inputs() + inputs.update(image=torch.rand(1, 3, 8, 8), generation_mode="editing") + output = pipe(**inputs).images + assert output.shape == (1, *self.output_shape) + assert not torch.isnan(output).any() + + @pytest.mark.parametrize("generation_mode", ["unknown", "editing"]) + def test_invalid_generation_mode_inputs(self, generation_mode): + pipe = self.get_pipeline().to(torch_device) + inputs = self.get_dummy_inputs() + inputs["generation_mode"] = generation_mode + with pytest.raises(ValueError): + pipe(**inputs) + + +class TestLLaDAImagePipelineMemory(LLaDAImagePipelineTesterConfig, MemoryTesterMixin): + pass