cleanup
This commit is contained in:
+88
-85
@@ -1,4 +1,5 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
from comfy.utils import common_upscale
|
||||
from comfy import model_management
|
||||
@@ -35,7 +36,7 @@ class WanVideoImageResizeToClosest:
|
||||
DESCRIPTION = "Resizes image to the closest supported resolution based on aspect ratio and max pixels, according to the original code"
|
||||
|
||||
def process(self, image, generation_width, generation_height, aspect_ratio_preservation ):
|
||||
|
||||
|
||||
H, W = image.shape[1], image.shape[2]
|
||||
max_area = generation_width * generation_height
|
||||
|
||||
@@ -47,7 +48,7 @@ class WanVideoImageResizeToClosest:
|
||||
aspect_ratio = generation_height / generation_width
|
||||
if aspect_ratio_preservation == "crop_to_new":
|
||||
crop = "center"
|
||||
|
||||
|
||||
lat_h = round(
|
||||
np.sqrt(max_area * aspect_ratio) // VAE_STRIDE[1] //
|
||||
PATCH_SIZE[1] * PATCH_SIZE[1])
|
||||
@@ -135,27 +136,27 @@ class WanVideoVACEStartToEndFrame:
|
||||
# Convert negative end_index to positive
|
||||
if end_index < 0:
|
||||
end_index = num_frames + end_index
|
||||
|
||||
|
||||
# Create output batch with empty frames
|
||||
out_batch = torch.ones((num_frames, H, W, 3), device=device) * empty_frame_level
|
||||
|
||||
|
||||
# Create mask tensor with proper dimensions
|
||||
masks = torch.ones((num_frames, H, W), device=device)
|
||||
|
||||
|
||||
# Pre-process all images at once to avoid redundant work
|
||||
if end_image is not None and (end_image.shape[1] != H or end_image.shape[2] != W):
|
||||
end_image = common_upscale(end_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1)
|
||||
|
||||
|
||||
if control_images is not None and (control_images.shape[1] != H or control_images.shape[2] != W):
|
||||
control_images = common_upscale(control_images.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1)
|
||||
|
||||
|
||||
# Place start image at start_index
|
||||
if start_image is not None:
|
||||
frames_to_copy = min(start_image.shape[0], num_frames - start_index)
|
||||
if frames_to_copy > 0:
|
||||
out_batch[start_index:start_index + frames_to_copy] = start_image[:frames_to_copy]
|
||||
masks[start_index:start_index + frames_to_copy] = 0
|
||||
|
||||
|
||||
# Place end image at end_index
|
||||
if end_image is not None:
|
||||
# Calculate where to start placing end images
|
||||
@@ -163,28 +164,28 @@ class WanVideoVACEStartToEndFrame:
|
||||
if end_start < 0: # Handle case where end images won't all fit
|
||||
end_image = end_image[abs(end_start):]
|
||||
end_start = 0
|
||||
|
||||
|
||||
frames_to_copy = min(end_image.shape[0], num_frames - end_start)
|
||||
if frames_to_copy > 0:
|
||||
out_batch[end_start:end_start + frames_to_copy] = end_image[:frames_to_copy]
|
||||
masks[end_start:end_start + frames_to_copy] = 0
|
||||
|
||||
|
||||
# Apply control images to remaining frames that don't have start or end images
|
||||
if control_images is not None:
|
||||
# Create a mask of frames that are still empty (mask == 1)
|
||||
empty_frames = masks.sum(dim=(1, 2)) > 0.5 * H * W
|
||||
|
||||
|
||||
if empty_frames.any():
|
||||
# Only apply control images where they exist
|
||||
control_length = control_images.shape[0]
|
||||
for frame_idx in range(num_frames):
|
||||
if empty_frames[frame_idx] and frame_idx < control_length:
|
||||
out_batch[frame_idx] = control_images[frame_idx]
|
||||
|
||||
|
||||
# Apply inpaint mask if provided
|
||||
if inpaint_mask is not None:
|
||||
inpaint_mask = common_upscale(inpaint_mask.unsqueeze(1), W, H, "nearest-exact", "disabled").squeeze(1).to(device)
|
||||
|
||||
|
||||
# Handle different mask lengths efficiently
|
||||
if inpaint_mask.shape[0] > num_frames:
|
||||
inpaint_mask = inpaint_mask[:num_frames]
|
||||
@@ -226,31 +227,31 @@ class CreateCFGScheduleFloatList:
|
||||
cfg_list = [1.0] * steps
|
||||
start_idx = min(int(steps * start_percent), steps - 1)
|
||||
end_idx = min(int(steps * end_percent), steps - 1)
|
||||
|
||||
|
||||
for i in range(start_idx, end_idx + 1):
|
||||
if i >= steps:
|
||||
break
|
||||
|
||||
|
||||
if end_idx == start_idx:
|
||||
t = 0
|
||||
else:
|
||||
t = (i - start_idx) / (end_idx - start_idx)
|
||||
|
||||
|
||||
if interpolation == "linear":
|
||||
factor = t
|
||||
elif interpolation == "ease_in":
|
||||
factor = t * t
|
||||
elif interpolation == "ease_out":
|
||||
factor = t * (2 - t)
|
||||
|
||||
|
||||
cfg_list[i] = round(cfg_scale_start + factor * (cfg_scale_end - cfg_scale_start), 2)
|
||||
|
||||
|
||||
# If start_percent > 0, always include the first step
|
||||
if start_percent > 0:
|
||||
cfg_list[0] = 1.0
|
||||
|
||||
if unique_id and PromptServer is not None:
|
||||
try:
|
||||
try:
|
||||
PromptServer.instance.send_progress_text(
|
||||
f"{cfg_list}",
|
||||
unique_id
|
||||
@@ -259,7 +260,7 @@ class CreateCFGScheduleFloatList:
|
||||
pass
|
||||
|
||||
return (cfg_list,)
|
||||
|
||||
|
||||
class CreateScheduleFloatList:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -289,16 +290,16 @@ class CreateScheduleFloatList:
|
||||
cfg_list = [default_value] * steps
|
||||
start_idx = min(int(steps * start_percent), steps - 1)
|
||||
end_idx = min(int(steps * end_percent), steps - 1)
|
||||
|
||||
|
||||
for i in range(start_idx, end_idx + 1):
|
||||
if i >= steps:
|
||||
break
|
||||
|
||||
|
||||
if end_idx == start_idx:
|
||||
t = 0
|
||||
else:
|
||||
t = (i - start_idx) / (end_idx - start_idx)
|
||||
|
||||
|
||||
if interpolation == "linear":
|
||||
factor = t
|
||||
elif interpolation == "ease_in":
|
||||
@@ -313,7 +314,7 @@ class CreateScheduleFloatList:
|
||||
cfg_list[0] = default_value
|
||||
|
||||
if unique_id and PromptServer is not None:
|
||||
try:
|
||||
try:
|
||||
PromptServer.instance.send_progress_text(
|
||||
f"{cfg_list}",
|
||||
unique_id
|
||||
@@ -322,7 +323,7 @@ class CreateScheduleFloatList:
|
||||
pass
|
||||
|
||||
return (cfg_list,)
|
||||
|
||||
|
||||
|
||||
class DummyComfyWanModelObject:
|
||||
@classmethod
|
||||
@@ -348,7 +349,7 @@ class DummyComfyWanModelObject:
|
||||
return model_sampling
|
||||
return None
|
||||
return (DummyModel(),)
|
||||
|
||||
|
||||
class WanVideoLatentReScale:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -405,7 +406,7 @@ class WanVideoLatentReScale:
|
||||
samples["samples"] = latents
|
||||
|
||||
return (samples,)
|
||||
|
||||
|
||||
class WanVideoSigmaToStep:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -422,7 +423,7 @@ class WanVideoSigmaToStep:
|
||||
|
||||
def convert(self, sigma):
|
||||
return (sigma,)
|
||||
|
||||
|
||||
class NormalizeAudioLoudness:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -437,11 +438,11 @@ class NormalizeAudioLoudness:
|
||||
FUNCTION = "normalize"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def normalize(self, audio, lufs):
|
||||
def normalize(self, audio, lufs):
|
||||
audio_input = audio["waveform"]
|
||||
sample_rate = audio["sample_rate"]
|
||||
if audio_input.dim() == 3:
|
||||
audio_input = audio_input.squeeze(0)
|
||||
audio_input = audio_input.squeeze(0)
|
||||
audio_input_np = audio_input.detach().transpose(0, 1).numpy().astype(np.float32)
|
||||
audio_input_np = np.ascontiguousarray(audio_input_np)
|
||||
normalized_audio = self.loudness_norm(audio_input_np, sr=sample_rate, lufs=lufs)
|
||||
@@ -449,7 +450,7 @@ class NormalizeAudioLoudness:
|
||||
out_audio = {"waveform": torch.from_numpy(normalized_audio).transpose(0, 1).unsqueeze(0).float(), "sample_rate": sample_rate}
|
||||
|
||||
return (out_audio, )
|
||||
|
||||
|
||||
def loudness_norm(self, audio_array, sr=16000, lufs=-23):
|
||||
try:
|
||||
import pyloudnorm
|
||||
@@ -461,7 +462,7 @@ class NormalizeAudioLoudness:
|
||||
return audio_array
|
||||
normalized_audio = pyloudnorm.normalize.loudness(audio_array, loudness, lufs)
|
||||
return normalized_audio
|
||||
|
||||
|
||||
class WanVideoPassImagesFromSamples:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -505,12 +506,12 @@ class FaceMaskFromPoseKeypoints:
|
||||
for i, pose_frame in enumerate(pose_frames):
|
||||
selected_idx, prev_center = self.select_closest_person(pose_frame, person_index if i == 0 else prev_center)
|
||||
np_frames.append(self.draw_kps(pose_frame, selected_idx))
|
||||
|
||||
|
||||
if not np_frames:
|
||||
# Handle case where no frames were processed
|
||||
log.warning("No valid pose frames found, returning empty mask")
|
||||
return (torch.zeros((1, 64, 64), dtype=torch.float32),)
|
||||
|
||||
|
||||
np_frames = np.stack(np_frames, axis=0)
|
||||
tensor = torch.from_numpy(np_frames).float() / 255.
|
||||
print("tensor.shape:", tensor.shape)
|
||||
@@ -521,41 +522,41 @@ class FaceMaskFromPoseKeypoints:
|
||||
people = pose_frame["people"]
|
||||
if not people:
|
||||
return -1, None
|
||||
|
||||
|
||||
centers = []
|
||||
valid_people_indices = []
|
||||
|
||||
|
||||
for idx, person in enumerate(people):
|
||||
# Check if face keypoints exist and are valid
|
||||
if "face_keypoints_2d" not in person or not person["face_keypoints_2d"]:
|
||||
continue
|
||||
|
||||
|
||||
kps = np.array(person["face_keypoints_2d"])
|
||||
if len(kps) == 0:
|
||||
continue
|
||||
|
||||
|
||||
n = len(kps) // 3
|
||||
if n == 0:
|
||||
continue
|
||||
|
||||
|
||||
facial_kps = rearrange(kps, "(n c) -> n c", n=n, c=3)[:, :2]
|
||||
|
||||
|
||||
# Check if we have valid coordinates (not all zeros)
|
||||
if np.all(facial_kps == 0):
|
||||
continue
|
||||
|
||||
|
||||
center = facial_kps.mean(axis=0)
|
||||
|
||||
|
||||
# Check if center is valid (not NaN or infinite)
|
||||
if np.isnan(center).any() or np.isinf(center).any():
|
||||
continue
|
||||
|
||||
|
||||
centers.append(center)
|
||||
valid_people_indices.append(idx)
|
||||
|
||||
|
||||
if not centers:
|
||||
return -1, None
|
||||
|
||||
|
||||
if isinstance(prev_center_or_index, (int, np.integer)):
|
||||
# First frame: use person_index, but map to valid people
|
||||
if 0 <= prev_center_or_index < len(valid_people_indices):
|
||||
@@ -587,58 +588,58 @@ class FaceMaskFromPoseKeypoints:
|
||||
width, height = pose_frame["canvas_width"], pose_frame["canvas_height"]
|
||||
canvas = np.zeros((height, width, 3), dtype=np.uint8)
|
||||
people = pose_frame["people"]
|
||||
|
||||
|
||||
if person_index < 0 or person_index >= len(people):
|
||||
return canvas # Out of bounds, return blank
|
||||
|
||||
|
||||
person = people[person_index]
|
||||
|
||||
|
||||
# Check if face keypoints exist and are valid
|
||||
if "face_keypoints_2d" not in person or not person["face_keypoints_2d"]:
|
||||
return canvas # No face keypoints, return blank
|
||||
|
||||
|
||||
face_kps_data = person["face_keypoints_2d"]
|
||||
if len(face_kps_data) == 0:
|
||||
return canvas # Empty keypoints, return blank
|
||||
|
||||
|
||||
n = len(face_kps_data) // 3
|
||||
if n < 17: # Need at least 17 points for outer contour
|
||||
return canvas # Not enough keypoints, return blank
|
||||
|
||||
|
||||
facial_kps = rearrange(np.array(face_kps_data), "(n c) -> n c", n=n, c=3)[:, :2]
|
||||
|
||||
|
||||
# Check if we have valid coordinates (not all zeros)
|
||||
if np.all(facial_kps == 0):
|
||||
return canvas # All keypoints are zero, return blank
|
||||
|
||||
|
||||
# Check for NaN or infinite values
|
||||
if np.isnan(facial_kps).any() or np.isinf(facial_kps).any():
|
||||
return canvas # Invalid coordinates, return blank
|
||||
|
||||
|
||||
# Check for negative coordinates or coordinates that would create streaks
|
||||
if np.any(facial_kps < 0):
|
||||
return canvas # Negative coordinates, likely bad detection
|
||||
|
||||
|
||||
# Check if coordinates are reasonable (not too close to edges which might indicate bad detection)
|
||||
min_margin = 5 # Minimum distance from edges
|
||||
if (np.any(facial_kps[:, 0] < min_margin) or
|
||||
np.any(facial_kps[:, 1] < min_margin) or
|
||||
np.any(facial_kps[:, 0] > width - min_margin) or
|
||||
if (np.any(facial_kps[:, 0] < min_margin) or
|
||||
np.any(facial_kps[:, 1] < min_margin) or
|
||||
np.any(facial_kps[:, 0] > width - min_margin) or
|
||||
np.any(facial_kps[:, 1] > height - min_margin)):
|
||||
# Check if this looks like a streak to corner (many points near 0,0)
|
||||
corner_points = np.sum((facial_kps[:, 0] < min_margin) & (facial_kps[:, 1] < min_margin))
|
||||
if corner_points > 3: # Too many points near corner, likely bad detection
|
||||
return canvas
|
||||
|
||||
|
||||
facial_kps = facial_kps.astype(np.int32)
|
||||
|
||||
|
||||
# Ensure coordinates are within canvas bounds
|
||||
facial_kps[:, 0] = np.clip(facial_kps[:, 0], 0, width - 1)
|
||||
facial_kps[:, 1] = np.clip(facial_kps[:, 1], 0, height - 1)
|
||||
|
||||
|
||||
part_color = (255, 255, 255)
|
||||
outer_contour = facial_kps[:17]
|
||||
|
||||
|
||||
# Additional validation for the contour before drawing
|
||||
# Check if contour points are too spread out (indicating bad detection)
|
||||
if len(outer_contour) >= 3:
|
||||
@@ -647,11 +648,11 @@ class FaceMaskFromPoseKeypoints:
|
||||
max_x, max_y = np.max(outer_contour, axis=0)
|
||||
contour_width = max_x - min_x
|
||||
contour_height = max_y - min_y
|
||||
|
||||
|
||||
# If contour spans more than 80% of canvas, likely bad detection
|
||||
if (contour_width > 0.8 * width or contour_height > 0.8 * height):
|
||||
return canvas
|
||||
|
||||
|
||||
# Check if we have a valid contour (at least 3 unique points)
|
||||
unique_points = np.unique(outer_contour, axis=0)
|
||||
if len(unique_points) >= 3:
|
||||
@@ -659,11 +660,11 @@ class FaceMaskFromPoseKeypoints:
|
||||
# Calculate area to see if it's too large or too small
|
||||
contour_area = cv2.contourArea(outer_contour)
|
||||
canvas_area = width * height
|
||||
|
||||
|
||||
# If contour is less than 0.1% or more than 50% of canvas, skip
|
||||
if 0.001 * canvas_area <= contour_area <= 0.5 * canvas_area:
|
||||
cv2.fillPoly(canvas, pts=[outer_contour], color=part_color)
|
||||
|
||||
|
||||
return canvas
|
||||
|
||||
|
||||
@@ -679,7 +680,7 @@ class DrawGaussianNoiseOnImage:
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("IMAGE", )
|
||||
RETURN_NAMES = ("images",)
|
||||
FUNCTION = "apply"
|
||||
@@ -694,66 +695,66 @@ class DrawGaussianNoiseOnImage:
|
||||
|
||||
in_masks = mask.clone().to(processing_device)
|
||||
in_images = image.clone().to(processing_device)
|
||||
|
||||
|
||||
# Resize mask to match image dimensions
|
||||
if HM != H or WM != W:
|
||||
in_masks = F.interpolate(mask.unsqueeze(1), size=(H, W), mode='nearest-exact').squeeze(1)
|
||||
|
||||
|
||||
# Match batch sizes
|
||||
if B > BM:
|
||||
in_masks = in_masks.repeat((B + BM - 1) // BM, 1, 1)[:B]
|
||||
elif BM > B:
|
||||
in_masks = in_masks[:B]
|
||||
|
||||
|
||||
output_images = []
|
||||
|
||||
|
||||
# Set random seed for reproducibility
|
||||
generator = torch.Generator(device=processing_device).manual_seed(seed)
|
||||
|
||||
|
||||
for i in tqdm(range(B), desc="DrawGaussianNoiseOnImage batch"):
|
||||
curr_mask = in_masks[i]
|
||||
img_idx = min(i, B - 1)
|
||||
curr_image = in_images[img_idx]
|
||||
|
||||
|
||||
# Expand mask to 3 channels
|
||||
mask_expanded = curr_mask.unsqueeze(-1).expand(-1, -1, 3)
|
||||
|
||||
|
||||
# Calculate mean and std per channel from the subject region (where mask is 1)
|
||||
subject_mask = mask_expanded > 0.5
|
||||
|
||||
|
||||
# Initialize noise tensor
|
||||
noise = torch.zeros_like(curr_image)
|
||||
|
||||
|
||||
for c in range(C):
|
||||
channel = curr_image[:, :, c]
|
||||
channel_mask = subject_mask[:, :, c]
|
||||
|
||||
|
||||
if channel_mask.sum() > 0:
|
||||
# Get subject pixels
|
||||
subject_pixels = channel[channel_mask]
|
||||
|
||||
|
||||
# Calculate statistics
|
||||
mean = subject_pixels.mean()
|
||||
std = subject_pixels.std()
|
||||
|
||||
|
||||
# Generate Gaussian noise for this channel
|
||||
noise[:, :, c] = torch.normal(mean=mean.item(), std=std.item(),
|
||||
size=(H, W), generator=generator,
|
||||
noise[:, :, c] = torch.normal(mean=mean.item(), std=std.item(),
|
||||
size=(H, W), generator=generator,
|
||||
device=processing_device)
|
||||
|
||||
|
||||
# Clamp noise to valid range
|
||||
noise = torch.clamp(noise, 0.0, 1.0)
|
||||
|
||||
|
||||
# Apply: keep subject, fill background with noise
|
||||
masked_image = curr_image * mask_expanded + noise * (1 - mask_expanded)
|
||||
output_images.append(masked_image)
|
||||
|
||||
|
||||
# If no masks were processed, return empty tensor
|
||||
if not output_images:
|
||||
return (torch.zeros((0, H, W, 3), dtype=image.dtype),)
|
||||
|
||||
out_rgb = torch.stack(output_images, dim=0).cpu()
|
||||
|
||||
|
||||
return (out_rgb, )
|
||||
|
||||
|
||||
@@ -773,6 +774,8 @@ class WanVideoPreviewEmbeds:
|
||||
def get(self, embeds):
|
||||
latents = embeds.get("image_embeds", None)
|
||||
mask = embeds.get("mask", None)
|
||||
if mask is not None:
|
||||
mask = mask[0].float().cpu()
|
||||
return ({"samples": latents.unsqueeze(0)}, mask)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
|
||||
Reference in New Issue
Block a user