From e5a326c9811514f2c08c89bccea9a7c731d9a503 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 21 Apr 2025 01:33:51 +0300 Subject: [PATCH] Use penultimate_hidden_states instead when using native clip vision The new Skyreels 1.3B I2V did not work otherwise, this seems to make it match the wrapper clip vision output better --- nodes.py | 28 +++++++++++++++++----------- 1 file changed, 17 insertions(+), 11 deletions(-) diff --git a/nodes.py b/nodes.py index 3e15b0e..6689482 100644 --- a/nodes.py +++ b/nodes.py @@ -620,7 +620,7 @@ class WanVideoModelLoader: "teacache_coefficients": teacache_coefficients_map[model_variant], "vace_layers": vace_layers, "vace_in_dim": vace_in_dim - } + } with init_empty_weights(): transformer = WanModel(**TRANSFORMER_CONFIG) @@ -662,10 +662,9 @@ class WanVideoModelLoader: total=param_count, leave=True): dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype - if "modulation" in name: + if "modulation" in name or "time_" in name: dtype_to_use = torch.float32 set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) - comfy_model.diffusion_model = transformer comfy_model.load_device = transformer_load_device @@ -1502,6 +1501,7 @@ class WanVideoClipVisionEncode: image = image_1 clip_vision.model.to(device) + negative_clip_embeds = None if tiles > 0: @@ -1511,16 +1511,17 @@ class WanVideoClipVisionEncode: negative_clip_embeds = clip_encode_image_tiled(clip_vision, negative_image.to(device), tiles=tiles, ratio=ratio) else: if isinstance(clip_vision, ClipVisionModel): - clip_embeds = clip_vision.encode_image(image).last_hidden_state.to(device) + clip_embeds = clip_vision.encode_image(image).penultimate_hidden_states.to(device) if negative_image is not None: - negative_clip_embeds = clip_vision.encode_image(negative_image).last_hidden_state.to(device) + negative_clip_embeds = clip_vision.encode_image(negative_image).penultimate_hidden_states.to(device) else: pixel_values = clip_preprocess(image.to(device), size=224, mean=image_mean, std=image_std, crop=(not crop == "disabled")).float() clip_embeds = clip_vision.visual(pixel_values) if negative_image is not None: pixel_values = clip_preprocess(negative_image.to(device), size=224, mean=image_mean, std=image_std, crop=(not crop == "disabled")).float() negative_clip_embeds = clip_vision.visual(pixel_values) - log.info(f"Clip embeds shape: {clip_embeds.shape}") + + log.info(f"Clip embeds shape: {clip_embeds.shape}, dtype: {clip_embeds.dtype}") weighted_embeds = [] weighted_embeds.append(clip_embeds[0:1] * strength_1) @@ -1542,6 +1543,8 @@ class WanVideoClipVisionEncode: clip_embeds = torch.cat(weighted_embeds, dim=1) elif combine_embeds == "batch": clip_embeds = torch.cat(weighted_embeds, dim=0) + else: + clip_embeds = weighted_embeds[0] log.info(f"Combined clip embeds shape: {clip_embeds.shape}") @@ -2139,7 +2142,7 @@ class WanVideoExperimentalArgs: return {"required": { "video_attention_split_steps": ("STRING", {"default": "", "tooltip": "Steps to split self attention when using multiple prompts"}), "cfg_zero_star": ("BOOLEAN", {"default": False, "tooltip": "https://github.com/WeichenFan/CFG-Zero-star"}), - "use_zero_init": ("BOOLEAN", {"default": True}), + "use_zero_init": ("BOOLEAN", {"default": False}), "zero_star_steps": ("INT", {"default": 0, "min": 0, "tooltip": "Steps to split self attention when using multiple prompts"}), "use_fresca": ("BOOLEAN", {"default": False, "tooltip": "https://github.com/WikiChao/FreSca"}), "fresca_scale_low": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), @@ -2689,7 +2692,6 @@ class WanVideoSampler: 'freqs': freqs, 't': timestep, 'current_step': idx, - 'y': [image_cond_input] if image_cond_input is not None else None, 'control_lora_enabled': control_lora_enabled, 'vace_data': vace_data, 'camera_embed': camera_embed, @@ -2704,7 +2706,8 @@ class WanVideoSampler: if not batched_cfg: #cond noise_pred_cond, teacache_state_cond = transformer( - [z], context=positive_embeds, clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage, + [z], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None, + clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage, pred_id=teacache_state[0] if teacache_state else None, **base_params ) @@ -2720,7 +2723,8 @@ class WanVideoSampler: return noise_pred_cond, [teacache_state_cond] #uncond noise_pred_uncond, teacache_state_uncond = transformer( - [z], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea, + [z], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea, + y=[image_cond_input] if image_cond_input is not None else None, is_uncond=True, current_step_percentage=current_step_percentage, pred_id=teacache_state[1] if teacache_state else None, **base_params @@ -2730,7 +2734,9 @@ class WanVideoSampler: else: teacache_state_uncond = None [noise_pred_cond, noise_pred_uncond], teacache_state_cond = transformer( - [z] + [z], context=positive_embeds + negative_embeds, clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage, + [z] + [z], context=positive_embeds + negative_embeds, + y=[image_cond_input] + [image_cond_input] if image_cond_input is not None else None, + clip_fea=clip_fea.repeat(2,1,1), is_uncond=False, current_step_percentage=current_step_percentage, pred_id=teacache_state[0] if teacache_state else None, **base_params )