diff --git a/CHANGELOG.md b/CHANGELOG.md
index 73c7025c5..b3ae5858a 100644
--- a/CHANGELOG.md
+++ b/CHANGELOG.md
@@ -16,6 +16,10 @@ All-about-optimizations:
- **Models**
- [Anima 2.9B Preview v1](https://huggingface.co/yeoj34760/Anima-2.9B)
expanded version of Anima 2B
+ - [inclusionAI LLaDA-Image](https://huggingface.co/inclusionAI/LLaDA-Image) in *base* and *turbo* variants
+ LLaDA-Image is a 6.5B transformer with massive 16.3B fully-custom MoE text-encoder and optional 1.3B SigVQ conditioning model
+ with support for text-to-image, vq-conditioned text-to-image and image-editing workflows
+ *note* model is extremely quantization sensitive so minimum allowed quant type is `uint8`
- **LoRA**
- *TODO*: see [LoRA docs](https://vladmandic.github.io/sdnext-docs/LoRA) for all of the improvements and usage instructions
*note*: lora now has its own settings section in *settings -> lora*
diff --git a/modules/modeldata.py b/modules/modeldata.py
index ee8122b15..4e7568436 100644
--- a/modules/modeldata.py
+++ b/modules/modeldata.py
@@ -119,6 +119,8 @@ def get_model_type(pipe):
model_type = 'longcat'
elif 'GlmImage' in name:
model_type = 'glmimage'
+ elif 'LLaDAImage' in name:
+ model_type = 'lladaimage'
elif 'Step1XEdit' in name:
model_type = 'step1x_edit'
elif 'JoyImageEdit' in name:
diff --git a/modules/sd_detect.py b/modules/sd_detect.py
index bf3327be1..a133c8e21 100644
--- a/modules/sd_detect.py
+++ b/modules/sd_detect.py
@@ -167,6 +167,8 @@ def guess_by_name(fn, current_guess):
new_guess = 'OvisImage'
elif 'glm-image' in fn.lower():
new_guess = 'GLMImage'
+ elif 'llada' in fn.lower():
+ new_guess = 'LLaDAImage'
elif 'sdxs-1b' in fn.lower():
new_guess = 'SDXS'
elif 'step1x-edit' in fn.lower():
diff --git a/modules/sd_models.py b/modules/sd_models.py
index 37ad35067..bbaf5d008 100644
--- a/modules/sd_models.py
+++ b/modules/sd_models.py
@@ -55,6 +55,7 @@ pipe_switch_task_exclude = [
'Kandinsky5I2IPipeline',
'GoogleNanoBananaPipeline',
'Step1XEditPipeline',
+ 'LLaDAImagePipeline',
'BooguImagePipeline',
'BooguImageTurboPipeline',
]
@@ -610,6 +611,10 @@ def load_diffuser_force(detected_model_type: str, checkpoint_info: CheckpointInf
from pipelines.model_glm import load_glm_image
sd_model = load_glm_image(checkpoint_info, diffusers_load_config)
allow_post_quant = False
+ elif model_type in ['LLaDAImage']:
+ from pipelines.model_llada import load_llada_image
+ sd_model = load_llada_image(checkpoint_info, diffusers_load_config)
+ allow_post_quant = False
elif model_type in ['SDXS']:
from pipelines.model_sdxs import load_sdxs
sd_model = load_sdxs(checkpoint_info, diffusers_load_config)
diff --git a/modules/vae/sd_vae_taesd.py b/modules/vae/sd_vae_taesd.py
index a898c3410..2cffe6041 100644
--- a/modules/vae/sd_vae_taesd.py
+++ b/modules/vae/sd_vae_taesd.py
@@ -73,7 +73,7 @@ def get_model(model_cls, variant=None):
elif model_cls in {'f1', 'h1', 'zimage', 'lumina2', 'chroma', 'longcat', 'omnigen2', 'flite', 'ovis', 'kandinsky5', 'glmimage', 'cogview3', 'cogview4', 'ultraflux'}:
model_cls = 'f1'
variant = 'TAE FLUX.1'
- elif model_cls in {'f2', 'ernieimage', 'lens', 'ideogram4'}:
+ elif model_cls in {'f2', 'ernieimage', 'lens', 'ideogram4', 'lladaimage'}:
model_cls = 'f2'
variant = 'TAE FLUX.2'
elif model_cls in {'sd3'}:
diff --git a/pipelines/llada/__init__.py b/pipelines/llada/__init__.py
new file mode 100644
index 000000000..af8b8951d
--- /dev/null
+++ b/pipelines/llada/__init__.py
@@ -0,0 +1,8 @@
+from .pipeline_llada_image import LLaDAImagePipeline
+from .pipeline_output import LLaDAImagePipelineOutput
+from .transformer_llada_image import LLaDAImageQueryFormerModel
+from .transformer_llada_image import LLaDAImageSigVQModel
+from .transformer_llada_image import LLaDAImageTextProjectionModel
+from .transformer_llada_image import LLaDAImageTransformer2DModel
+
+__all__ = ['LLaDAImagePipeline', 'LLaDAImagePipelineOutput', 'LLaDAImageQueryFormerModel', 'LLaDAImageSigVQModel', 'LLaDAImageTextProjectionModel', 'LLaDAImageTransformer2DModel']
diff --git a/pipelines/llada/configuration_llada2uni_moe.py b/pipelines/llada/configuration_llada2uni_moe.py
new file mode 100644
index 000000000..fdd7a8845
--- /dev/null
+++ b/pipelines/llada/configuration_llada2uni_moe.py
@@ -0,0 +1,133 @@
+# Copyright 2025 Antgroup and The HuggingFace Inc. 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.
+"""LLaDA2 MoE model configuration."""
+
+from transformers.configuration_utils import PretrainedConfig
+
+
+class LLaDA2MoeConfig(PretrainedConfig):
+ r"""
+ Configuration class for the LLaDA2 MoE model.
+
+ ```python
+ >>> from configuration_llada2uni_moe import LLaDA2MoeConfig
+ >>> config = LLaDA2MoeConfig()
+ ```
+ """
+
+ # Keep the original value because it selects the fused-expert implementation.
+ model_type = "llada2_moe_veomni"
+
+ def __init__(
+ self,
+ vocab_size=30592,
+ hidden_size=1024,
+ intermediate_size=None,
+ num_hidden_layers=24,
+ num_attention_heads=16,
+ num_key_value_heads=0,
+ head_dim=None,
+ hidden_act="silu",
+ use_qkv_bias=False,
+ use_qk_norm=True,
+ use_bias=True,
+ rms_norm_eps=1e-05,
+ tie_word_embeddings=False,
+ attention_dropout=0.1,
+ initializer_range=0.02,
+ max_position_embeddings=16384,
+ rope_theta=10000.0,
+ rope_parameters=None,
+ rope_scaling=None,
+ partial_rotary_factor=0.5,
+ use_cache=True,
+ sliding_window=None,
+ pad_token_id=126081,
+ # Image
+ image_token_offset=157184,
+ # MoE
+ num_experts=16,
+ num_shared_experts=0,
+ num_experts_per_tok=2,
+ n_group=8,
+ topk_group=4,
+ routed_scaling_factor=2.5,
+ moe_router_enable_expert_bias=True,
+ norm_topk_prob=True,
+ router_dtype="fp32",
+ score_function="sigmoid",
+ moe_intermediate_size=None,
+ first_k_dense_replace=0,
+ output_router_logits=False,
+ **kwargs,
+ ):
+ self.vocab_size = vocab_size
+ self.hidden_size = hidden_size
+ self.intermediate_size = intermediate_size
+ self.num_hidden_layers = num_hidden_layers
+ self.num_attention_heads = num_attention_heads
+ self.num_key_value_heads = num_key_value_heads
+ self.head_dim = head_dim or hidden_size // num_attention_heads
+ self.hidden_act = hidden_act
+ self.use_qkv_bias = use_qkv_bias
+ self.use_qk_norm = use_qk_norm
+ self.use_bias = use_bias
+ self.rms_norm_eps = rms_norm_eps
+ self.attention_dropout = attention_dropout
+ self.initializer_range = initializer_range
+ self.max_position_embeddings = max_position_embeddings
+ self.rope_theta = rope_theta
+ self.rope_scaling = rope_scaling
+ self.partial_rotary_factor = partial_rotary_factor
+ self.use_cache = use_cache
+ self.sliding_window = sliding_window
+
+ # Image token offset: VQ codebook indices are shifted by this amount in the vocabulary
+ self.image_token_offset = image_token_offset
+
+ # RoPE parameters dict — used by LLaDA2MoeRotaryEmbedding
+ if rope_parameters is None:
+ rope_parameters = {
+ "rope_type": "default",
+ "rope_theta": rope_theta,
+ "partial_rotary_factor": partial_rotary_factor,
+ }
+ self.rope_parameters = rope_parameters
+
+ # MoE
+ self.num_experts = num_experts
+ self.num_shared_experts = num_shared_experts
+ self.num_experts_per_tok = num_experts_per_tok
+ self.n_group = n_group
+ self.topk_group = topk_group
+ self.routed_scaling_factor = routed_scaling_factor
+ self.moe_router_enable_expert_bias = moe_router_enable_expert_bias
+ self.norm_topk_prob = norm_topk_prob
+ self.router_dtype = router_dtype
+ self.score_function = score_function
+ self.moe_intermediate_size = moe_intermediate_size
+ self.first_k_dense_replace = first_k_dense_replace
+ self.output_router_logits = output_router_logits
+
+ # FP8 quantization flag — set to True to use FP8Linear for experts
+ self.use_fp8_experts = kwargs.pop("use_fp8_experts", False)
+
+ super().__init__(
+ pad_token_id=pad_token_id,
+ tie_word_embeddings=tie_word_embeddings,
+ **kwargs,
+ )
+
+
+__all__ = ["LLaDA2MoeConfig"]
diff --git a/pipelines/llada/fused_moe_ops.py b/pipelines/llada/fused_moe_ops.py
new file mode 100644
index 000000000..30296fcbe
--- /dev/null
+++ b/pipelines/llada/fused_moe_ops.py
@@ -0,0 +1,375 @@
+# Copyright 2025 Bytedance Ltd. and/or its affiliates
+#
+# 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.
+
+"""Standalone, inference-only VeOmni v0.1.0 fused-MoE compatibility shim.
+
+This module preserves the ``veomni.ops.fused_moe_forward`` call signature used
+by VeOmni v0.1.0 while removing VeOmni's training, Expert Parallelism (EP), NPU,
+and Seed-kernel dependencies. It is intended for single-device inference only.
+
+The CUDA fast path uses a small Triton grouped-linear kernel. If Triton is not
+available, the tensors are not on CUDA, or ``LLADA_MOE_BACKEND=eager`` is set,
+the implementation falls back to ordinary PyTorch operations.
+
+Replace the original model-code import with, for example,
+``from .fused_moe_v010 import fused_moe_forward``.
+
+Derived from ByteDance-Seed/VeOmni v0.1.0.post1:
+https://github.com/ByteDance-Seed/VeOmni/tree/v0.1.0.post1
+"""
+
+from __future__ import annotations
+
+import os
+
+import torch
+import torch.nn.functional as F
+
+try:
+ import triton
+ import triton.language as tl
+except ImportError: # The eager fallback does not require Triton.
+ triton = None
+ tl = None
+
+
+_SUPPORTED_TRITON_DTYPES = (torch.float16, torch.bfloat16)
+
+
+if triton is not None:
+
+ @triton.jit
+ def _grouped_linear_kernel(
+ input_ptr,
+ weight_ptr,
+ output_ptr,
+ expert_cumsum_ptr,
+ N: tl.constexpr,
+ K: tl.constexpr,
+ BLOCK_M: tl.constexpr,
+ BLOCK_N: tl.constexpr,
+ BLOCK_K: tl.constexpr,
+ ):
+ """Compute per-expert ``input @ weight.T`` for contiguous tensors."""
+ block_m = tl.program_id(axis=0)
+ block_n = tl.program_id(axis=1)
+ expert = tl.program_id(axis=2)
+
+ expert_start = tl.load(expert_cumsum_ptr + expert - 1, mask=expert > 0, other=0)
+ expert_end = tl.load(expert_cumsum_ptr + expert)
+ expert_tokens = expert_end - expert_start
+
+ if block_m * BLOCK_M >= expert_tokens:
+ return
+
+ row_offsets = block_m * BLOCK_M + tl.arange(0, BLOCK_M)
+ col_offsets = block_n * BLOCK_N + tl.arange(0, BLOCK_N)
+ k_offsets = tl.arange(0, BLOCK_K)
+
+ input_ptrs = (
+ input_ptr
+ + (expert_start + row_offsets[:, None]) * K
+ + k_offsets[None, :]
+ )
+ weight_ptrs = (
+ weight_ptr
+ + expert * N * K
+ + col_offsets[None, :] * K
+ + k_offsets[:, None]
+ )
+
+ accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
+ for k_block in range(0, tl.cdiv(K, BLOCK_K)):
+ remaining_k = K - k_block * BLOCK_K
+ inputs = tl.load(
+ input_ptrs,
+ mask=(row_offsets[:, None] < expert_tokens) & (k_offsets[None, :] < remaining_k),
+ other=0.0,
+ )
+ weights = tl.load(
+ weight_ptrs,
+ mask=(col_offsets[None, :] < N) & (k_offsets[:, None] < remaining_k),
+ other=0.0,
+ )
+ accumulator += tl.dot(inputs, weights)
+ input_ptrs += BLOCK_K
+ weight_ptrs += BLOCK_K
+
+ output_ptrs = (
+ output_ptr
+ + (expert_start + row_offsets[:, None]) * N
+ + col_offsets[None, :]
+ )
+ tl.store(
+ output_ptrs,
+ accumulator,
+ mask=(row_offsets[:, None] < expert_tokens) & (col_offsets[None, :] < N),
+ )
+
+
+def _validate_inputs(
+ num_experts: int,
+ routing_weights: torch.Tensor,
+ selected_experts: torch.Tensor,
+ hidden_states: torch.Tensor,
+ fc1_1_weight: torch.Tensor,
+ fc1_2_weight: torch.Tensor,
+ fc2_weight: torch.Tensor,
+) -> None:
+ if num_experts <= 0:
+ raise ValueError(f"num_experts must be positive, got {num_experts}")
+ if torch.is_grad_enabled():
+ raise RuntimeError(
+ "This standalone fused_moe_forward is inference-only. Call it under "
+ "torch.no_grad() or torch.inference_mode()."
+ )
+ if hidden_states.ndim != 2:
+ raise ValueError(f"hidden_states must have shape [tokens, hidden], got {tuple(hidden_states.shape)}")
+ if routing_weights.ndim != 2 or selected_experts.shape != routing_weights.shape:
+ raise ValueError(
+ "routing_weights and selected_experts must have the same [tokens, top_k] shape, got "
+ f"{tuple(routing_weights.shape)} and {tuple(selected_experts.shape)}"
+ )
+ if routing_weights.shape[1] == 0:
+ raise ValueError("top_k must be positive")
+ if routing_weights.shape[0] != hidden_states.shape[0]:
+ raise ValueError("routing_weights and hidden_states must contain the same number of tokens")
+ if selected_experts.dtype not in (torch.int32, torch.int64):
+ raise TypeError(f"selected_experts must be int32 or int64, got {selected_experts.dtype}")
+ if fc1_1_weight.ndim != 3 or fc1_2_weight.ndim != 3 or fc2_weight.ndim != 3:
+ raise ValueError("expert weights must be rank-3 tensors")
+ if fc1_1_weight.shape != fc1_2_weight.shape:
+ raise ValueError("fc1_1_weight and fc1_2_weight must have identical shapes")
+
+ experts, intermediate_size, hidden_size = fc1_1_weight.shape
+ expected_fc2_shape = (experts, hidden_size, intermediate_size)
+ if experts != num_experts:
+ raise ValueError(f"num_experts={num_experts}, but the weights contain {experts} experts")
+ if hidden_states.shape[1] != hidden_size:
+ raise ValueError(f"hidden size is {hidden_states.shape[1]}, but the weights expect {hidden_size}")
+ if tuple(fc2_weight.shape) != expected_fc2_shape:
+ raise ValueError(f"fc2_weight must have shape {expected_fc2_shape}, got {tuple(fc2_weight.shape)}")
+ if selected_experts.numel():
+ # These scalar checks synchronize CUDA once, before launching harder-to-debug kernels.
+ min_expert = int(selected_experts.min().item())
+ max_expert = int(selected_experts.max().item())
+ if min_expert < 0 or max_expert >= num_experts:
+ raise ValueError(f"selected expert IDs must be in [0, {num_experts}), got [{min_expert}, {max_expert}]")
+
+ devices = {
+ hidden_states.device,
+ routing_weights.device,
+ selected_experts.device,
+ fc1_1_weight.device,
+ fc1_2_weight.device,
+ fc2_weight.device,
+ }
+ if len(devices) != 1:
+ raise ValueError(f"all inputs and weights must be on one device, got {sorted(map(str, devices))}")
+
+
+def _route_tokens(
+ num_experts: int,
+ routing_weights: torch.Tensor,
+ selected_experts: torch.Tensor,
+ hidden_states: torch.Tensor,
+) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
+ """Sort routed token copies by expert and return the inverse permutation."""
+ top_k = selected_experts.shape[1]
+ flat_experts = selected_experts.reshape(-1).to(torch.int64)
+ order = torch.argsort(flat_experts, stable=True)
+ sorted_hidden_states = hidden_states[torch.div(order, top_k, rounding_mode="floor")].contiguous()
+ sorted_routing_weights = routing_weights.reshape(-1)[order].contiguous()
+ tokens_per_expert = torch.bincount(flat_experts, minlength=num_experts)
+ expert_cumsum = torch.cumsum(tokens_per_expert, dim=0, dtype=torch.int32).contiguous()
+ return sorted_hidden_states, sorted_routing_weights, expert_cumsum, order
+
+
+def _unroute_tokens(
+ sorted_outputs: torch.Tensor,
+ order: torch.Tensor,
+ num_tokens: int,
+ top_k: int,
+) -> torch.Tensor:
+ restored = torch.empty_like(sorted_outputs)
+ restored[order] = sorted_outputs
+ # VeOmni's v0.1.0 gather kernel accumulates the top-k outputs in FP32.
+ return restored.view(num_tokens, top_k, -1).sum(dim=1, dtype=torch.float32).to(sorted_outputs.dtype)
+
+
+def _grouped_linear_triton(
+ inputs: torch.Tensor,
+ weights: torch.Tensor,
+ expert_cumsum: torch.Tensor,
+) -> torch.Tensor:
+ if triton is None: # pragma: no cover - guarded by the caller
+ raise RuntimeError("Triton is not available")
+ if not inputs.is_contiguous() or not weights.is_contiguous():
+ raise ValueError("the Triton path requires contiguous inputs and expert weights")
+
+ num_experts, output_size, input_size = weights.shape
+ if inputs.shape[1] != input_size:
+ raise ValueError(f"input width is {inputs.shape[1]}, but the weights expect {input_size}")
+
+ output = torch.empty((inputs.shape[0], output_size), dtype=inputs.dtype, device=inputs.device)
+ block_m, block_n, block_k = 128, 128, 32
+ grid = (
+ triton.cdiv(inputs.shape[0], block_m),
+ triton.cdiv(output_size, block_n),
+ num_experts,
+ )
+ with torch.cuda.device(inputs.device):
+ _grouped_linear_kernel[grid](
+ inputs,
+ weights,
+ output,
+ expert_cumsum,
+ N=output_size,
+ K=input_size,
+ BLOCK_M=block_m,
+ BLOCK_N=block_n,
+ BLOCK_K=block_k,
+ num_warps=8,
+ num_stages=3,
+ )
+ return output
+
+
+def _triton_moe_forward(
+ num_experts: int,
+ routing_weights: torch.Tensor,
+ selected_experts: torch.Tensor,
+ hidden_states: torch.Tensor,
+ fc1_1_weight: torch.Tensor,
+ fc1_2_weight: torch.Tensor,
+ fc2_weight: torch.Tensor,
+) -> torch.Tensor:
+ sorted_hidden, sorted_routing, expert_cumsum, order = _route_tokens(
+ num_experts, routing_weights, selected_experts, hidden_states
+ )
+ gate = _grouped_linear_triton(sorted_hidden, fc1_1_weight, expert_cumsum)
+ up = _grouped_linear_triton(sorted_hidden, fc1_2_weight, expert_cumsum)
+ intermediate = F.silu(gate) * up
+ intermediate.mul_(sorted_routing.unsqueeze(-1))
+ sorted_outputs = _grouped_linear_triton(intermediate.contiguous(), fc2_weight, expert_cumsum)
+ return _unroute_tokens(sorted_outputs, order, hidden_states.shape[0], selected_experts.shape[1])
+
+
+def _eager_moe_forward(
+ num_experts: int,
+ routing_weights: torch.Tensor,
+ selected_experts: torch.Tensor,
+ hidden_states: torch.Tensor,
+ fc1_1_weight: torch.Tensor,
+ fc1_2_weight: torch.Tensor,
+ fc2_weight: torch.Tensor,
+) -> torch.Tensor:
+ sorted_hidden, sorted_routing, expert_cumsum, order = _route_tokens(
+ num_experts, routing_weights, selected_experts, hidden_states
+ )
+ expert_ends = expert_cumsum.to(device="cpu", dtype=torch.int64).tolist()
+ outputs: list[torch.Tensor] = []
+ start = 0
+ for expert, end in enumerate(expert_ends):
+ if end > start:
+ expert_inputs = sorted_hidden[start:end]
+ gate = F.linear(expert_inputs, fc1_1_weight[expert])
+ up = F.linear(expert_inputs, fc1_2_weight[expert])
+ intermediate = F.silu(gate) * up
+ intermediate.mul_(sorted_routing[start:end].unsqueeze(-1))
+ outputs.append(F.linear(intermediate, fc2_weight[expert]))
+ start = end
+
+ sorted_outputs = torch.cat(outputs, dim=0) if outputs else hidden_states.new_empty((0, hidden_states.shape[1]))
+ return _unroute_tokens(sorted_outputs, order, hidden_states.shape[0], selected_experts.shape[1])
+
+
+def fused_moe_forward(
+ module: torch.nn.Module,
+ num_experts: int,
+ routing_weights: torch.Tensor,
+ selected_experts: torch.Tensor,
+ hidden_states: torch.Tensor,
+ fc1_1_weight: torch.Tensor,
+ fc1_2_weight: torch.Tensor,
+ fc2_weight: torch.Tensor,
+) -> torch.Tensor:
+ """Run the VeOmni v0.1.0 split-weight MoE operation for inference.
+
+ ``module`` is retained for call-site compatibility. Like VeOmni's original
+ non-EP implementation, this function does not use it.
+
+ Set ``LLADA_MOE_BACKEND`` to ``auto`` (default), ``triton``, or ``eager``.
+ The ``triton`` setting fails loudly if its requirements are not met;
+ ``auto`` falls back to the PyTorch implementation.
+ """
+ del module
+ _validate_inputs(
+ num_experts,
+ routing_weights,
+ selected_experts,
+ hidden_states,
+ fc1_1_weight,
+ fc1_2_weight,
+ fc2_weight,
+ )
+
+ backend = os.getenv("LLADA_MOE_BACKEND", "auto").lower()
+ if backend not in {"auto", "triton", "eager"}:
+ raise ValueError(f"LLADA_MOE_BACKEND must be auto, triton, or eager; got {backend!r}")
+
+ compute_dtype = fc1_1_weight.dtype
+ if fc1_2_weight.dtype != compute_dtype or fc2_weight.dtype != compute_dtype:
+ raise TypeError("all expert weights must have the same dtype")
+ hidden_states = hidden_states.to(dtype=compute_dtype)
+ routing_weights = routing_weights.to(dtype=compute_dtype)
+
+ if hidden_states.shape[0] == 0:
+ return hidden_states
+
+ can_use_triton = (
+ triton is not None
+ and hidden_states.is_cuda
+ and compute_dtype in _SUPPORTED_TRITON_DTYPES
+ and fc1_1_weight.is_contiguous()
+ and fc1_2_weight.is_contiguous()
+ and fc2_weight.is_contiguous()
+ )
+ if backend == "triton" and not can_use_triton:
+ raise RuntimeError(
+ "The Triton backend requires Triton, CUDA tensors, contiguous expert weights, "
+ "and float16 or bfloat16 weights."
+ )
+ if backend != "eager" and can_use_triton:
+ return _triton_moe_forward(
+ num_experts,
+ routing_weights,
+ selected_experts,
+ hidden_states,
+ fc1_1_weight,
+ fc1_2_weight,
+ fc2_weight,
+ )
+ return _eager_moe_forward(
+ num_experts,
+ routing_weights,
+ selected_experts,
+ hidden_states,
+ fc1_1_weight,
+ fc1_2_weight,
+ fc2_weight,
+ )
+
+
+__all__ = ["fused_moe_forward"]
diff --git a/pipelines/llada/modeling_llada2uni_moe.py b/pipelines/llada/modeling_llada2uni_moe.py
new file mode 100644
index 000000000..77eb547b3
--- /dev/null
+++ b/pipelines/llada/modeling_llada2uni_moe.py
@@ -0,0 +1,1284 @@
+# coding=utf-8
+# Copyright 2025 Antgroup and The HuggingFace Inc. team. All rights reserved.
+#
+# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
+# and OPT implementations in this library. It has been modified from its
+# original forms to accommodate minor architectural differences compared
+# to GPT-NeoX and OPT used by the Meta AI team that trained the model.
+#
+# 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.
+"""PyTorch implementation of the fused LLaDA2 MoE model."""
+
+from dataclasses import dataclass
+import math
+from typing import List, Optional, Tuple, Union
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+from transformers.activations import ACT2FN
+from transformers.cache_utils import Cache
+from transformers.generation import GenerationMixin
+from transformers.modeling_attn_mask_utils import (
+ _prepare_4d_attention_mask,
+ _prepare_4d_causal_attention_mask,
+ _prepare_4d_causal_attention_mask_for_sdpa,
+)
+from transformers.modeling_outputs import ModelOutput, MoeModelOutputWithPast
+from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS, dynamic_rope_update
+from transformers.modeling_utils import PreTrainedModel
+from transformers.pytorch_utils import ALL_LAYERNORM_LAYERS
+from transformers.utils import logging
+from .fused_moe_ops import fused_moe_forward
+
+from .configuration_llada2uni_moe import LLaDA2MoeConfig
+
+logger = logging.get_logger(__name__)
+
+
+class LLaDA2MoeRMSNorm(nn.Module):
+ """RMSNorm used by the LLaDA2 model."""
+
+ def __init__(self, hidden_size, eps=1e-6):
+ super().__init__()
+ self.weight = nn.Parameter(torch.ones(hidden_size))
+ self.variance_epsilon = eps
+
+ def forward(self, hidden_states):
+ input_dtype = hidden_states.dtype
+ hidden_states = hidden_states.to(torch.float32)
+ variance = hidden_states.pow(2).mean(-1, keepdim=True)
+ hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
+ return self.weight * hidden_states.to(input_dtype)
+
+
+# Preserve the historical spelling used by the original implementation.
+LLaDA2MoERMSNorm = LLaDA2MoeRMSNorm
+ALL_LAYERNORM_LAYERS.append(LLaDA2MoeRMSNorm)
+
+
+class LLaDA2MoePreTrainedModel(PreTrainedModel):
+ config_class = LLaDA2MoeConfig
+ base_model_prefix = "model"
+ supports_gradient_checkpointing = True
+ _no_split_modules = ["LLaDA2MoeDecoderLayer"]
+ _skip_keys_device_placement = "past_key_values"
+ _supports_flash_attn_2 = True
+ _supports_sdpa = True
+ _supports_cache_class = True
+ _supports_flash_attn = True
+ _can_compile_fullgraph = True
+ _supports_attention_backend = True
+
+ def _init_weights(self, module):
+ std = self.config.initializer_range
+ if isinstance(module, nn.Linear):
+ module.weight.data.normal_(mean=0.0, std=std)
+ if module.bias is not None:
+ module.bias.data.zero_()
+ elif isinstance(module, nn.Embedding):
+ module.weight.data.normal_(mean=0.0, std=std)
+ if module.padding_idx is not None:
+ module.weight.data[module.padding_idx].zero_()
+
+
+def rotate_half(hidden_states):
+ first, second = hidden_states.chunk(2, dim=-1)
+ return torch.cat((-second, first), dim=-1)
+
+
+def apply_rotary_pos_emb(query, key, cos, sin, position_ids=None, unsqueeze_dim=1): # pylint: disable=unused-argument
+ """Apply RoPE to the rotary part of query and key states."""
+ cos = cos.unsqueeze(unsqueeze_dim)
+ sin = sin.unsqueeze(unsqueeze_dim)
+ rotary_dim = cos.shape[-1]
+ query_rotary, query_pass = query[..., :rotary_dim], query[..., rotary_dim:]
+ key_rotary, key_pass = key[..., :rotary_dim], key[..., rotary_dim:]
+ query_rotary = query_rotary * cos + rotate_half(query_rotary) * sin
+ key_rotary = key_rotary * cos + rotate_half(key_rotary) * sin
+ return torch.cat((query_rotary, query_pass), dim=-1), torch.cat(
+ (key_rotary, key_pass), dim=-1
+ )
+
+
+def _compute_default_rope_parameters(config, device=None, **kwargs): # pylint: disable=unused-argument
+ """Compute the unscaled RoPE frequencies removed from Transformers 5.x."""
+ head_dim = getattr(config, "head_dim", None) or config.hidden_size // config.num_attention_heads
+ partial_rotary_factor = getattr(config, "partial_rotary_factor", 1.0)
+ dim = int(head_dim * partial_rotary_factor)
+ base = getattr(config, "rope_theta", 10000.0)
+ inv_freq = 1.0 / (
+ base ** (torch.arange(0, dim, 2, dtype=torch.int64, device=device).float() / dim)
+ )
+ return inv_freq, 1.0
+
+
+class LLaDA2MoeRotaryEmbedding(nn.Module):
+ def __init__(self, config: LLaDA2MoeConfig, device=None):
+ super().__init__()
+ # BC: "rope_type" was originally "type"
+ if hasattr(config, "rope_scaling") and config.rope_scaling is not None:
+ self.rope_type = config.rope_scaling.get(
+ "rope_type", config.rope_scaling.get("type")
+ )
+ else:
+ self.rope_type = "default"
+ self.max_seq_len_cached = config.max_position_embeddings
+ self.original_max_seq_len = config.max_position_embeddings
+
+ self.config = config
+ self.rope_init_fn = ROPE_INIT_FUNCTIONS.get(self.rope_type)
+ if self.rope_init_fn is None and self.rope_type == "default":
+ self.rope_init_fn = _compute_default_rope_parameters
+ if self.rope_init_fn is None:
+ raise KeyError(f"Unsupported RoPE type: {self.rope_type}")
+
+ inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device)
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
+ self.original_inv_freq = self.inv_freq
+
+ @torch.no_grad()
+ @dynamic_rope_update # power user: used with advanced RoPE types (e.g. dynamic rope)
+ def forward(self, x, position_ids):
+ inv_freq_expanded = (
+ self.inv_freq[None, :, None]
+ .float()
+ .expand(position_ids.shape[0], -1, 1)
+ .to(x.device)
+ )
+ position_ids_expanded = position_ids[:, None, :].float()
+
+ device_type = (
+ x.device.type
+ if isinstance(x.device.type, str) and x.device.type != "mps"
+ else "cpu"
+ )
+ with torch.autocast(device_type=device_type, enabled=False): # Force float32
+ freqs = (
+ inv_freq_expanded.float() @ position_ids_expanded.float()
+ ).transpose(1, 2)
+ emb = torch.cat((freqs, freqs), dim=-1)
+ cos = emb.cos() * self.attention_scaling
+ sin = emb.sin() * self.attention_scaling
+
+ return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
+
+
+class LLaDA2MoeMLP(nn.Module):
+ def __init__(self, config: LLaDA2MoeConfig, intermediate_size: int):
+ super().__init__()
+ self.config = config
+ self.hidden_size = config.hidden_size
+ self.intermediate_size = intermediate_size
+
+ self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
+ self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
+ self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
+ self.act_fn = ACT2FN[config.hidden_act]
+
+ def forward(self, x):
+ return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
+
+
+class LLaDA2MoeGate(nn.Module):
+ def __init__(self, config):
+ super().__init__()
+ self.config = config
+ self.top_k = config.num_experts_per_tok
+ self.num_experts = config.num_experts
+
+ self.n_group = config.n_group
+ self.topk_group = config.topk_group
+
+ self.gating_dim = config.hidden_size
+ self.weight = nn.Parameter(torch.empty((self.num_experts, self.gating_dim)))
+ self.routed_scaling_factor = config.routed_scaling_factor
+
+ self.register_buffer("expert_bias", torch.zeros((self.num_experts)))
+ self.reset_parameters()
+
+ def reset_parameters(self) -> None:
+ import torch.nn.init as init
+
+ init.kaiming_uniform_(self.weight, a=math.sqrt(5))
+
+ def group_limited_topk(
+ self,
+ scores: torch.Tensor,
+ ):
+ num_tokens, _ = scores.size()
+ group_scores = (
+ scores.view(num_tokens, self.n_group, -1).topk(2, dim=-1)[0].sum(dim=-1)
+ )
+ group_idx = torch.topk(group_scores, k=self.topk_group, dim=-1, sorted=False)[1]
+ group_mask = torch.zeros_like(group_scores)
+ group_mask.scatter_(1, group_idx, 1)
+
+ score_mask = (
+ group_mask.unsqueeze(-1)
+ .expand(num_tokens, self.n_group, self.num_experts // self.n_group)
+ .reshape(num_tokens, -1)
+ )
+
+ masked_scores = scores.masked_fill(~score_mask.bool(), float("-inf"))
+ probs, top_indices = torch.topk(masked_scores, k=self.top_k, dim=-1)
+
+ return probs, top_indices
+
+ def forward(self, hidden_states):
+ hidden_states = hidden_states.view(-1, hidden_states.shape[-1])
+ logits = F.linear(
+ hidden_states.type(torch.float32), self.weight.type(torch.float32)
+ )
+
+ scores = torch.sigmoid(logits.float()).type_as(logits)
+
+ scores_for_routing = scores + self.expert_bias
+ _, topk_idx = self.group_limited_topk(scores_for_routing)
+
+ scores = torch.gather(scores, dim=1, index=topk_idx).type_as(logits)
+
+ topk_weight = (
+ scores / (scores.sum(dim=-1, keepdim=True) + 1e-20)
+ if self.top_k > 1
+ else scores
+ )
+ topk_weight = topk_weight * self.routed_scaling_factor
+
+ return topk_idx, topk_weight, logits
+
+
+class LLaDA2MoeExperts(nn.Module):
+ def __init__(self, config):
+ super().__init__()
+ self.num_experts = config.num_experts
+ self.hidden_dim = config.hidden_size
+ self.intermediate_size = config.moe_intermediate_size
+ self.gate_proj = torch.nn.Parameter(
+ torch.empty(self.num_experts, self.intermediate_size, self.hidden_dim),
+ requires_grad=True,
+ )
+ self.up_proj = torch.nn.Parameter(
+ torch.empty(self.num_experts, self.intermediate_size, self.hidden_dim),
+ requires_grad=True,
+ )
+ self.down_proj = torch.nn.Parameter(
+ torch.empty(self.num_experts, self.hidden_dim, self.intermediate_size),
+ requires_grad=True,
+ )
+
+ def forward(self, hidden_states, routing_weights, selected_experts):
+ return fused_moe_forward(
+ module=self,
+ num_experts=self.num_experts,
+ routing_weights=routing_weights,
+ selected_experts=selected_experts,
+ hidden_states=hidden_states,
+ fc1_1_weight=self.gate_proj,
+ fc1_2_weight=self.up_proj,
+ fc2_weight=self.down_proj,
+ )
+
+ def reset_parameters(self):
+ """
+ Initialize the parameters of all expert networks.
+ Uses different initialization strategies for different projection layers.
+ """
+ for expert_id in range(self.num_experts):
+ nn.init.kaiming_uniform_(self.gate_proj[expert_id], a=math.sqrt(5))
+ nn.init.kaiming_uniform_(self.up_proj[expert_id], a=math.sqrt(5))
+ nn.init.xavier_uniform_(self.down_proj[expert_id])
+
+
+class LLaDA2MoeSparseMoeBlock(nn.Module):
+ """Fused routed experts plus a shared expert."""
+
+ def __init__(self, config: LLaDA2MoeConfig):
+ super().__init__()
+ self.config = config
+ self.experts = LLaDA2MoeExperts(config)
+ self.gate = LLaDA2MoeGate(config)
+ if config.num_shared_experts is not None:
+ self.shared_experts = LLaDA2MoeMLP(
+ config=config,
+ intermediate_size=config.moe_intermediate_size
+ * config.num_shared_experts,
+ )
+
+ def forward(self, hidden_states):
+ identity = hidden_states
+ bsz, seq_len, h = hidden_states.shape
+ topk_idx, topk_weight, router_logits = self.gate(hidden_states)
+ hidden_states = hidden_states.view(-1, hidden_states.shape[-1])
+ y = self.experts(
+ hidden_states, routing_weights=topk_weight, selected_experts=topk_idx
+ ).reshape(bsz, seq_len, h)
+ if self.config.num_shared_experts is not None:
+ y = y + self.shared_experts(identity)
+ return y, (
+ router_logits.view(bsz, seq_len, -1),
+ topk_idx.view(bsz, seq_len, -1),
+ )
+
+
+def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
+ """
+ This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
+ num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
+ """
+ batch, num_key_value_heads, slen, head_dim = hidden_states.shape
+ if n_rep == 1:
+ return hidden_states
+ hidden_states = hidden_states[:, :, None, :, :].expand(
+ batch, num_key_value_heads, n_rep, slen, head_dim
+ )
+ return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
+
+
+class LLaDA2MoeAttention(nn.Module):
+ """Multi-headed attention from 'Attention Is All You Need' paper"""
+
+ def __init__(self, config: LLaDA2MoeConfig, layer_idx: Optional[int] = None):
+ super().__init__()
+ self.config = config
+ self.layer_idx = layer_idx
+ if layer_idx is None:
+ logger.warning_once(
+ f"Instantiating {self.__class__.__name__} without passing `layer_idx` is not recommended and will "
+ "to errors during the forward call, if caching is used. Please make sure to provide a `layer_idx` "
+ "when creating this class."
+ )
+
+ self.attention_dropout = config.attention_dropout
+ self.hidden_size = config.hidden_size
+ self.num_heads = config.num_attention_heads
+ self.head_dim = config.head_dim or self.hidden_size // self.num_heads
+ partial_rotary_factor = (
+ config.partial_rotary_factor
+ if hasattr(config, "partial_rotary_factor")
+ else 1.0
+ )
+ self.rope_dim = int(self.head_dim * partial_rotary_factor)
+ self.num_key_value_heads = config.num_key_value_heads
+ self.num_key_value_groups = self.num_heads // self.num_key_value_heads
+ self.max_position_embeddings = config.max_position_embeddings
+ self.rope_theta = config.rope_theta
+ self.is_causal = False
+
+ self.query_key_value = nn.Linear(
+ self.hidden_size,
+ (self.num_heads + 2 * self.num_key_value_heads) * self.head_dim,
+ bias=config.use_qkv_bias,
+ )
+
+ self.query_layernorm = LLaDA2MoERMSNorm(self.head_dim, eps=config.rms_norm_eps)
+ self.key_layernorm = LLaDA2MoERMSNorm(self.head_dim, eps=config.rms_norm_eps)
+ self.dense = nn.Linear(
+ self.num_heads * self.head_dim, self.hidden_size, bias=config.use_bias
+ )
+
+ def forward( # pylint: disable=unused-argument
+ self,
+ hidden_states: torch.Tensor,
+ attention_mask: Optional[torch.Tensor] = None,
+ position_ids: Optional[torch.LongTensor] = None,
+ past_key_value: Optional[Cache] = None,
+ output_attentions: bool = False,
+ use_cache: bool = False,
+ position_embeddings: Optional[
+ Tuple[torch.Tensor, torch.Tensor]
+ ] = None, # necessary, but kept here for BC
+ **kwargs,
+ ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
+ bsz, q_len, _ = hidden_states.size()
+
+ qkv = self.query_key_value(hidden_states)
+ qkv = qkv.view(
+ bsz, q_len, self.num_heads + 2 * self.num_key_value_heads, self.head_dim
+ )
+
+ query_states, key_states, value_states = qkv.split(
+ [self.num_heads, self.num_key_value_heads, self.num_key_value_heads], dim=-2
+ )
+ query_states = query_states.transpose(1, 2)
+ key_states = key_states.transpose(1, 2)
+ value_states = value_states.transpose(1, 2)
+
+ query_states = self.query_layernorm(query_states)
+ key_states = self.key_layernorm(key_states)
+
+ kv_seq_len = key_states.shape[-2]
+ if past_key_value is not None:
+ if self.layer_idx is None:
+ raise ValueError(
+ f"The cache structure has changed since version v4.36. If you are using {self.__class__.__name__} "
+ "for auto-regressive decoding with k/v caching, please make sure to initialize the attention class "
+ "with a layer index."
+ )
+ kv_seq_len += past_key_value.get_seq_length(self.layer_idx)
+ cos, sin = position_embeddings
+ query_states, key_states = apply_rotary_pos_emb(
+ query_states, key_states, cos, sin, position_ids
+ )
+
+ if past_key_value is not None:
+ cache_kwargs = {"sin": sin, "cos": cos} # Specific to RoPE models
+ key_states, value_states = past_key_value.update(
+ key_states, value_states, self.layer_idx, cache_kwargs
+ )
+
+ key_states = repeat_kv(key_states, self.num_key_value_groups)
+ value_states = repeat_kv(value_states, self.num_key_value_groups)
+
+ attn_weights = torch.matmul(
+ query_states, key_states.transpose(2, 3)
+ ) / math.sqrt(self.head_dim)
+
+ if attn_weights.size() != (bsz, self.num_heads, q_len, kv_seq_len):
+ raise ValueError(
+ f"Attention weights should be of size {(bsz, self.num_heads, q_len, kv_seq_len)}, but is"
+ f" {attn_weights.size()}"
+ )
+ if attention_mask is not None:
+ if attention_mask.size() != (bsz, 1, q_len, kv_seq_len):
+ raise ValueError(
+ f"Attention mask should be of size {(bsz, 1, q_len, kv_seq_len)}, but is {attention_mask.size()}"
+ )
+ attn_weights = attn_weights + attention_mask
+
+ attn_weights = nn.functional.softmax(
+ attn_weights, dim=-1, dtype=torch.float32
+ ).to(query_states.dtype)
+ attn_weights = nn.functional.dropout(
+ attn_weights, p=self.attention_dropout, training=self.training
+ )
+ attn_output = torch.matmul(attn_weights, value_states)
+
+ if attn_output.size() != (bsz, self.num_heads, q_len, self.head_dim):
+ raise ValueError(
+ f"`attn_output` should be of size {(bsz, self.num_heads, q_len, self.head_dim)}, but is"
+ f" {attn_output.size()}"
+ )
+
+ attn_output = attn_output.transpose(1, 2).contiguous()
+
+ attn_output = attn_output.reshape(bsz, q_len, -1)
+
+ attn_output = self.dense(attn_output)
+
+ if not output_attentions:
+ attn_weights = None
+
+ return attn_output, attn_weights, past_key_value
+
+
+class LLaDA2MoeSdpaAttention(LLaDA2MoeAttention):
+ """
+ LLaDA2Moe attention module using torch.nn.functional.scaled_dot_product_attention. This module inherits from
+ `LLaDA2MoeAttention` as the weights of the module stays untouched. The only changes are on the forward pass to adapt to
+ SDPA API.
+ """
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ attention_mask: Optional[torch.Tensor] = None,
+ position_ids: Optional[torch.LongTensor] = None,
+ past_key_value: Optional[Cache] = None,
+ output_attentions: bool = False,
+ use_cache: bool = False,
+ position_embeddings: Optional[
+ Tuple[torch.Tensor, torch.Tensor]
+ ] = None, # necessary, but kept here for BC
+ **kwargs,
+ ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
+ if output_attentions:
+ logger.warning_once(
+ "LLaDA2MoeModel is using LLaDA2MoeSdpaAttention, but `torch.nn.functional.scaled_dot_product_attention` does not support `output_attentions=True`. Falling back to the manual attention implementation, "
+ 'but specifying the manual implementation will be required from Transformers version v5.0.0 onwards. This warning can be removed using the argument `attn_implementation="eager"` when loading the model.'
+ )
+ return super().forward(
+ hidden_states=hidden_states,
+ attention_mask=attention_mask,
+ position_ids=position_ids,
+ past_key_value=past_key_value,
+ output_attentions=output_attentions,
+ use_cache=use_cache,
+ )
+
+ bsz, q_len, _ = hidden_states.size()
+
+ qkv = self.query_key_value(hidden_states)
+ qkv = qkv.view(
+ bsz, q_len, self.num_heads + 2 * self.num_key_value_heads, self.head_dim
+ )
+
+ query_states, key_states, value_states = qkv.split(
+ [self.num_heads, self.num_key_value_heads, self.num_key_value_heads], dim=-2
+ )
+ query_states = query_states.transpose(1, 2)
+ key_states = key_states.transpose(1, 2)
+ value_states = value_states.transpose(1, 2)
+
+ query_states = self.query_layernorm(query_states)
+ key_states = self.key_layernorm(key_states)
+
+ kv_seq_len = key_states.shape[-2]
+ if past_key_value is not None:
+ kv_seq_len += past_key_value.get_seq_length(self.layer_idx)
+ cos, sin = position_embeddings
+
+ query_states, key_states = apply_rotary_pos_emb(
+ query_states, key_states, cos, sin, position_ids
+ )
+
+ if past_key_value is not None:
+ cache_kwargs = {"sin": sin, "cos": cos} # Specific to RoPE models
+ key_states, value_states = past_key_value.update(
+ key_states, value_states, self.layer_idx, cache_kwargs
+ )
+
+ key_states = repeat_kv(key_states, self.num_key_value_groups)
+ value_states = repeat_kv(value_states, self.num_key_value_groups)
+
+ if attention_mask is not None:
+ if attention_mask.size() != (bsz, 1, q_len, kv_seq_len):
+ raise ValueError(
+ f"Attention mask should be of size {(bsz, 1, q_len, kv_seq_len)}, but is {attention_mask.size()}"
+ )
+
+ # SDPA with memory-efficient backend is currently (torch==2.1.2) bugged with non-contiguous inputs with custom attn_mask,
+ # Reference: https://github.com/pytorch/pytorch/issues/112577.
+ if query_states.device.type == "cuda" and attention_mask is not None:
+ query_states = query_states.contiguous()
+ key_states = key_states.contiguous()
+ value_states = value_states.contiguous()
+
+ attn_output = torch.nn.functional.scaled_dot_product_attention(
+ query_states,
+ key_states,
+ value_states,
+ attn_mask=attention_mask,
+ dropout_p=self.attention_dropout if self.training else 0.0,
+ # The q_len > 1 is necessary to match with AttentionMaskConverter.to_causal_4d that does not create a causal mask in case q_len == 1.
+ is_causal=self.is_causal and attention_mask is None and q_len > 1,
+ )
+
+ attn_output = attn_output.transpose(1, 2).contiguous()
+ attn_output = attn_output.reshape(bsz, q_len, -1)
+
+ attn_output = self.dense(attn_output)
+
+ return attn_output, None, past_key_value
+
+
+ATTENTION_CLASSES = {
+ "eager": LLaDA2MoeSdpaAttention,
+ "flash_attention_2": LLaDA2MoeSdpaAttention,
+ "sdpa": LLaDA2MoeSdpaAttention,
+}
+
+
+class LLaDA2MoeDecoderLayer(nn.Module):
+ def __init__(self, config: LLaDA2MoeConfig, layer_idx: int):
+ super().__init__()
+ self.hidden_size = config.hidden_size
+
+ self.attention = ATTENTION_CLASSES[config._attn_implementation](
+ config=config, layer_idx=layer_idx
+ )
+
+ self.mlp = (
+ LLaDA2MoeSparseMoeBlock(config)
+ if (
+ config.num_experts is not None
+ and layer_idx >= config.first_k_dense_replace
+ )
+ else LLaDA2MoeMLP(config=config, intermediate_size=config.intermediate_size)
+ )
+ self.input_layernorm = LLaDA2MoERMSNorm(
+ config.hidden_size, eps=config.rms_norm_eps
+ )
+ self.post_attention_layernorm = LLaDA2MoERMSNorm(
+ config.hidden_size, eps=config.rms_norm_eps
+ )
+
+ def forward(
+ self, # pylint: disable=unused-argument
+ hidden_states: torch.Tensor,
+ attention_mask: Optional[torch.Tensor] = None,
+ position_ids: Optional[torch.LongTensor] = None,
+ past_key_value: Optional[Tuple[torch.Tensor]] = None,
+ output_attentions: Optional[bool] = False,
+ output_router_logits: Optional[bool] = False,
+ use_cache: Optional[bool] = False,
+ position_embeddings: Optional[
+ Tuple[torch.Tensor, torch.Tensor]
+ ] = None, # necessary, but kept here for BC
+ **kwargs,
+ ) -> Tuple[
+ torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]
+ ]:
+ """
+ Args:
+ hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`
+ attention_mask (`torch.FloatTensor`, *optional*):
+ attention mask of size `(batch_size, sequence_length)` if flash attention is used or `(batch_size, 1,
+ query_sequence_length, key_sequence_length)` if default attention is used.
+ position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
+ Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
+ config.n_positions - 1]`.
+ past_key_value (`Tuple(torch.FloatTensor)`, *optional*):
+ cached past key and value projection states
+ output_attentions (`bool`, *optional*):
+ Whether to return the attentions tensors of all attention layers. See `attentions` under
+ returned tensors for more detail.
+ output_router_logits (`bool`, *optional*):
+ Whether or not to return the logits of all the routers. They are useful for computing the router loss,
+ and should not be returned during inference.
+ use_cache (`bool`, *optional*):
+ If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding
+ (see `past_key_values`).
+ """
+ residual = hidden_states
+
+ hidden_states = self.input_layernorm(hidden_states)
+
+ hidden_states, self_attn_weights, present_key_value = self.attention(
+ hidden_states=hidden_states,
+ attention_mask=attention_mask,
+ position_ids=position_ids,
+ past_key_value=past_key_value,
+ output_attentions=output_attentions,
+ position_embeddings=position_embeddings,
+ use_cache=use_cache,
+ )
+ hidden_states = residual + hidden_states
+
+ residual = hidden_states
+ hidden_states = self.post_attention_layernorm(hidden_states)
+ hidden_states = self.mlp(hidden_states)
+ if isinstance(hidden_states, tuple):
+ hidden_states, router_logits = hidden_states
+ else:
+ router_logits = None
+ hidden_states = residual + hidden_states.to(residual.device)
+
+ outputs = (hidden_states,)
+
+ if output_attentions:
+ outputs += (self_attn_weights,)
+
+ if use_cache:
+ outputs += (present_key_value,)
+
+ if output_router_logits:
+ outputs += (router_logits,)
+
+ return outputs
+
+
+def calculate_pack_position_ids(
+ input_ids: Optional[torch.Tensor] = None,
+ inputs_embeds: Optional[torch.Tensor] = None,
+ position_ids: Optional[torch.Tensor] = None,
+ past_key_values_length: int = 0,
+ cu_lengths_list: Optional[List[torch.Tensor]] = None,
+):
+ """Build continuous or per-sequence packed position IDs."""
+ if position_ids is not None:
+ return position_ids
+
+ if input_ids is not None:
+ device = input_ids.device
+ batch_size, seq_length = input_ids.shape
+ elif inputs_embeds is not None:
+ device = inputs_embeds.device
+ batch_size, seq_length, _ = inputs_embeds.shape
+ else:
+ raise ValueError("You have to specify either input_ids or inputs_embeds")
+
+ if cu_lengths_list is not None:
+ all_position_ids = []
+
+ for i in range(batch_size):
+ cu_seqlens = cu_lengths_list[i].to(device)
+ starts = cu_seqlens[:-1]
+ lengths = cu_seqlens[1:] - cu_seqlens[:-1]
+ total_len = cu_seqlens[-1].item()
+ global_positions = torch.arange(total_len, device=device, dtype=torch.long)
+ subtraction_mask = torch.repeat_interleave(starts, lengths)
+ current_pos_ids = global_positions - subtraction_mask
+ all_position_ids.append(current_pos_ids)
+
+ position_ids = torch.nn.utils.rnn.pad_sequence(
+ all_position_ids,
+ batch_first=True,
+ padding_value=0,
+ )
+
+ if position_ids.shape[1] < seq_length:
+ pad_right = seq_length - position_ids.shape[1]
+ position_ids = F.pad(position_ids, (0, pad_right), "constant", 0)
+ else:
+ position_ids = torch.arange(
+ past_key_values_length,
+ seq_length + past_key_values_length,
+ dtype=torch.long,
+ device=device,
+ )
+ position_ids = position_ids.unsqueeze(0).expand(batch_size, -1)
+
+ return position_ids
+
+
+class LLaDA2MoeModel(LLaDA2MoePreTrainedModel):
+ """
+ Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`LLaDA2MoeDecoderLayer`]
+ Args:
+ config: LLaDA2MoeConfig
+ """
+
+ def __init__(self, config: LLaDA2MoeConfig):
+ super().__init__(config)
+ self.padding_idx = config.pad_token_id
+ self.vocab_size = config.vocab_size
+
+ self.word_embeddings = nn.Embedding(
+ config.vocab_size, config.hidden_size, self.padding_idx
+ )
+ self.layers = nn.ModuleList(
+ [
+ LLaDA2MoeDecoderLayer(config, layer_idx)
+ for layer_idx in range(config.num_hidden_layers)
+ ]
+ )
+
+ self._use_sdpa = config._attn_implementation == "sdpa"
+ self._use_flash_attention_2 = config._attn_implementation == "flash_attention_2"
+ self.norm = LLaDA2MoERMSNorm(config.hidden_size, eps=config.rms_norm_eps)
+ self.rotary_emb = LLaDA2MoeRotaryEmbedding(config=config)
+ self.gradient_checkpointing = False
+
+ self.post_init()
+
+ def get_input_embeddings(self):
+ return self.word_embeddings
+
+ def set_input_embeddings(self, value):
+ self.word_embeddings = value
+
+ def forward(
+ self, # pylint: disable=unused-argument
+ input_ids: torch.LongTensor = None,
+ attention_mask: Optional[torch.Tensor] = None,
+ position_ids: Optional[torch.LongTensor] = None,
+ past_key_values: Optional[List[torch.FloatTensor]] = None,
+ inputs_embeds: Optional[torch.FloatTensor] = None,
+ use_cache: Optional[bool] = None,
+ output_attentions: Optional[bool] = None,
+ output_hidden_states: Optional[bool] = None,
+ output_router_logits: Optional[bool] = None,
+ cu_lengths_list: Optional[List] = None,
+ return_dict: Optional[bool] = None,
+ **kwargs,
+ ) -> Union[Tuple, MoeModelOutputWithPast]:
+ output_attentions = (
+ output_attentions
+ if output_attentions is not None
+ else self.config.output_attentions
+ )
+ output_hidden_states = (
+ output_hidden_states
+ if output_hidden_states is not None
+ else self.config.output_hidden_states
+ )
+ output_router_logits = (
+ output_router_logits
+ if output_router_logits is not None
+ else self.config.output_router_logits
+ )
+ use_cache = use_cache if use_cache is not None else self.config.use_cache
+
+ return_dict = (
+ return_dict if return_dict is not None else self.config.use_return_dict
+ )
+
+ if input_ids is not None and inputs_embeds is not None:
+ raise ValueError(
+ "You cannot specify both input_ids and inputs_embeds at the same time"
+ )
+ elif input_ids is not None:
+ batch_size, seq_length = input_ids.shape[:2]
+ elif inputs_embeds is not None:
+ batch_size, seq_length = inputs_embeds.shape[:2]
+ else:
+ raise ValueError("You have to specify either input_ids or inputs_embeds")
+
+ if self.gradient_checkpointing and self.training:
+ if use_cache:
+ logger.warning_once(
+ "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`transformers."
+ )
+ use_cache = False
+
+ past_key_values_length = past_key_values.get_seq_length() if use_cache and past_key_values is not None else 0
+
+ if position_ids is None:
+ position_ids = calculate_pack_position_ids(
+ input_ids,
+ inputs_embeds,
+ position_ids,
+ past_key_values_length,
+ cu_lengths_list,
+ )
+
+ if inputs_embeds is None:
+ inputs_embeds = self.word_embeddings(input_ids)
+
+ if hasattr(attention_mask, "dim") and attention_mask.dim() == 2:
+ if self._use_sdpa and not output_attentions:
+ # output_attentions=True can not be supported when using SDPA, and we fall back on
+ # the manual implementation that requires a 4D causal mask in all cases.
+ attention_mask = _prepare_4d_causal_attention_mask_for_sdpa(
+ attention_mask,
+ (batch_size, seq_length),
+ inputs_embeds,
+ past_key_values_length,
+ )
+ else:
+ if attention_mask is not None:
+ attention_mask = _prepare_4d_attention_mask(
+ attention_mask, inputs_embeds.dtype
+ )
+ else:
+ attention_mask = _prepare_4d_causal_attention_mask(
+ attention_mask,
+ (batch_size, seq_length),
+ inputs_embeds,
+ past_key_values_length,
+ )
+ hidden_states = inputs_embeds
+
+ position_embeddings = self.rotary_emb(hidden_states, position_ids)
+
+ all_hidden_states = () if output_hidden_states else None
+ all_self_attns = () if output_attentions else None
+ all_router_logits = () if output_router_logits else None
+ next_decoder_cache = None
+
+ for decoder_layer in self.layers:
+ if output_hidden_states:
+ all_hidden_states += (hidden_states,)
+
+ if self.gradient_checkpointing and self.training:
+ layer_outputs = self._gradient_checkpointing_func(
+ decoder_layer.__call__,
+ hidden_states,
+ attention_mask,
+ position_ids,
+ past_key_values,
+ output_attentions,
+ output_router_logits,
+ use_cache,
+ position_embeddings,
+ )
+ else:
+ layer_outputs = decoder_layer(
+ hidden_states,
+ attention_mask=attention_mask,
+ position_ids=position_ids,
+ past_key_value=past_key_values,
+ output_attentions=output_attentions,
+ output_router_logits=output_router_logits,
+ use_cache=use_cache,
+ position_embeddings=position_embeddings,
+ )
+ hidden_states = layer_outputs[0]
+
+ if use_cache:
+ next_decoder_cache = layer_outputs[2 if output_attentions else 1]
+
+ if output_attentions:
+ all_self_attns += (layer_outputs[1],)
+
+ if output_router_logits and layer_outputs[-1] is not None:
+ all_router_logits += (layer_outputs[-1],)
+
+ hidden_states = self.norm(hidden_states)
+
+ if output_hidden_states:
+ all_hidden_states += (hidden_states,)
+
+ next_cache = next_decoder_cache if use_cache else None
+ if not return_dict:
+ return tuple(
+ v
+ for v in [
+ hidden_states,
+ next_cache,
+ all_hidden_states,
+ all_self_attns,
+ all_router_logits,
+ ]
+ if v is not None
+ )
+ return MoeModelOutputWithPast(
+ last_hidden_state=hidden_states,
+ past_key_values=next_cache,
+ hidden_states=all_hidden_states,
+ attentions=all_self_attns,
+ router_logits=all_router_logits,
+ )
+
+
+@dataclass
+class LLaDA2MoeCausalLMOutputWithPast(ModelOutput):
+ r"""
+ loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
+ Training loss.
+ logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):
+ Prediction scores for each vocabulary token before SoftMax.
+ past_key_values (`Cache`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
+ It is a [`~cache_utils.Cache`] instance. For more details, see our [kv cache guide](https://huggingface.co/docs/transformers/en/kv_cache).
+
+ Contains pre-computed hidden-states (key and values in the self-attention blocks) that can be used (see
+ `past_key_values` input) to speed up sequential decoding.
+ rope_deltas (`torch.LongTensor` of shape `(batch_size, )`, *optional*):
+ The offset between the sequence length and rotary position indices.
+ """
+
+ loss: Optional[torch.FloatTensor] = None
+ z_loss: Optional[torch.FloatTensor] = None
+ logits: Optional[torch.FloatTensor] = None
+ past_key_values: Optional[Cache] = None
+ hidden_states: Optional[tuple[torch.FloatTensor]] = None
+ attentions: Optional[tuple[torch.FloatTensor]] = None
+ rope_deltas: Optional[torch.LongTensor] = None
+
+
+class LLaDA2MoeBackbone(nn.Module):
+ """Container for the model backbone and output head."""
+
+ def __init__(self, config: LLaDA2MoeConfig):
+ super().__init__()
+ self.language_model = LLaDA2MoeModel(config)
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
+
+ def get_input_embeddings(self):
+ return self.language_model.get_input_embeddings()
+
+ def set_input_embeddings(self, value):
+ self.language_model.set_input_embeddings(value)
+
+ def forward(self, *args, **kwargs):
+ return self.language_model(*args, **kwargs)
+
+
+class LLaDA2MoeModelLM(LLaDA2MoePreTrainedModel, GenerationMixin):
+ """Fused LLaDA2 MoE model."""
+
+ accepts_loss_kwargs = False
+
+ def __init__(self, config: LLaDA2MoeConfig):
+ super().__init__(config)
+ self.model = LLaDA2MoeBackbone(config)
+ self.img_token_id = 157184
+ self.img_start_id = 157185
+ self.img_end_id = 157186
+ self.img_pad_id = 157187
+ self.post_init()
+
+ @property
+ def language_model(self):
+ return self.model.language_model
+
+ def get_input_embeddings(self):
+ return self.model.get_input_embeddings()
+
+ def set_input_embeddings(self, value):
+ self.model.set_input_embeddings(value)
+
+ def get_output_embeddings(self):
+ return self.model.lm_head
+
+ def set_output_embeddings(self, value):
+ self.model.lm_head = value
+
+ def get_decoder(self):
+ return self.model.language_model
+
+ def set_decoder(self, decoder):
+ self.model.language_model = decoder
+
+ def forward(
+ self,
+ input_ids: Optional[torch.LongTensor] = None,
+ attention_mask: Optional[torch.Tensor] = None,
+ position_ids: Optional[torch.LongTensor] = None,
+ past_key_values: Optional[Cache] = None,
+ inputs_embeds: Optional[torch.FloatTensor] = None,
+ labels: Optional[torch.LongTensor] = None,
+ use_cache: Optional[bool] = None,
+ output_attentions: Optional[bool] = None,
+ output_router_logits: Optional[bool] = None,
+ output_hidden_states: Optional[bool] = None,
+ return_dict: Optional[bool] = None,
+ logits_to_keep: Union[int, torch.Tensor] = 0,
+ cu_lengths_list: Optional[List] = None,
+ **kwargs,
+ ) -> Union[tuple, LLaDA2MoeCausalLMOutputWithPast]:
+ return_dict = (
+ return_dict if return_dict is not None else self.config.use_return_dict
+ )
+ if inputs_embeds is None:
+ if input_ids is None:
+ raise ValueError("Provide either input_ids or inputs_embeds")
+ inputs_embeds = self.get_input_embeddings()(input_ids)
+
+ outputs = self.model(
+ input_ids=None,
+ attention_mask=attention_mask,
+ position_ids=position_ids,
+ past_key_values=past_key_values,
+ inputs_embeds=inputs_embeds,
+ use_cache=use_cache,
+ output_attentions=output_attentions,
+ output_router_logits=output_router_logits,
+ output_hidden_states=output_hidden_states,
+ return_dict=True,
+ cu_lengths_list=cu_lengths_list,
+ **kwargs,
+ )
+ hidden_states = outputs.last_hidden_state
+ indices = (
+ slice(-logits_to_keep, None)
+ if isinstance(logits_to_keep, int)
+ else logits_to_keep
+ )
+ logits = self.model.lm_head(hidden_states[:, indices, :])
+
+ loss = None
+ if labels is not None:
+ loss = F.cross_entropy(
+ logits.reshape(-1, logits.shape[-1]), labels.reshape(-1)
+ )
+
+ if not return_dict:
+ result = (
+ logits,
+ outputs.past_key_values,
+ outputs.hidden_states,
+ outputs.attentions,
+ )
+ return ((loss,) + result) if loss is not None else result
+ return LLaDA2MoeCausalLMOutputWithPast(
+ loss=loss,
+ z_loss=None,
+ logits=logits,
+ past_key_values=outputs.past_key_values,
+ hidden_states=outputs.hidden_states,
+ attentions=outputs.attentions,
+ rope_deltas=None,
+ )
+
+ @staticmethod
+ def _top_k_logits(logits, k):
+ if k is None or k <= 0:
+ return logits
+ values, _ = torch.topk(logits, min(k, logits.shape[-1]))
+ return torch.where(logits < values[..., -1, None], -torch.inf, logits)
+
+ @staticmethod
+ def _top_p_logits(logits, p):
+ if p is None or p >= 1.0:
+ return logits
+ sorted_logits, sorted_indices = torch.sort(logits, descending=True)
+ cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
+ sorted_mask = cumulative_probs > p
+ sorted_mask[..., 1:] = sorted_mask[..., :-1].clone()
+ sorted_mask[..., 0] = False
+ mask = torch.zeros_like(sorted_mask).scatter(-1, sorted_indices, sorted_mask)
+ return logits.masked_fill(mask, -torch.inf)
+
+ def _sample_with_temperature_topk_topp(
+ self, logits, temperature=1.0, top_k=0, top_p=1.0
+ ):
+ original_shape = logits.shape[:-1]
+ logits = logits.reshape(-1, logits.shape[-1])
+ if (
+ temperature == 0.0
+ and (top_k in (None, 0))
+ and (top_p is None or top_p >= 1.0)
+ ):
+ probs = F.softmax(logits, dim=-1)
+ token = logits.argmax(dim=-1, keepdim=True)
+ token_prob = probs.gather(-1, token)
+ return token.view(*original_shape), token_prob.view(*original_shape)
+ if temperature > 0 and temperature != 1.0:
+ logits = logits / temperature
+ logits = self._top_k_logits(logits, top_k)
+ logits = self._top_p_logits(logits, top_p)
+ probs = F.softmax(logits, dim=-1)
+ token = torch.multinomial(probs, num_samples=1)
+ token_prob = probs.gather(-1, token)
+ return token.view(*original_shape), token_prob.view(*original_shape)
+
+ @staticmethod
+ def _get_num_transfer_tokens(block_length, steps):
+ if steps == 0:
+ return torch.empty(0, dtype=torch.int64)
+ schedule = torch.full((steps,), block_length // steps, dtype=torch.int64)
+ schedule[: block_length % steps] += 1
+ return schedule
+
+ @torch.no_grad()
+ def generate_bd_image_logic(
+ self,
+ data: Optional[dict] = None,
+ temperature: float = 0.0,
+ block_length: int = 32,
+ steps: int = 32,
+ gen_length: int = 2048,
+ top_p: Optional[float] = None,
+ top_k: Optional[int] = None,
+ eos_early_stop: bool = True,
+ minimal_topk: int = 1,
+ threshold: float = 0.95,
+ eos_id: int = 156892,
+ mask_id: int = 156895,
+ cfg_scale: float = 1.0,
+ mode: str = "eoi",
+ ):
+ """Generate discrete image tokens with the original block-diffusion logic."""
+ if data is None or "input_ids" not in data:
+ raise ValueError("data must contain input_ids")
+ steps = min(steps, gen_length // minimal_topk)
+ input_ids = data["input_ids"]
+ eoi_id = 156902
+ prompt_length = input_ids.shape[1]
+ num_blocks = (prompt_length + gen_length + block_length - 1) // block_length
+ total_length = num_blocks * block_length
+
+ block_mask = torch.tril(torch.ones(num_blocks, num_blocks, device=self.device))
+ full_attention_mask = (
+ block_mask.repeat_interleave(block_length, 0)
+ .repeat_interleave(block_length, 1)[None, None]
+ .bool()
+ )
+ position_ids = torch.arange(total_length, device=self.device).unsqueeze(0)
+ x = torch.full((1, total_length), mask_id, dtype=torch.long, device=self.device)
+ x[:, :prompt_length] = input_ids
+ prefill_blocks = prompt_length // block_length
+ schedule = self._get_num_transfer_tokens(block_length, steps)
+ use_cfg = cfg_scale != 1.0
+
+ if use_cfg:
+ uncond_ids = data.get("uncond_ids", [27, 411, 19483, 29])
+ if torch.is_tensor(uncond_ids):
+ uncond_ids = uncond_ids.flatten().tolist()
+ pad_len = prompt_length - len(uncond_ids)
+ if pad_len < 0:
+ raise ValueError(
+ "The unconditional prompt is longer than the conditional prompt"
+ )
+ uncond_input = torch.full(
+ (1, prompt_length), mask_id, dtype=torch.long, device=self.device
+ )
+ uncond_input[0, -len(uncond_ids) :] = torch.tensor(
+ uncond_ids, device=self.device
+ )
+ uncond_attention_mask = full_attention_mask.clone()
+ uncond_attention_mask[:, :, :, :pad_len] = False
+ uncond_position_ids = torch.cat(
+ [
+ torch.zeros(pad_len, device=self.device, dtype=torch.long),
+ torch.arange(total_length - pad_len, device=self.device),
+ ]
+ ).unsqueeze(0)
+
+ for block_index in range(prefill_blocks, num_blocks):
+ window_end = (block_index + 1) * block_length
+ current = x[:, :window_end]
+ current_mask = full_attention_mask[:, :, :window_end, :window_end]
+ current_positions = position_ids[:, :window_end]
+
+ for step_index in range(steps):
+ active = current[:, -block_length:] == mask_id
+ if not active.any():
+ break
+ if use_cfg:
+ unconditional = current.clone()
+ unconditional[:, :prompt_length] = uncond_input
+ combined_ids = torch.cat([current, unconditional], dim=0)
+ combined_positions = torch.cat(
+ [current_positions, uncond_position_ids[:, :window_end]], dim=0
+ )
+ combined_mask = torch.cat(
+ [
+ current_mask,
+ uncond_attention_mask[:, :, :window_end, :window_end],
+ ],
+ dim=0,
+ )
+ logits = self(
+ input_ids=combined_ids,
+ attention_mask=combined_mask,
+ position_ids=combined_positions,
+ ).logits
+ conditional_logits, unconditional_logits = logits.chunk(2, dim=0)
+ active_logits = unconditional_logits[
+ :, -block_length:
+ ] + cfg_scale * (
+ conditional_logits[:, -block_length:]
+ - unconditional_logits[:, -block_length:]
+ )
+ else:
+ active_logits = self(
+ input_ids=current,
+ attention_mask=current_mask,
+ position_ids=current_positions,
+ ).logits[:, -block_length:]
+
+ tokens, confidence = self._sample_with_temperature_topk_topp(
+ active_logits, temperature=temperature, top_k=top_k, top_p=top_p
+ )
+ count = schedule[step_index].item()
+ scores = torch.where(active, confidence, -torch.inf)
+ selected = torch.zeros_like(tokens, dtype=torch.bool)
+ high_confidence = scores[0] > threshold
+ if high_confidence.sum().item() >= count:
+ selected[0] = high_confidence
+ else:
+ _, indices = torch.topk(
+ scores[0], k=min(count, active.sum().item())
+ )
+ selected[0, indices] = True
+ current[:, -block_length:][selected] = tokens[selected]
+
+ stop_token = eoi_id if mode == "eoi" else eos_id
+ positions = (current[0, prompt_length:] == stop_token).nonzero(
+ as_tuple=True
+ )[0]
+ if eos_early_stop and len(positions) > 0:
+ stop_position = positions[0].item() + prompt_length
+ if (current[0, prompt_length:stop_position] != mask_id).all():
+ x[:, :window_end] = current
+ return x[:, : stop_position + 1]
+
+ x[:, :window_end] = current
+ return x[:, : prompt_length + gen_length]
+
+
+__all__ = ["LLaDA2MoeModelLM", "LLaDA2MoeModel", "LLaDA2MoePreTrainedModel"]
diff --git a/pipelines/llada/pipeline_llada_image.py b/pipelines/llada/pipeline_llada_image.py
new file mode 100644
index 000000000..eabe87622
--- /dev/null
+++ b/pipelines/llada/pipeline_llada_image.py
@@ -0,0 +1,687 @@
+# 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
+from pathlib import Path
+
+import torch
+import torch.nn.functional as F
+from transformers import AutoModel, AutoTokenizer, PreTrainedModel, PreTrainedTokenizerBase
+
+from diffusers.image_processor import PipelineImageInput, VaeImageProcessor
+from diffusers.models import AutoencoderKLFlux2
+from diffusers.pipelines.pipeline_utils import DiffusionPipeline
+from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
+from diffusers.utils import logging
+from diffusers.utils.torch_utils import randn_tensor
+
+from .transformer_llada_image import (
+ LLaDAImageQueryFormerModel,
+ LLaDAImageSigVQModel,
+ LLaDAImageTextProjectionModel,
+ LLaDAImageTransformer2DModel,
+)
+from .pipeline_output import LLaDAImagePipelineOutput
+
+
+logger = logging.get_logger(__name__)
+
+
+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"
+ _callback_tensor_inputs = ["latents", "noise_pred"]
+
+ 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)
+
+ @classmethod
+ def from_pretrained( # pylint: disable=arguments-differ
+ cls,
+ pretrained_model_name_or_path: str | Path,
+ *,
+ torch_dtype: torch.dtype | None = None,
+ device: torch.device | str | None = None,
+ cache_dir: str | None = None,
+ scheduler=None,
+ vae=None,
+ text_encoder=None,
+ tokenizer=None,
+ queryformer=None,
+ text_projection=None,
+ sigvq=None,
+ transformer=None,
+ **kwargs,
+ ) -> "LLaDAImagePipeline":
+ """Load all LLaDA-Image components from a converted model directory or Hugging Face repository.
+
+ This model stores a LLaDA2 text encoder that requires `trust_remote_code=True`, so its components are loaded
+ explicitly instead of relying on the generic Diffusers pipeline resolver.
+ """
+ model_path = Path(pretrained_model_name_or_path)
+ if not model_path.is_dir():
+ from huggingface_hub import snapshot_download
+
+ ignore_patterns = ["assets/**"]
+ if text_encoder is not None:
+ ignore_patterns.append("text_encoder/**")
+ if transformer is not None:
+ ignore_patterns.append("transformer/**")
+ model_path = Path(
+ snapshot_download(
+ repo_id=str(pretrained_model_name_or_path),
+ cache_dir=cache_dir,
+ ignore_patterns=ignore_patterns or None,
+ )
+ )
+
+ if not (model_path / "model_index.json").is_file():
+ raise ValueError(
+ "Expected a converted LLaDA-Image model directory containing `model_index.json`, got "
+ f"{model_path}."
+ )
+
+ if scheduler is None:
+ scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(model_path / "scheduler", cache_dir=cache_dir)
+ if vae is None:
+ vae = AutoencoderKLFlux2.from_pretrained(model_path / "vae", torch_dtype=torch_dtype, cache_dir=cache_dir)
+ if text_encoder is None:
+ text_encoder_kwargs = {"dtype": torch_dtype, "trust_remote_code": True, "cache_dir": cache_dir}
+ if device is not None:
+ text_encoder_kwargs["device_map"] = {"": device}
+ text_encoder = AutoModel.from_pretrained(model_path / "text_encoder", **text_encoder_kwargs)
+ if tokenizer is None:
+ tokenizer = AutoTokenizer.from_pretrained(model_path / "tokenizer", cache_dir=cache_dir)
+ if queryformer is None:
+ queryformer = LLaDAImageQueryFormerModel.from_pretrained(
+ model_path / "queryformer", torch_dtype=torch_dtype, cache_dir=cache_dir
+ )
+ if text_projection is None:
+ text_projection = LLaDAImageTextProjectionModel.from_pretrained(
+ model_path / "text_projection", torch_dtype=torch_dtype, cache_dir=cache_dir
+ )
+ if sigvq is None:
+ sigvq = LLaDAImageSigVQModel.from_pretrained(
+ model_path / "sigvq", torch_dtype=torch_dtype, cache_dir=cache_dir
+ )
+ if transformer is None:
+ transformer = LLaDAImageTransformer2DModel.from_pretrained(
+ model_path / "transformer", torch_dtype=torch_dtype, cache_dir=cache_dir
+ )
+
+ if device is not None:
+ vae = vae.to(device)
+ queryformer = queryformer.to(device)
+ text_projection = text_projection.to(device)
+ sigvq = sigvq.to(device)
+ transformer = transformer.to(device)
+
+ return cls(
+ scheduler=scheduler,
+ vae=vae,
+ text_encoder=text_encoder,
+ tokenizer=tokenizer,
+ queryformer=queryformer,
+ text_projection=text_projection,
+ sigvq=sigvq,
+ transformer=transformer,
+ )
+
+ @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",
+ )
+ input_ids = text_inputs.input_ids.to(self.text_encoder.device)
+ attention_mask = text_inputs.attention_mask.to(input_ids.device).bool()
+ inputs_embeds = self.text_encoder.get_input_embeddings()(input_ids)
+ text_encoder_device = inputs_embeds.device
+ attention_mask = attention_mask.to(text_encoder_device)
+
+ query_embeds = self.queryformer(
+ inputs_embeds.to(device=self.queryformer.device, dtype=self.queryformer.dtype),
+ attention_mask.to(self.queryformer.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.text_projection.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:
+ 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 = []
+
+ 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.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)
+
+ 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.sigvq.device, dtype=self.sigvq.dtype)
+ ).semantic_features
+
+ source_latents = self.vae.encode(image.to(device=self.vae.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.transformer.device
+ self._guidance_scale = guidance_scale # pylint: disable=attribute-defined-outside-init
+ 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.sigvq.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) # pylint: disable=attribute-defined-outside-init
+
+ cond_cap_feats = [
+ embeds[mask].to(device=self.transformer.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=self.transformer.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=self.transformer.device, dtype=self.transformer.dtype)
+ for features in semantic_features
+ ]
+ if source_latents is not None:
+ source_latent_list = [
+ latent.unsqueeze(1).to(device=self.transformer.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=self.transformer.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]
+ )
+
+ noise_pred = 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
+ noise_pred = -torch.stack(noise_pred, dim=0).squeeze(2).float()
+
+ if do_classifier_free_guidance:
+ conditional_output, unconditional_output = noise_pred.chunk(2)
+ noise_pred = unconditional_output + self.guidance_scale * (
+ conditional_output - unconditional_output
+ )
+
+ latents = self.scheduler.step(noise_pred, 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=self.vae.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/pipelines/llada/pipeline_output.py b/pipelines/llada/pipeline_output.py
new file mode 100644
index 000000000..f79efaa72
--- /dev/null
+++ b/pipelines/llada/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/pipelines/llada/transformer_llada_image.py b/pipelines/llada/transformer_llada_image.py
new file mode 100644
index 000000000..d1aa783d2
--- /dev/null
+++ b/pipelines/llada/transformer_llada_image.py
@@ -0,0 +1,1874 @@
+# 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 diffusers.configuration_utils import ConfigMixin, register_to_config
+from diffusers.models.attention import AttentionMixin, AttentionModuleMixin, FeedForward
+from diffusers.models.attention_dispatch import dispatch_attention_fn
+from diffusers.models.modeling_outputs import Transformer2DModelOutput
+from diffusers.models.modeling_utils import ModelMixin
+from diffusers.models.normalization import RMSNorm
+from diffusers.utils import BaseOutput
+from diffusers.utils.torch_utils import maybe_allow_in_graph
+
+
+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 self.freqs_cis is None or self.freqs_cis[0].device != position_ids.device:
+ self.freqs_cis = self._create_frequencies(position_ids.device)
+
+ frequencies = []
+ for axis, axis_frequencies in enumerate(self.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/pipelines/model_llada.py b/pipelines/model_llada.py
new file mode 100644
index 000000000..647cd059c
--- /dev/null
+++ b/pipelines/model_llada.py
@@ -0,0 +1,80 @@
+import diffusers
+from modules import shared, devices, sd_models, sd_hijack_te, sd_hijack_vae
+from modules.logger import log
+from pipelines import generic
+
+
+def load_llada_image(checkpoint_info, diffusers_load_config=None):
+ if diffusers_load_config is None:
+ diffusers_load_config = {}
+ repo_id = sd_models.path_to_repo(checkpoint_info)
+ sd_models.hf_auth_check(checkpoint_info)
+ log.debug(f'Load model: type=LLaDAImage repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}')
+
+ from pipelines.llada import LLaDAImagePipeline
+ from pipelines.llada.transformer_llada_image import LLaDAImageTransformer2DModel
+ from pipelines.llada.modeling_llada2uni_moe import LLaDA2MoeModelLM
+
+ generic.set_pipeline('LLaDAImage', LLaDAImagePipeline)
+ if repo_id is None or repo_id.lower() == 'none':
+ return None
+
+ if 'Model' in shared.opts.sdnq_quantize_weights:
+ if any(x in shared.opts.sdnq_quantize_weights_mode for x in ['2', '3', '4', '5', '6']):
+ shared.opts.sdnq_quantize_weights_mode = 'uint8'
+ log.warning('LLaDAImage: cls=LLaDAImageTransformer2DModel quant=uint8 override')
+ if 'TE' in shared.opts.sdnq_quantize_weights:
+ if any(x in shared.opts.sdnq_quantize_weights_mode_te for x in ['2', '3', '4', '5', '6']):
+ shared.opts.sdnq_quantize_weights_mode_te = 'uint8'
+ log.warning('LLaDAImage: cls=LLaDA2MoeModelLM quant=uint8 override')
+ if shared.opts.sdnq_quantize_matmul_mode_te != 'disabled':
+ shared.opts.sdnq_quantize_matmul_mode_te = 'disabled'
+ log.warning('LLaDAImage: cls=LLaDA2MoeModelLM matmul=disabled override')
+
+ transformer = generic.load_transformer(
+ repo_id,
+ cls_name=LLaDAImageTransformer2DModel,
+ load_config=diffusers_load_config,
+ modules_to_not_convert=[
+ 'all_x_embedder',
+ 'all_final_layer',
+ 't_embedder',
+ 'cap_embedder',
+ 'semantic_embedder',
+ 'sigvq_embedder',
+ ],
+ )
+ text_encoder = generic.load_text_encoder(
+ repo_id,
+ cls_name=LLaDA2MoeModelLM,
+ load_config=diffusers_load_config,
+ allow_shared=False,
+ trust_remote_code=True,
+ modules_to_not_convert=[
+ '.model.language_model.word_embeddings',
+ '.model.language_model.norm',
+ '.model.lm_head',
+ ],
+ )
+
+ diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING['llada-image'] = LLaDAImagePipeline
+ diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING['llada-image'] = LLaDAImagePipeline
+
+ pipe = LLaDAImagePipeline.from_pretrained(
+ repo_id,
+ cache_dir=shared.opts.diffusers_dir,
+ torch_dtype=devices.dtype,
+ transformer=transformer,
+ text_encoder=text_encoder,
+ )
+ pipe.task_args = {
+ 'output_type': 'np',
+ 'generation_mode': 'text',
+ }
+ # generation_mode = "text", "vq", "editing"
+
+ del transformer, text_encoder
+ sd_hijack_te.init_hijack(pipe)
+ sd_hijack_vae.init_hijack(pipe)
+ devices.torch_gc(force=True, reason='load')
+ return pipe