Experimental clip embed enhancement through tiling

Allows more detail in the clip embeds as well as details from non square images without stretching

Idea and most of the code is from Matteo:

https://github.com/cubiq/ComfyUI_IPAdapter_plus/blob/9d076a3df0d2763cef5510ec5ab807f6632c39f5/utils.py#L181
This commit is contained in:
kijai
2025-03-21 19:41:49 +02:00
parent 5a2383621a
commit c3e96fd6f9
4 changed files with 148 additions and 21 deletions
+19 -10
View File
@@ -1,7 +1,7 @@
import os
import torch
import gc
from .utils import log, print_memory, apply_lora
from .utils import log, print_memory, apply_lora, clip_encode_image_tiled
import numpy as np
import math
from tqdm import tqdm
@@ -1056,6 +1056,7 @@ class WanVideoImageClipEncode:
return (image_embeds,)
#region clip vision
class WanVideoClipVisionEncode:
@classmethod
def INPUT_TYPES(s):
@@ -1070,6 +1071,8 @@ class WanVideoClipVisionEncode:
},
"optional": {
"image_2": ("IMAGE", ),
"tiles": ("INT", {"default": 0, "min": 0, "max": 16, "step": 2, "tooltip": "Use matteo's tiled image encoding for improved accuracy"}),
"ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Ratio of the tile average"}),
}
}
@@ -1078,26 +1081,32 @@ class WanVideoClipVisionEncode:
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, clip_vision, image_1, strength_1, strength_2, force_offload, crop, combine_embeds, image_2=None):
def process(self, clip_vision, image_1, strength_1, strength_2, force_offload, crop, combine_embeds, image_2=None, tiles=0, ratio=1.0):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
self.image_mean = [0.48145466, 0.4578275, 0.40821073]
self.image_std = [0.26862954, 0.26130258, 0.27577711]
clip_vision.model.to(device)
image_mean = [0.48145466, 0.4578275, 0.40821073]
image_std = [0.26862954, 0.26130258, 0.27577711]
if image_2 is not None:
image = torch.cat([image_1, image_2], dim=0)
else:
image = image_1
if isinstance(clip_vision, ClipVisionModel):
clip_embeds = clip_vision.encode_image(image).last_hidden_state.to(device)
clip_vision.model.to(device)
image = image.to(device)
if tiles > 0:
log.info("Using tiled image encoding")
clip_embeds = clip_encode_image_tiled(clip_vision, image, tiles=tiles, ratio=ratio)
else:
pixel_values = clip_preprocess(image.to(device), size=224, mean=self.image_mean, std=self.image_std, crop=(not crop == "disabled")).float()
clip_embeds = clip_vision.visual(pixel_values)
if isinstance(clip_vision, ClipVisionModel):
clip_embeds = clip_vision.encode_image(image).last_hidden_state.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)
log.info(f"Clip embeds shape: {clip_embeds.shape}")
if clip_embeds.shape[0] > 1:
embed_1 = clip_embeds[0:1] * strength_1
+85
View File
@@ -75,4 +75,89 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
#print("param.device", param.device)
set_module_tensor_to_device(model.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[name])
return model
# from https://github.com/cubiq/ComfyUI_IPAdapter_plus/blob/9d076a3df0d2763cef5510ec5ab807f6632c39f5/utils.py#L181
def split_tiles(embeds, num_split):
_, H, W, _ = embeds.shape
out = []
for x in embeds:
x = x.unsqueeze(0)
h, w = H // num_split, W // num_split
x_split = torch.cat([x[:, i*h:(i+1)*h, j*w:(j+1)*w, :] for i in range(num_split) for j in range(num_split)], dim=0)
out.append(x_split)
x_split = torch.stack(out, dim=0)
return x_split
def merge_hiddenstates(x, tiles):
chunk_size = tiles*tiles
x = x.split(chunk_size)
out = []
for embeds in x:
num_tiles = embeds.shape[0]
tile_size = int((embeds.shape[1]-1) ** 0.5)
grid_size = int(num_tiles ** 0.5)
# Extract class tokens
class_tokens = embeds[:, 0, :] # Save class tokens: [num_tiles, embeds[-1]]
avg_class_token = class_tokens.mean(dim=0, keepdim=True).unsqueeze(0) # Average token, shape: [1, 1, embeds[-1]]
patch_embeds = embeds[:, 1:, :] # Shape: [num_tiles, tile_size^2, embeds[-1]]
reshaped = patch_embeds.reshape(grid_size, grid_size, tile_size, tile_size, embeds.shape[-1])
merged = torch.cat([torch.cat([reshaped[i, j] for j in range(grid_size)], dim=1)
for i in range(grid_size)], dim=0)
merged = merged.unsqueeze(0) # Shape: [1, grid_size*tile_size, grid_size*tile_size, embeds[-1]]
# Pool to original size
pooled = torch.nn.functional.adaptive_avg_pool2d(merged.permute(0, 3, 1, 2), (tile_size, tile_size)).permute(0, 2, 3, 1)
flattened = pooled.reshape(1, tile_size*tile_size, embeds.shape[-1])
# Add back the class token
with_class = torch.cat([avg_class_token, flattened], dim=1) # Shape: original shape
out.append(with_class)
out = torch.cat(out, dim=0)
return out
from comfy.clip_vision import clip_preprocess, ClipVisionModel
def clip_encode_image_tiled(clip_vision, image, tiles=1, ratio=1.0):
embeds = encode_image_(clip_vision, image)
tiles = min(tiles, 16)
if tiles > 1:
# split in tiles
image_split = split_tiles(image, tiles)
# get the embeds for each tile
embeds_split = {}
for i in image_split:
encoded = encode_image_(clip_vision, i)
if not hasattr(embeds_split, "last_hidden_state"):
embeds_split["last_hidden_state"] = encoded
else:
embeds_split["last_hidden_state"] = torch.cat(embeds_split["last_hidden_state"], encoded, dim=0)
embeds_split['last_hidden_state'] = merge_hiddenstates(embeds_split['last_hidden_state'], tiles)
if embeds.shape[0] > 1: # if we have more than one image we need to average the embeddings for consistency
embeds = embeds * ratio + embeds_split['last_hidden_state']*(1-ratio)
else: # otherwise we can concatenate them, they can be averaged later
embeds = torch.cat([embeds * ratio, embeds_split['last_hidden_state']])
return embeds
def encode_image_(clip_vision, image):
if isinstance(clip_vision, ClipVisionModel):
out = clip_vision.encode_image(image).last_hidden_state
else:
pixel_values = clip_preprocess(image, size=224, crop=True).float()
out = clip_vision.visual(pixel_values)
return out
+35 -6
View File
@@ -198,11 +198,12 @@ class VisionTransformer(nn.Module):
kernel_size=patch_size,
stride=patch_size,
bias=not pre_norm)
if pool_type in ('token', 'token_fc'):
self.cls_embedding = nn.Parameter(gain * torch.randn(1, 1, dim))
self.pos_embedding = nn.Parameter(gain * torch.randn(
1, self.num_patches +
(1 if pool_type in ('token', 'token_fc') else 0), dim))
(1 if pool_type in ('token', 'token_fc') else 0), dim)) #torch.Size([1, 257, 1280])
self.dropout = nn.Dropout(embedding_dropout)
# transformer
@@ -222,13 +223,41 @@ class VisionTransformer(nn.Module):
def forward(self, x, interpolation=False, use_31_block=False):
b = x.size(0)
height, width = x.size(2), x.size(3)
# embeddings
x = self.patch_embedding(x).flatten(2).permute(0, 2, 1)
x_patch = self.patch_embedding(x)
B, C, H, W = x_patch.shape
x = x_patch.flatten(2).permute(0, 2, 1)
if self.pool_type in ('token', 'token_fc'):
x = torch.cat([self.cls_embedding.expand(b, -1, -1), x], dim=1)
if interpolation:
e = pos_interpolate(self.pos_embedding, x.size(1))
embeddings = x
num_patches = embeddings.shape[1] - 1
position_embedding = self.pos_embedding
num_positions = position_embedding.shape[1] - 1
class_pos_embed = position_embedding[:, :1]
patch_pos_embed = position_embedding[:, 1:]
dim = embeddings.shape[-1]
new_height = height // self.patch_size
new_width = width // self.patch_size
sqrt_num_positions = int(num_positions**0.5)
patch_pos_embed = patch_pos_embed.reshape(1, sqrt_num_positions, sqrt_num_positions, dim)
patch_pos_embed = patch_pos_embed.permute(0, 3, 1, 2)
patch_pos_embed = nn.functional.interpolate(
patch_pos_embed,
size=(new_height, new_width),
mode="bicubic",
align_corners=False,
)
patch_pos_embed = patch_pos_embed.permute(0, 2, 3, 1).view(1, -1, dim)
e = torch.cat((class_pos_embed, patch_pos_embed), dim=1)
else:
e = self.pos_embedding
x = self.dropout(x + e)
@@ -426,8 +455,8 @@ class CLIPModel:
for name, param in self.model.named_parameters():
set_module_tensor_to_device(self.model, name, device=device, dtype=dtype, value=state_dict[name])
def visual(self, image):
def visual(self, image, interpolation=False):
# forward
with torch.autocast(device_type=mm.get_autocast_device(self.device), dtype=self.dtype):
out = self.model.visual(image, use_31_block=True)
out = self.model.visual(image, interpolation=interpolation, use_31_block=True)
return out
+9 -5
View File
@@ -329,15 +329,15 @@ class WanI2VCrossAttention(WanSelfAttention):
self.norm_k_img = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.attention_mode = attention_mode
def forward(self, x, context, context_lens):
def forward(self, x, context, context_lens, clip_fea_tokens=257):
r"""
Args:
x(Tensor): Shape [B, L1, C]
context(Tensor): Shape [B, L2, C]
context_lens(Tensor): Shape [B]
"""
context_img = context[:, :257]
context = context[:, 257:]
context_img = context[:, :clip_fea_tokens]
context = context[:, clip_fea_tokens:]
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
@@ -417,6 +417,7 @@ class WanAttentionBlock(nn.Module):
context,
context_lens,
rope_func = "default",
clip_fea_tokens=257,
):
r"""
Args:
@@ -437,13 +438,13 @@ class WanAttentionBlock(nn.Module):
x = x.to(torch.float32) + (y.to(torch.float32) * e[2].to(torch.float32))
# cross-attention & ffn function
def cross_attn_ffn(x, context, context_lens, e):
def cross_attn_ffn(x, context, context_lens, e, clip_fea_tokens=clip_fea_tokens):
x = x + self.cross_attn(self.norm3(x), context, context_lens)
y = self.ffn(self.norm2(x).float() * (1 + e[4]) + e[3])
x = x.to(torch.float32) + (y.to(torch.float32) * e[5].to(torch.float32))
return x
x = cross_attn_ffn(x, context, context_lens, e)
x = cross_attn_ffn(x, context, context_lens, e, clip_fea_tokens=clip_fea_tokens)
return x
@@ -789,6 +790,8 @@ class WanModel(ModelMixin, ConfigMixin):
self.text_embedding.to(self.offload_device, non_blocking=self.use_non_blocking)
if clip_fea is not None:
clip_fea_tokens = clip_fea.shape[1]
clip_fea = clip_fea.to(self.main_device)
if self.offload_img_emb:
self.img_emb.to(self.main_device)
context_clip = self.img_emb(clip_fea) # bs x 257 x dim
@@ -846,6 +849,7 @@ class WanModel(ModelMixin, ConfigMixin):
freqs=freqs,
context=context,
context_lens=context_lens,
clip_fea_tokens=clip_fea_tokens,
rope_func=rope_func
)