FoundationVision/LlamaGen
64
1# Modified from:2# VQGAN: https://github.com/CompVis/taming-transformers/blob/master/taming/modules/transformer/mingpt.py3# DiT: https://github.com/facebookresearch/DiT/blob/main/models.py 4# nanoGPT: https://github.com/karpathy/nanoGPT/blob/master/model.py5# llama: https://github.com/facebookresearch/llama/blob/main/llama/model.py6# gpt-fast: https://github.com/pytorch-labs/gpt-fast/blob/main/model.py7# PixArt: https://github.com/PixArt-alpha/PixArt-alpha/blob/master/diffusion/model/nets/PixArt_blocks.py8from dataclasses import dataclass9from typing import Optional, List10 11 12import torch13import torch.nn as nn14from torch.nn import functional as F15 16 17def find_multiple(n: int, k: int):18 if n % k == 0:19 return n20 return n + k - (n % k)21 22@dataclass23class ModelArgs:24 dim: int = 409625 n_layer: int = 3226 n_head: int = 3227 n_kv_head: Optional[int] = None28 multiple_of: int = 256 # make SwiGLU hidden layer size multiple of large power of 229 ffn_dim_multiplier: Optional[float] = None30 rope_base: float = 1000031 norm_eps: float = 1e-532 initializer_range: float = 0.0233 34 token_dropout_p: float = 0.135 attn_dropout_p: float = 0.036 resid_dropout_p: float = 0.137 ffn_dropout_p: float = 0.138 drop_path_rate: float = 0.039 40 num_classes: int = 100041 caption_dim: int = 204842 class_dropout_prob: float = 0.143 model_type: str = 'c2i'44 45 vocab_size: int = 1638446 cls_token_num: int = 147 block_size: int = 25648 max_batch_size: int = 3249 max_seq_len: int = 204850 51 52#################################################################################53# Embedding Layers for Class Labels #54#################################################################################55class LabelEmbedder(nn.Module):56 """57 Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.58 """59 def __init__(self, num_classes, hidden_size, dropout_prob):60 super().__init__()61 use_cfg_embedding = dropout_prob > 062 self.embedding_table = nn.Embedding(num_classes + use_cfg_embedding, hidden_size)63 self.num_classes = num_classes64 self.dropout_prob = dropout_prob65 66 def token_drop(self, labels, force_drop_ids=None):67 """68 Drops labels to enable classifier-free guidance.69 """70 if force_drop_ids is None:71 drop_ids = torch.rand(labels.shape[0], device=labels.device) < self.dropout_prob72 else:73 drop_ids = force_drop_ids == 174 labels = torch.where(drop_ids, self.num_classes, labels)75 return labels76 77 def forward(self, labels, train, force_drop_ids=None):78 use_dropout = self.dropout_prob > 079 if (train and use_dropout) or (force_drop_ids is not None):80 labels = self.token_drop(labels, force_drop_ids)81 embeddings = self.embedding_table(labels).unsqueeze(1)82 return embeddings83 84 85#################################################################################86# Embedding Layers for Text Feature #87#################################################################################88class CaptionEmbedder(nn.Module):89 """90 Embeds text caption into vector representations. Also handles label dropout for classifier-free guidance.91 """92 def __init__(self, in_channels, hidden_size, uncond_prob, token_num=120):93 super().__init__()94 self.cap_proj = MLP(in_features=in_channels, hidden_features=hidden_size, out_features=hidden_size)95 self.register_buffer("uncond_embedding", nn.Parameter(torch.randn(token_num, in_channels) / in_channels ** 0.5))96 self.uncond_prob = uncond_prob97 98 def token_drop(self, caption, force_drop_ids=None):99 """100 Drops labels to enable classifier-free guidance.101 """102 if force_drop_ids is None:103 drop_ids = torch.rand(caption.shape[0], device=caption.device) < self.uncond_prob104 else:105 drop_ids = force_drop_ids == 1106 caption = torch.where(drop_ids[:, None, None], self.uncond_embedding, caption)107 return caption108 109 def forward(self, caption, train, force_drop_ids=None):110 use_dropout = self.uncond_prob > 0111 if (train and use_dropout) or (force_drop_ids is not None):112 caption = self.token_drop(caption, force_drop_ids)113 embeddings = self.cap_proj(caption)114 return embeddings115 116 117class MLP(nn.Module):118 def __init__(self, in_features, hidden_features, out_features):119 super().__init__()120 out_features = out_features or in_features121 hidden_features = hidden_features or in_features122 self.fc1 = nn.Linear(in_features, hidden_features, bias=False)123 self.act = nn.GELU(approximate='tanh')124 self.fc2 = nn.Linear(hidden_features, out_features, bias=False)125 126 def forward(self, x):127 x = self.fc1(x)128 x = self.act(x)129 x = self.fc2(x)130 return x131 132 133#################################################################################134# GPT Model #135#################################################################################136class RMSNorm(torch.nn.Module):137 def __init__(self, dim: int, eps: float = 1e-5):138 super().__init__()139 self.eps = eps140 self.weight = nn.Parameter(torch.ones(dim))141 142 def _norm(self, x):143 return x * torch.rsqrt(torch.mean(x * x, dim=-1, keepdim=True) + self.eps)144 145 def forward(self, x):146 output = self._norm(x.float()).type_as(x)147 return output * self.weight148 149 150class FeedForward(nn.Module):151 def __init__(self, config: ModelArgs):152 super().__init__()153 hidden_dim = 4 * config.dim154 hidden_dim = int(2 * hidden_dim / 3)155 # custom dim factor multiplier156 if config.ffn_dim_multiplier is not None:157 hidden_dim = int(config.ffn_dim_multiplier * hidden_dim)158 hidden_dim = find_multiple(hidden_dim, config.multiple_of)159 160 self.w1 = nn.Linear(config.dim, hidden_dim, bias=False)161 self.w3 = nn.Linear(config.dim, hidden_dim, bias=False)162 self.w2 = nn.Linear(hidden_dim, config.dim, bias=False)163 self.ffn_dropout = nn.Dropout(config.ffn_dropout_p)164 165 def forward(self, x):166 return self.ffn_dropout(self.w2(F.silu(self.w1(x)) * self.w3(x)))167 168 169class KVCache(nn.Module):170 def __init__(self, max_batch_size, max_seq_length, n_head, head_dim, dtype):171 super().__init__()172 cache_shape = (max_batch_size, n_head, max_seq_length, head_dim)173 self.register_buffer('k_cache', torch.zeros(cache_shape, dtype=dtype))174 self.register_buffer('v_cache', torch.zeros(cache_shape, dtype=dtype))175 176 def update(self, input_pos, k_val, v_val):177 # input_pos: [S], k_val: [B, H, S, D]178 assert input_pos.shape[0] == k_val.shape[2]179 k_out = self.k_cache180 v_out = self.v_cache181 k_out[:, :, input_pos] = k_val182 v_out[:, :, input_pos] = v_val183 184 return k_out, v_out185 186 187class Attention(nn.Module):188 def __init__(self, config: ModelArgs):189 super().__init__()190 assert config.dim % config.n_head == 0191 self.dim = config.dim192 self.head_dim = config.dim // config.n_head193 self.n_head = config.n_head194 self.n_kv_head = config.n_kv_head if config.n_kv_head is not None else config.n_head195 total_kv_dim = (self.n_head + 2 * self.n_kv_head) * self.head_dim196 197 # key, query, value projections for all heads, but in a batch198 self.wqkv = nn.Linear(config.dim, total_kv_dim, bias=False)199 self.wo = nn.Linear(config.dim, config.dim, bias=False)200 self.kv_cache = None201 202 # regularization203 self.attn_dropout_p = config.attn_dropout_p204 self.resid_dropout = nn.Dropout(config.resid_dropout_p)205 206 def forward(207 self, x: torch.Tensor, freqs_cis: torch.Tensor = None, 208 input_pos: Optional[torch.Tensor] = None, 209 mask: Optional[torch.Tensor] = None210 ):211 bsz, seqlen, _ = x.shape212 kv_size = self.n_kv_head * self.head_dim213 xq, xk, xv = self.wqkv(x).split([self.dim, kv_size, kv_size], dim=-1)214 215 xq = xq.view(bsz, seqlen, self.n_head, self.head_dim)216 xk = xk.view(bsz, seqlen, self.n_kv_head, self.head_dim)217 xv = xv.view(bsz, seqlen, self.n_kv_head, self.head_dim)218 219 xq = apply_rotary_emb(xq, freqs_cis)220 xk = apply_rotary_emb(xk, freqs_cis)221 222 xq, xk, xv = map(lambda x: x.transpose(1, 2), (xq, xk, xv))223 224 if self.kv_cache is not None:225 keys, values = self.kv_cache.update(input_pos, xk, xv)226 else:227 keys, values = xk, xv228 keys = keys.repeat_interleave(self.n_head // self.n_kv_head, dim=1)229 values = values.repeat_interleave(self.n_head // self.n_kv_head, dim=1)230 231 output = F.scaled_dot_product_attention(232 xq, keys, values, 233 attn_mask=mask, 234 is_causal=True if mask is None else False, # is_causal=False is for KV cache235 dropout_p=self.attn_dropout_p if self.training else 0) 236 237 output = output.transpose(1, 2).contiguous().view(bsz, seqlen, self.dim)238 239 output = self.resid_dropout(self.wo(output))240 return output241 242 243class TransformerBlock(nn.Module):244 def __init__(self, config: ModelArgs, drop_path: float):245 super().__init__()246 self.attention = Attention(config)247 self.feed_forward = FeedForward(config)248 self.attention_norm = RMSNorm(config.dim, eps=config.norm_eps)249 self.ffn_norm = RMSNorm(config.dim, eps=config.norm_eps)250 251 def forward(252 self, x: torch.Tensor, freqs_cis: torch.Tensor, start_pos: int, mask: Optional[torch.Tensor] = None):253 h = x + self.attention(self.attention_norm(x), freqs_cis, start_pos, mask)254 out = h + self.feed_forward(self.ffn_norm(h))255 return out256 257 258class Transformer(nn.Module):259 def __init__(self, config: ModelArgs):260 super().__init__()261 self.config = config262 self.vocab_size = config.vocab_size263 self.n_layer = config.n_layer264 self.block_size = config.block_size265 self.num_classes = config.num_classes266 self.model_type = config.model_type267 self.cls_token_num = config.cls_token_num268 if self.model_type == 'c2i':269 self.cls_embedding = LabelEmbedder(config.num_classes, config.dim, config.class_dropout_prob)270 elif self.model_type == 't2i':271 self.cls_embedding = CaptionEmbedder(config.caption_dim, config.dim, config.class_dropout_prob)272 else:273 raise Exception("please check model type")274 self.tok_embeddings = nn.Embedding(config.vocab_size, config.dim)275 self.tok_dropout = nn.Dropout(config.token_dropout_p)276 277 # transformer blocks278 dpr = [x.item() for x in torch.linspace(0, config.drop_path_rate, config.n_layer)]279 self.layers = torch.nn.ModuleList()280 for layer_id in range(config.n_layer):281 self.layers.append(TransformerBlock(config, dpr[layer_id]))282 283 # output layer284 self.norm = RMSNorm(config.dim, eps=config.norm_eps)285 self.output = nn.Linear(config.dim, config.vocab_size, bias=False)286 287 # 2d rotary pos embedding288 grid_size = int(self.block_size ** 0.5)289 assert grid_size * grid_size == self.block_size290 self.freqs_cis = precompute_freqs_cis_2d(grid_size, self.config.dim // self.config.n_head, self.config.rope_base, self.cls_token_num)291 292 # KVCache293 self.max_batch_size = -1294 self.max_seq_length = -1295 296 self.initialize_weights()297 298 def initialize_weights(self): 299 # Initialize nn.Linear and nn.Embedding300 self.apply(self._init_weights)301 302 # Zero-out output layers:303 nn.init.constant_(self.output.weight, 0)304 305 def _init_weights(self, module):306 std = self.config.initializer_range307 if isinstance(module, nn.Linear):308 module.weight.data.normal_(mean=0.0, std=std)309 if module.bias is not None:310 module.bias.data.zero_()311 elif isinstance(module, nn.Embedding):312 module.weight.data.normal_(mean=0.0, std=std)313 314 def setup_caches(self, max_batch_size, max_seq_length, dtype):315 # if self.max_seq_length >= max_seq_length and self.max_batch_size >= max_batch_size:316 # return317 head_dim = self.config.dim // self.config.n_head318 max_seq_length = find_multiple(max_seq_length, 8)319 self.max_seq_length = max_seq_length320 self.max_batch_size = max_batch_size321 for b in self.layers:322 b.attention.kv_cache = KVCache(max_batch_size, max_seq_length, self.config.n_head, head_dim, dtype)323 324 causal_mask = torch.tril(torch.ones(self.max_seq_length, self.max_seq_length, dtype=torch.bool))325 self.causal_mask = causal_mask.unsqueeze(0).repeat(self.max_batch_size, 1, 1)326 grid_size = int(self.config.block_size ** 0.5)327 assert grid_size * grid_size == self.block_size328 self.freqs_cis = precompute_freqs_cis_2d(grid_size, self.config.dim // self.config.n_head, self.config.rope_base, self.cls_token_num)329 330 def forward(331 self, 332 idx: torch.Tensor, 333 cond_idx: torch.Tensor, # cond_idx_or_embed334 input_pos: Optional[torch.Tensor] = None, 335 targets: Optional[torch.Tensor] = None,336 mask: Optional[torch.Tensor] = None,337 valid: Optional[torch.Tensor] = None,338 ):339 if idx is not None and cond_idx is not None: # training or naive inference340 cond_embeddings = self.cls_embedding(cond_idx, train=self.training)[:,:self.cls_token_num]341 token_embeddings = self.tok_embeddings(idx)342 token_embeddings = torch.cat((cond_embeddings, token_embeddings), dim=1)343 h = self.tok_dropout(token_embeddings)344 self.freqs_cis = self.freqs_cis.to(h.device)345 else:346 if cond_idx is not None: # prefill in inference347 token_embeddings = self.cls_embedding(cond_idx, train=self.training)[:,:self.cls_token_num]348 else: # decode_n_tokens(kv cache) in inference349 token_embeddings = self.tok_embeddings(idx)350 351 bs = token_embeddings.shape[0]352 mask = self.causal_mask[:bs, None, input_pos]353 h = self.tok_dropout(token_embeddings)354 self.freqs_cis = self.freqs_cis355 356 if self.training:357 freqs_cis = self.freqs_cis[:token_embeddings.shape[1]]358 else:359 freqs_cis = self.freqs_cis[input_pos]360 # transformer blocks361 for layer in self.layers:362 h = layer(h, freqs_cis, input_pos, mask)363 364 # output layers365 h = self.norm(h)366 logits = self.output(h).float()367 368 if self.training:369 logits = logits[:, self.cls_token_num - 1:].contiguous()370 371 # if we are given some desired targets also calculate the loss372 loss = None373 if valid is not None:374 loss_all = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1), reduction='none')375 valid_all = valid[:,None].repeat(1, targets.shape[1]).view(-1)376 loss = (loss_all * valid_all).sum() / max(valid_all.sum(), 1)377 elif targets is not None:378 loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1))379 380 return logits, loss381 382 383 def get_fsdp_wrap_module_list(self) -> List[nn.Module]:384 return list(self.layers)385 386 387 388#################################################################################389# Rotary Positional Embedding Functions #390#################################################################################391# https://github.com/pytorch-labs/gpt-fast/blob/main/model.py 392def precompute_freqs_cis(seq_len: int, n_elem: int, base: int = 10000, cls_token_num=120):393 freqs = 1.0 / (base ** (torch.arange(0, n_elem, 2)[: (n_elem // 2)].float() / n_elem))394 t = torch.arange(seq_len, device=freqs.device)395 freqs = torch.outer(t, freqs) # (seq_len, head_dim // 2)396 freqs_cis = torch.polar(torch.ones_like(freqs), freqs)397 cache = torch.stack([freqs_cis.real, freqs_cis.imag], dim=-1) # (cls_token_num+seq_len, head_dim // 2, 2)398 cond_cache = torch.cat([torch.zeros(cls_token_num, n_elem // 2, 2), cache]) # (cls_token_num+seq_len, head_dim // 2, 2)399 return cond_cache 400 401 402def precompute_freqs_cis_2d(grid_size: int, n_elem: int, base: int = 10000, cls_token_num=120):403 # split the dimension into half, one for x and one for y404 half_dim = n_elem // 2405 freqs = 1.0 / (base ** (torch.arange(0, half_dim, 2)[: (half_dim // 2)].float() / half_dim))406 t = torch.arange(grid_size, device=freqs.device)407 freqs = torch.outer(t, freqs) # (grid_size, head_dim // 2)408 freqs_grid = torch.concat([409 freqs[:, None, :].expand(-1, grid_size, -1),410 freqs[None, :, :].expand(grid_size, -1, -1),411 ], dim=-1) # (grid_size, grid_size, head_dim // 2)412 cache_grid = torch.stack([torch.cos(freqs_grid), torch.sin(freqs_grid)], dim=-1) # (grid_size, grid_size, head_dim // 2, 2)413 cache = cache_grid.flatten(0, 1)414 cond_cache = torch.cat([torch.zeros(cls_token_num, n_elem // 2, 2), cache]) # (cls_token_num+grid_size**2, head_dim // 2, 2)415 return cond_cache 416 417 418def apply_rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor):419 # x: (bs, seq_len, n_head, head_dim)420 # freqs_cis (seq_len, head_dim // 2, 2)421 xshaped = x.float().reshape(*x.shape[:-1], -1, 2) # (bs, seq_len, n_head, head_dim//2, 2)422 freqs_cis = freqs_cis.view(1, xshaped.size(1), 1, xshaped.size(3), 2) # (1, seq_len, 1, head_dim//2, 2)423 x_out2 = torch.stack([424 xshaped[..., 0] * freqs_cis[..., 0] - xshaped[..., 1] * freqs_cis[..., 1],425 xshaped[..., 1] * freqs_cis[..., 0] + xshaped[..., 0] * freqs_cis[..., 1],426 ], dim=-1)427 x_out2 = x_out2.flatten(3)428 return x_out2.type_as(x)429 430 431 432#################################################################################433# GPT Configs #434#################################################################################435### text-conditional436def GPT_7B(**kwargs):437 return Transformer(ModelArgs(n_layer=32, n_head=32, dim=4096, **kwargs)) # 6.6B438 439def GPT_3B(**kwargs):440 return Transformer(ModelArgs(n_layer=24, n_head=32, dim=3200, **kwargs)) # 3.1B441 442def GPT_1B(**kwargs):443 return Transformer(ModelArgs(n_layer=22, n_head=32, dim=2048, **kwargs)) # 1.2B444 445### class-conditional446def GPT_XXXL(**kwargs):447 return Transformer(ModelArgs(n_layer=48, n_head=40, dim=2560, **kwargs)) # 3.9B448 449def GPT_XXL(**kwargs):450 return Transformer(ModelArgs(n_layer=48, n_head=24, dim=1536, **kwargs)) # 1.4B451 452def GPT_XL(**kwargs):453 return Transformer(ModelArgs(n_layer=36, n_head=20, dim=1280, **kwargs)) # 775M454 455def GPT_L(**kwargs):456 return Transformer(ModelArgs(n_layer=24, n_head=16, dim=1024, **kwargs)) # 343M457 458def GPT_B(**kwargs):459 return Transformer(ModelArgs(n_layer=12, n_head=12, dim=768, **kwargs)) # 111M460 461 462GPT_models = {463 'GPT-B': GPT_B, 'GPT-L': GPT_L, 'GPT-XL': GPT_XL, 'GPT-XXL': GPT_XXL, 'GPT-XXXL': GPT_XXXL,464 'GPT-1B': GPT_1B, 'GPT-3B': GPT_3B, 'GPT-7B': GPT_7B, 465}