Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
wan_video_text_encoder.py270 linesDownload Raw Back to models
1import math2 3import torch4import torch.nn as nn5import torch.nn.functional as F6 7 8def fp16_clamp(x):9    if x.dtype == torch.float16 and torch.isinf(x).any():10        clamp = torch.finfo(x.dtype).max - 100011        x = torch.clamp(x, min=-clamp, max=clamp)12    return x13 14 15class GELU(nn.Module):16 17    def forward(self, x):18        return 0.5 * x * (1.0 + torch.tanh(19            math.sqrt(2.0 / math.pi) * (x + 0.044715 * torch.pow(x, 3.0))))20 21 22class T5LayerNorm(nn.Module):23 24    def __init__(self, dim, eps=1e-6):25        super(T5LayerNorm, self).__init__()26        self.dim = dim27        self.eps = eps28        self.weight = nn.Parameter(torch.ones(dim))29 30    def forward(self, x):31        x = x * torch.rsqrt(x.float().pow(2).mean(dim=-1, keepdim=True) +32                            self.eps)33        if self.weight.dtype in [torch.float16, torch.bfloat16]:34            x = x.type_as(self.weight)35        return self.weight * x36 37 38class T5Attention(nn.Module):39 40    def __init__(self, dim, dim_attn, num_heads, dropout=0.1):41        assert dim_attn % num_heads == 042        super(T5Attention, self).__init__()43        self.dim = dim44        self.dim_attn = dim_attn45        self.num_heads = num_heads46        self.head_dim = dim_attn // num_heads47 48        # layers49        self.q = nn.Linear(dim, dim_attn, bias=False)50        self.k = nn.Linear(dim, dim_attn, bias=False)51        self.v = nn.Linear(dim, dim_attn, bias=False)52        self.o = nn.Linear(dim_attn, dim, bias=False)53        self.dropout = nn.Dropout(dropout)54 55    def forward(self, x, context=None, mask=None, pos_bias=None):56        """57        x:          [B, L1, C].58        context:    [B, L2, C] or None.59        mask:       [B, L2] or [B, L1, L2] or None.60        """61        # check inputs62        context = x if context is None else context63        b, n, c = x.size(0), self.num_heads, self.head_dim64 65        # compute query, key, value66        q = self.q(x).view(b, -1, n, c)67        k = self.k(context).view(b, -1, n, c)68        v = self.v(context).view(b, -1, n, c)69 70        # attention bias71        attn_bias = x.new_zeros(b, n, q.size(1), k.size(1))72        if pos_bias is not None:73            attn_bias += pos_bias74        if mask is not None:75            assert mask.ndim in [2, 3]76            mask = mask.view(b, 1, 1,77                             -1) if mask.ndim == 2 else mask.unsqueeze(1)78            attn_bias.masked_fill_(mask == 0, torch.finfo(x.dtype).min)79 80        # compute attention (T5 does not use scaling)81        attn = torch.einsum('binc,bjnc->bnij', q, k) + attn_bias82        attn = F.softmax(attn.float(), dim=-1).type_as(attn)83        x = torch.einsum('bnij,bjnc->binc', attn, v)84 85        # output86        x = x.reshape(b, -1, n * c)87        x = self.o(x)88        x = self.dropout(x)89        return x90 91 92class T5FeedForward(nn.Module):93 94    def __init__(self, dim, dim_ffn, dropout=0.1):95        super(T5FeedForward, self).__init__()96        self.dim = dim97        self.dim_ffn = dim_ffn98 99        # layers100        self.gate = nn.Sequential(nn.Linear(dim, dim_ffn, bias=False), GELU())101        self.fc1 = nn.Linear(dim, dim_ffn, bias=False)102        self.fc2 = nn.Linear(dim_ffn, dim, bias=False)103        self.dropout = nn.Dropout(dropout)104 105    def forward(self, x):106        x = self.fc1(x) * self.gate(x)107        x = self.dropout(x)108        x = self.fc2(x)109        x = self.dropout(x)110        return x111 112 113class T5SelfAttention(nn.Module):114 115    def __init__(self,116                 dim,117                 dim_attn,118                 dim_ffn,119                 num_heads,120                 num_buckets,121                 shared_pos=True,122                 dropout=0.1):123        super(T5SelfAttention, self).__init__()124        self.dim = dim125        self.dim_attn = dim_attn126        self.dim_ffn = dim_ffn127        self.num_heads = num_heads128        self.num_buckets = num_buckets129        self.shared_pos = shared_pos130 131        # layers132        self.norm1 = T5LayerNorm(dim)133        self.attn = T5Attention(dim, dim_attn, num_heads, dropout)134        self.norm2 = T5LayerNorm(dim)135        self.ffn = T5FeedForward(dim, dim_ffn, dropout)136        self.pos_embedding = None if shared_pos else T5RelativeEmbedding(137            num_buckets, num_heads, bidirectional=True)138 139    def forward(self, x, mask=None, pos_bias=None):140        e = pos_bias if self.shared_pos else self.pos_embedding(141            x.size(1), x.size(1))142        x = fp16_clamp(x + self.attn(self.norm1(x), mask=mask, pos_bias=e))143        x = fp16_clamp(x + self.ffn(self.norm2(x)))144        return x145 146 147class T5RelativeEmbedding(nn.Module):148 149    def __init__(self, num_buckets, num_heads, bidirectional, max_dist=128):150        super(T5RelativeEmbedding, self).__init__()151        self.num_buckets = num_buckets152        self.num_heads = num_heads153        self.bidirectional = bidirectional154        self.max_dist = max_dist155 156        # layers157        self.embedding = nn.Embedding(num_buckets, num_heads)158 159    def forward(self, lq, lk):160        device = self.embedding.weight.device161        # rel_pos = torch.arange(lk).unsqueeze(0).to(device) - \162        #     torch.arange(lq).unsqueeze(1).to(device)163        rel_pos = torch.arange(lk, device=device).unsqueeze(0) - \164            torch.arange(lq, device=device).unsqueeze(1)165        rel_pos = self._relative_position_bucket(rel_pos)166        rel_pos_embeds = self.embedding(rel_pos)167        rel_pos_embeds = rel_pos_embeds.permute(2, 0, 1).unsqueeze(168            0)  # [1, N, Lq, Lk]169        return rel_pos_embeds.contiguous()170 171    def _relative_position_bucket(self, rel_pos):172        # preprocess173        if self.bidirectional:174            num_buckets = self.num_buckets // 2175            rel_buckets = (rel_pos > 0).long() * num_buckets176            rel_pos = torch.abs(rel_pos)177        else:178            num_buckets = self.num_buckets179            rel_buckets = 0180            rel_pos = -torch.min(rel_pos, torch.zeros_like(rel_pos))181 182        # embeddings for small and large positions183        max_exact = num_buckets // 2184        rel_pos_large = max_exact + (torch.log(rel_pos.float() / max_exact) /185                                     math.log(self.max_dist / max_exact) *186                                     (num_buckets - max_exact)).long()187        rel_pos_large = torch.min(188            rel_pos_large, torch.full_like(rel_pos_large, num_buckets - 1))189        rel_buckets += torch.where(rel_pos < max_exact, rel_pos, rel_pos_large)190        return rel_buckets191 192def init_weights(m):193    if isinstance(m, T5LayerNorm):194        nn.init.ones_(m.weight)195    elif isinstance(m, T5FeedForward):196        nn.init.normal_(m.gate[0].weight, std=m.dim**-0.5)197        nn.init.normal_(m.fc1.weight, std=m.dim**-0.5)198        nn.init.normal_(m.fc2.weight, std=m.dim_ffn**-0.5)199    elif isinstance(m, T5Attention):200        nn.init.normal_(m.q.weight, std=(m.dim * m.dim_attn)**-0.5)201        nn.init.normal_(m.k.weight, std=m.dim**-0.5)202        nn.init.normal_(m.v.weight, std=m.dim**-0.5)203        nn.init.normal_(m.o.weight, std=(m.num_heads * m.dim_attn)**-0.5)204    elif isinstance(m, T5RelativeEmbedding):205        nn.init.normal_(206            m.embedding.weight, std=(2 * m.num_buckets * m.num_heads)**-0.5)207 208 209class WanTextEncoder(torch.nn.Module):210 211    def __init__(self,212                 vocab=256384,213                 dim=4096,214                 dim_attn=4096,215                 dim_ffn=10240,216                 num_heads=64,217                 num_layers=24,218                 num_buckets=32,219                 shared_pos=False,220                 dropout=0.1):221        super(WanTextEncoder, self).__init__()222        self.dim = dim223        self.dim_attn = dim_attn224        self.dim_ffn = dim_ffn225        self.num_heads = num_heads226        self.num_layers = num_layers227        self.num_buckets = num_buckets228        self.shared_pos = shared_pos229 230        # layers231        self.token_embedding = vocab if isinstance(vocab, nn.Embedding) \232            else nn.Embedding(vocab, dim)233        self.pos_embedding = T5RelativeEmbedding(234            num_buckets, num_heads, bidirectional=True) if shared_pos else None235        self.dropout = nn.Dropout(dropout)236        self.blocks = nn.ModuleList([237            T5SelfAttention(dim, dim_attn, dim_ffn, num_heads, num_buckets,238                            shared_pos, dropout) for _ in range(num_layers)239        ])240        self.norm = T5LayerNorm(dim)241 242        # initialize weights243        self.apply(init_weights)244 245    def forward(self, ids, mask=None):246        x = self.token_embedding(ids)247        x = self.dropout(x)248        e = self.pos_embedding(x.size(1),249                               x.size(1)) if self.shared_pos else None250        for block in self.blocks:251            x = block(x, mask, pos_bias=e)252        x = self.norm(x)253        x = self.dropout(x)254        return x255    256    @staticmethod257    def state_dict_converter():258        return WanTextEncoderStateDictConverter()259    260    261class WanTextEncoderStateDictConverter:262    def __init__(self):263        pass264 265    def from_diffusers(self, state_dict):266        return state_dict267    268    def from_civitai(self, state_dict):269        return state_dict270