@@ -33,6 +33,10 @@ class EXM_HYDiT_Model(comfy.model_base.BaseModel):
|
|||||||
for name in ["context_t5", "context_mask", "context_t5_mask"]:
|
for name in ["context_t5", "context_mask", "context_t5_mask"]:
|
||||||
out[name] = comfy.conds.CONDRegular(kwargs[name])
|
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
|
return out
|
||||||
|
|
||||||
def load_hydit(model_path, model_conf):
|
def load_hydit(model_path, model_conf):
|
||||||
|
|||||||
@@ -212,6 +212,7 @@ class HunYuanDiT(nn.Module):
|
|||||||
self.extra_in_dim = 256 * 6 + hidden_size
|
self.extra_in_dim = 256 * 6 + hidden_size
|
||||||
|
|
||||||
# Text embedding for `add`
|
# Text embedding for `add`
|
||||||
|
self.last_size = input_size
|
||||||
self.x_embedder = PatchEmbed(input_size, patch_size, in_channels, hidden_size)
|
self.x_embedder = PatchEmbed(input_size, patch_size, in_channels, hidden_size)
|
||||||
self.t_embedder = TimestepEmbedder(hidden_size)
|
self.t_embedder = TimestepEmbedder(hidden_size)
|
||||||
self.extra_in_dim += 1024
|
self.extra_in_dim += 1024
|
||||||
@@ -358,7 +359,7 @@ class HunYuanDiT(nn.Module):
|
|||||||
rope = get_2d_rotary_pos_embed(self.head_size, *sub_args)
|
rope = get_2d_rotary_pos_embed(self.head_size, *sub_args)
|
||||||
return rope
|
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
|
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)
|
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)
|
# context_t5_mask = torch.zeros(x.shape[0], 256, device=x.device)
|
||||||
|
|
||||||
# style
|
# 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
|
# image size - todo separate for cond/uncond when batched
|
||||||
width = x.shape[2]
|
if torch.is_tensor(src_size_cond):
|
||||||
height = x.shape[3]
|
src_size_cond = (int(src_size_cond[0][0]), int(src_size_cond[0][1]))
|
||||||
src_size_cond = (width//2*16, height//2*16)
|
|
||||||
size_cond = list(src_size_cond) + [width*8, height*8, 0, 0]
|
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)
|
image_meta_size = torch.as_tensor([size_cond] * x.shape[0], device=x.device)
|
||||||
|
|
||||||
# RoPE
|
# 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
|
# Run original forward pass
|
||||||
out = self.forward_raw(
|
out = self.forward_raw(
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
import os
|
import os
|
||||||
import folder_paths
|
import folder_paths
|
||||||
|
from copy import deepcopy
|
||||||
|
|
||||||
from .conf import hydit_conf
|
from .conf import hydit_conf
|
||||||
from .loader import load_hydit
|
from .loader import load_hydit
|
||||||
@@ -163,9 +164,35 @@ class HYDiTTextEncodeSimple(HYDiTTextEncode):
|
|||||||
def encode_simple(self, text, **args):
|
def encode_simple(self, text, **args):
|
||||||
return self.encode(text=text, text_t5=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 = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"HYDiTCheckpointLoader": HYDiTCheckpointLoader,
|
"HYDiTCheckpointLoader": HYDiTCheckpointLoader,
|
||||||
"HYDiTTextEncoderLoader": HYDiTTextEncoderLoader,
|
"HYDiTTextEncoderLoader": HYDiTTextEncoderLoader,
|
||||||
"HYDiTTextEncode": HYDiTTextEncode,
|
"HYDiTTextEncode": HYDiTTextEncode,
|
||||||
"HYDiTTextEncodeSimple": HYDiTTextEncodeSimple,
|
"HYDiTTextEncodeSimple": HYDiTTextEncodeSimple,
|
||||||
|
"HYDiTSrcSizeCond": HYDiTSrcSizeCond,
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user