Team Ai
Modelpublic

diffusion-reasoning/gdsd_code_llada

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes22downloads
modeling_llada.py1515 linesDownload Raw Back to root
1from __future__ import annotations2 3import logging4import math5import sys6from abc import abstractmethod7from collections import defaultdict8from functools import partial9from typing import (10    Callable,11    Dict,12    Iterable,13    List,14    NamedTuple,15    Optional,16    Sequence,17    Set,18    Tuple,19    cast,20)21from dataclasses import fields22from typing import List, Optional, Tuple, Union23 24import torch25import torch.backends.cuda26import torch.nn as nn27import torch.nn.functional as F28from torch import einsum29from transformers import PreTrainedModel30from transformers.modeling_outputs import CausalLMOutputWithPast31from transformers.models.auto import AutoModel32from transformers.cache_utils import Cache33 34from .configuration_llada import (35    LLaDAConfig,36    StrEnum,37    InitFnType,38    ActivationType,39    BlockType,40    LayerNormType,41    ModelConfig,42    ActivationCheckpointingStrategy,43)44 45if sys.version_info.minor > 8:46    from collections.abc import MutableMapping47elif sys.version_info.minor == 8:48    from typing import MutableMapping49else:50    raise SystemExit("This script supports Python 3.8 or higher")51 52__all__ = [53    "LayerNormBase",54    "LayerNorm",55    "RMSLayerNorm",56    "GemmaRMSLayerNorm",57    "RotaryEmbedding",58    "Activation",59    "GELU",60    "ReLU",61    "SwiGLU",62    "LLaDABlock",63    "LLaDASequentialBlock",64    "LLaDAModel",65    "LLaDAOutput",66    "LLaDAGenerateOutput",67]68 69 70log = logging.getLogger(__name__)71 72 73class ModuleType(StrEnum):74    in_module = "in"75    out_module = "out"76    emb = "emb"77    final_out = "final_out"78 79 80def init_weights(81    config: ModelConfig,82    module: Union[nn.Linear, nn.Embedding],83    d: Optional[int] = None,84    layer_id: Optional[int] = None,85    std_factor: float = 1.0,86    type_of_module: Optional[ModuleType] = None,87) -> None:88    """89    Initialize weights of a linear or embedding module.90 91    :param config: The model config.92    :param module: The linear or embedding submodule to initialize.93    :param d: The effective input dimensionality of the weights. This could be smaller than the actual dimensions94        for fused layers.95    :param layer_id: When set, the standard deviation for the "mitchell" method will be adjusted by96        ``1 / sqrt(2 * (layer_id + 1))``.97    """98    d = d if d is not None else config.d_model99    if config.init_fn == InitFnType.normal:100        std = config.init_std * std_factor101        if config.init_cutoff_factor is not None:102            cutoff_value = config.init_cutoff_factor * std103            nn.init.trunc_normal_(module.weight, mean=0.0, std=std, a=-cutoff_value, b=cutoff_value)104        else:105            nn.init.normal_(module.weight, mean=0.0, std=std)106    elif config.init_fn == InitFnType.mitchell:107        std = std_factor / math.sqrt(d)108        if layer_id is not None:109            std = std / math.sqrt(2 * (layer_id + 1))110        nn.init.trunc_normal_(module.weight, mean=0.0, std=std, a=-3 * std, b=3 * std)111    elif config.init_fn == InitFnType.kaiming_normal:112        nn.init.kaiming_normal_(module.weight, nonlinearity="relu")113    elif config.init_fn == InitFnType.fan_in:114        std = std_factor / math.sqrt(d)115        nn.init.normal_(module.weight, mean=0.0, std=std)116    elif config.init_fn == InitFnType.full_megatron:117        if type_of_module is None:118            raise RuntimeError(f"When using the {InitFnType.full_megatron} init, every module must have a type.")119 120        cutoff_factor = config.init_cutoff_factor121        if cutoff_factor is None:122            cutoff_factor = 3123 124        if type_of_module == ModuleType.in_module:125            # for att_proj (same as QKV), ff_proj126            std = config.init_std127        elif type_of_module == ModuleType.out_module:128            # for attn_out, ff_out129            std = config.init_std / math.sqrt(2.0 * config.n_layers)130        elif type_of_module == ModuleType.emb:131            # positional embeddings (wpe)132            # token embeddings (wte)133            std = config.init_std134        elif type_of_module == ModuleType.final_out:135            # final output (ff_out)136            std = config.d_model**-0.5137        else:138            raise RuntimeError(f"Unknown module type '{type_of_module}'")139        nn.init.trunc_normal_(140            module.weight,141            mean=0.0,142            std=std,143            a=-cutoff_factor * std,144            b=cutoff_factor * std,145        )146    else:147        raise NotImplementedError(config.init_fn)148 149    if isinstance(module, nn.Linear):150        if module.bias is not None:151            nn.init.zeros_(module.bias)152 153        if config.init_fn == InitFnType.normal and getattr(module, "_is_residual", False):154            with torch.no_grad():155                module.weight.div_(math.sqrt(2 * config.n_layers))156 157 158def ensure_finite_(x: torch.Tensor, check_neg_inf: bool = True, check_pos_inf: bool = False):159    """160    Modify ``x`` in place to replace ``float("-inf")`` with the minimum value of the dtype when ``check_neg_inf``161    is ``True`` and to replace ``float("inf")`` with the maximum value of the dtype when ``check_pos_inf`` is ``True``.162    """163    if check_neg_inf:164        x.masked_fill_(x == float("-inf"), torch.finfo(x.dtype).min)165    if check_pos_inf:166        x.masked_fill_(x == float("inf"), torch.finfo(x.dtype).max)167 168 169def activation_checkpoint_function(cfg: ModelConfig):170    preserve_rng_state = (171        (cfg.attention_dropout == 0.0) and (cfg.embedding_dropout == 0.0) and (cfg.residual_dropout == 0.0)172    )173    from torch.utils.checkpoint import checkpoint174 175    return partial(176        checkpoint,177        preserve_rng_state=preserve_rng_state,178        use_reentrant=False,179    )180 181 182class BufferCache(dict, MutableMapping[str, torch.Tensor]):183    """184    Cache for attention biases and other things that would normally be stored as buffers.185    We avoid using buffers because we've run into various issues doing so with FSDP.186    In general it appears the way FSDP handles buffers is not well-defined.187    It doesn't shard them but apparently it does synchronize them across processes, which we want to avoid188    since (A) it isn't necessary, and (B) we sometimes have `-inf` in these biases which might get turned into189    NaNs when they're synchronized due to casting or some other issue.190    """191 192 193def _non_meta_init_device(config: ModelConfig) -> torch.device:194    if config.init_device is not None and config.init_device != "meta":195        return torch.device(config.init_device)196    else:197        return torch.device("cuda" if torch.cuda.is_available() else "cpu")198 199 200class Dropout(nn.Dropout):201    def forward(self, input: torch.Tensor) -> torch.Tensor:202        if self.p == 0.0:203            return input204        else:205            return F.dropout(input, self.p, self.training, self.inplace)206 207 208class LayerNormBase(nn.Module):209    def __init__(210        self,211        config: ModelConfig,212        *,213        size: Optional[int] = None,214        elementwise_affine: Optional[bool] = True,215        eps: float = 1e-05,216    ):217        super().__init__()218        self.config = config219        self.eps = eps220        self.normalized_shape = (size or config.d_model,)221        if elementwise_affine or (elementwise_affine is None and self.config.layer_norm_with_affine):222            self.weight = nn.Parameter(torch.ones(self.normalized_shape, device=config.init_device))223            use_bias = self.config.bias_for_layer_norm224            if use_bias is None:225                use_bias = self.config.include_bias226            if use_bias:227                self.bias = nn.Parameter(torch.zeros(self.normalized_shape, device=config.init_device))228            else:229                self.register_parameter("bias", None)230        else:231            self.register_parameter("bias", None)232            self.register_parameter("weight", None)233 234    @abstractmethod235    def forward(self, x: torch.Tensor) -> torch.Tensor:236        raise NotImplementedError237 238    @classmethod239    def build(cls, config: ModelConfig, size: Optional[int] = None, **kwargs) -> LayerNormBase:240        if config.layer_norm_type == LayerNormType.default:241            return LayerNorm(config, size=size, low_precision=False, **kwargs)242        elif config.layer_norm_type == LayerNormType.low_precision:243            return LayerNorm(config, size=size, low_precision=True, **kwargs)244        elif config.layer_norm_type == LayerNormType.rms:245            return RMSLayerNorm(config, size=size, **kwargs)246        elif config.layer_norm_type == LayerNormType.gemma_rms:247            return GemmaRMSLayerNorm(config, size=size, **kwargs)248        else:249            raise NotImplementedError(f"Unknown LayerNorm type: '{config.layer_norm_type}'")250 251    def _cast_if_autocast_enabled(self, tensor: torch.Tensor, dtype: Optional[torch.dtype] = None) -> torch.Tensor:252        # NOTE: `is_autocast_enabled()` only checks for CUDA autocast, so we use the separate function253        # `is_autocast_cpu_enabled()` for CPU autocast.254        # See https://github.com/pytorch/pytorch/issues/110966.255        if tensor.device.type == "cuda" and torch.is_autocast_enabled():256            return tensor.to(dtype=dtype if dtype is not None else torch.get_autocast_gpu_dtype())257        elif tensor.device.type == "cpu" and torch.is_autocast_cpu_enabled():258            return tensor.to(dtype=dtype if dtype is not None else torch.get_autocast_cpu_dtype())259        else:260            return tensor261 262    def reset_parameters(self):263        if self.weight is not None:264            torch.nn.init.ones_(self.weight)  # type: ignore265        if self.bias is not None:266            torch.nn.init.zeros_(self.bias)  # type: ignore267 268 269class LayerNorm(LayerNormBase):270    """271    The default :class:`LayerNorm` implementation which can optionally run in low precision.272    """273 274    def __init__(275        self,276        config: ModelConfig,277        size: Optional[int] = None,278        low_precision: bool = False,279        elementwise_affine: Optional[bool] = None,280        eps: float = 1e-05,281    ):282        super().__init__(config, size=size, elementwise_affine=elementwise_affine, eps=eps)283        self.low_precision = low_precision284 285    def forward(self, x: torch.Tensor) -> torch.Tensor:286        if self.low_precision:287            module_device = x.device288            downcast_x = self._cast_if_autocast_enabled(x)289            downcast_weight = (290                self._cast_if_autocast_enabled(self.weight) if self.weight is not None else self.weight291            )292            downcast_bias = self._cast_if_autocast_enabled(self.bias) if self.bias is not None else self.bias293            with torch.autocast(enabled=False, device_type=module_device.type):294                return F.layer_norm(295                    downcast_x, self.normalized_shape, weight=downcast_weight, bias=downcast_bias, eps=self.eps296                )297        else:298            return F.layer_norm(x, self.normalized_shape, weight=self.weight, bias=self.bias, eps=self.eps)299 300 301class RMSLayerNorm(LayerNormBase):302    """303    RMS layer norm, a simplified :class:`LayerNorm` implementation304    """305 306    def __init__(307        self,308        config: ModelConfig,309        size: Optional[int] = None,310        elementwise_affine: Optional[bool] = None,311        eps: float = 1e-5,312    ):313        super().__init__(config, size=size, elementwise_affine=elementwise_affine, eps=config.rms_norm_eps)314 315    def forward(self, x: torch.Tensor) -> torch.Tensor:316        with torch.autocast(enabled=False, device_type=x.device.type):317            og_dtype = x.dtype318            x = x.to(torch.float32)319            variance = x.pow(2).mean(-1, keepdim=True)320            x = x * torch.rsqrt(variance + self.eps)321            x = x.to(og_dtype)322 323        if self.weight is not None:324            if self.bias is not None:325                return self.weight * x + self.bias326            else:327                return self.weight * x328        else:329            return x330 331 332class GemmaRMSLayerNorm(LayerNormBase):333    """334    Gemma RMS layer norm, a simplified :class:`LayerNorm` implementation335    """336 337    def __init__(338        self,339        config: ModelConfig,340        size: Optional[int] = None,341        elementwise_affine: Optional[bool] = None,342        eps: float = 1e-5,343    ):344        super().__init__(config, size=size, elementwise_affine=elementwise_affine, eps=config.rms_norm_eps)345 346    def forward(self, x: torch.Tensor) -> torch.Tensor:347        with torch.autocast(enabled=False, device_type=x.device.type):348            og_dtype = x.dtype349            x = x.to(torch.float32)350            variance = x.pow(2).mean(-1, keepdim=True)351            x = x * torch.rsqrt(variance + self.eps)352            x = x.to(og_dtype)353 354        if self.weight is not None:355            if self.bias is not None:356                return x * (1 + self.weight) + self.bias357            else:358                return x * (1 + self.weight)359        else:360            return x361 362 363class RotaryEmbedding(nn.Module):364    """365    [Rotary positional embeddings (RoPE)](https://arxiv.org/abs/2104.09864).366    """367 368    def __init__(self, config: ModelConfig, cache: BufferCache):369        super().__init__()370        self.config = config371        self.__cache = cache372        # Warm up cache.373        self.rope_theta = config.rope_theta374        self.get_rotary_embedding(config.max_sequence_length, _non_meta_init_device(config))375 376    def get_rotary_embedding(self, seq_len: int, device: torch.device) -> Tuple[torch.Tensor, torch.Tensor]:377        if (378            (pos_sin := self.__cache.get("rope_pos_sin")) is not None379            and (pos_cos := self.__cache.get("rope_pos_cos")) is not None380            and pos_sin.shape[-2] >= seq_len381            and pos_cos.shape[-2] >= seq_len382        ):383            if pos_sin.device != device:384                pos_sin = pos_sin.to(device)385                self.__cache["rope_pos_sin"] = pos_sin386            if pos_cos.device != device:387                pos_cos = pos_cos.to(device)388                self.__cache["rope_pos_cos"] = pos_cos389            return pos_sin[:, :, :seq_len, :], pos_cos[:, :, :seq_len, :]390 391        with torch.autocast(device.type, enabled=False):392            dim = self.config.d_model // self.config.n_heads393            inv_freq = 1.0 / (self.rope_theta ** (torch.arange(0, dim, 2, device=device, dtype=torch.float) / dim))394            seq = torch.arange(seq_len, device=device, dtype=torch.float)395            freqs = einsum("i , j -> i j", seq, inv_freq)396            positions = torch.cat((freqs, freqs), dim=-1)397            pos_sin, pos_cos = positions.sin()[None, None, :, :], positions.cos()[None, None, :, :]398        self.__cache["rope_pos_sin"] = pos_sin399        self.__cache["rope_pos_cos"] = pos_cos400        return pos_sin, pos_cos401 402    def rotate_half(self, x: torch.Tensor) -> torch.Tensor:403        B, nh, T, hs = x.size()404        x = x.view(B, nh, T, 2, hs // 2)405        x1, x2 = x.unbind(dim=-2)406        return torch.cat((-x2, x1), dim=-1)407 408    def apply_rotary_pos_emb(self, pos_sin: torch.Tensor, pos_cos: torch.Tensor, t: torch.Tensor) -> torch.Tensor:409        return ((t * pos_cos) + (self.rotate_half(t) * pos_sin)).to(t.dtype)410 411    def forward(self, q: torch.Tensor, k: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:412        if self.config.rope_full_precision:413            q_, k_ = q.float(), k.float()414        else:415            q_, k_ = q, k416 417        with torch.autocast(q.device.type, enabled=False):418            query_len, key_len = q_.shape[-2], k_.shape[-2]  # could be different if layer_past not None419            pos_sin, pos_cos = self.get_rotary_embedding(key_len, q_.device)420            pos_sin = pos_sin.type_as(q_)421            pos_cos = pos_cos.type_as(q_)422            q_ = self.apply_rotary_pos_emb(423                pos_sin[:, :, key_len - query_len : key_len, :],424                pos_cos[:, :, key_len - query_len : key_len, :],425                q_,426            )427            k_ = self.apply_rotary_pos_emb(pos_sin, pos_cos, k_)428        return q_.type_as(q), k_.type_as(k)429 430 431class Activation(nn.Module):432    def __init__(self, config: ModelConfig):433        super().__init__()434        self.config = config435 436    @abstractmethod437    def forward(self, x: torch.Tensor) -> torch.Tensor:438        raise NotImplementedError439 440    @property441    @abstractmethod442    def output_multiplier(self) -> float:443        raise NotImplementedError444 445    @classmethod446    def build(cls, config: ModelConfig) -> Activation:447        if config.activation_type == ActivationType.gelu:448            return cast(Activation, GELU(approximate="none"))449        elif config.activation_type == ActivationType.relu:450            return cast(Activation, ReLU(inplace=False))451        elif config.activation_type == ActivationType.silu:452            return cast(Activation, SiLU(inplace=False))453        elif config.activation_type == ActivationType.swiglu:454            return SwiGLU(config)455        else:456            raise NotImplementedError(f"Unknown activation: '{config.activation_type}'")457 458 459class GELU(nn.GELU):460    @property461    def output_multiplier(self) -> float:462        return 1.0463 464 465class ReLU(nn.ReLU):466    @property467    def output_multiplier(self) -> float:468        return 1.0469 470class SiLU(nn.SiLU):471    @property472    def output_multiplier(self) -> float:473        return 1.0474 475class SwiGLU(Activation):476    def forward(self, x: torch.Tensor) -> torch.Tensor:477        x, gate = x.chunk(2, dim=-1)478        return F.silu(gate) * x479 480    @property481    def output_multiplier(self) -> float:482        return 0.5483 484 485def causal_attention_bias(seq_len: int, device: torch.device) -> torch.FloatTensor:486    att_bias = torch.triu(487        torch.ones(seq_len, seq_len, device=device, dtype=torch.float),488        diagonal=1,489    )490    att_bias.masked_fill_(att_bias == 1, torch.finfo(att_bias.dtype).min)491    return att_bias.view(1, 1, seq_len, seq_len)  # type: ignore492 493 494def get_causal_attention_bias(cache: BufferCache, seq_len: int, device: torch.device) -> torch.Tensor:495    if (causal_bias := cache.get("causal_attention_bias")) is not None and causal_bias.shape[-1] >= seq_len:496        if causal_bias.device != device:497            causal_bias = causal_bias.to(device)498            cache["causal_attention_bias"] = causal_bias499        return causal_bias500    with torch.autocast(device.type, enabled=False):501        causal_bias = causal_attention_bias(seq_len, device)502    cache["causal_attention_bias"] = causal_bias503    return causal_bias504 505 506def alibi_attention_bias(seq_len: int, config: ModelConfig, device: torch.device) -> torch.FloatTensor:507    alibi_bias = torch.arange(1 - seq_len, 1, dtype=torch.float, device=device).view(1, 1, 1, seq_len)508 509    # shape: (1, 1, seq_len, seq_len)510    alibi_bias = alibi_bias - torch.arange(1 - seq_len, 1, dtype=torch.float, device=device).view(1, 1, seq_len, 1)511    alibi_bias.abs_().mul_(-1)512 513    # shape: (n_heads,)514    m = torch.arange(1, config.n_heads + 1, dtype=torch.float, device=device)515    m.mul_(config.alibi_bias_max / config.n_heads)516 517    # shape: (1, n_heads, seq_len, seq_len)518    return alibi_bias * (1.0 / (2 ** m.view(1, config.n_heads, 1, 1)))  # type: ignore519 520 521class LLaDABlock(nn.Module):522    """523    A base class for transformer block implementations.524    """525 526    def __init__(self, layer_id: int, config: ModelConfig, cache: BufferCache):527        super().__init__()528        self.layer_id = layer_id529        self.config = config530        self.hidden_size = (531            config.mlp_hidden_size if config.mlp_hidden_size is not None else config.mlp_ratio * config.d_model532        )533        self.__cache = cache534        assert config.d_model % config.n_heads == 0535 536        self._activation_checkpoint_fn = None537 538        # Dropout.539        self.dropout = Dropout(config.residual_dropout)540 541        # Layer norms.542        self.k_norm: Optional[LayerNormBase] = None543        self.q_norm: Optional[LayerNormBase] = None544        if config.attention_layer_norm:545            self.k_norm = LayerNormBase.build(546                config,547                size=(config.d_model // config.n_heads) * config.effective_n_kv_heads,548                elementwise_affine=config.attention_layer_norm_with_affine,549            )550            self.q_norm = LayerNormBase.build(config, elementwise_affine=config.attention_layer_norm_with_affine)551 552        # Activation function.553        self.act = Activation.build(config)554        assert (self.act.output_multiplier * self.hidden_size) % 1 == 0555 556        # Attention output projection.557        self.attn_out = nn.Linear(558            config.d_model, config.d_model, bias=config.include_bias, device=config.init_device559        )560 561        # Feed-forward output projection.562        self.ff_out = nn.Linear(563            int(self.act.output_multiplier * self.hidden_size),564            config.d_model,565            bias=config.include_bias,566            device=config.init_device,567        )568        self.ff_out._is_residual = True  # type: ignore569 570        # Rotary embeddings.571        if self.config.rope:572            self.rotary_emb = RotaryEmbedding(config, self.__cache)573 574        self.flash_attn_func = None575        if config.flash_attention:576            try:577                from flash_attn import flash_attn_func  # type: ignore578 579                self.flash_attn_func = flash_attn_func580            except ModuleNotFoundError:581                pass582 583    def reset_parameters(self):584        if self.k_norm is not None:585            self.k_norm.reset_parameters()586        if self.q_norm is not None:587            self.q_norm.reset_parameters()588        init_weights(589            self.config,590            self.attn_out,591            d=self.config.d_model,592            layer_id=self.layer_id,593            type_of_module=ModuleType.out_module,594        )595        init_weights(596            self.config,597            self.ff_out,598            d=self.ff_out.in_features,599            layer_id=self.layer_id,600            type_of_module=ModuleType.out_module,601        )602 603    def set_activation_checkpointing(self, strategy: Optional[ActivationCheckpointingStrategy]):604        if strategy == ActivationCheckpointingStrategy.fine_grained:605            self._activation_checkpoint_fn = activation_checkpoint_function(self.config)606        else:607            self._activation_checkpoint_fn = None608 609    @classmethod610    def _cast_attn_bias(cls, bias: torch.Tensor, input_dtype: torch.dtype) -> torch.Tensor:611        target_dtype = input_dtype612        # NOTE: `is_autocast_enabled()` only checks for CUDA autocast, so we use the separate function613        # `is_autocast_cpu_enabled()` for CPU autocast.614        # See https://github.com/pytorch/pytorch/issues/110966.615        if bias.device.type == "cuda" and torch.is_autocast_enabled():616            target_dtype = torch.get_autocast_gpu_dtype()617        elif bias.device.type == "cpu" and torch.is_autocast_cpu_enabled():618            target_dtype = torch.get_autocast_cpu_dtype()619        if bias.dtype != target_dtype:620            bias = bias.to(target_dtype)621            ensure_finite_(bias, check_neg_inf=True, check_pos_inf=False)622        return bias623 624    def _scaled_dot_product_attention(625        self,626        q: torch.Tensor,627        k: torch.Tensor,628        v: torch.Tensor,629        attn_mask: Optional[torch.Tensor] = None,630        dropout_p: float = 0.0,631        is_causal: bool = False,632    ) -> torch.Tensor:633        """634        Computes scaled dot product attention on query, key and value tensors, using an optional635        attention mask if passed, and applying dropout if a probability greater than 0.0 is specified.636        """637        if self.flash_attn_func is not None and attn_mask is None:638            r = self.flash_attn_func(639                q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), dropout_p=dropout_p, causal=False640            )641            return r.transpose(1, 2)642        else:643            # torch's sdpa doesn't support GQA, so we're doing this644            assert k.size(1) == v.size(1)645            num_kv_heads = k.size(1)646            num_q_heads = q.size(1)647            if num_q_heads != num_kv_heads:648                assert num_q_heads % num_kv_heads == 0649                k = k.repeat_interleave(num_q_heads // num_kv_heads, dim=1, output_size=num_q_heads)650                v = v.repeat_interleave(num_q_heads // num_kv_heads, dim=1, output_size=num_q_heads)651 652            # Modify: MDM set causal to False.653            return F.scaled_dot_product_attention(654                q,655                k,656                v,657                attn_mask=attn_mask,658                dropout_p=dropout_p,659                is_causal=False,660            )661 662    def attention(663        self,664        q: torch.Tensor,665        k: torch.Tensor,666        v: torch.Tensor,667        attention_bias: Optional[torch.Tensor] = None,668        layer_past: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,669        use_cache: bool = False,670    ) -> Tuple[torch.Tensor, Optional[Tuple[torch.Tensor, torch.Tensor]]]:671        B, T, C = q.size()  # batch size, sequence length, d_model672        dtype = k.dtype673 674        # Optionally apply layer norm to keys and queries.675        if self.q_norm is not None and self.k_norm is not None:676            q = self.q_norm(q).to(dtype=dtype)677            k = self.k_norm(k).to(dtype=dtype)678 679        # Move head forward to be next to the batch dim.680        # shape: (B, nh, T, hs)681        q = q.view(B, T, self.config.n_heads, C // self.config.n_heads).transpose(1, 2)682        # shape: (B, n_kv_h, T, hs)683        k = k.view(B, T, self.config.effective_n_kv_heads, C // self.config.n_heads).transpose(1, 2)684        # shape: (B, n_kv_h, T, hs)685        v = v.view(B, T, self.config.effective_n_kv_heads, C // self.config.n_heads).transpose(1, 2)686 687        if layer_past is not None:688            past_key, past_value = layer_past689            k = torch.cat((past_key, k), dim=-2)690            v = torch.cat((past_value, v), dim=-2)691 692        present = (k, v) if use_cache else None693        query_len, key_len = q.shape[-2], k.shape[-2]  # could be different if layer_past not None694 695        if self.config.rope:696            # Apply rotary embeddings.697            q, k = self.rotary_emb(q, k)698 699        if attention_bias is not None:700            # Resize and cast attention bias.701            # The current dtype of the attention bias might not match the dtype that the SDP attn function will702            # run in if AMP is enabled, and this can be a problem if some tokens are masked out due to padding703            # as down-casting the attention bias to the autocast precision will result in -infs, which will704            # cause the SDP attn function to produce NaNs.705            attention_bias = self._cast_attn_bias(706                attention_bias[:, :, key_len - query_len : key_len, :key_len], dtype707            )708 709        # Get the attention scores.710        # shape: (B, nh, T, hs)711        att = self._scaled_dot_product_attention(712            q,713            k,714            v,715            attn_mask=attention_bias,716            dropout_p=0.0 if not self.training else self.config.attention_dropout,717            is_causal=False,718        )719 720        # Re-assemble all head outputs side-by-side.721        att = att.transpose(1, 2).contiguous().view(B, T, C)722 723        # Apply output projection.724        return self.attn_out(att), present725 726    @abstractmethod727    def forward(728        self,729        x: torch.Tensor,730        attention_bias: Optional[torch.FloatTensor] = None,731        layer_past: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,732        use_cache: bool = False,733    ) -> Tuple[torch.Tensor, Optional[Tuple[torch.Tensor, torch.Tensor]]]:734        raise NotImplementedError735 736    @classmethod737    def build(cls, layer_id: int, config: ModelConfig, cache: BufferCache) -> LLaDABlock:738        if config.block_type == BlockType.sequential:739            return LLaDASequentialBlock(layer_id, config, cache)740        elif config.block_type == BlockType.llama:741            return LLaDALlamaBlock(layer_id, config, cache)742        else:743            raise NotImplementedError(f"Unknown block type: '{config.block_type}'")744 745 746class LLaDASequentialBlock(LLaDABlock):747    """748    This is a typical transformer block where the output is computed as ``MLP(LN(x + Attention(LN(x))))``749    (plus another skip connection).750    """751 752    def __init__(self, layer_id: int, config: ModelConfig, cache: BufferCache):753        super().__init__(layer_id, config, cache)754        # Layer norms.755        self.attn_norm = LayerNorm.build(config)756        self.ff_norm = LayerNorm.build(config)757        # Attention input projection. Projects x -> (q, k, v)758        head_dim = config.d_model // config.n_heads759        self.fused_dims = (760            config.d_model,761            config.effective_n_kv_heads * head_dim,762            config.effective_n_kv_heads * head_dim,763        )764        self.att_proj = nn.Linear(765            config.d_model, sum(self.fused_dims), bias=config.include_bias | config.include_qkv_bias, device=config.init_device766        )767        # Feed-forward input projection.768        self.ff_proj = nn.Linear(769            config.d_model, self.hidden_size, bias=config.include_bias, device=config.init_device770        )771 772    def reset_parameters(self):773        super().reset_parameters()774        self.attn_norm.reset_parameters()775        self.ff_norm.reset_parameters()776        # NOTE: the standard deviation for these weights does not depend on the layer.777        init_weights(778            self.config, self.att_proj, d=self.config.d_model, layer_id=None, type_of_module=ModuleType.in_module779        )780        init_weights(781            self.config, self.ff_proj, d=self.config.d_model, layer_id=None, type_of_module=ModuleType.in_module782        )783 784    def forward(785        self,786        x: torch.Tensor,787        attention_bias: Optional[torch.Tensor] = None,788        layer_past: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,789        use_cache: bool = False,790    ) -> Tuple[torch.Tensor, Optional[Tuple[torch.Tensor, torch.Tensor]]]:791        # Get query, key, value projections.792        # shape:793        #  - for regular attn q, k, v: (batch_size, seq_len, d_model)794        #  - for multi-query attn q: (batch_size, seq_len, d_model)795        #                      k, v: (batch_size, seq_len, d_model // n_heads)796        #  - for group query attn q: (batch_size, seq_len, d_model)797        #                      k, v: (batch_size, seq_len, d_model // n_kv_heads)798        if self._activation_checkpoint_fn is not None:799            q, k, v = self.att_proj(self._activation_checkpoint_fn(self.attn_norm, x)).split(800                self.fused_dims, dim=-1801            )802        else:803            q, k, v = self.att_proj(self.attn_norm(x)).split(self.fused_dims, dim=-1)804 805        # Get attention scores.806        if self._activation_checkpoint_fn is not None:807            att, cache = self._activation_checkpoint_fn(  # type: ignore808                self.attention, q, k, v, attention_bias, layer_past=layer_past, use_cache=use_cache809            )810        else:811            att, cache = self.attention(q, k, v, attention_bias, layer_past=layer_past, use_cache=use_cache)812 813        # Add attention scores.814        # shape: (B, T, C)815        x = x + self.dropout(att)816 817        # Add feed-forward projection.818        # shape: (batch_size, seq_len, d_model)819        og_x = x820        if self._activation_checkpoint_fn is not None:821            x = self._activation_checkpoint_fn(self.ff_norm, x)  # type: ignore822        else:823            x = self.ff_norm(x)824        x = self.ff_proj(x)825        if self._activation_checkpoint_fn is not None:826            x = self._activation_checkpoint_fn(self.act, x)  # type: ignore827        else:828            x = self.act(x)829        x = self.ff_out(x)830        x = self.dropout(x)831        x = og_x + x832 833        return x, cache834 835 836class LLaDALlamaBlock(LLaDABlock):837    """838    This is a transformer block where the output is computed as ``MLP(LN(x + Attention(LN(x))))``839    (plus another skip connection). This block is similar to `LLaDASequentialBlock`840    but some operations have slightly different implementations to imitate the841    behavior of Llama.842    """843 844    def __init__(self, layer_id: int, config: ModelConfig, cache: BufferCache):845        super().__init__(layer_id, config, cache)846        # Layer norms.847        self.attn_norm = LayerNorm.build(config)848        self.ff_norm = LayerNorm.build(config)849        self.__cache = cache850 851        # Attention input projection. Projects x -> (q, k, v)852        head_dim = config.d_model // config.n_heads853        q_proj_out_dim = config.d_model854        k_proj_out_dim = config.effective_n_kv_heads * head_dim855        v_proj_out_dim = config.effective_n_kv_heads * head_dim856        self.q_proj = nn.Linear(857            config.d_model, q_proj_out_dim, bias=config.include_bias | config.include_qkv_bias, device=config.init_device858        )859        self.k_proj = nn.Linear(860            config.d_model, k_proj_out_dim, bias=config.include_bias | config.include_qkv_bias, device=config.init_device861        )862        self.v_proj = nn.Linear(863            config.d_model, v_proj_out_dim, bias=config.include_bias | config.include_qkv_bias, device=config.init_device864        )865 866        # Feed-forward input projection.867        self.ff_proj = nn.Linear(868            config.d_model, self.hidden_size, bias=config.include_bias, device=config.init_device869        )870        # new add871        self.up_proj = nn.Linear(872            config.d_model, self.hidden_size, bias=config.include_bias, device=config.init_device873        )874 875    def reset_parameters(self):876        super().reset_parameters()877        self.attn_norm.reset_parameters()878        self.ff_norm.reset_parameters()879        # NOTE: the standard deviation for these weights does not depend on the layer.880        init_weights(self.config, self.q_proj, d=self.config.d_model, layer_id=None)881        init_weights(self.config, self.k_proj, d=self.config.d_model, layer_id=None)882        init_weights(self.config, self.v_proj, d=self.config.d_model, layer_id=None)883        init_weights(self.config, self.ff_proj, d=self.config.d_model, layer_id=None)884        init_weights(self.config, self.up_proj, d=self.config.d_model, layer_id=None)  # new add885 886    def forward(887        self,888        x: torch.Tensor,889        attention_bias: Optional[torch.Tensor] = None,890        layer_past: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,891        use_cache: bool = False,892    ) -> Tuple[torch.Tensor, Optional[Tuple[torch.Tensor, torch.Tensor]]]:893        # Get query, key, value projections.894        # shape:895        #  - for regular attn q, k, v: (batch_size, seq_len, d_model)896        #  - for multi-query attn q: (batch_size, seq_len, d_model)897        #                      k, v: (batch_size, seq_len, d_model // n_heads)898        #  - for group query attn q: (batch_size, seq_len, d_model)899        #                      k, v: (batch_size, seq_len, d_model // n_kv_heads)900        x_normed = self.attn_norm(x)901        q = self.q_proj(x_normed)902        k = self.k_proj(x_normed)903        v = self.v_proj(x_normed)904 905        # Get attention scores.906        if self._activation_checkpoint_fn is not None:907            att, cache = self._activation_checkpoint_fn(  # type: ignore908                self.attention, q, k, v, attention_bias, layer_past=layer_past, use_cache=use_cache909            )910        else:911            att, cache = self.attention(q, k, v, attention_bias, layer_past=layer_past, use_cache=use_cache)912 913        # Add attention scores.914        # shape: (B, T, C)915        x = x + self.dropout(att)916 917        # Add feed-forward projection.918        # shape: (batch_size, seq_len, d_model)919        og_x = x920        if self._activation_checkpoint_fn is not None:921            x = self._activation_checkpoint_fn(self.ff_norm, x)  # type: ignore922        else:923            x = self.ff_norm(x)924        x, x_up = self.ff_proj(x), self.up_proj(x) # new add925        if self._activation_checkpoint_fn is not None:926            x = self._activation_checkpoint_fn(self.act, x)  # type: ignore927        else:928            x = self.act(x)929        x = x * x_up # new add930        x = self.ff_out(x)931        x = self.dropout(x)932        x = og_x + x933 934        return x, cache935 936 937class LLaDAOutput(NamedTuple):938    logits: torch.FloatTensor939    """940    A tensor of shape `(batch_size, seq_len, vocab_size)` representing the log probabilities941    for the next token *before* normalization via (log) softmax.942    """943 944    attn_key_values: Optional[List[Tuple[torch.Tensor, torch.Tensor]]]945    """946    Attention keys and values from each block.947    """948 949    hidden_states: Optional[Tuple[torch.Tensor]]950    """951    Hidden states from each block.952    """953 954 955class LLaDAGenerateOutput(NamedTuple):956    token_ids: torch.LongTensor957    """958    The generated token IDs, a tensor of shape `(batch_size, beam_size, max_steps)`.959    These do *not* include the original input IDs.960    """961 962    scores: torch.FloatTensor963    """964    The scores of the generated sequences, a tensor of shape `(batch_size, beam_size)`.965    """966 967 968class LLaDABlockGroup(nn.ModuleList):969    def __init__(self, config: ModelConfig, layer_offset: int, modules: Optional[Iterable[nn.Module]] = None):970        super().__init__(modules)971        self.config = config972        self.layer_offset = layer_offset973        self.activation_checkpointing_strategy: Optional[ActivationCheckpointingStrategy] = None974        self._activation_checkpoint_fn = activation_checkpoint_function(self.config)975 976    def forward(977        self,978        x: torch.Tensor,979        attention_bias: Optional[torch.FloatTensor] = None,980        layers_past: Optional[List[Tuple[torch.Tensor, torch.Tensor]]] = None,981        use_cache: bool = False,982    ) -> Tuple[torch.Tensor, Optional[List[Tuple[torch.Tensor, torch.Tensor]]]]:983        attn_key_values: Optional[List[Tuple[torch.Tensor, torch.Tensor]]] = [] if use_cache else None984        for block_idx, block in enumerate(self):985            layer_past = None if layers_past is None else layers_past[block_idx]986            block_idx += self.layer_offset987            if (988                (self.activation_checkpointing_strategy == ActivationCheckpointingStrategy.whole_layer)989                or (990                    self.activation_checkpointing_strategy == ActivationCheckpointingStrategy.one_in_two991                    and block_idx % 2 == 0992                )993                or (994                    self.activation_checkpointing_strategy == ActivationCheckpointingStrategy.one_in_three995                    and block_idx % 3 == 0996                )997                or (998                    self.activation_checkpointing_strategy == ActivationCheckpointingStrategy.one_in_four999                    and block_idx % 4 == 01000                )1001            ):1002                # shape: (batch_size, seq_len, d_model)1003                x, cache = self._activation_checkpoint_fn(  # type: ignore1004                    block, x, attention_bias=attention_bias, layer_past=layer_past, use_cache=use_cache1005                )1006            else:1007                # shape: (batch_size, seq_len, d_model)1008                x, cache = block(x, attention_bias=attention_bias, layer_past=layer_past, use_cache=use_cache)1009            if attn_key_values is not None:1010                assert cache is not None1011                attn_key_values.append(cache)1012        return x, attn_key_values1013 1014    def reset_parameters(self):1015        for block in self:1016            block.reset_parameters()1017 1018    def set_activation_checkpointing(self, strategy: Optional[ActivationCheckpointingStrategy]):1019        self.activation_checkpointing_strategy = strategy1020        for block in self:1021            block.set_activation_checkpointing(strategy)1022 1023 1024class LLaDAModel(nn.Module):1025    def __init__(self, config: ModelConfig, init_params: bool = True):1026        super().__init__()1027        self.config = config1028        self.__cache = BufferCache()1029 1030        # Validate config.1031        if self.config.alibi and self.config.flash_attention:1032            raise Exception("ALiBi is currently not supported with FlashAttention")1033 1034        if self.config.alibi and self.config.rope:1035            raise Exception("ALiBi and RoPE are mutually exclusive")1036 1037        if self.config.embedding_size is not None and self.config.embedding_size != self.config.vocab_size:1038            if self.config.embedding_size < self.config.vocab_size:1039                raise Exception("embedding size should be at least as big as vocab size")1040            elif self.config.embedding_size % 128 != 0:1041                import warnings1042 1043                warnings.warn(1044                    "Embedding size is not a multiple of 128! This could hurt throughput performance.", UserWarning1045                )1046 1047        self.activation_checkpointing_strategy: Optional[ActivationCheckpointingStrategy] = None1048        self._activation_checkpoint_fn: Callable = activation_checkpoint_function(self.config)1049 1050        if not (1051            0 < self.config.block_group_size <= self.config.n_layers1052            and self.config.n_layers % self.config.block_group_size == 01053        ):1054            raise Exception("n layers must be divisible by block group size")1055 1056        torch.backends.cuda.enable_flash_sdp(True)1057        torch.backends.cuda.enable_mem_efficient_sdp(False)  # this is super slow so make sure torch won't use it1058 1059        self.transformer = nn.ModuleDict(1060            dict(1061                wte=nn.Embedding(1062                    config.embedding_size or config.vocab_size, config.d_model, device=config.init_device1063                ),1064                emb_drop=Dropout(config.embedding_dropout),1065                ln_f=LayerNorm.build(config),1066            )1067        )1068 1069        blocks = [LLaDABlock.build(i, config, self.__cache) for i in range(config.n_layers)]1070        if self.config.block_group_size > 1:1071            block_groups = [1072                LLaDABlockGroup(config, i, blocks[i : i + config.block_group_size])1073                for i in range(0, config.n_layers, config.block_group_size)1074            ]1075            self.transformer.update({"block_groups": nn.ModuleList(block_groups)})1076        else:1077            self.transformer.update({"blocks": nn.ModuleList(blocks)})1078 1079        if not (self.config.alibi or self.config.rope):1080            self.transformer.update(1081                {"wpe": nn.Embedding(config.max_sequence_length, config.d_model, device=config.init_device)}1082            )1083        if not config.weight_tying:1084            self.transformer.update(1085                {1086                    "ff_out": nn.Linear(1087                        config.d_model,1088                        config.embedding_size or config.vocab_size,1089                        bias=config.include_bias,1090                        device=config.init_device,1091                    )1092                }1093            )1094        # When `init_device="meta"` FSDP will call `reset_parameters()` to initialize weights.1095        if init_params and self.config.init_device != "meta":1096            self.reset_parameters()1097        self.__num_fwd_flops: Optional[int] = None1098 1099        # Warm up cache.1100        if self.config.alibi:1101            get_causal_attention_bias(self.__cache, config.max_sequence_length, _non_meta_init_device(config))1102            self.get_alibi_attention_bias(config.max_sequence_length, _non_meta_init_device(config))1103 1104    def set_activation_checkpointing(self, strategy: Optional[ActivationCheckpointingStrategy]):1105        self.activation_checkpointing_strategy = strategy1106        if self.config.block_group_size != 1:1107            for block_group in self.transformer.block_groups:1108                block_group.set_activation_checkpointing(strategy)1109        else:1110            for block in self.transformer.blocks:1111                block.set_activation_checkpointing(strategy)1112 1113    @property1114    def device(self) -> torch.device:1115        device: torch.device = self.transformer.wte.weight.device  # type: ignore1116        if device.type == "meta":1117            return _non_meta_init_device(self.config)1118        else:1119            return device1120 1121    def reset_parameters(self):1122        log.info("Initializing model parameters...")1123        # Top-level embeddings / linear layers.1124        init_weights(1125            self.config,1126            self.transformer.wte,  # type: ignore1127            std_factor=(0.5 * math.sqrt(self.config.d_model)) if self.config.scale_logits else 1.0,1128            type_of_module=ModuleType.emb,1129        )1130        if hasattr(self.transformer, "wpe"):1131            init_weights(self.config, self.transformer.wpe, type_of_module=ModuleType.emb)  # type: ignore1132 1133        # Top-level layer norm.1134        self.transformer.ln_f.reset_parameters()  # type: ignore1135 1136        # Output weights.1137        if hasattr(self.transformer, "ff_out"):1138            init_weights(self.config, self.transformer.ff_out, type_of_module=ModuleType.final_out)  # type: ignore1139 1140        # Let the blocks handle themselves.1141        if self.config.block_group_size == 1:1142            for block in self.transformer.blocks:1143                block.reset_parameters()1144        else:1145            for block_group in self.transformer.block_groups:1146                block_group.reset_parameters()1147 1148    def get_alibi_attention_bias(self, seq_len: int, device: torch.device) -> torch.Tensor:1149        if (alibi_bias := self.__cache.get("alibi_attention_bias")) is not None and alibi_bias.shape[1150            -11151        ] >= seq_len:1152            if alibi_bias.device != device:1153                alibi_bias = alibi_bias.to(device)1154                self.__cache["alibi_attention_bias"] = alibi_bias1155            return alibi_bias1156        with torch.autocast(device.type, enabled=False):1157            alibi_bias = alibi_attention_bias(seq_len, self.config, device)1158        self.__cache["alibi_attention_bias"] = alibi_bias1159        return alibi_bias1160    1161    def get_bidirectional_attention_bias(self, seq_len: int, device: torch.device) -> torch.Tensor:1162        if (bidirectional_bias := self.__cache.get("bidirectional_attention_bias")) is not None and bidirectional_bias.shape[1163            -11164        ] >= seq_len:1165            if bidirectional_bias.device != device:1166                bidirectional_bias = bidirectional_bias.to(device)1167                self.__cache["bidirectional_attention_bias"] = bidirectional_bias1168            return bidirectional_bias1169        with torch.autocast(device.type, enabled=False):1170            bidirectional_bias = torch.zeros((1, 1, seq_len, seq_len), device=device, dtype=torch.float)1171        self.__cache["bidirectional_attention_bias"] = bidirectional_bias1172        return bidirectional_bias1173    1174    def forward(1175        self,1176        input_ids: torch.LongTensor,1177        input_embeddings: Optional[torch.FloatTensor] = None,1178        attention_mask: Optional[torch.Tensor] = None,1179        attention_bias: Optional[torch.Tensor] = None,1180        past_key_values: Optional[Sequence[Tuple[torch.Tensor, torch.Tensor]]] = None,1181        use_cache: bool = False,1182        last_logits_only: bool = False,1183        output_hidden_states: Optional[bool] = None,1184    ) -> LLaDAOutput:1185        """1186        :param input_ids: A tensor of shape `(batch_size, seq_len)`.1187        :param input_embeddings: A tensor of shape `(batch_size, seq_len, d_model)` with input1188            embeddings. When provided, it is treated as the output of the input embedding layer.1189        :param attention_mask: A tensor of shape `(batch_size, seq_len)` that indicates1190            which input IDs are masked. A `1` value in the mask means that1191            the corresponding input ID should *not* be ignored. A `0` means1192            that the corresponding input ID is masked.1193 1194            This has the same meaning as the `attention_mask` in HuggingFace's `transformers`1195            library.1196        :param attention_bias: A tensor of shape `(batch_size, 1, seq_len, seq_len)`,1197            `(1, 1, seq_len, seq_len)`, or `(seq_len, seq_len)`. This is used1198            to introduce causal or other biases.1199 1200            If the tensor is a bool or byte tensor, a `True` or `1` at `attention_bias[:, :, i, j]`

Showing the first 1,200 of 1515 lines. Download the file for the rest.