From 96cc89babb401d720e02a2bed15d883d23f84221 Mon Sep 17 00:00:00 2001 From: Haoming Date: Mon, 8 Dec 2025 16:15:49 +0800 Subject: [PATCH] lumina --- backend/nn/lumina.py | 72 ++++++++++++------- .../packages/huggingface_guess/model_list.py | 2 + 2 files changed, 48 insertions(+), 26 deletions(-) diff --git a/backend/nn/lumina.py b/backend/nn/lumina.py index b89628eb..4464aa50 100644 --- a/backend/nn/lumina.py +++ b/backend/nn/lumina.py @@ -1,4 +1,4 @@ -# https://github.com/comfyanonymous/ComfyUI/blob/v0.3.75/comfy/ldm/lumina/model.py +# https://github.com/comfyanonymous/ComfyUI/blob/v0.3.77/comfy/ldm/lumina/model.py # Reference: https://github.com/Alpha-VLLM/Lumina-Image-2.0 import math @@ -15,8 +15,8 @@ if xformers_enabled(): else: from backend.attention import attention_pytorch as attention_function -from backend.args import dynamic_args -from backend.nn.flux import EmbedND +from backend.nn.flux import EmbedND, apply_rope +from backend.utils import fp16_fix as clamp_fp16 from backend.utils import pad_to_patch_size @@ -73,12 +73,6 @@ class JointAttention(nn.Module): else: self.q_norm = self.k_norm = nn.Identity() - @staticmethod - def apply_rotary_emb(x_in: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor: - t_ = x_in.reshape(*x_in.shape[:-1], -1, 1, 2) - t_out = freqs_cis[..., 0] * t_[..., 0] + freqs_cis[..., 1] * t_[..., 1] - return t_out.reshape(*x_in.shape) - def forward(self, x: torch.Tensor, x_mask: torch.Tensor, freqs_cis: torch.Tensor, transformer_options={}) -> torch.Tensor: bsz, seqlen, _ = x.shape @@ -98,8 +92,7 @@ class JointAttention(nn.Module): xq = self.q_norm(xq) xk = self.k_norm(xk) - xq = JointAttention.apply_rotary_emb(xq, freqs_cis=freqs_cis) - xk = JointAttention.apply_rotary_emb(xk, freqs_cis=freqs_cis) + xq, xk = apply_rope(xq, xk, freqs_cis) n_rep = self.n_local_heads // self.n_local_kv_heads if n_rep >= 1: @@ -122,7 +115,7 @@ class FeedForward(nn.Module): self.w3 = nn.Linear(dim, hidden_dim, bias=False) def _forward_silu_gating(self, x1, x3): - return F.silu(x1) * x3 + return clamp_fp16(F.silu(x1) * x3) def forward(self, x): return self.w2(self._forward_silu_gating(self.w1(x), self.w3(x))) @@ -165,26 +158,32 @@ class JointTransformerBlock(nn.Module): scale_msa, gate_msa, scale_mlp, gate_mlp = self.adaLN_modulation(adaln_input).chunk(4, dim=1) x = x + gate_msa.unsqueeze(1).tanh() * self.attention_norm2( - self.attention( - modulate(self.attention_norm1(x), scale_msa), - x_mask, - freqs_cis, - transformer_options=transformer_options, + clamp_fp16( + self.attention( + modulate(self.attention_norm1(x), scale_msa), + x_mask, + freqs_cis, + transformer_options=transformer_options, + ) ) ) x = x + gate_mlp.unsqueeze(1).tanh() * self.ffn_norm2( - self.feed_forward( - modulate(self.ffn_norm1(x), scale_mlp), + clamp_fp16( + self.feed_forward( + modulate(self.ffn_norm1(x), scale_mlp), + ) ) ) else: assert adaln_input is None x = x + self.attention_norm2( - self.attention( - self.attention_norm1(x), - x_mask, - freqs_cis, - transformer_options=transformer_options, + clamp_fp16( + self.attention( + self.attention_norm1(x), + x_mask, + freqs_cis, + transformer_options=transformer_options, + ) ) ) x = x + self.ffn_norm2( @@ -241,6 +240,7 @@ class NextDiT(nn.Module): z_image_modulation: bool = False, time_scale: float = 1.0, pad_tokens_multiple: int = None, + clip_text_dim=None, ): super().__init__() @@ -292,6 +292,19 @@ class NextDiT(nn.Module): nn.Linear(cap_feat_dim, dim, bias=True), ) + self.clip_text_pooled_proj = None + + if clip_text_dim is not None: + self.clip_text_dim = clip_text_dim + self.clip_text_pooled_proj = nn.Sequential( + nn.RMSNorm(clip_text_dim, eps=norm_eps, elementwise_affine=True), + nn.Linear(clip_text_dim, clip_text_dim, bias=True), + ) + self.time_text_embed = nn.Sequential( + nn.SiLU(), + nn.Linear(min(dim, 1024) + clip_text_dim, min(dim, 1024), bias=True), + ) + self.layers = nn.ModuleList( [ JointTransformerBlock( @@ -379,7 +392,7 @@ class NextDiT(nn.Module): l_effective_cap_len = [cap_feats.shape[1]] * bsz return padded_full_embed, mask, img_sizes, l_effective_cap_len, freqs_cis - def forward(self, x, timesteps, context, num_tokens=None, attention_mask=None, **kwargs): + def forward(self, x, timesteps, context, num_tokens=None, attention_mask=None, transformer_options={}, **kwargs): t = 1.0 - timesteps cap_feats = context cap_mask = attention_mask @@ -391,7 +404,14 @@ class NextDiT(nn.Module): cap_feats = self.cap_embedder(cap_feats) # (N, L, D) # todo check if able to batchify w.o. redundant compute - transformer_options = kwargs.get("transformer_options", {}) + if self.clip_text_pooled_proj is not None: + pooled = kwargs.get("clip_text_pooled", None) + if pooled is not None: + pooled = self.clip_text_pooled_proj(pooled) + else: + pooled = torch.zeros((1, self.clip_text_dim), device=x.device, dtype=x.dtype) + + adaln_input = self.time_text_embed(torch.cat((t, pooled), dim=-1)) x_is_tensor = isinstance(x, torch.Tensor) x, mask, img_size, cap_size, freqs_cis = self.patchify_and_embed(x, cap_feats, cap_mask, t, num_tokens, transformer_options=transformer_options) diff --git a/modules_forge/packages/huggingface_guess/model_list.py b/modules_forge/packages/huggingface_guess/model_list.py index 7a213aa5..93b3fa25 100644 --- a/modules_forge/packages/huggingface_guess/model_list.py +++ b/modules_forge/packages/huggingface_guess/model_list.py @@ -367,6 +367,8 @@ class ZImage(Lumina2): memory_usage_factor = 1.7 + supported_inference_dtypes = [torch.bfloat16, torch.float16, torch.float32] + def clip_target(self, state_dict={}): return {"qwen3_4b.transformer": "text_encoder"}