Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
wan_video_image_encoder.py903 linesDownload Raw Back to models
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