diff --git a/nodes.py b/nodes.py index 99377b9..9e8dcad 100644 --- a/nodes.py +++ b/nodes.py @@ -39,7 +39,7 @@ class WanVideoBlockSwap: def INPUT_TYPES(s): return { "required": { - "blocks_to_swap": ("INT", {"default": 20, "min": 0, "max": 40, "step": 1, "tooltip": "Number of double blocks to swap"}), + "blocks_to_swap": ("INT", {"default": 20, "min": 0, "max": 40, "step": 1, "tooltip": "Number of transformer blocks to swap, the 14B model has 40, while the 1.3B model has 30 blocks"}), "offload_img_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload img_emb to offload_device"}), "offload_txt_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload time_emb to offload_device"}), }, diff --git a/pyproject.toml b/pyproject.toml index a39e49d..1f2251b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "ComfyUI-WanVideoWrapper" description = "ComfyUI diffusers wrapper nodes for WanVideo" -version = "1.0.3" +version = "1.0.4" license = {file = "LICENSE"} dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.31.0", "ftfy"] diff --git a/utils.py b/utils.py index ac263e9..ce7dbe9 100644 --- a/utils.py +++ b/utils.py @@ -21,4 +21,11 @@ def print_memory(device): log.info(f"Max allocated memory: {max_memory=:.3f} GB") log.info(f"Max reserved memory: {max_reserved=:.3f} GB") #memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False) - #log.info(f"Memory Summary:\n{memory_summary}") \ No newline at end of file + #log.info(f"Memory Summary:\n{memory_summary}") + +def get_module_memory_mb(module): + memory = 0 + for param in module.parameters(): + if param.data is not None: + memory += param.nelement() * param.element_size() + return memory / (1024 * 1024) # Convert to MB \ No newline at end of file diff --git a/wanvideo/modules/clip.py b/wanvideo/modules/clip.py index bd9da3f..304fbb3 100644 --- a/wanvideo/modules/clip.py +++ b/wanvideo/modules/clip.py @@ -9,8 +9,6 @@ import torch.nn.functional as F import torchvision.transforms as T from .attention import attention -from .tokenizers import HuggingfaceTokenizer -from .xlm_roberta import XLMRoberta __all__ = [ 'XLMRobertaCLIP', @@ -155,60 +153,6 @@ class AttentionBlock(nn.Module): x = x + self.mlp(self.norm2(x)) return x - -class AttentionPool(nn.Module): - - def __init__(self, - dim, - mlp_ratio, - num_heads, - activation='gelu', - proj_dropout=0.0, - norm_eps=1e-5): - assert dim % num_heads == 0 - super().__init__() - self.dim = dim - self.mlp_ratio = mlp_ratio - self.num_heads = num_heads - self.head_dim = dim // num_heads - self.proj_dropout = proj_dropout - self.norm_eps = norm_eps - - # layers - gain = 1.0 / math.sqrt(dim) - self.cls_embedding = nn.Parameter(gain * torch.randn(1, 1, dim)) - self.to_q = nn.Linear(dim, dim) - self.to_kv = nn.Linear(dim, dim * 2) - self.proj = nn.Linear(dim, dim) - self.norm = LayerNorm(dim, eps=norm_eps) - self.mlp = nn.Sequential( - nn.Linear(dim, int(dim * mlp_ratio)), - QuickGELU() if activation == 'quick_gelu' else nn.GELU(), - nn.Linear(int(dim * mlp_ratio), dim), nn.Dropout(proj_dropout)) - - def forward(self, x): - """ - x: [B, L, C]. - """ - b, s, c, n, d = *x.size(), self.num_heads, self.head_dim - - # compute query, key, value - q = self.to_q(self.cls_embedding).view(1, 1, n, d).expand(b, -1, -1, -1) - k, v = self.to_kv(x).view(b, s, 2, n, d).unbind(2) - - # compute attention - x = flash_attention(q, k, v, version=2) - x = x.reshape(b, 1, c) - - # output - x = self.proj(x) - x = F.dropout(x, self.proj_dropout, self.training) - - # mlp - x = x + self.mlp(self.norm(x)) - return x[:, 0] - - class VisionTransformer(nn.Module): def __init__(self, @@ -275,9 +219,6 @@ class VisionTransformer(nn.Module): self.head = nn.Parameter(gain * torch.randn(dim, out_dim)) elif pool_type == 'token_fc': self.head = nn.Linear(dim, out_dim) - elif pool_type == 'attn_pool': - self.head = AttentionPool(dim, mlp_ratio, num_heads, activation, - proj_dropout, norm_eps) def forward(self, x, interpolation=False, use_31_block=False): b = x.size(0) @@ -303,31 +244,6 @@ class VisionTransformer(nn.Module): return x -class XLMRobertaWithHead(XLMRoberta): - - def __init__(self, **kwargs): - self.out_dim = kwargs.pop('out_dim') - super().__init__(**kwargs) - - # head - mid_dim = (self.dim + self.out_dim) // 2 - self.head = nn.Sequential( - nn.Linear(self.dim, mid_dim, bias=False), nn.GELU(), - nn.Linear(mid_dim, self.out_dim, bias=False)) - - def forward(self, ids): - # xlm-roberta - x = super().forward(ids) - - # average pooling - mask = ids.ne(self.pad_id).unsqueeze(-1).to(x) - x = (x * mask).sum(dim=1) / mask.sum(dim=1) - - # head - x = self.head(x) - return x - - class XLMRobertaCLIP(nn.Module): def __init__(self, diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index d9a8bf3..73f28c2 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -15,7 +15,7 @@ __all__ = ['WanModel'] from tqdm import tqdm -from ...utils import log +from ...utils import log, get_module_memory_mb def poly1d(coefficients, x): result = torch.zeros_like(x) @@ -547,12 +547,27 @@ class WanModel(ModelMixin, ConfigMixin): self.blocks_to_swap = blocks_to_swap self.offload_img_emb = offload_img_emb self.offload_txt_emb = offload_txt_emb + + total_offload_memory = 0 + total_main_memory = 0 for b, block in tqdm(enumerate(self.blocks), total=len(self.blocks), desc="Initializing block swap"): + block_memory = get_module_memory_mb(block) + if b > self.blocks_to_swap: block.to(self.main_device) + total_main_memory += block_memory else: block.to(self.offload_device) + total_offload_memory += block_memory + + #print(f"Block {b}: {block_memory:.2f}MB on {block.parameters().__next__().device}") + log.info("----------------------") + log.info(f"Block swap memory summary:") + log.info(f"Transformer blocks on {self.offload_device}: {total_offload_memory:.2f}MB") + log.info(f"Transformer blocks on {self.main_device}: {total_main_memory:.2f}MB") + log.info(f"Total Memory: {(total_offload_memory + total_main_memory):.2f}MB") + log.info("----------------------") def forward( self, diff --git a/wanvideo/modules/xlm_roberta.py b/wanvideo/modules/xlm_roberta.py deleted file mode 100644 index 4bd38c1..0000000 --- a/wanvideo/modules/xlm_roberta.py +++ /dev/null @@ -1,170 +0,0 @@ -# Modified from transformers.models.xlm_roberta.modeling_xlm_roberta -# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved. -import torch -import torch.nn as nn -import torch.nn.functional as F - -__all__ = ['XLMRoberta', 'xlm_roberta_large'] - - -class SelfAttention(nn.Module): - - def __init__(self, dim, num_heads, dropout=0.1, eps=1e-5): - assert dim % num_heads == 0 - super().__init__() - self.dim = dim - self.num_heads = num_heads - self.head_dim = dim // num_heads - self.eps = eps - - # layers - self.q = nn.Linear(dim, dim) - self.k = nn.Linear(dim, dim) - self.v = nn.Linear(dim, dim) - self.o = nn.Linear(dim, dim) - self.dropout = nn.Dropout(dropout) - - def forward(self, x, mask): - """ - x: [B, L, C]. - """ - b, s, c, n, d = *x.size(), self.num_heads, self.head_dim - - # compute query, key, value - q = self.q(x).reshape(b, s, n, d).permute(0, 2, 1, 3) - k = self.k(x).reshape(b, s, n, d).permute(0, 2, 1, 3) - v = self.v(x).reshape(b, s, n, d).permute(0, 2, 1, 3) - - # compute attention - p = self.dropout.p if self.training else 0.0 - x = F.scaled_dot_product_attention(q, k, v, mask, p) - x = x.permute(0, 2, 1, 3).reshape(b, s, c) - - # output - x = self.o(x) - x = self.dropout(x) - return x - - -class AttentionBlock(nn.Module): - - def __init__(self, dim, num_heads, post_norm, dropout=0.1, eps=1e-5): - super().__init__() - self.dim = dim - self.num_heads = num_heads - self.post_norm = post_norm - self.eps = eps - - # layers - self.attn = SelfAttention(dim, num_heads, dropout, eps) - self.norm1 = nn.LayerNorm(dim, eps=eps) - self.ffn = nn.Sequential( - nn.Linear(dim, dim * 4), nn.GELU(), nn.Linear(dim * 4, dim), - nn.Dropout(dropout)) - self.norm2 = nn.LayerNorm(dim, eps=eps) - - def forward(self, x, mask): - if self.post_norm: - x = self.norm1(x + self.attn(x, mask)) - x = self.norm2(x + self.ffn(x)) - else: - x = x + self.attn(self.norm1(x), mask) - x = x + self.ffn(self.norm2(x)) - return x - - -class XLMRoberta(nn.Module): - """ - XLMRobertaModel with no pooler and no LM head. - """ - - def __init__(self, - vocab_size=250002, - max_seq_len=514, - type_size=1, - pad_id=1, - dim=1024, - num_heads=16, - num_layers=24, - post_norm=True, - dropout=0.1, - eps=1e-5): - super().__init__() - self.vocab_size = vocab_size - self.max_seq_len = max_seq_len - self.type_size = type_size - self.pad_id = pad_id - self.dim = dim - self.num_heads = num_heads - self.num_layers = num_layers - self.post_norm = post_norm - self.eps = eps - - # embeddings - self.token_embedding = nn.Embedding(vocab_size, dim, padding_idx=pad_id) - self.type_embedding = nn.Embedding(type_size, dim) - self.pos_embedding = nn.Embedding(max_seq_len, dim, padding_idx=pad_id) - self.dropout = nn.Dropout(dropout) - - # blocks - self.blocks = nn.ModuleList([ - AttentionBlock(dim, num_heads, post_norm, dropout, eps) - for _ in range(num_layers) - ]) - - # norm layer - self.norm = nn.LayerNorm(dim, eps=eps) - - def forward(self, ids): - """ - ids: [B, L] of torch.LongTensor. - """ - b, s = ids.shape - mask = ids.ne(self.pad_id).long() - - # embeddings - x = self.token_embedding(ids) + \ - self.type_embedding(torch.zeros_like(ids)) + \ - self.pos_embedding(self.pad_id + torch.cumsum(mask, dim=1) * mask) - if self.post_norm: - x = self.norm(x) - x = self.dropout(x) - - # blocks - mask = torch.where( - mask.view(b, 1, 1, s).gt(0), 0.0, - torch.finfo(x.dtype).min) - for block in self.blocks: - x = block(x, mask) - - # output - if not self.post_norm: - x = self.norm(x) - return x - - -def xlm_roberta_large(pretrained=False, - return_tokenizer=False, - device='cpu', - **kwargs): - """ - XLMRobertaLarge adapted from Huggingface. - """ - # params - cfg = dict( - vocab_size=250002, - max_seq_len=514, - type_size=1, - pad_id=1, - dim=1024, - num_heads=16, - num_layers=24, - post_norm=True, - dropout=0.1, - eps=1e-5) - cfg.update(**kwargs) - - # init a model on device - with torch.device(device): - model = XLMRoberta(**cfg) - return model