diffusion-reasoning/gdsd_code_llada
022
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]`