Merge branch 'main' into longcat

This commit is contained in:
kijai
2025-10-29 02:33:37 +02:00
3 changed files with 30 additions and 27 deletions
+26
View File
@@ -3080,10 +3080,36 @@ class WanVideoSampler:
"samples": callback_latent.unsqueeze(0).cpu() if callback is not None else None,
})
class WanVideoSamplerSettings(WanVideoSampler):
RETURN_TYPES = ("SAMPLER_ARGS",)
RETURN_NAMES = ("sampler_inputs", )
DESCRIPTION = "Node to output all settings and inputs for the WanVideoSamplerFromSettings -node"
def process(self, *args, **kwargs):
import inspect
params = inspect.signature(WanVideoSampler.process).parameters
args_dict = {name: kwargs.get(name, param.default if param.default is not inspect.Parameter.empty else None)
for name, param in params.items() if name != "self"}
return args_dict,
class WanVideoSamplerFromSettings(WanVideoSampler):
DESCRIPTION = "Utility node with no other functionality than to look cleaner, useful for the live preview as the main sampler node has become a messy monster"
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"sampler_inputs": ("SAMPLER_ARGS",),},
}
def process(self, sampler_inputs):
return super().process(**sampler_inputs)
NODE_CLASS_MAPPINGS = {
"WanVideoSampler": WanVideoSampler,
"WanVideoSamplerSettings": WanVideoSamplerSettings,
"WanVideoSamplerFromSettings": WanVideoSamplerFromSettings,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoSampler": "WanVideo Sampler",
"WanVideoSamplerSettings": "WanVideo Sampler Settings",
"WanVideoSamplerFromSettings": "WanVideo Sampler From Settings",
}
+1 -1
View File
@@ -2398,7 +2398,7 @@ class WanModel(torch.nn.Module):
block.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=F+1)
if self.ref_conv is not None and fun_ref is not None:
fun_ref = self.ref_conv(fun_ref).flatten(2).transpose(1, 2)
fun_ref = self.ref_conv(fun_ref.to(self.ref_conv.weight.dtype)).flatten(2).transpose(1, 2)
grid_sizes = torch.stack([torch.tensor([u[0] + 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
seq_len += fun_ref.size(1)
F += 1
+3 -26
View File
@@ -5,19 +5,11 @@ import torch.nn as nn
import torch.nn.functional as F
from tqdm import tqdm
from comfy.utils import ProgressBar
import comfy.ops
ops = comfy.ops.disable_weight_init
CACHE_T = 2
# Workaround for increased memory usage in Conv3D with bfloat16 in torch 2.9.0 stable and up
try:
torch_cudnn_bug = (
hasattr(torch.backends.cudnn, 'version') and
torch.backends.cudnn.version() >= 90800 and
torch.__version__ in ["2.9.0+cu126", "2.9.0+cu128", "2.9.0+cu130"] or torch.__version__.startswith("2.10.")
)
except:
torch_cudnn_bug = False
def check_is_instance(model, module_class):
if isinstance(model, module_class):
return True
@@ -26,7 +18,7 @@ def check_is_instance(model, module_class):
return False
class CausalConv3d(nn.Conv3d):
class CausalConv3d(ops.Conv3d):
"""
Causal 3d convolusion.
"""
@@ -45,21 +37,6 @@ class CausalConv3d(nn.Conv3d):
padding[4] -= cache_x.shape[2]
x = F.pad(x, padding)
# Convert to float32 only if this would trigger cuDNN bug
if (
torch_cudnn_bug and
x.dtype in (torch.bfloat16, torch.half) and
len(self.weight.shape) == 5 and
any(self.weight.shape[i] != 1 for i in range(2, 5))
):
self.weight.data = self.weight.data.float()
if self.bias is not None:
self.bias.data = self.bias.data.float()
result = super().forward(x.float())
return result.to(x.dtype)
return super().forward(x)