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