From b773f348e688fac3ba112eeb3634584643679ec5 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Tue, 28 May 2024 00:18:44 +0200 Subject: [PATCH] HunYuanDiT fix image size cond logic #54 --- HunYuanDiT/loader.py | 4 ++++ HunYuanDiT/models/models.py | 25 +++++++++++++++++-------- HunYuanDiT/nodes.py | 27 +++++++++++++++++++++++++++ 3 files changed, 48 insertions(+), 8 deletions(-) diff --git a/HunYuanDiT/loader.py b/HunYuanDiT/loader.py index 942b267..4d9726f 100644 --- a/HunYuanDiT/loader.py +++ b/HunYuanDiT/loader.py @@ -33,6 +33,10 @@ class EXM_HYDiT_Model(comfy.model_base.BaseModel): for name in ["context_t5", "context_mask", "context_t5_mask"]: out[name] = comfy.conds.CONDRegular(kwargs[name]) + src_size_cond = kwargs.get("src_size_cond", None) + if src_size_cond is not None: + out["src_size_cond"] = comfy.conds.CONDRegular(torch.tensor(src_size_cond)) + return out def load_hydit(model_path, model_conf): diff --git a/HunYuanDiT/models/models.py b/HunYuanDiT/models/models.py index f0309a9..2d54e31 100644 --- a/HunYuanDiT/models/models.py +++ b/HunYuanDiT/models/models.py @@ -212,6 +212,7 @@ class HunYuanDiT(nn.Module): self.extra_in_dim = 256 * 6 + hidden_size # Text embedding for `add` + self.last_size = input_size self.x_embedder = PatchEmbed(input_size, patch_size, in_channels, hidden_size) self.t_embedder = TimestepEmbedder(hidden_size) self.extra_in_dim += 1024 @@ -358,7 +359,7 @@ class HunYuanDiT(nn.Module): rope = get_2d_rotary_pos_embed(self.head_size, *sub_args) return rope - def forward(self, x, timesteps, context, context_mask=None, context_t5=None, context_t5_mask=None, image_meta_size=None, **kwargs): + def forward(self, x, timesteps, context, context_mask=None, context_t5=None, context_t5_mask=None, src_size_cond=(1024,1024), **kwargs): """ Forward pass that adapts comfy input to original forward function x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images) @@ -370,17 +371,25 @@ class HunYuanDiT(nn.Module): # context_t5_mask = torch.zeros(x.shape[0], 256, device=x.device) # style - style = torch.as_tensor([0, 0] * (x.shape[0]//2), device=x.device) + style = torch.as_tensor([0] * (x.shape[0]), device=x.device) - # image size - todo: from cond - width = x.shape[2] - height = x.shape[3] - src_size_cond = (width//2*16, height//2*16) - size_cond = list(src_size_cond) + [width*8, height*8, 0, 0] + # image size - todo separate for cond/uncond when batched + if torch.is_tensor(src_size_cond): + src_size_cond = (int(src_size_cond[0][0]), int(src_size_cond[0][1])) + + image_size = (x.shape[2]//2*16, x.shape[3]//2*16) + size_cond = list(src_size_cond) + [image_size[0], image_size[1], 0, 0] image_meta_size = torch.as_tensor([size_cond] * x.shape[0], device=x.device) # RoPE - rope = self.calc_rope(*src_size_cond) + rope = self.calc_rope(*image_size) + + # Update x_embedder if image size changed + if self.last_size != image_size: + from tqdm import tqdm + tqdm.write(f"HyDiT: New image size {image_size}") + self.x_embedder.update_image_size(image_size) + self.last_size = image_size # Run original forward pass out = self.forward_raw( diff --git a/HunYuanDiT/nodes.py b/HunYuanDiT/nodes.py index 8bc39f5..a9e1b6a 100644 --- a/HunYuanDiT/nodes.py +++ b/HunYuanDiT/nodes.py @@ -1,5 +1,6 @@ import os import folder_paths +from copy import deepcopy from .conf import hydit_conf from .loader import load_hydit @@ -163,9 +164,35 @@ class HYDiTTextEncodeSimple(HYDiTTextEncode): def encode_simple(self, text, **args): return self.encode(text=text, text_t5=text, **args) +class HYDiTSrcSizeCond: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "cond": ("CONDITIONING", ), + "width": ("INT", {"default": 1024.0, "min": 0, "max": 8192, "step": 16}), + "height": ("INT", {"default": 1024.0, "min": 0, "max": 8192, "step": 16}), + } + } + + RETURN_TYPES = ("CONDITIONING",) + RETURN_NAMES = ("cond",) + FUNCTION = "add_cond" + CATEGORY = "ExtraModels/HunyuanDiT" + TITLE = "Hunyuan DiT Size Conditioning (advanced)" + + def add_cond(self, cond, width, height): + cond = deepcopy(cond) + for c in range(len(cond)): + cond[c][1].update({ + "src_size_cond": [[height, width]], + }) + return (cond,) + NODE_CLASS_MAPPINGS = { "HYDiTCheckpointLoader": HYDiTCheckpointLoader, "HYDiTTextEncoderLoader": HYDiTTextEncoderLoader, "HYDiTTextEncode": HYDiTTextEncode, "HYDiTTextEncodeSimple": HYDiTTextEncodeSimple, + "HYDiTSrcSizeCond": HYDiTSrcSizeCond, }