Add a warning for possible custom node version conflict from having multiple WanVideoWrappers installed
This commit is contained in:
+17
-6
@@ -1,3 +1,14 @@
|
||||
try:
|
||||
from .utils import check_duplicate_nodes, log
|
||||
duplicate_dirs = check_duplicate_nodes()
|
||||
if duplicate_dirs:
|
||||
warning_msg = f"⚠️ Found {len(duplicate_dirs)} other WanVideoWrapper directories: ⚠️\n"
|
||||
for dir_path in duplicate_dirs:
|
||||
warning_msg += f" - {dir_path}\n"
|
||||
log.warning(warning_msg + "Please remove duplicates to avoid possible conflicts.")
|
||||
except:
|
||||
pass
|
||||
|
||||
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .recammaster.nodes import NODE_CLASS_MAPPINGS as RECAM_MASTER_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .skyreels.nodes import NODE_CLASS_MAPPINGS as SKYREELS_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as SKYREELS_NODE_DISPLAY_NAME_MAPPINGS
|
||||
@@ -17,36 +28,36 @@ from .s2v.nodes import NODE_CLASS_MAPPINGS as S2V_NODE_CLASS_MAPPINGS, NODE_DISP
|
||||
try:
|
||||
from .qwen.qwen import NODE_CLASS_MAPPINGS as QWEN_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as QWEN_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as e:
|
||||
print(f"Qwen nodes not available due to error in importing them: {e}")
|
||||
log.warning(f"⚠️ WanVideoWrapper: Qwen nodes not available due to error in importing them: {e} ⚠️")
|
||||
QWEN_NODE_CLASS_MAPPINGS = {}
|
||||
QWEN_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
|
||||
|
||||
try:
|
||||
from .fantasyportrait.nodes import NODE_CLASS_MAPPINGS as FANTASYPORTRAIT_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FANTASYPORTRAIT_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as e:
|
||||
print(f"FantasyPortrait nodes not available due to error in importing them: {e}")
|
||||
log.warning(f"⚠️ WanVideoWrapper: FantasyPortrait nodes not available due to error in importing them: {e} ⚠️")
|
||||
FANTASYPORTRAIT_NODE_CLASS_MAPPINGS = {}
|
||||
FANTASYPORTRAIT_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
try:
|
||||
from .unianimate.nodes import NODE_CLASS_MAPPINGS as UNIANIMATE_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as e:
|
||||
print(f"UniAnimate nodes not available due to error in importing them: {e}")
|
||||
log.warning(f"⚠️ WanVideoWrapper: UniAnimate nodes not available due to error in importing them: {e} ⚠️")
|
||||
UNIANIMATE_NODE_CLASS_MAPPINGS = {}
|
||||
UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
try:
|
||||
from .MTV.nodes import NODE_CLASS_MAPPINGS as MTV_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MTV_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as e:
|
||||
print(f"MTV nodes not available due to error in importing them: {e}")
|
||||
log.warning(f"⚠️ WanVideoWrapper: MTV nodes not available due to error in importing them: {e} ⚠️")
|
||||
MTV_NODE_CLASS_MAPPINGS = {}
|
||||
MTV_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
try:
|
||||
from .HuMo.nodes import NODE_CLASS_MAPPINGS as HUMO_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as HUMO_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as e:
|
||||
print(f"HuMo nodes not available due to error in importing them: {e}")
|
||||
log.warning(f"⚠️ WanVideoWrapper: HuMo nodes not available due to error in importing them: {e} ⚠️")
|
||||
HUMO_NODE_CLASS_MAPPINGS = {}
|
||||
HUMO_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
|
||||
@@ -63,7 +63,7 @@ try:
|
||||
from comfy.latent_formats import Wan21, Wan22
|
||||
latent_format = Wan21
|
||||
except: #for backwards compatibility
|
||||
log.warning("Wan21 latent format not found, update ComfyUI for better livepreview")
|
||||
log.warning("⚠️ Wan21 latent format not found, update ComfyUI for better livepreview ⚠️")
|
||||
from comfy.latent_formats import HunyuanVideo
|
||||
latent_format = HunyuanVideo
|
||||
|
||||
@@ -1092,7 +1092,7 @@ class WanVideoModelLoader:
|
||||
major, minor = torch.cuda.get_device_capability(device)
|
||||
log.info(f"CUDA Compute Capability: {major}.{minor}")
|
||||
if compile_args is not None and "e4" in quantization and (major, minor) < (8, 9):
|
||||
log.warning("Torch.compile with fp8_e4m3fn weights on CUDA compute capability < 8.9 is not supported. Please use fp8_e5m2, GGUF or higher precision instead.")
|
||||
log.warning("⚠️Torch.compile with fp8_e4m3fn weights on CUDA compute capability < 8.9 is not supported. Please use fp8_e5m2, GGUF or higher precision instead.⚠️")
|
||||
|
||||
if "scaled_fp8" in sd and "scaled" not in quantization:
|
||||
raise ValueError("The model is a scaled fp8 model, please set quantization to '_scaled'")
|
||||
|
||||
@@ -3,13 +3,15 @@ import torch
|
||||
import logging
|
||||
import math
|
||||
from tqdm import tqdm
|
||||
from copy import deepcopy
|
||||
from pathlib import Path
|
||||
import os
|
||||
import types, collections
|
||||
from comfy.utils import ProgressBar, copy_to_param, set_attr_param
|
||||
from comfy.model_patcher import get_key_weight, string_to_seed
|
||||
from comfy.lora import calculate_weight
|
||||
from comfy.model_management import cast_to_device
|
||||
from comfy.float import stochastic_rounding
|
||||
import folder_paths
|
||||
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
@@ -130,7 +132,7 @@ def get_tensor_memory(tensor):
|
||||
def patch_weight_to_device(self, key, device_to=None, inplace_update=False, backup_keys=False, scale_weight=None):
|
||||
if key not in self.patches:
|
||||
return
|
||||
|
||||
|
||||
weight, set_func, convert_func = get_key_weight(self.model, key)
|
||||
inplace_update = self.weight_inplace_update or inplace_update
|
||||
|
||||
@@ -148,7 +150,7 @@ def patch_weight_to_device(self, key, device_to=None, inplace_update=False, back
|
||||
temp_weight = temp_weight * scale_weight.to(temp_weight.device, temp_weight.dtype)
|
||||
|
||||
out_weight = calculate_weight(self.patches[key], temp_weight, key)
|
||||
|
||||
|
||||
if set_func is None:
|
||||
out_weight = stochastic_rounding(out_weight, weight.dtype, seed=string_to_seed(key))
|
||||
if inplace_update:
|
||||
@@ -225,7 +227,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
|
||||
param_fp32.mul_(scale.to(param.device, torch.float32))
|
||||
param.copy_(param_fp32.to(param.dtype))
|
||||
model.scale_weights_applied = True
|
||||
|
||||
|
||||
model.current_weight_patches_uuid = model.patches_uuid
|
||||
if low_mem_load:
|
||||
for name, param in model.model.diffusion_model.named_parameters():
|
||||
@@ -247,11 +249,11 @@ def split_tiles(embeds, num_split):
|
||||
for x in embeds:
|
||||
x = x.unsqueeze(0)
|
||||
h, w = H // num_split, W // num_split
|
||||
x_split = torch.cat([x[:, i*h:(i+1)*h, j*w:(j+1)*w, :] for i in range(num_split) for j in range(num_split)], dim=0)
|
||||
x_split = torch.cat([x[:, i*h:(i+1)*h, j*w:(j+1)*w, :] for i in range(num_split) for j in range(num_split)], dim=0)
|
||||
out.append(x_split)
|
||||
|
||||
|
||||
x_split = torch.stack(out, dim=0)
|
||||
|
||||
|
||||
return x_split
|
||||
|
||||
def merge_hiddenstates(x, tiles):
|
||||
@@ -271,19 +273,19 @@ def merge_hiddenstates(x, tiles):
|
||||
patch_embeds = embeds[:, 1:, :] # Shape: [num_tiles, tile_size^2, embeds[-1]]
|
||||
reshaped = patch_embeds.reshape(grid_size, grid_size, tile_size, tile_size, embeds.shape[-1])
|
||||
|
||||
merged = torch.cat([torch.cat([reshaped[i, j] for j in range(grid_size)], dim=1)
|
||||
merged = torch.cat([torch.cat([reshaped[i, j] for j in range(grid_size)], dim=1)
|
||||
for i in range(grid_size)], dim=0)
|
||||
|
||||
|
||||
merged = merged.unsqueeze(0) # Shape: [1, grid_size*tile_size, grid_size*tile_size, embeds[-1]]
|
||||
|
||||
|
||||
# Pool to original size
|
||||
pooled = torch.nn.functional.adaptive_avg_pool2d(merged.permute(0, 3, 1, 2), (tile_size, tile_size)).permute(0, 2, 3, 1)
|
||||
flattened = pooled.reshape(1, tile_size*tile_size, embeds.shape[-1])
|
||||
|
||||
|
||||
# Add back the class token
|
||||
with_class = torch.cat([avg_class_token, flattened], dim=1) # Shape: original shape
|
||||
out.append(with_class)
|
||||
|
||||
|
||||
out = torch.cat(out, dim=0)
|
||||
|
||||
return out
|
||||
@@ -384,7 +386,7 @@ def is_image_black(image, threshold=1e-3):
|
||||
return torch.all(image < threshold).item()
|
||||
|
||||
def add_noise_to_reference_video(image, ratio=None):
|
||||
sigma = torch.ones((image.shape[0],)).to(image.device, image.dtype) * ratio
|
||||
sigma = torch.ones((image.shape[0],)).to(image.device, image.dtype) * ratio
|
||||
image_noise = torch.randn_like(image) * sigma[:, None, None, None]
|
||||
image_noise = torch.where(image==-1, torch.zeros_like(image), image_noise)
|
||||
image = image + image_noise
|
||||
@@ -400,7 +402,7 @@ def optimized_scale(positive_flat, negative_flat):
|
||||
|
||||
# st_star = v_cond^T * v_uncond / ||v_uncond||^2
|
||||
st_star = dot_product / squared_norm
|
||||
|
||||
|
||||
return st_star
|
||||
|
||||
def find_closest_valid_dim(fixed_dim, var_dim, block_size):
|
||||
@@ -468,7 +470,7 @@ def setup_radial_attention(transformer, transformer_options, latent, seq_len, la
|
||||
block.dense_attention_mode = dense_attention_mode
|
||||
block.dense_timesteps = dense_timesteps
|
||||
block.self_attn.decay_factor = decay_factor
|
||||
|
||||
|
||||
log.info(f"Radial attention mode enabled.")
|
||||
log.info(f"dense_attention_mode: {dense_attention_mode}, dense_timesteps: {dense_timesteps}, decay_factor: {decay_factor}")
|
||||
log.info(f"dense_blocks: {[i for i, block in enumerate(transformer.blocks) if getattr(block, 'dense_block', False)]})")
|
||||
@@ -567,4 +569,22 @@ def tensor_pingpong_pad(video, target_len):
|
||||
padded_video = video[:, :, indices, :, :]
|
||||
if in_dims == 4:
|
||||
padded_video = padded_video.squeeze(0)
|
||||
return padded_video
|
||||
return padded_video
|
||||
|
||||
|
||||
def check_duplicate_nodes():
|
||||
"""Check ComfyUI custom_nodes directory for duplicate installations"""
|
||||
custom_nodes_dir = Path(folder_paths.folder_names_and_paths["custom_nodes"][0][0])
|
||||
current_path = Path(__file__).parent
|
||||
|
||||
wanvideo_dirs = []
|
||||
|
||||
# Check all directories in custom_nodes
|
||||
for path in custom_nodes_dir.iterdir():
|
||||
if (path.is_dir() and
|
||||
path != current_path and
|
||||
'wanvideo' in path.name.lower() and
|
||||
'wrapper' in path.name.lower()):
|
||||
wanvideo_dirs.append(str(path))
|
||||
|
||||
return wanvideo_dirs
|
||||
|
||||
Reference in New Issue
Block a user