From cb0aea8eb5deadce32c95d05cf53f91b7aae2400 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 5 Sep 2026 13:01:36 +0200 Subject: [PATCH] llada t2i Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 4 + modules/modeldata.py | 2 + modules/sd_detect.py | 2 + modules/sd_models.py | 5 + modules/vae/sd_vae_taesd.py | 2 +- pipelines/llada/__init__.py | 8 + .../llada/configuration_llada2uni_moe.py | 133 ++ pipelines/llada/fused_moe_ops.py | 375 ++++ pipelines/llada/modeling_llada2uni_moe.py | 1284 +++++++++++ pipelines/llada/pipeline_llada_image.py | 687 ++++++ pipelines/llada/pipeline_output.py | 34 + pipelines/llada/transformer_llada_image.py | 1874 +++++++++++++++++ pipelines/model_llada.py | 80 + 13 files changed, 4489 insertions(+), 1 deletion(-) create mode 100644 pipelines/llada/__init__.py create mode 100644 pipelines/llada/configuration_llada2uni_moe.py create mode 100644 pipelines/llada/fused_moe_ops.py create mode 100644 pipelines/llada/modeling_llada2uni_moe.py create mode 100644 pipelines/llada/pipeline_llada_image.py create mode 100644 pipelines/llada/pipeline_output.py create mode 100644 pipelines/llada/transformer_llada_image.py create mode 100644 pipelines/model_llada.py 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