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:
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user