Binary file not shown.
Binary file not shown.
@@ -0,0 +1,134 @@
|
||||
import json
|
||||
import torch
|
||||
import torchvision.transforms.functional as TF
|
||||
from ..utils import log
|
||||
from .trajectory import create_pos_feature_map, draw_tracks_on_video
|
||||
import os
|
||||
from comfy import model_management as mm
|
||||
device = mm.get_torch_device()
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
VAE_STRIDE = (4, 8, 8) # t, h, w
|
||||
|
||||
class WanVideoWanDrawWanMoveTracks:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"tracks": ("WANMOVETRACKS",),
|
||||
"width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the image to encode"}),
|
||||
"height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the image to encode"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("image",)
|
||||
FUNCTION = "execute"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def execute(self, tracks, width, height):
|
||||
track = tracks["tracks"].unsqueeze(0)
|
||||
track_visibility = tracks["track_visibility"].unsqueeze(0)
|
||||
|
||||
video = torch.zeros((track.shape[1], height, width, 3), device=device)
|
||||
track_video = draw_tracks_on_video(video, track, track_visibility)
|
||||
track_video = torch.stack([TF.to_tensor(frame) for frame in track_video], dim=0).movedim(1, -1)
|
||||
|
||||
return (track_video.float().cpu(), )
|
||||
|
||||
|
||||
class WanVideoAddWanMoveTracks:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"image_embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"track_coords": ("STRING", {"forceInput": True, "tooltip": "JSON string or list of JSON strings representing the tracks"}),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
|
||||
},
|
||||
"optional": {
|
||||
"track_mask": ("MASK",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "WANMOVETRACKS")
|
||||
RETURN_NAMES = ("image_embeds", "tracks")
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, image_embeds, track_coords, strength, track_mask=None):
|
||||
updated = dict(image_embeds)
|
||||
|
||||
target_shape = image_embeds.get("target_shape")
|
||||
if target_shape is not None:
|
||||
height = target_shape[2] * VAE_STRIDE[1]
|
||||
width = target_shape[3] * VAE_STRIDE[2]
|
||||
else:
|
||||
height = image_embeds["lat_h"] * VAE_STRIDE[1]
|
||||
width = image_embeds["lat_w"] * VAE_STRIDE[2]
|
||||
num_frames = image_embeds["num_frames"]
|
||||
|
||||
tracks_data = parse_json_tracks(track_coords)
|
||||
track_list = [
|
||||
[[track[frame]['x'], track[frame]['y']] for track in tracks_data]
|
||||
for frame in range(len(tracks_data[0]))
|
||||
]
|
||||
track = torch.tensor(track_list, dtype=torch.float32, device=device) # shape: (frames, num_tracks, 2)
|
||||
track = track[:num_frames]
|
||||
|
||||
num_tracks = track.shape[-2]
|
||||
if track_mask is None:
|
||||
track_visibility = torch.ones((num_frames, num_tracks), dtype=torch.bool, device=device)
|
||||
else:
|
||||
track_visibility = (track_mask > 0).any(dim=(1, 2)).unsqueeze(-1)
|
||||
|
||||
feature_map, track_pos = create_pos_feature_map(track, track_visibility, VAE_STRIDE, height, width, 16, track_num=1, device=device)
|
||||
|
||||
updated.setdefault("wanmove_embeds", {})
|
||||
updated["wanmove_embeds"]["track_pos"] = track_pos * strength
|
||||
|
||||
tracks_dict = {
|
||||
"tracks": track,
|
||||
"track_visibility": track_visibility,
|
||||
}
|
||||
|
||||
return (updated, tracks_dict,)
|
||||
|
||||
|
||||
def parse_json_tracks(tracks):
|
||||
tracks_data = []
|
||||
try:
|
||||
# If tracks is a string, try to parse it as JSON
|
||||
if isinstance(tracks, str):
|
||||
parsed = json.loads(tracks.replace("'", '"'))
|
||||
tracks_data.extend(parsed)
|
||||
else:
|
||||
# If tracks is a list of strings, parse each one
|
||||
for track_str in tracks:
|
||||
parsed = json.loads(track_str.replace("'", '"'))
|
||||
tracks_data.append(parsed)
|
||||
|
||||
# Check if we have a single track (dict with x,y) or a list of tracks
|
||||
if tracks_data and isinstance(tracks_data[0], dict) and 'x' in tracks_data[0]:
|
||||
# Single track detected, wrap it in a list
|
||||
tracks_data = [tracks_data]
|
||||
elif tracks_data and isinstance(tracks_data[0], list) and tracks_data[0] and isinstance(tracks_data[0][0], dict) and 'x' in tracks_data[0][0]:
|
||||
# Already a list of tracks, nothing to do
|
||||
pass
|
||||
else:
|
||||
# Unexpected format
|
||||
log.warning(f"Warning: Unexpected track format: {type(tracks_data[0])}")
|
||||
|
||||
except json.JSONDecodeError as e:
|
||||
log.warning(f"Error parsing tracks JSON: {e}")
|
||||
tracks_data = []
|
||||
|
||||
return tracks_data
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAddWanMoveTracks": WanVideoAddWanMoveTracks,
|
||||
"WanVideoWanDrawWanMoveTracks": WanVideoWanDrawWanMoveTracks,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoAddWanMoveTracks": "WanVideo Add WanMove Tracks",
|
||||
"WanVideoWanDrawWanMoveTracks": "WanVideo Draw WanMove Tracks",
|
||||
}
|
||||
@@ -0,0 +1,324 @@
|
||||
# https://github.com/ali-vilab/Wan-Move/blob/main/wan/modules/trajectory.py
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image, ImageDraw
|
||||
|
||||
SKIP_ZERO = False
|
||||
|
||||
def get_pos_emb(
|
||||
pos_k: torch.Tensor,
|
||||
pos_emb_dim: int,
|
||||
theta_func: callable = lambda i, d: torch.pow(10000, torch.mul(2, torch.div(i.to(torch.float32), d))),
|
||||
device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu"),
|
||||
dtype: torch.dtype = torch.float32,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Generate batch position embeddings.
|
||||
|
||||
Args:
|
||||
pos_k (torch.Tensor): A 1D tensor containing positions for which to generate embeddings.
|
||||
pos_emb_dim (int): The dimension of position embeddings.
|
||||
theta_func (callable): Function to compute thetas based on position and embedding dimensions.
|
||||
device (torch.device): Device to store the position embeddings.
|
||||
dtype (torch.dtype): Desired data type for computations.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: The position embeddings with shape (batch_size, pos_emb_dim).
|
||||
"""
|
||||
assert pos_emb_dim % 2 == 0, "The dimension of position embeddings must be even."
|
||||
pos_k = pos_k.to(device, dtype)
|
||||
if SKIP_ZERO:
|
||||
pos_k = pos_k + 1
|
||||
batch_size = pos_k.size(0)
|
||||
|
||||
denominator = torch.arange(0, pos_emb_dim // 2, device=device, dtype=dtype)
|
||||
# Expand denominator to match the shape needed for broadcasting
|
||||
denominator_expanded = denominator.view(1, -1).expand(batch_size, -1)
|
||||
|
||||
thetas = theta_func(denominator_expanded, pos_emb_dim)
|
||||
|
||||
# Ensure pos_k is in the correct shape for broadcasting
|
||||
pos_k_expanded = pos_k.view(-1, 1).to(dtype)
|
||||
sin_thetas = torch.sin(torch.div(pos_k_expanded, thetas))
|
||||
cos_thetas = torch.cos(torch.div(pos_k_expanded, thetas))
|
||||
|
||||
# Concatenate sine and cosine embeddings along the last dimension
|
||||
pos_emb = torch.cat([sin_thetas, cos_thetas], dim=-1)
|
||||
|
||||
return pos_emb
|
||||
|
||||
def create_pos_feature_map(
|
||||
pred_tracks: torch.Tensor, # [T, N, 2]
|
||||
pred_visibility: torch.Tensor, # [T, N]
|
||||
downsample_ratios: list[int],
|
||||
height: int,
|
||||
width: int,
|
||||
pos_emb_dim: int,
|
||||
track_num: int = -1,
|
||||
t_down_strategy: str = "sample",
|
||||
device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu"),
|
||||
dtype: torch.dtype = torch.float32,
|
||||
):
|
||||
"""
|
||||
Create a feature map from the predicted tracks.
|
||||
|
||||
Args:
|
||||
- pred_tracks: torch.Tensor, the predicted tracks, [T, N, 2]
|
||||
- pred_visibility: torch.Tensor, the predicted visibility, [T, N]
|
||||
- downsample_ratios: list[int], the ratios for downsampling time, height, and width
|
||||
- height: int, the height of the feature map
|
||||
- width: int, the width of the feature map
|
||||
- pos_emb_dim: int, the dimension of the position embeddings
|
||||
- track_num: int, the number of tracks to use
|
||||
- t_down_strategy: str, the strategy for downsampling time dimension
|
||||
- device: torch.device, the device
|
||||
- dtype: torch.dtype, the data type
|
||||
|
||||
Returns:
|
||||
- feature_map: torch.Tensor, the feature map, [T', H', W', pos_emb_dim]
|
||||
- track_pos: torch.Tensor, the position embeddings, [N, T', 2], 2 = height, width
|
||||
"""
|
||||
|
||||
assert t_down_strategy in ["sample", "average"], "Invalid strategy for downsampling time dimension."
|
||||
|
||||
t, n, _ = pred_tracks.shape
|
||||
t_down, h_down, w_down = downsample_ratios
|
||||
feature_map = torch.zeros((t-1) // t_down + 1, height // h_down, width // w_down, pos_emb_dim, device=device, dtype=dtype)
|
||||
track_pos = - torch.ones(n, (t-1) // t_down + 1, 2, dtype=torch.long)
|
||||
|
||||
if track_num == -1:
|
||||
track_num = n
|
||||
|
||||
tracks_idx = torch.randperm(n)[:track_num]
|
||||
tracks = pred_tracks[:, tracks_idx]
|
||||
visibility = pred_visibility[:, tracks_idx]
|
||||
tracks_embs = get_pos_emb(torch.randperm(n)[:track_num], pos_emb_dim, device=device, dtype=dtype)
|
||||
|
||||
for t_idx in range(0, t, t_down):
|
||||
if t_down_strategy == "sample" or t_idx == 0:
|
||||
cur_tracks = tracks[t_idx] # [N, 2]
|
||||
cur_visibility = visibility[t_idx] # [N]
|
||||
else:
|
||||
cur_tracks = tracks[t_idx:t_idx+t_down].mean(dim=0)
|
||||
cur_visibility = torch.any(visibility[t_idx:t_idx+t_down], dim=0)
|
||||
|
||||
for i in range(track_num):
|
||||
if not cur_visibility[i] or cur_tracks[i][0] < 0 or cur_tracks[i][1] < 0 or cur_tracks[i][0] >= width or cur_tracks[i][1] >= height:
|
||||
continue
|
||||
x, y = cur_tracks[i]
|
||||
x, y = int(x // w_down), int(y // h_down)
|
||||
feature_map[t_idx // t_down, y, x] += tracks_embs[i]
|
||||
track_pos[i, t_idx // t_down, 0], track_pos[i, t_idx // t_down, 1] = y, x
|
||||
|
||||
return feature_map, track_pos
|
||||
|
||||
|
||||
def replace_feature(
|
||||
vae_feature: torch.Tensor, # [B, C', T', H', W']
|
||||
track_pos: torch.Tensor, # [B, N, T', 2]
|
||||
) -> torch.Tensor:
|
||||
b, _, t, h, w = vae_feature.shape
|
||||
assert b == track_pos.shape[0], "Batch size mismatch."
|
||||
n = track_pos.shape[1]
|
||||
|
||||
# Shuffle the trajectory order
|
||||
track_pos = track_pos[:, torch.randperm(n), :, :]
|
||||
|
||||
# Extract coordinates at time steps ≥ 1 and generate a valid mask
|
||||
current_pos = track_pos[:, :, 1:, :] # [B, N, T-1, 2]
|
||||
mask = (current_pos[..., 0] >= 0) & (current_pos[..., 1] >= 0) # [B, N, T-1]
|
||||
|
||||
# Get all valid indices
|
||||
valid_indices = mask.nonzero(as_tuple=False) # [num_valid, 3]
|
||||
num_valid = valid_indices.shape[0]
|
||||
|
||||
if num_valid == 0:
|
||||
return vae_feature
|
||||
|
||||
# Decompose valid indices into each dimension
|
||||
batch_idx = valid_indices[:, 0]
|
||||
track_idx = valid_indices[:, 1]
|
||||
t_rel = valid_indices[:, 2]
|
||||
t_target = t_rel + 1 # Convert to original time step indices
|
||||
|
||||
# Extract target position coordinates
|
||||
h_target = current_pos[batch_idx, track_idx, t_rel, 0].long() # Ensure integer indices
|
||||
w_target = current_pos[batch_idx, track_idx, t_rel, 1].long()
|
||||
|
||||
# Extract source position coordinates (t=0)
|
||||
h_source = track_pos[batch_idx, track_idx, 0, 0].long()
|
||||
w_source = track_pos[batch_idx, track_idx, 0, 1].long()
|
||||
|
||||
# Get source features and assign to target positions
|
||||
src_features = vae_feature[batch_idx, :, 0, h_source, w_source]
|
||||
vae_feature[batch_idx, :, t_target, h_target, w_target] = src_features
|
||||
|
||||
return vae_feature
|
||||
|
||||
def get_video_track_video(
|
||||
model,
|
||||
video_tensor: torch.Tensor, # [T, C, H, W]
|
||||
downsample_ratios: list[int],
|
||||
pos_emb_dim: int,
|
||||
grid_size: int = 32,
|
||||
track_num: int = -1,
|
||||
t_down_strategy: str = "sample",
|
||||
device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu"),
|
||||
dtype: torch.dtype = torch.float32,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Get the track video from the video tensor.
|
||||
|
||||
Args:
|
||||
- model: torch.nn.Module, the model for tracking, CoTracker
|
||||
- video_tensor: torch.Tensor, the video tensor, [T, C, H, W]
|
||||
- downsample_ratios: list[int], the ratios for downsampling time, height, and width
|
||||
- height: int, the height of the feature map
|
||||
- width: int, the width of the feature map
|
||||
- pos_emb_dim: int, the dimension of the position embeddings
|
||||
- grid_size: int, the size of the grid
|
||||
- track_num: int, the number of tracks to use
|
||||
- t_down_strategy: str, the strategy for downsampling time dimension
|
||||
- device: torch.device, the device
|
||||
- dtype: torch.dtype, the data type
|
||||
|
||||
Returns:
|
||||
- track_video: torch.Tensor, the track video, [pos_emb_dim, T', H', W']
|
||||
- track_pos: torch.Tensor, the position embeddings, [N, T', 2], 2 = height, width
|
||||
- pred_tracks: the predicted point trajectories
|
||||
- pred_visibility: visibility of the predicted point trajectories
|
||||
"""
|
||||
|
||||
t, c, height, width = video_tensor.shape
|
||||
with (
|
||||
torch.autocast(device_type=device.type, dtype=dtype),
|
||||
torch.no_grad(),
|
||||
):
|
||||
pred_tracks, pred_visibility = model(
|
||||
video_tensor.unsqueeze(0),
|
||||
grid_size=grid_size,
|
||||
backward_tracking=False,
|
||||
)
|
||||
|
||||
track_video, track_pos = create_pos_feature_map(
|
||||
pred_tracks[0], pred_visibility[0], downsample_ratios, height, width, pos_emb_dim, track_num, t_down_strategy, device, dtype
|
||||
)
|
||||
|
||||
return track_video.permute(3, 0, 1, 2), track_pos, pred_tracks, pred_visibility
|
||||
|
||||
# ---------------------------
|
||||
# Visualize functions
|
||||
# --------------------------
|
||||
|
||||
def draw_overall_gradient_polyline_on_image(image, line_width, points, start_color):
|
||||
"""
|
||||
- image (Image): target image to draw on.
|
||||
- line_width (int): initial line width.
|
||||
- points (list of tuples): list of points forming the polyline, each point is (x, y).
|
||||
- start_color (tuple): starting color of the line (R, G, B).
|
||||
|
||||
Return:
|
||||
- Image: original image with the gradient polyline drawn.
|
||||
"""
|
||||
|
||||
def get_distance(p1, p2):
|
||||
return ((p2[0] - p1[0]) ** 2 + (p2[1] - p1[1]) ** 2) ** 0.5
|
||||
|
||||
# Create a new image with the same size as the original
|
||||
new_image = Image.new('RGBA', image.size)
|
||||
draw = ImageDraw.Draw(new_image, 'RGBA')
|
||||
points = points[::-1]
|
||||
|
||||
# Compute total length
|
||||
total_length = sum(get_distance(points[i], points[i+1]) for i in range(len(points)-1))
|
||||
|
||||
# Accumulated length
|
||||
accumulated_length = 0
|
||||
|
||||
# Draw the gradient polyline
|
||||
for start_point, end_point in zip(points[:-1], points[1:]):
|
||||
segment_length = get_distance(start_point, end_point)
|
||||
steps = int(segment_length)
|
||||
|
||||
for i in range(steps):
|
||||
# Current accumulated length
|
||||
current_length = accumulated_length + (i / steps) * segment_length
|
||||
|
||||
# Alpha from fully opaque to fully transparent
|
||||
alpha = int(255 * (1 - current_length / total_length))
|
||||
color = (*start_color, alpha)
|
||||
|
||||
# Interpolated coordinates
|
||||
x = int(start_point[0] + (end_point[0] - start_point[0]) * i / steps)
|
||||
y = int(start_point[1] + (end_point[1] - start_point[1]) * i / steps)
|
||||
|
||||
# Dynamic line width, decreasing from initial width to 1
|
||||
dynamic_line_width = int(line_width * (1 - (current_length / total_length)))
|
||||
dynamic_line_width = max(dynamic_line_width, 1) # minimum width is 1 to avoid 0
|
||||
|
||||
draw.line([(x, y), (x + 1, y)], fill=color, width=dynamic_line_width)
|
||||
|
||||
accumulated_length += segment_length
|
||||
|
||||
return new_image
|
||||
|
||||
def add_weighted(rgb, track):
|
||||
rgb = np.array(rgb) # [H, W, C] "RGB"
|
||||
track = np.array(track) # [H, W, C] "RGBA"
|
||||
|
||||
# Compute weights from the alpha channel
|
||||
alpha = track[:, :, 3] / 255.0
|
||||
|
||||
# Expand alpha to 3 channels to match RGB
|
||||
alpha = np.stack([alpha] * 3, axis=-1)
|
||||
|
||||
# Blend the two images
|
||||
blend_img = track[:, :, :3] * alpha + rgb * (1 - alpha)
|
||||
|
||||
return Image.fromarray(blend_img.astype(np.uint8))
|
||||
|
||||
def draw_tracks_on_video(video, tracks, visibility=None, track_frame=24):
|
||||
color_map = [
|
||||
(102, 153, 255),
|
||||
(0, 255, 255),
|
||||
(255, 255, 0),
|
||||
(255, 102, 204),
|
||||
(0, 255, 0)
|
||||
]
|
||||
circle_size = 12
|
||||
line_width = 16
|
||||
|
||||
video = video.byte().cpu().numpy() # (81, 480, 832, 3)
|
||||
tracks = tracks[0].long().detach().cpu().numpy()
|
||||
if visibility is not None:
|
||||
visibility = visibility[0].detach().cpu().numpy()
|
||||
# print(video.shape, tracks.shape)
|
||||
|
||||
output_frames = []
|
||||
# Process the video
|
||||
for t in range(video.shape[0]):
|
||||
# Extract current frame
|
||||
frame = video[t]
|
||||
frame = Image.fromarray(frame).convert("RGB")
|
||||
|
||||
# Draw tracks
|
||||
for n in range(tracks.shape[1]):
|
||||
if visibility is not None and visibility[t, n] == 0:
|
||||
continue
|
||||
|
||||
# Track coordinate at current frame
|
||||
track_coord = tracks[t, n]
|
||||
tracks_coord = tracks[max(t-track_frame, 0):t+1, n]
|
||||
|
||||
# Draw a circle
|
||||
draw = ImageDraw.Draw(frame)
|
||||
draw.ellipse((track_coord[0] - circle_size, track_coord[1] - circle_size, track_coord[0] + circle_size, track_coord[1] + circle_size), fill=color_map[n % len(color_map)])
|
||||
# Draw the polyline
|
||||
track_image = draw_overall_gradient_polyline_on_image(frame, line_width, tracks_coord, color_map[n % len(color_map)])
|
||||
frame = add_weighted(frame, track_image)
|
||||
|
||||
# Save current frame
|
||||
output_frames.append(frame.convert("RGB"))
|
||||
|
||||
return output_frames
|
||||
+12
-2
@@ -291,7 +291,7 @@ class WanVideoSampler:
|
||||
else:
|
||||
cfg = [cfg] * (steps + 1)
|
||||
|
||||
control_latents = control_camera_latents = clip_fea = clip_fea_neg = end_image = recammaster = camera_embed = unianim_data = mocha_embeds = None
|
||||
control_latents = control_camera_latents = clip_fea = clip_fea_neg = end_image = recammaster = camera_embed = unianim_data = mocha_embeds = image_cond_neg =None
|
||||
vace_data = vace_context = vace_scale = None
|
||||
fun_or_fl2v_model = has_ref = drop_last = False
|
||||
phantom_latents = fun_ref_image = ATI_tracks = None
|
||||
@@ -301,6 +301,15 @@ class WanVideoSampler:
|
||||
#I2V
|
||||
image_cond = image_embeds.get("image_embeds", None)
|
||||
if image_cond is not None:
|
||||
# WanMove
|
||||
wanmove_embeds = image_embeds.get("wanmove_embeds", None)
|
||||
if wanmove_embeds is not None:
|
||||
from .WanMove.trajectory import replace_feature
|
||||
track_pos = wanmove_embeds["track_pos"]
|
||||
if any(not math.isclose(c, 1.0) for c in cfg):
|
||||
image_cond_neg = torch.cat([image_embeds["mask"], image_cond])
|
||||
image_cond = replace_feature(image_cond.unsqueeze(0), track_pos.unsqueeze(0))[0]
|
||||
|
||||
if transformer.in_dim == 16:
|
||||
raise ValueError("T2V (text to video) model detected, encoded images only work with I2V (Image to video) models")
|
||||
elif transformer.in_dim not in [48, 32]: # fun 2.1 models don't use the mask
|
||||
@@ -1118,7 +1127,7 @@ class WanVideoSampler:
|
||||
lynx_embeds=lynx_embeds
|
||||
)
|
||||
log.info(f"Extracted {len(lynx_ref_buffer)} cond ref buffers")
|
||||
if not math.isclose(cfg[0], 1.0):
|
||||
if any(not math.isclose(c, 1.0) for c in cfg):
|
||||
log.info("Extracting Lynx ref uncond buffer...")
|
||||
if transformer.in_dim == 36:
|
||||
lynx_ref_input_uncond = torch.cat([lynx_ref_latent_uncond, empty_image_cond], dim=0)
|
||||
@@ -1536,6 +1545,7 @@ class WanVideoSampler:
|
||||
base_params['is_uncond'] = True
|
||||
base_params['clip_fea'] = clip_fea_neg if clip_fea_neg is not None else clip_fea
|
||||
base_params["add_text_emb"] = qwenvl_embeds_neg.to(device) if qwenvl_embeds_neg is not None else None # QwenVL embeddings for Bindweave
|
||||
base_params['y'] = image_cond_neg if image_cond_neg is not None else base_params['y']
|
||||
if wananim_face_pixels is not None:
|
||||
base_params['wananim_face_pixel_values'] = torch.zeros_like(wananim_face_pixels).to(device, torch.float32) - 1
|
||||
if humo_audio_input_neg is not None:
|
||||
|
||||
Reference in New Issue
Block a user