From 7c81d7ce103f1a5d90a5ada86d76112391d988c8 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 16 Apr 2025 19:40:04 +0300 Subject: [PATCH] Better support for the official FirstLastFrame2Video -model (FLF2V) https://huggingface.co/Kijai/WanVideo_comfy/blob/main/Wan2_1-FLF2V-14B-720P_fp16.safetensors https://huggingface.co/Kijai/WanVideo_comfy/blob/main/Wan2_1-FLF2V-14B-720P_fp8_e4m3fn.safetensors --- nodes.py | 37 ++++++++++++++++++++++--------------- wanvideo/modules/model.py | 11 +++++++---- 2 files changed, 29 insertions(+), 19 deletions(-) diff --git a/nodes.py b/nodes.py index cf66229..03c7419 100644 --- a/nodes.py +++ b/nodes.py @@ -545,9 +545,12 @@ class WanVideoModelLoader: raise ValueError("Invalid WanVideo model selected") dim = sd["patch_embedding.weight"].shape[0] in_channels = sd["patch_embedding.weight"].shape[1] - print("in_channels: ", in_channels) + log.info(f"Detected model in_channels: {in_channels}") ffn_dim = sd["blocks.0.ffn.0.bias"].shape[0] - if in_channels in [36, 48]: + + if "model_type.Wan2_1-FLF2V-14B-720P" in sd or "flf2v" in model.lower(): + model_type = "fl2v" + elif in_channels in [36, 48]: model_type = "i2v" elif in_channels == 16: model_type = "t2v" @@ -583,13 +586,17 @@ class WanVideoModelLoader: "e0": [8.10705460e+03, 2.13393892e+03, -3.72934672e+02, 1.66203073e+01, -4.17769401e-02], }, } - if model_type == "i2v": + + if model_type == "i2v" or model_type == "fl2v": if "480" in model or "fun" in model.lower() or "a2" in model.lower(): #just a guess for the Fun model for now... model_variant = "i2v_480" elif "720" in model: model_variant = "i2v_720" elif model_type == "t2v": model_variant = "14B" + else: + model_variant = "14B" #default to this + log.warning("Model variant not detected, defaulting TeaCache coefficients to 14B") if dim == 1536: model_variant = "1_3B" log.info(f"Model variant detected: {model_variant}") @@ -1561,7 +1568,7 @@ class WanVideoImageToVideoEncode: "start_image": ("IMAGE", {"tooltip": "Image to encode"}), "end_image": ("IMAGE", {"tooltip": "end frame"}), "control_embeds": ("WANVIDIMAGE_EMBEDS", {"tooltip": "Control signal for the Fun -model"}), - "fun_model": ("BOOLEAN", {"default": False, "tooltip": "Enable when using Fun model"}), + "fun_or_fl2v_model": ("BOOLEAN", {"default": True, "tooltip": "Enable when using official FLF2V or Fun model"}), "temporal_mask": ("MASK", {"tooltip": "mask"}), "extra_latents": ("LATENT", {"tooltip": "Extra latents to add to the input front, used for Skyreels A2 reference images"}), } @@ -1573,7 +1580,7 @@ class WanVideoImageToVideoEncode: CATEGORY = "WanVideoWrapper" def process(self, vae, width, height, num_frames, clip_embeds, force_offload, noise_aug_strength, - start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_model=False, temporal_mask=None, extra_latents=None): + start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False, temporal_mask=None, extra_latents=None): device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -1589,7 +1596,7 @@ class WanVideoImageToVideoEncode: num_frames = ((num_frames - 1) // 4) * 4 + 1 two_ref_images = start_image is not None and end_image is not None - base_frames = num_frames + (1 if two_ref_images and not fun_model else 0) + base_frames = num_frames + (1 if two_ref_images and not fun_or_fl2v_model else 0) if temporal_mask is None: mask = torch.zeros(1, base_frames, lat_h, lat_w, device=device) if start_image is not None: @@ -1606,7 +1613,7 @@ class WanVideoImageToVideoEncode: # Repeat first frame and optionally end frame start_mask_repeated = torch.repeat_interleave(mask[:, 0:1], repeats=4, dim=1) # T, C, H, W - if end_image is not None and not fun_model: + if end_image is not None and not fun_or_fl2v_model: end_mask_repeated = torch.repeat_interleave(mask[:, -1:], repeats=4, dim=1) # T, C, H, W mask = torch.cat([start_mask_repeated, mask[:, 1:-1], end_mask_repeated], dim=1) else: @@ -1642,7 +1649,7 @@ class WanVideoImageToVideoEncode: elif start_image is None and end_image is None: concatenated = torch.zeros(3, num_frames, H, W, device=device) else: - if fun_model: + if fun_or_fl2v_model: zero_frames = torch.zeros(3, num_frames-(start_image.shape[0]+end_image.shape[0]), H, W, device=device) else: zero_frames = torch.zeros(3, num_frames-1, H, W, device=device) @@ -1651,7 +1658,7 @@ class WanVideoImageToVideoEncode: temporal_mask = common_upscale(temporal_mask.unsqueeze(1), W, H, "nearest", "disabled").squeeze(1) concatenated = resized_start_image[:,:num_frames] * temporal_mask[:num_frames].unsqueeze(0) - y = vae.encode([concatenated.to(device=device, dtype=vae.dtype)], device, end_=(end_image is not None and not fun_model))[0] + y = vae.encode([concatenated.to(device=device, dtype=vae.dtype)], device, end_=(end_image is not None and not fun_or_fl2v_model))[0] has_ref = False if extra_latents is not None: samples = extra_latents["samples"].squeeze(0) @@ -1673,7 +1680,7 @@ class WanVideoImageToVideoEncode: # Calculate maximum sequence length patches_per_frame = lat_h * lat_w // (patch_size[1] * patch_size[2]) - frames_per_stride = (num_frames - 1) // 4 + (2 if end_image is not None and not fun_model else 1) + frames_per_stride = (num_frames - 1) // 4 + (2 if end_image is not None and not fun_or_fl2v_model else 1) max_seq_len = frames_per_stride * patches_per_frame vae.model.clear_cache() @@ -1692,7 +1699,7 @@ class WanVideoImageToVideoEncode: "lat_w": lat_w, "control_embeds": control_embeds["control_embeds"] if control_embeds is not None else None, "end_image": resized_end_image if end_image is not None else None, - "fun_model": fun_model, + "fun_or_fl2v_model": fun_or_fl2v_model, "has_ref": has_ref, } @@ -2242,7 +2249,7 @@ class WanVideoSampler: control_latents, clip_fea, clip_fea_neg, end_image, recammaster, camera_embed = None, None, None, None, None, None vace_data, vace_context, vace_scale = None, None, None - fun_model, has_ref, drop_last = False, False, False + fun_or_fl2v_model, has_ref, drop_last = False, False, False image_cond = image_embeds.get("image_embeds", None) @@ -2252,10 +2259,10 @@ class WanVideoSampler: lat_w = image_embeds.get("lat_w", None) if lat_h is None or lat_w is None: raise ValueError("Clip encoded image embeds must be provided for I2V (Image to Video) model") - fun_model = image_embeds.get("fun_model", False) + fun_or_fl2v_model = image_embeds.get("fun_or_fl2v_model", False) noise = torch.randn( 16, - (image_embeds["num_frames"] - 1) // 4 + (2 if end_image is not None and not fun_model else 1), + (image_embeds["num_frames"] - 1) // 4 + (2 if end_image is not None and not fun_or_fl2v_model else 1), lat_h, lat_w, dtype=torch.float32, @@ -3029,7 +3036,7 @@ class WanVideoSampler: pass return ({ - "samples": x0.unsqueeze(0).cpu(), "looped": is_looped, "end_image": end_image if not fun_model else None, "has_ref": has_ref, "drop_last": drop_last, + "samples": x0.unsqueeze(0).cpu(), "looped": is_looped, "end_image": end_image if not fun_or_fl2v_model else None, "has_ref": has_ref, "drop_last": drop_last, }, ) class WindowTracker: diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 057b7f5..c19be23 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -665,15 +665,19 @@ class Head(nn.Module): class MLPProj(torch.nn.Module): - def __init__(self, in_dim, out_dim): + def __init__(self, in_dim, out_dim, fl_pos_emb=False): super().__init__() self.proj = torch.nn.Sequential( torch.nn.LayerNorm(in_dim), torch.nn.Linear(in_dim, in_dim), torch.nn.GELU(), torch.nn.Linear(in_dim, out_dim), torch.nn.LayerNorm(out_dim)) + if fl_pos_emb: # NOTE: we only use this for `fl2v` + self.emb_pos = nn.Parameter(torch.zeros(1, 257 * 2, 1280)) def forward(self, image_embeds): + if hasattr(self, 'emb_pos'): + image_embeds = image_embeds + self.emb_pos clip_extra_context_tokens = self.proj(image_embeds) return clip_extra_context_tokens @@ -750,7 +754,6 @@ class WanModel(ModelMixin, ConfigMixin): super().__init__() - assert model_type in ['t2v', 'i2v'] self.model_type = model_type self.patch_size = patch_size @@ -861,8 +864,8 @@ class WanModel(ModelMixin, ConfigMixin): assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0 - if model_type == 'i2v': - self.img_emb = MLPProj(1280, dim) + if model_type == 'i2v' or model_type == 'fl2v': + self.img_emb = MLPProj(1280, dim, fl_pos_emb=model_type == 'fl2v') def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False, vace_blocks_to_swap=None): log.info(f"Swapping {blocks_to_swap + 1} transformer blocks")