mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-03 22:27:13 +08:00
Support all latent resolutions.
This commit is contained in:
parent
83491965da
commit
9c90edf403
@ -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
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user