From c3e96fd6f9de1807f03d72bffffe2697ecb3e0e6 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 21 Mar 2025 19:41:49 +0200 Subject: [PATCH] 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 --- nodes.py | 29 ++++++++----- utils.py | 85 +++++++++++++++++++++++++++++++++++++++ wanvideo/modules/clip.py | 41 ++++++++++++++++--- wanvideo/modules/model.py | 14 ++++--- 4 files changed, 148 insertions(+), 21 deletions(-) diff --git a/nodes.py b/nodes.py index 40eb4e5..ea15fbd 100644 --- a/nodes.py +++ b/nodes.py @@ -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 diff --git a/utils.py b/utils.py index 9799f7b..2b6234b 100644 --- a/utils.py +++ b/utils.py @@ -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 \ No newline at end of file diff --git a/wanvideo/modules/clip.py b/wanvideo/modules/clip.py index 304fbb3..eb39f36 100644 --- a/wanvideo/modules/clip.py +++ b/wanvideo/modules/clip.py @@ -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 diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index a9f7da3..bf39b4c 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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 )