HunYuanDiT fix image size cond logic

#54
This commit is contained in:
City
2024-05-28 00:18:44 +02:00
parent 014e067483
commit b773f348e6
3 changed files with 48 additions and 8 deletions
+4
View File
@@ -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):
+17 -8
View File
@@ -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(
+27
View File
@@ -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,
}