Support all latent resolutions.

This commit is contained in:
comfyanonymous 2025-06-25 19:32:07 -04:00
parent 83491965da
commit 9c90edf403

View File

@ -10,6 +10,7 @@ from comfy.ldm.lightricks.model import Timesteps
from comfy.ldm.flux.layers import EmbedND from comfy.ldm.flux.layers import EmbedND
from comfy.ldm.modules.attention import optimized_attention_masked from comfy.ldm.modules.attention import optimized_attention_masked
import comfy.model_management import comfy.model_management
import comfy.ldm.common_dit
def apply_rotary_emb(x, freqs_cis): def apply_rotary_emb(x, freqs_cis):
@ -362,6 +363,7 @@ class OmniGen2Transformer2DModel(nn.Module):
l_effective_img_len = [(H // p) * (W // p) for (H, W) in img_sizes] l_effective_img_len = [(H // p) * (W // p) for (H, W) in img_sizes]
if ref_image_hidden_states is not None: if ref_image_hidden_states is not None:
ref_image_hidden_states = list(map(lambda ref: comfy.ldm.common_dit.pad_to_patch_size(ref, (p, p)), ref_image_hidden_states))
ref_img_sizes = [[(imgs.size(2), imgs.size(3)) if imgs is not None else None for imgs in ref_image_hidden_states]] * batch_size ref_img_sizes = [[(imgs.size(2), imgs.size(3)) if imgs is not None else None for imgs in ref_image_hidden_states]] * batch_size
l_effective_ref_img_len = [[(ref_img_size[0] // p) * (ref_img_size[1] // p) for ref_img_size in _ref_img_sizes] if _ref_img_sizes is not None else [0] for _ref_img_sizes in ref_img_sizes] l_effective_ref_img_len = [[(ref_img_size[0] // p) * (ref_img_size[1] // p) for ref_img_size in _ref_img_sizes] if _ref_img_sizes is not None else [0] for _ref_img_sizes in ref_img_sizes]
else: else:
@ -415,7 +417,8 @@ class OmniGen2Transformer2DModel(nn.Module):
def forward(self, x, timesteps, context, num_tokens, ref_latents=None, attention_mask=None, **kwargs): def forward(self, x, timesteps, context, num_tokens, ref_latents=None, attention_mask=None, **kwargs):
B, C, H, W = x.shape B, C, H, W = x.shape
hidden_states = x hidden_states = comfy.ldm.common_dit.pad_to_patch_size(x, (self.patch_size, self.patch_size))
_, _, H_padded, W_padded = hidden_states.shape
timestep = 1.0 - timesteps timestep = 1.0 - timesteps
text_hidden_states = context text_hidden_states = context
text_attention_mask = attention_mask text_attention_mask = attention_mask
@ -461,6 +464,6 @@ class OmniGen2Transformer2DModel(nn.Module):
hidden_states = self.norm_out(hidden_states, temb) hidden_states = self.norm_out(hidden_states, temb)
p = self.patch_size p = self.patch_size
output = rearrange(hidden_states[:, -img_len:], 'b (h w) (p1 p2 c) -> b c (h p1) (w p2)', h=H // p, w=W// p, p1=p, p2=p) output = rearrange(hidden_states[:, -img_len:], 'b (h w) (p1 p2 c) -> b c (h p1) (w p2)', h=H_padded // p, w=W_padded// p, p1=p, p2=p)[:, :, :H, :W]
return -output return -output