Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
transformer_arch.py274 linesDownload Raw Back to archs
1import math2 3import numpy as np4import torch5import torch.nn as nn6import torch.nn.functional as F7 8 9class CausalSelfAttention(nn.Module):10    """11    A vanilla multi-head masked self-attention layer with a projection at the end.12    It is possible to use torch.nn.MultiheadAttention here but I am including an13    explicit implementation here to show that there is nothing too scary here.14    """15 16    def __init__(self, bert_n_emb, bert_n_head, attn_pdrop, resid_pdrop,17                 latent_shape, sampler):18        super().__init__()19        assert bert_n_emb % bert_n_head == 020        # key, query, value projections for all heads21        self.key = nn.Linear(bert_n_emb, bert_n_emb)22        self.query = nn.Linear(bert_n_emb, bert_n_emb)23        self.value = nn.Linear(bert_n_emb, bert_n_emb)24        # regularization25        self.attn_drop = nn.Dropout(attn_pdrop)26        self.resid_drop = nn.Dropout(resid_pdrop)27        # output projection28        self.proj = nn.Linear(bert_n_emb, bert_n_emb)29        self.n_head = bert_n_head30        self.causal = True if sampler == 'autoregressive' else False31        if self.causal:32            block_size = np.prod(latent_shape)33            mask = torch.tril(torch.ones(block_size, block_size))34            self.register_buffer("mask", mask.view(1, 1, block_size,35                                                   block_size))36 37    def forward(self, x, layer_past=None):38        B, T, C = x.size()39 40        # calculate query, key, values for all heads in batch and move head forward to be the batch dim41        k = self.key(x).view(B, T, self.n_head,42                             C // self.n_head).transpose(1,43                                                         2)  # (B, nh, T, hs)44        q = self.query(x).view(B, T, self.n_head,45                               C // self.n_head).transpose(1,46                                                           2)  # (B, nh, T, hs)47        v = self.value(x).view(B, T, self.n_head,48                               C // self.n_head).transpose(1,49                                                           2)  # (B, nh, T, hs)50 51        present = torch.stack((k, v))52        if self.causal and layer_past is not None:53            past_key, past_value = layer_past54            k = torch.cat((past_key, k), dim=-2)55            v = torch.cat((past_value, v), dim=-2)56 57        # causal self-attention; Self-attend: (B, nh, T, hs) x (B, nh, hs, T) -> (B, nh, T, T)58        att = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(k.size(-1)))59 60        if self.causal and layer_past is None:61            att = att.masked_fill(self.mask[:, :, :T, :T] == 0, float('-inf'))62 63        att = F.softmax(att, dim=-1)64        att = self.attn_drop(att)65        y = att @ v  # (B, nh, T, T) x (B, nh, T, hs) -> (B, nh, T, hs)66        # re-assemble all head outputs side by side67        y = y.transpose(1, 2).contiguous().view(B, T, C)68 69        # output projection70        y = self.resid_drop(self.proj(y))71        return y, present72 73 74class Block(nn.Module):75    """ an unassuming Transformer block """76 77    def __init__(self, bert_n_emb, resid_pdrop, bert_n_head, attn_pdrop,78                 latent_shape, sampler):79        super().__init__()80        self.ln1 = nn.LayerNorm(bert_n_emb)81        self.ln2 = nn.LayerNorm(bert_n_emb)82        self.attn = CausalSelfAttention(bert_n_emb, bert_n_head, attn_pdrop,83                                        resid_pdrop, latent_shape, sampler)84        self.mlp = nn.Sequential(85            nn.Linear(bert_n_emb, 4 * bert_n_emb),86            nn.GELU(),  # nice87            nn.Linear(4 * bert_n_emb, bert_n_emb),88            nn.Dropout(resid_pdrop),89        )90 91    def forward(self, x, layer_past=None, return_present=False):92 93        attn, present = self.attn(self.ln1(x), layer_past)94        x = x + attn95        x = x + self.mlp(self.ln2(x))96 97        if layer_past is not None or return_present:98            return x, present99        return x100 101 102class Transformer(nn.Module):103    """  the full GPT language model, with a context size of block_size """104 105    def __init__(self,106                 codebook_size,107                 segm_codebook_size,108                 bert_n_emb,109                 bert_n_layers,110                 bert_n_head,111                 block_size,112                 latent_shape,113                 embd_pdrop,114                 resid_pdrop,115                 attn_pdrop,116                 sampler='absorbing'):117        super().__init__()118 119        self.vocab_size = codebook_size + 1120        self.n_embd = bert_n_emb121        self.block_size = block_size122        self.n_layers = bert_n_layers123        self.codebook_size = codebook_size124        self.segm_codebook_size = segm_codebook_size125        self.causal = sampler == 'autoregressive'126        if self.causal:127            self.vocab_size = codebook_size128 129        self.tok_emb = nn.Embedding(self.vocab_size, self.n_embd)130        self.pos_emb = nn.Parameter(131            torch.zeros(1, self.block_size, self.n_embd))132        self.segm_emb = nn.Embedding(self.segm_codebook_size, self.n_embd)133        self.start_tok = nn.Parameter(torch.zeros(1, 1, self.n_embd))134        self.drop = nn.Dropout(embd_pdrop)135 136        # transformer137        self.blocks = nn.Sequential(*[138            Block(bert_n_emb, resid_pdrop, bert_n_head, attn_pdrop,139                  latent_shape, sampler) for _ in range(self.n_layers)140        ])141        # decoder head142        self.ln_f = nn.LayerNorm(self.n_embd)143        self.head = nn.Linear(self.n_embd, self.codebook_size, bias=False)144 145    def get_block_size(self):146        return self.block_size147 148    def _init_weights(self, module):149        if isinstance(module, (nn.Linear, nn.Embedding)):150            module.weight.data.normal_(mean=0.0, std=0.02)151            if isinstance(module, nn.Linear) and module.bias is not None:152                module.bias.data.zero_()153        elif isinstance(module, nn.LayerNorm):154            module.bias.data.zero_()155            module.weight.data.fill_(1.0)156 157    def forward(self, idx, segm_tokens, t=None):158        # each index maps to a (learnable) vector159        token_embeddings = self.tok_emb(idx)160 161        segm_embeddings = self.segm_emb(segm_tokens)162 163        if self.causal:164            token_embeddings = torch.cat((self.start_tok.repeat(165                token_embeddings.size(0), 1, 1), token_embeddings),166                                         dim=1)167 168        t = token_embeddings.shape[1]169        assert t <= self.block_size, "Cannot forward, model block size is exhausted."170        # each position maps to a (learnable) vector171 172        position_embeddings = self.pos_emb[:, :t, :]173 174        x = token_embeddings + position_embeddings + segm_embeddings175        x = self.drop(x)176        for block in self.blocks:177            x = block(x)178        x = self.ln_f(x)179        logits = self.head(x)180 181        return logits182 183 184class TransformerMultiHead(nn.Module):185    """  the full GPT language model, with a context size of block_size """186 187    def __init__(self,188                 codebook_size,189                 segm_codebook_size,190                 texture_codebook_size,191                 bert_n_emb,192                 bert_n_layers,193                 bert_n_head,194                 block_size,195                 latent_shape,196                 embd_pdrop,197                 resid_pdrop,198                 attn_pdrop,199                 num_head,200                 sampler='absorbing'):201        super().__init__()202 203        self.vocab_size = codebook_size + 1204        self.n_embd = bert_n_emb205        self.block_size = block_size206        self.n_layers = bert_n_layers207        self.codebook_size = codebook_size208        self.segm_codebook_size = segm_codebook_size209        self.texture_codebook_size = texture_codebook_size210        self.causal = sampler == 'autoregressive'211        if self.causal:212            self.vocab_size = codebook_size213 214        self.tok_emb = nn.Embedding(self.vocab_size, self.n_embd)215        self.pos_emb = nn.Parameter(216            torch.zeros(1, self.block_size, self.n_embd))217        self.segm_emb = nn.Embedding(self.segm_codebook_size, self.n_embd)218        self.texture_emb = nn.Embedding(self.texture_codebook_size,219                                        self.n_embd)220        self.start_tok = nn.Parameter(torch.zeros(1, 1, self.n_embd))221        self.drop = nn.Dropout(embd_pdrop)222 223        # transformer224        self.blocks = nn.Sequential(*[225            Block(bert_n_emb, resid_pdrop, bert_n_head, attn_pdrop,226                  latent_shape, sampler) for _ in range(self.n_layers)227        ])228        # decoder head229        self.num_head = num_head230        self.head_class_num = codebook_size // self.num_head231        self.ln_f = nn.LayerNorm(self.n_embd)232        self.head_list = nn.ModuleList([233            nn.Linear(self.n_embd, self.head_class_num, bias=False)234            for _ in range(self.num_head)235        ])236 237    def get_block_size(self):238        return self.block_size239 240    def _init_weights(self, module):241        if isinstance(module, (nn.Linear, nn.Embedding)):242            module.weight.data.normal_(mean=0.0, std=0.02)243            if isinstance(module, nn.Linear) and module.bias is not None:244                module.bias.data.zero_()245        elif isinstance(module, nn.LayerNorm):246            module.bias.data.zero_()247            module.weight.data.fill_(1.0)248 249    def forward(self, idx, segm_tokens, texture_tokens, t=None):250        # each index maps to a (learnable) vector251        token_embeddings = self.tok_emb(idx)252        segm_embeddings = self.segm_emb(segm_tokens)253        texture_embeddings = self.texture_emb(texture_tokens)254 255        if self.causal:256            token_embeddings = torch.cat((self.start_tok.repeat(257                token_embeddings.size(0), 1, 1), token_embeddings),258                                         dim=1)259 260        t = token_embeddings.shape[1]261        assert t <= self.block_size, "Cannot forward, model block size is exhausted."262        # each position maps to a (learnable) vector263 264        position_embeddings = self.pos_emb[:, :t, :]265 266        x = token_embeddings + position_embeddings + segm_embeddings + texture_embeddings267        x = self.drop(x)268        for block in self.blocks:269            x = block(x)270        x = self.ln_f(x)271        logits_list = [self.head_list[i](x) for i in range(self.num_head)]272 273        return logits_list274