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
This commit is contained in:
kijai
2025-04-21 01:33:51 +03:00
parent 7d005201a2
commit e5a326c981
+17 -11
View File
@@ -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
)