radames/Text2Human-API
1
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 