适配comfyui版本0.5.1

This commit is contained in:
刘雪峰
2026-01-12 12:48:33 +08:00
parent ebff62ede8
commit 190e79fb07
2 changed files with 16 additions and 4 deletions
+8 -2
View File
@@ -3,7 +3,10 @@ import math
import torch
from torch import Tensor
from comfy.ldm.lightricks.model import precompute_freqs_cis
try:
from comfy.ldm.lightricks.model import LTXBaseModel
except:
from comfy.ldm.lightricks.model import precompute_freqs_cis
from comfy.ldm.lightricks.symmetric_patchifier import latent_to_pixel_coords
from ..patch_util import PatchKeys
@@ -60,7 +63,10 @@ def ltx_forward_orig(
if attention_mask is not None and not torch.is_floating_point(attention_mask):
attention_mask = (attention_mask - 1).to(x.dtype).reshape((attention_mask.shape[0], 1, -1, attention_mask.shape[-1])) * torch.finfo(x.dtype).max
pe = precompute_freqs_cis(fractional_coords, dim=self.inner_dim, out_dtype=x.dtype)
try:
pe = LTXBaseModel.precompute_freqs_cis(fractional_coords, dim=self.inner_dim, out_dtype=x.dtype)
except:
pe = precompute_freqs_cis(fractional_coords, dim=self.inner_dim, out_dtype=x.dtype)
batch_size = x.shape[0]
timestep, embedded_timestep = self.adaln_single(
+8 -2
View File
@@ -3,7 +3,10 @@ import math
import torch
from torch import Tensor
from comfy.ldm.lightricks.model import precompute_freqs_cis
try:
from comfy.ldm.lightricks.model import LTXBaseModel
except:
from comfy.ldm.lightricks.model import precompute_freqs_cis
from ...patch_util import PatchKeys
# changed in comfyui hash commit 93fedd92fe0eb67a09e29069b05adebb40678639 (between comfyui version 0.3.19 and 1.3.20)
@@ -77,7 +80,10 @@ def ltx_forward_orig(
if attention_mask is not None and not torch.is_floating_point(attention_mask):
attention_mask = (attention_mask - 1).to(x.dtype).reshape((attention_mask.shape[0], 1, -1, attention_mask.shape[-1])) * torch.finfo(x.dtype).max
pe = precompute_freqs_cis(indices_grid, dim=self.inner_dim, out_dtype=x.dtype)
try:
pe = LTXBaseModel.precompute_freqs_cis(indices_grid, dim=self.inner_dim, out_dtype=x.dtype)
except:
pe = precompute_freqs_cis(indices_grid, dim=self.inner_dim, out_dtype=x.dtype)
batch_size = x.shape[0]
timestep, embedded_timestep = self.adaln_single(