hugging-apps/echo-memory
0
1"""2Concise re-implementation of3``https://github.com/openai/CLIP'' and4``https://github.com/mlfoundations/open_clip''.5"""6import math7import torch8import torch.nn as nn9import torch.nn.functional as F10import torchvision.transforms as T11from .wan_video_dit import flash_attention12 13 14class SelfAttention(nn.Module):15 16 def __init__(self, dim, num_heads, dropout=0.1, eps=1e-5):17 assert dim % num_heads == 018 super().__init__()19 self.dim = dim20 self.num_heads = num_heads21 self.head_dim = dim // num_heads22 self.eps = eps23 24 # layers25 self.q = nn.Linear(dim, dim)26 self.k = nn.Linear(dim, dim)27 self.v = nn.Linear(dim, dim)28 self.o = nn.Linear(dim, dim)29 self.dropout = nn.Dropout(dropout)30 31 def forward(self, x, mask):32 """33 x: [B, L, C].34 """35 b, s, c, n, d = *x.size(), self.num_heads, self.head_dim36 37 # compute query, key, value38 q = self.q(x).reshape(b, s, n, d).permute(0, 2, 1, 3)39 k = self.k(x).reshape(b, s, n, d).permute(0, 2, 1, 3)40 v = self.v(x).reshape(b, s, n, d).permute(0, 2, 1, 3)41 42 # compute attention43 p = self.dropout.p if self.training else 0.044 x = F.scaled_dot_product_attention(q, k, v, mask, p)45 x = x.permute(0, 2, 1, 3).reshape(b, s, c)46 47 # output48 x = self.o(x)49 x = self.dropout(x)50 return x51 52 53class AttentionBlock(nn.Module):54 55 def __init__(self, dim, num_heads, post_norm, dropout=0.1, eps=1e-5):56 super().__init__()57 self.dim = dim58 self.num_heads = num_heads59 self.post_norm = post_norm60 self.eps = eps61 62 # layers63 self.attn = SelfAttention(dim, num_heads, dropout, eps)64 self.norm1 = nn.LayerNorm(dim, eps=eps)65 self.ffn = nn.Sequential(66 nn.Linear(dim, dim * 4), nn.GELU(), nn.Linear(dim * 4, dim),67 nn.Dropout(dropout))68 self.norm2 = nn.LayerNorm(dim, eps=eps)69 70 def forward(self, x, mask):71 if self.post_norm:72 x = self.norm1(x + self.attn(x, mask))73 x = self.norm2(x + self.ffn(x))74 else:75 x = x + self.attn(self.norm1(x), mask)76 x = x + self.ffn(self.norm2(x))77 return x78 79 80class XLMRoberta(nn.Module):81 """82 XLMRobertaModel with no pooler and no LM head.83 """84 85 def __init__(self,86 vocab_size=250002,87 max_seq_len=514,88 type_size=1,89 pad_id=1,90 dim=1024,91 num_heads=16,92 num_layers=24,93 post_norm=True,94 dropout=0.1,95 eps=1e-5):96 super().__init__()97 self.vocab_size = vocab_size98 self.max_seq_len = max_seq_len99 self.type_size = type_size100 self.pad_id = pad_id101 self.dim = dim102 self.num_heads = num_heads103 self.num_layers = num_layers104 self.post_norm = post_norm105 self.eps = eps106 107 # embeddings108 self.token_embedding = nn.Embedding(vocab_size, dim, padding_idx=pad_id)109 self.type_embedding = nn.Embedding(type_size, dim)110 self.pos_embedding = nn.Embedding(max_seq_len, dim, padding_idx=pad_id)111 self.dropout = nn.Dropout(dropout)112 113 # blocks114 self.blocks = nn.ModuleList([115 AttentionBlock(dim, num_heads, post_norm, dropout, eps)116 for _ in range(num_layers)117 ])118 119 # norm layer120 self.norm = nn.LayerNorm(dim, eps=eps)121 122 def forward(self, ids):123 """124 ids: [B, L] of torch.LongTensor.125 """126 b, s = ids.shape127 mask = ids.ne(self.pad_id).long()128 129 # embeddings130 x = self.token_embedding(ids) + \131 self.type_embedding(torch.zeros_like(ids)) + \132 self.pos_embedding(self.pad_id + torch.cumsum(mask, dim=1) * mask)133 if self.post_norm:134 x = self.norm(x)135 x = self.dropout(x)136 137 # blocks138 mask = torch.where(139 mask.view(b, 1, 1, s).gt(0), 0.0,140 torch.finfo(x.dtype).min)141 for block in self.blocks:142 x = block(x, mask)143 144 # output145 if not self.post_norm:146 x = self.norm(x)147 return x148 149 150def xlm_roberta_large(pretrained=False,151 return_tokenizer=False,152 device='cpu',153 **kwargs):154 """155 XLMRobertaLarge adapted from Huggingface.156 """157 # params158 cfg = dict(159 vocab_size=250002,160 max_seq_len=514,161 type_size=1,162 pad_id=1,163 dim=1024,164 num_heads=16,165 num_layers=24,166 post_norm=True,167 dropout=0.1,168 eps=1e-5)169 cfg.update(**kwargs)170 171 # init model172 if pretrained:173 from sora import DOWNLOAD_TO_CACHE174 175 # init a meta model176 with torch.device('meta'):177 model = XLMRoberta(**cfg)178 179 # load checkpoint180 model.load_state_dict(181 torch.load(182 DOWNLOAD_TO_CACHE('models/xlm_roberta/xlm_roberta_large.pth'),183 map_location=device),184 assign=True)185 else:186 # init a model on device187 with torch.device(device):188 model = XLMRoberta(**cfg)189 190 # init tokenizer191 if return_tokenizer:192 from sora.data import HuggingfaceTokenizer193 tokenizer = HuggingfaceTokenizer(194 name='xlm-roberta-large',195 seq_len=model.text_len,196 clean='whitespace')197 return model, tokenizer198 else:199 return model200 201 202 203def pos_interpolate(pos, seq_len):204 if pos.size(1) == seq_len:205 return pos206 else:207 src_grid = int(math.sqrt(pos.size(1)))208 tar_grid = int(math.sqrt(seq_len))209 n = pos.size(1) - src_grid * src_grid210 return torch.cat([211 pos[:, :n],212 F.interpolate(213 pos[:, n:].float().reshape(1, src_grid, src_grid, -1).permute(214 0, 3, 1, 2),215 size=(tar_grid, tar_grid),216 mode='bicubic',217 align_corners=False).flatten(2).transpose(1, 2)218 ],219 dim=1)220 221 222class QuickGELU(nn.Module):223 224 def forward(self, x):225 return x * torch.sigmoid(1.702 * x)226 227 228class LayerNorm(nn.LayerNorm):229 230 def forward(self, x):231 return super().forward(x).type_as(x)232 233 234class SelfAttention(nn.Module):235 236 def __init__(self,237 dim,238 num_heads,239 causal=False,240 attn_dropout=0.0,241 proj_dropout=0.0):242 assert dim % num_heads == 0243 super().__init__()244 self.dim = dim245 self.num_heads = num_heads246 self.head_dim = dim // num_heads247 self.causal = causal248 self.attn_dropout = attn_dropout249 self.proj_dropout = proj_dropout250 251 # layers252 self.to_qkv = nn.Linear(dim, dim * 3)253 self.proj = nn.Linear(dim, dim)254 255 def forward(self, x):256 """257 x: [B, L, C].258 """259 # compute query, key, value260 q, k, v = self.to_qkv(x).chunk(3, dim=-1)261 262 # compute attention263 x = flash_attention(q, k, v, num_heads=self.num_heads, compatibility_mode=True)264 265 # output266 x = self.proj(x)267 x = F.dropout(x, self.proj_dropout, self.training)268 return x269 270 271class SwiGLU(nn.Module):272 273 def __init__(self, dim, mid_dim):274 super().__init__()275 self.dim = dim276 self.mid_dim = mid_dim277 278 # layers279 self.fc1 = nn.Linear(dim, mid_dim)280 self.fc2 = nn.Linear(dim, mid_dim)281 self.fc3 = nn.Linear(mid_dim, dim)282 283 def forward(self, x):284 x = F.silu(self.fc1(x)) * self.fc2(x)285 x = self.fc3(x)286 return x287 288 289class AttentionBlock(nn.Module):290 291 def __init__(self,292 dim,293 mlp_ratio,294 num_heads,295 post_norm=False,296 causal=False,297 activation='quick_gelu',298 attn_dropout=0.0,299 proj_dropout=0.0,300 norm_eps=1e-5):301 assert activation in ['quick_gelu', 'gelu', 'swi_glu']302 super().__init__()303 self.dim = dim304 self.mlp_ratio = mlp_ratio305 self.num_heads = num_heads306 self.post_norm = post_norm307 self.causal = causal308 self.norm_eps = norm_eps309 310 # layers311 self.norm1 = LayerNorm(dim, eps=norm_eps)312 self.attn = SelfAttention(dim, num_heads, causal, attn_dropout,313 proj_dropout)314 self.norm2 = LayerNorm(dim, eps=norm_eps)315 if activation == 'swi_glu':316 self.mlp = SwiGLU(dim, int(dim * mlp_ratio))317 else:318 self.mlp = nn.Sequential(319 nn.Linear(dim, int(dim * mlp_ratio)),320 QuickGELU() if activation == 'quick_gelu' else nn.GELU(),321 nn.Linear(int(dim * mlp_ratio), dim), nn.Dropout(proj_dropout))322 323 def forward(self, x):324 if self.post_norm:325 x = x + self.norm1(self.attn(x))326 x = x + self.norm2(self.mlp(x))327 else:328 x = x + self.attn(self.norm1(x))329 x = x + self.mlp(self.norm2(x))330 return x331 332 333class AttentionPool(nn.Module):334 335 def __init__(self,336 dim,337 mlp_ratio,338 num_heads,339 activation='gelu',340 proj_dropout=0.0,341 norm_eps=1e-5):342 assert dim % num_heads == 0343 super().__init__()344 self.dim = dim345 self.mlp_ratio = mlp_ratio346 self.num_heads = num_heads347 self.head_dim = dim // num_heads348 self.proj_dropout = proj_dropout349 self.norm_eps = norm_eps350 351 # layers352 gain = 1.0 / math.sqrt(dim)353 self.cls_embedding = nn.Parameter(gain * torch.randn(1, 1, dim))354 self.to_q = nn.Linear(dim, dim)355 self.to_kv = nn.Linear(dim, dim * 2)356 self.proj = nn.Linear(dim, dim)357 self.norm = LayerNorm(dim, eps=norm_eps)358 self.mlp = nn.Sequential(359 nn.Linear(dim, int(dim * mlp_ratio)),360 QuickGELU() if activation == 'quick_gelu' else nn.GELU(),361 nn.Linear(int(dim * mlp_ratio), dim), nn.Dropout(proj_dropout))362 363 def forward(self, x):364 """365 x: [B, L, C].366 """367 b, s, c, n, d = *x.size(), self.num_heads, self.head_dim368 369 # compute query, key, value370 q = self.to_q(self.cls_embedding).view(1, 1, n*d).expand(b, -1, -1)371 k, v = self.to_kv(x).chunk(2, dim=-1)372 373 # compute attention374 x = flash_attention(q, k, v, num_heads=self.num_heads, compatibility_mode=True)375 x = x.reshape(b, 1, c)376 377 # output378 x = self.proj(x)379 x = F.dropout(x, self.proj_dropout, self.training)380 381 # mlp382 x = x + self.mlp(self.norm(x))383 return x[:, 0]384 385 386class VisionTransformer(nn.Module):387 388 def __init__(self,389 image_size=224,390 patch_size=16,391 dim=768,392 mlp_ratio=4,393 out_dim=512,394 num_heads=12,395 num_layers=12,396 pool_type='token',397 pre_norm=True,398 post_norm=False,399 activation='quick_gelu',400 attn_dropout=0.0,401 proj_dropout=0.0,402 embedding_dropout=0.0,403 norm_eps=1e-5):404 if image_size % patch_size != 0:405 print(406 '[WARNING] image_size is not divisible by patch_size',407 flush=True)408 assert pool_type in ('token', 'token_fc', 'attn_pool')409 out_dim = out_dim or dim410 super().__init__()411 self.image_size = image_size412 self.patch_size = patch_size413 self.num_patches = (image_size // patch_size)**2414 self.dim = dim415 self.mlp_ratio = mlp_ratio416 self.out_dim = out_dim417 self.num_heads = num_heads418 self.num_layers = num_layers419 self.pool_type = pool_type420 self.post_norm = post_norm421 self.norm_eps = norm_eps422 423 # embeddings424 gain = 1.0 / math.sqrt(dim)425 self.patch_embedding = nn.Conv2d(426 3,427 dim,428 kernel_size=patch_size,429 stride=patch_size,430 bias=not pre_norm)431 if pool_type in ('token', 'token_fc'):432 self.cls_embedding = nn.Parameter(gain * torch.randn(1, 1, dim))433 self.pos_embedding = nn.Parameter(gain * torch.randn(434 1, self.num_patches +435 (1 if pool_type in ('token', 'token_fc') else 0), dim))436 self.dropout = nn.Dropout(embedding_dropout)437 438 # transformer439 self.pre_norm = LayerNorm(dim, eps=norm_eps) if pre_norm else None440 self.transformer = nn.Sequential(*[441 AttentionBlock(dim, mlp_ratio, num_heads, post_norm, False,442 activation, attn_dropout, proj_dropout, norm_eps)443 for _ in range(num_layers)444 ])445 self.post_norm = LayerNorm(dim, eps=norm_eps)446 447 # head448 if pool_type == 'token':449 self.head = nn.Parameter(gain * torch.randn(dim, out_dim))450 elif pool_type == 'token_fc':451 self.head = nn.Linear(dim, out_dim)452 elif pool_type == 'attn_pool':453 self.head = AttentionPool(dim, mlp_ratio, num_heads, activation,454 proj_dropout, norm_eps)455 456 def forward(self, x, interpolation=False, use_31_block=False):457 b = x.size(0)458 459 # embeddings460 x = self.patch_embedding(x).flatten(2).permute(0, 2, 1)461 if self.pool_type in ('token', 'token_fc'):462 x = torch.cat([self.cls_embedding.expand(b, -1, -1).to(dtype=x.dtype, device=x.device), x], dim=1)463 if interpolation:464 e = pos_interpolate(self.pos_embedding, x.size(1))465 else:466 e = self.pos_embedding467 e = e.to(dtype=x.dtype, device=x.device)468 x = self.dropout(x + e)469 if self.pre_norm is not None:470 x = self.pre_norm(x)471 472 # transformer473 if use_31_block:474 x = self.transformer[:-1](x)475 return x476 else:477 x = self.transformer(x)478 return x479 480 481class CLIP(nn.Module):482 483 def __init__(self,484 embed_dim=512,485 image_size=224,486 patch_size=16,487 vision_dim=768,488 vision_mlp_ratio=4,489 vision_heads=12,490 vision_layers=12,491 vision_pool='token',492 vision_pre_norm=True,493 vision_post_norm=False,494 vocab_size=49408,495 text_len=77,496 text_dim=512,497 text_mlp_ratio=4,498 text_heads=8,499 text_layers=12,500 text_causal=True,501 text_pool='argmax',502 text_head_bias=False,503 logit_bias=None,504 activation='quick_gelu',505 attn_dropout=0.0,506 proj_dropout=0.0,507 embedding_dropout=0.0,508 norm_eps=1e-5):509 super().__init__()510 self.embed_dim = embed_dim511 self.image_size = image_size512 self.patch_size = patch_size513 self.vision_dim = vision_dim514 self.vision_mlp_ratio = vision_mlp_ratio515 self.vision_heads = vision_heads516 self.vision_layers = vision_layers517 self.vision_pool = vision_pool518 self.vision_pre_norm = vision_pre_norm519 self.vision_post_norm = vision_post_norm520 self.vocab_size = vocab_size521 self.text_len = text_len522 self.text_dim = text_dim523 self.text_mlp_ratio = text_mlp_ratio524 self.text_heads = text_heads525 self.text_layers = text_layers526 self.text_causal = text_causal527 self.text_pool = text_pool528 self.text_head_bias = text_head_bias529 self.norm_eps = norm_eps530 531 # models532 self.visual = VisionTransformer(533 image_size=image_size,534 patch_size=patch_size,535 dim=vision_dim,536 mlp_ratio=vision_mlp_ratio,537 out_dim=embed_dim,538 num_heads=vision_heads,539 num_layers=vision_layers,540 pool_type=vision_pool,541 pre_norm=vision_pre_norm,542 post_norm=vision_post_norm,543 activation=activation,544 attn_dropout=attn_dropout,545 proj_dropout=proj_dropout,546 embedding_dropout=embedding_dropout,547 norm_eps=norm_eps)548 self.textual = TextTransformer(549 vocab_size=vocab_size,550 text_len=text_len,551 dim=text_dim,552 mlp_ratio=text_mlp_ratio,553 out_dim=embed_dim,554 num_heads=text_heads,555 num_layers=text_layers,556 causal=text_causal,557 pool_type=text_pool,558 head_bias=text_head_bias,559 activation=activation,560 attn_dropout=attn_dropout,561 proj_dropout=proj_dropout,562 embedding_dropout=embedding_dropout,563 norm_eps=norm_eps)564 self.log_scale = nn.Parameter(math.log(1 / 0.07) * torch.ones([]))565 if logit_bias is not None:566 self.logit_bias = nn.Parameter(logit_bias * torch.ones([]))567 568 # initialize weights569 self.init_weights()570 571 def forward(self, imgs, txt_ids):572 """573 imgs: [B, 3, H, W] of torch.float32.574 - mean: [0.48145466, 0.4578275, 0.40821073]575 - std: [0.26862954, 0.26130258, 0.27577711]576 txt_ids: [B, L] of torch.long. Encoded by data.CLIPTokenizer.577 """578 xi = self.visual(imgs)579 xt = self.textual(txt_ids)580 return xi, xt581 582 def init_weights(self):583 # embeddings584 nn.init.normal_(self.textual.token_embedding.weight, std=0.02)585 nn.init.normal_(self.visual.patch_embedding.weight, std=0.1)586 587 # attentions588 for modality in ['visual', 'textual']:589 dim = self.vision_dim if modality == 'visual' else self.text_dim590 transformer = getattr(self, modality).transformer591 proj_gain = (1.0 / math.sqrt(dim)) * (592 1.0 / math.sqrt(2 * len(transformer)))593 attn_gain = 1.0 / math.sqrt(dim)594 mlp_gain = 1.0 / math.sqrt(2.0 * dim)595 for block in transformer:596 nn.init.normal_(block.attn.to_qkv.weight, std=attn_gain)597 nn.init.normal_(block.attn.proj.weight, std=proj_gain)598 nn.init.normal_(block.mlp[0].weight, std=mlp_gain)599 nn.init.normal_(block.mlp[2].weight, std=proj_gain)600 601 def param_groups(self):602 groups = [{603 'params': [604 p for n, p in self.named_parameters()605 if 'norm' in n or n.endswith('bias')606 ],607 'weight_decay': 0.0608 }, {609 'params': [610 p for n, p in self.named_parameters()611 if not ('norm' in n or n.endswith('bias'))612 ]613 }]614 return groups615 616 617class XLMRobertaWithHead(XLMRoberta):618 619 def __init__(self, **kwargs):620 self.out_dim = kwargs.pop('out_dim')621 super().__init__(**kwargs)622 623 # head624 mid_dim = (self.dim + self.out_dim) // 2625 self.head = nn.Sequential(626 nn.Linear(self.dim, mid_dim, bias=False), nn.GELU(),627 nn.Linear(mid_dim, self.out_dim, bias=False))628 629 def forward(self, ids):630 # xlm-roberta631 x = super().forward(ids)632 633 # average pooling634 mask = ids.ne(self.pad_id).unsqueeze(-1).to(x)635 x = (x * mask).sum(dim=1) / mask.sum(dim=1)636 637 # head638 x = self.head(x)639 return x640 641 642class XLMRobertaCLIP(nn.Module):643 644 def __init__(self,645 embed_dim=1024,646 image_size=224,647 patch_size=14,648 vision_dim=1280,649 vision_mlp_ratio=4,650 vision_heads=16,651 vision_layers=32,652 vision_pool='token',653 vision_pre_norm=True,654 vision_post_norm=False,655 activation='gelu',656 vocab_size=250002,657 max_text_len=514,658 type_size=1,659 pad_id=1,660 text_dim=1024,661 text_heads=16,662 text_layers=24,663 text_post_norm=True,664 text_dropout=0.1,665 attn_dropout=0.0,666 proj_dropout=0.0,667 embedding_dropout=0.0,668 norm_eps=1e-5):669 super().__init__()670 self.embed_dim = embed_dim671 self.image_size = image_size672 self.patch_size = patch_size673 self.vision_dim = vision_dim674 self.vision_mlp_ratio = vision_mlp_ratio675 self.vision_heads = vision_heads676 self.vision_layers = vision_layers677 self.vision_pre_norm = vision_pre_norm678 self.vision_post_norm = vision_post_norm679 self.activation = activation680 self.vocab_size = vocab_size681 self.max_text_len = max_text_len682 self.type_size = type_size683 self.pad_id = pad_id684 self.text_dim = text_dim685 self.text_heads = text_heads686 self.text_layers = text_layers687 self.text_post_norm = text_post_norm688 self.norm_eps = norm_eps689 690 # models691 self.visual = VisionTransformer(692 image_size=image_size,693 patch_size=patch_size,694 dim=vision_dim,695 mlp_ratio=vision_mlp_ratio,696 out_dim=embed_dim,697 num_heads=vision_heads,698 num_layers=vision_layers,699 pool_type=vision_pool,700 pre_norm=vision_pre_norm,701 post_norm=vision_post_norm,702 activation=activation,703 attn_dropout=attn_dropout,704 proj_dropout=proj_dropout,705 embedding_dropout=embedding_dropout,706 norm_eps=norm_eps)707 self.textual = None708 self.log_scale = nn.Parameter(math.log(1 / 0.07) * torch.ones([]))709 710 def forward(self, imgs, txt_ids):711 """712 imgs: [B, 3, H, W] of torch.float32.713 - mean: [0.48145466, 0.4578275, 0.40821073]714 - std: [0.26862954, 0.26130258, 0.27577711]715 txt_ids: [B, L] of torch.long.716 Encoded by data.CLIPTokenizer.717 """718 xi = self.visual(imgs)719 xt = self.textual(txt_ids)720 return xi, xt721 722 def param_groups(self):723 groups = [{724 'params': [725 p for n, p in self.named_parameters()726 if 'norm' in n or n.endswith('bias')727 ],728 'weight_decay': 0.0729 }, {730 'params': [731 p for n, p in self.named_parameters()732 if not ('norm' in n or n.endswith('bias'))733 ]734 }]735 return groups736 737 738def _clip(pretrained=False,739 pretrained_name=None,740 model_cls=CLIP,741 return_transforms=False,742 return_tokenizer=False,743 tokenizer_padding='eos',744 dtype=torch.float32,745 device='cpu',746 **kwargs):747 # init model748 if pretrained and pretrained_name:749 from sora import BUCKET, DOWNLOAD_TO_CACHE750 751 # init a meta model752 with torch.device('meta'):753 model = model_cls(**kwargs)754 755 # checkpoint path756 checkpoint = f'models/clip/{pretrained_name}'757 if dtype in (torch.float16, torch.bfloat16):758 suffix = '-' + {759 torch.float16: 'fp16',760 torch.bfloat16: 'bf16'761 }[dtype]762 if object_exists(BUCKET, f'{checkpoint}{suffix}.pth'):763 checkpoint = f'{checkpoint}{suffix}'764 checkpoint += '.pth'765 766 # load767 model.load_state_dict(768 torch.load(DOWNLOAD_TO_CACHE(checkpoint), map_location=device),769 assign=True,770 strict=False)771 else:772 # init a model on device773 with torch.device(device):774 model = model_cls(**kwargs)775 776 # set device777 output = (model,)778 779 # init transforms780 if return_transforms:781 # mean and std782 if 'siglip' in pretrained_name.lower():783 mean, std = [0.5, 0.5, 0.5], [0.5, 0.5, 0.5]784 else:785 mean = [0.48145466, 0.4578275, 0.40821073]786 std = [0.26862954, 0.26130258, 0.27577711]787 788 # transforms789 transforms = T.Compose([790 T.Resize((model.image_size, model.image_size),791 interpolation=T.InterpolationMode.BICUBIC),792 T.ToTensor(),793 T.Normalize(mean=mean, std=std)794 ])795 output += (transforms,)796 797 # init tokenizer798 if return_tokenizer:799 from sora import data800 if 'siglip' in pretrained_name.lower():801 tokenizer = data.HuggingfaceTokenizer(802 name=f'timm/{pretrained_name}',803 seq_len=model.text_len,804 clean='canonicalize')805 elif 'xlm' in pretrained_name.lower():806 tokenizer = data.HuggingfaceTokenizer(807 name='xlm-roberta-large',808 seq_len=model.max_text_len - 2,809 clean='whitespace')810 elif 'mba' in pretrained_name.lower():811 tokenizer = data.HuggingfaceTokenizer(812 name='facebook/xlm-roberta-xl',813 seq_len=model.max_text_len - 2,814 clean='whitespace')815 else:816 tokenizer = data.CLIPTokenizer(817 seq_len=model.text_len, padding=tokenizer_padding)818 output += (tokenizer,)819 return output[0] if len(output) == 1 else output820 821 822def clip_xlm_roberta_vit_h_14(823 pretrained=False,824 pretrained_name='open-clip-xlm-roberta-large-vit-huge-14',825 **kwargs):826 cfg = dict(827 embed_dim=1024,828 image_size=224,829 patch_size=14,830 vision_dim=1280,831 vision_mlp_ratio=4,832 vision_heads=16,833 vision_layers=32,834 vision_pool='token',835 activation='gelu',836 vocab_size=250002,837 max_text_len=514,838 type_size=1,839 pad_id=1,840 text_dim=1024,841 text_heads=16,842 text_layers=24,843 text_post_norm=True,844 text_dropout=0.1,845 attn_dropout=0.0,846 proj_dropout=0.0,847 embedding_dropout=0.0)848 cfg.update(**kwargs)849 return _clip(pretrained, pretrained_name, XLMRobertaCLIP, **cfg)850 851 852class WanImageEncoder(torch.nn.Module):853 854 def __init__(self):855 super().__init__()856 # init model857 self.model, self.transforms = clip_xlm_roberta_vit_h_14(858 pretrained=False,859 return_transforms=True,860 return_tokenizer=False,861 dtype=torch.float32,862 device="cpu")863 864 def encode_image(self, videos):865 # preprocess866 size = (self.model.image_size,) * 2867 videos = torch.cat([868 F.interpolate(869 u,870 size=size,871 mode='bicubic',872 align_corners=False) for u in videos873 ])874 videos = self.transforms.transforms[-1](videos.mul_(0.5).add_(0.5))875 876 # forward877 dtype = next(iter(self.model.visual.parameters())).dtype878 videos = videos.to(dtype)879 out = self.model.visual(videos, use_31_block=True)880 return out881 882 @staticmethod883 def state_dict_converter():884 return WanImageEncoderStateDictConverter()885 886 887class WanImageEncoderStateDictConverter:888 def __init__(self):889 pass890 891 def from_diffusers(self, state_dict):892 return state_dict893 894 def from_civitai(self, state_dict):895 state_dict_ = {}896 for name, param in state_dict.items():897 if name.startswith("textual."):898 continue899 name = "model." + name900 state_dict_[name] = param901 return state_dict_902 903 