Merge branch 'main' into longcat
This commit is contained in:
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user