CoolFace
Apppublic

FoundationVision/LlamaGen

sourceHugging Facemitupdated 2y agoView on Hugging Face
64likes
gpt.py465 linesDownload Raw Back to models
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}