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 )