multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
7b592f7 verified
Raw
History Blame Contribute Delete
43 kB
# Copyright Lightning AI. Licensed under the Apache License 2.0, see LICENSE file.
"""Full definition of a decoder-only transformer-based language model, all of it in this single file.
Based on the nanoGPT implementation: https://github.com/karpathy/nanoGPT and
https://github.com/EleutherAI/gpt-neox/tree/main/megatron/model.
"""
import math
from typing import Any, Optional, Tuple, Union, List
from functools import partial
from transformers import AutoConfig, Qwen2_5OmniForConditionalGeneration
import torch
import torch.nn as nn
import torch.nn.functional as F
from typing_extensions import Self
import whisper
from transformers import Qwen2AudioEncoder, Qwen2AudioConfig
from src.audiointeraction.config import Config
def qkv_reassemble(
param: torch.Tensor, config: Config
) -> torch.Tensor:
"""Reassemble from a normal to an interleaved placement in a QKV matrix.
[Q, K, V, Q, K, V, ...] --> [Q, Q, ..., K, K, ..., V, V, ...]
"""
q_per_kv = config.n_head // config.n_query_groups
qs = []
ks = []
vs = []
for chunk in torch.chunk(param, config.n_query_groups):
split = torch.split(chunk, [config.head_size * q_per_kv, config.head_size, config.head_size])
qs.append(split[0])
ks.append(split[1])
vs.append(split[2])
q = torch.cat(qs)
k = torch.cat(ks)
v = torch.cat(vs)
return torch.cat((q, k, v))
class GPT(nn.Module):
def __init__(self, config: Config) -> None:
super().__init__()
assert config.padded_vocab_size is not None
self.config = config
self.lm_head = nn.Linear(
config.n_embd, config.padded_vocab_size, bias=config.lm_head_bias
)
self.transformer = nn.ModuleDict(
dict(
wte=nn.Embedding(config.padded_vocab_size, config.n_embd),
h=nn.ModuleList(
Block(config, block_idx)
for block_idx in range(config.n_layer)
),
ln_f=config.norm_class(config.n_embd, eps=config.norm_eps),
)
)
self.mask_cache: Optional[torch.Tensor] = None
self.max_seq_length = self.config.block_size
@property
def max_seq_length(self) -> int:
return self._max_seq_length
@max_seq_length.setter
def max_seq_length(self, value: int) -> None:
"""
When doing inference, the sequences used might be shorter than the model's context length.
This allows setting a smaller number to avoid allocating unused memory
"""
if value > self.config.block_size:
raise ValueError(
f"Cannot attend to {value}, block size is only {self.config.block_size}."
" This is likely because the input text exceeds the supported context length of this model."
)
self._max_seq_length = value
if not hasattr(self, "cos"):
# first call
cos, sin = self.rope_cache()
self.register_buffer("cos", cos, persistent=False)
self.register_buffer("sin", sin, persistent=False)
# override
elif value != self.cos.size(0):
self.cos, self.sin = self.rope_cache(device=self.cos.device)
# the mask and kv cache size will get updated on `set_kv_cache`. we cannot update it here because we don't know
# if the kv cache is expected
if self.mask_cache is not None and self.mask_cache.shape[-1] < value:
print(f"Warning: KV cache has length {self.mask_cache.shape[-1]} < {value} = max_seq_length. Call 'set_kv_cache' before doing any forwards!")
def reset_parameters(self) -> None:
# Trigger resetting the rope-cache
self.cos, self.sin = self.rope_cache(device=self.cos.device)
def _init_weights(self, module: nn.Module) -> None:
"""Meant to be used with `gpt.apply(gpt._init_weights)`."""
if isinstance(module, nn.Linear):
torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
if module.bias is not None:
torch.nn.init.zeros_(module.bias)
elif isinstance(module, nn.Embedding):
torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
def fill_in_audio_feature(self,
input_embeddings: torch.Tensor,
batch_size: int,
audio_feats_list,
audio_pos,
tasks) -> torch.Tensor:
"""Replace AUDIO_PAD positions in input_embeddings with precomputed audio features.
Two fill modes, dispatched per-sample by `tasks[batch_idx]`:
- "online": audio is streamed in fixed 10-frame chunks. `audio_pos[i]`
is a list of (start, end) tuples (each end-start == 10);
slice the feature tensor by 10 per chunk.
- "offline": audio is one contiguous block. `audio_pos[i]` is a single
(start, end) tuple covering `output_len` positions; place
the entire feature tensor in one shot.
"""
_, _, emb_dim = input_embeddings.shape
if not (batch_size == len(audio_feats_list) == len(audio_pos) == len(tasks)):
raise ValueError(
f"length mismatch: batch_size={batch_size}, "
f"feats={len(audio_feats_list)}, pos={len(audio_pos)}, tasks={len(tasks)}"
)
for batch_idx in range(batch_size):
audio_feats = audio_feats_list[batch_idx]
segments = audio_pos[batch_idx]
if segments is None or segments == -1:
continue
task = tasks[batch_idx]
if task == "offline":
# Single big block: place the whole feature tensor at the one segment.
start, end = segments[0]
if start >= self.max_seq_length:
continue
end = min(end, self.max_seq_length)
input_embeddings[batch_idx, start:end, :] = audio_feats[: end - start]
continue
# Online: per-10-frame chunk placement.
for seg_idx, (start, end) in enumerate(segments):
if start > self.max_seq_length:
continue
if end > self.max_seq_length:
input_embeddings[batch_idx, start:self.max_seq_length, :] = torch.zeros(
self.max_seq_length - start, emb_dim
)
else:
audio_feat = audio_feats[seg_idx * 10 : (seg_idx + 1) * 10]
seg_len, feat_dim = audio_feat.shape
expected_len = end - start
if seg_len != expected_len or feat_dim != emb_dim:
raise ValueError(
f"Loaded feature shape {audio_feat.shape} does not match expected "
f"({expected_len}, {emb_dim}) at batch {batch_idx}, segment {seg_idx}")
# Overwrite the embedding segment
input_embeddings[batch_idx, start:end, :] = audio_feat
return input_embeddings
def forward(
self,
idx: torch.Tensor,
tasks: Optional[List[str]],
batch_size: int,
audio_info: Optional[Union[dict, torch.Tensor]] = None,
input_pos: Optional[torch.Tensor] = None,
input_pos_maxp1: Optional[torch.Tensor] = None,
audio_tokens_per_chunk: int = 10,
lm_head_chunk_size: int = 0,
) -> Union[torch.Tensor, List[torch.Tensor]]:
"""
If `input_pos` is provided, the KV cache uses K and V vectors for
positions smaller than entries in `input_pos`. For efficiency, pass
`input_pos_maxp1` as `max(input_pos) + 1` if already available from
your forward algorithm. This slices the KV cache buffers and speeds
up multi-head attention.
Without `input_pos_maxp1`, the computation uses the full KV cache
(`max_seq_length`) with masking applied. Note that inferring
`input_pos_maxp1` from `input_pos` causes graph breaks and prevents
compilation.
Args:
idx: Token indices of input sequences, shape `(B, T)`, where `B`
is batch size.
input_pos: Optional. Positions of input tokens. The default is
`arange(T)`. Can have shape `(T,)` or `(B, T)` (batched index).
input_pos_maxp1: Optional. See above.
lm_head_chunk_size: Optional. If `lm_head_chunk_size > 0`, the final
`lm_head` computation is done in chunks of this size.
Returns:
Logit outputs, shape `(B, T, config.padded_vocab_size)`. If
`lm_head_chunk_size > 0`, this is a list of chunks of shape
`(B, lm_head_chunk_size, config.padded_vocab_size)`, the final
entry can be shorter.
"""
T = idx.size(1)
if self.max_seq_length < T:
raise ValueError(f"Cannot forward sequence of length {T}, max seq length is only {self.max_seq_length}.")
if input_pos is not None: # use the kv cache
if input_pos.dim() > 2:
# otherwise, things go wrong in `apply_rope`
raise ValueError(f"input_pos must have 1 or 2 dimensions, input_pos.shape = {input_pos.shape}")
if input_pos.shape[-1] != T:
raise ValueError(f"input_pos.shape[-1] = {input_pos.shape[-1]} != {T} = idx.shape[1], must be the same")
cos = batched_index_select(self.cos, 0, input_pos)
sin = batched_index_select(self.sin, 0, input_pos)
if input_pos.dim() == 1:
cos = cos.unsqueeze(0)
sin = sin.unsqueeze(0)
if self.mask_cache is None:
raise TypeError("You need to call `gpt.set_kv_cache()`")
mask = batched_index_select(self.mask_cache, 2, input_pos)
if mask.dim() > 4:
# the mask cache has a batch dim of 1 in addition to the one
# we get if input_pos has a batch dimension
mask = mask.view(*(mask.shape[0:1] + mask.shape[2:]))
if input_pos_maxp1 is not None:
# Shorten final dimension so it just covers all `input_pos` entries
if input_pos_maxp1 > self.max_seq_length:
raise ValueError(f"Positions in 'input_pos' must be in [0,{self.max_seq_length})")
mask = mask[..., :input_pos_maxp1]
else:
# unsqueeze to have a batch dimension
cos = self.cos[:T].unsqueeze(0)
sin = self.sin[:T].unsqueeze(0)
# `cos`, `sin` have shape (1, T, config.rope_n_elem)
mask = None # defaults to causal mask
input_pos_maxp1 = None
x = self.transformer.wte(idx) # token embeddings of shape (B, T, n_embd)
# Audio feature injection — dispatch on input type. Encoder features are
# already n_embd-dim (projected by audio_tower.proj), so we place them
# directly into the input embeddings.
# - dict : training path, segment-based fill from precomputed features
# - Tensor : inference path, streaming chunk replacement
# - None : no audio (e.g. text-only data or inter-token decoding step)
if isinstance(audio_info, dict):
# T_T (text-only) samples have audio_pos == None — nothing to fill.
if audio_info.get("audio_pos") is not None:
x = self.fill_in_audio_feature(
x, batch_size, audio_info["feats_paths"], audio_info["audio_pos"], tasks,
)
elif torch.is_tensor(audio_info):
if T > audio_tokens_per_chunk:
if x.size(0) != 1:
raise ValueError("inference mode, it is not supported for batch size > 1")
x[0, T - (audio_tokens_per_chunk + 1): T - 1, :] = audio_info
if self.config.scale_embeddings:
x = x * torch.tensor(self.config.n_embd ** 0.5, dtype=x.dtype)
for block in self.transformer.h:
x = block(x, cos, sin, mask, input_pos, input_pos_maxp1)
x = self.transformer.ln_f(x)
clamp_head = (
partial(do_softcapping, thresh=self.config.final_logit_softcapping)
if self.config.final_logit_softcapping is not None
else nn.Identity()
)
if lm_head_chunk_size > 0:
# chunk the lm head logits to reduce the peak memory used by autograd
return [
clamp_head(self.lm_head(x_i))
for x_i in x.split(lm_head_chunk_size, dim=1)
]
else:
return clamp_head(self.lm_head(x)) # (B, T, padded_vocab_size)
def rope_cache(self, device: Optional[torch.device] = None) -> Tuple[torch.Tensor, torch.Tensor]:
if self.config.rope_adjustments is None:
extra_config = None
else:
adjusted_params_required = ["factor", "low_freq_factor", "high_freq_factor", "original_max_seq_len"]
params_present = [param in self.config.rope_adjustments for param in adjusted_params_required]
num_params_present = sum(params_present)
if num_params_present == 0:
extra_config = None # uses standard RoPE
elif num_params_present == 4:
# These parameters should always be used together so that we don't interfere with standard rope
extra_config = {
name: self.config.rope_adjustments[name]
for name in adjusted_params_required
}
else:
# Some but not all parameters are specified; raise an error
missing_params = [
param for param, present in zip(adjusted_params_required, params_present) if not present
]
raise ValueError(
f"The following adjusted RoPE parameters are missing in rope_adjustments: {', '.join(missing_params)}. "
"All adjusted RoPE parameters must be specified together."
)
return build_rope_cache(
seq_len=self.max_seq_length,
n_elem=self.config.rope_n_elem,
device=device,
condense_ratio=self.config.rope_condense_ratio,
base=self.config.rope_base,
extra_config=extra_config,
)
def set_kv_cache(
self,
batch_size: int,
max_seq_length: Optional[int] = None,
rope_cache_length: Optional[int] = None,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
) -> None:
if rope_cache_length is None:
rope_cache_length = self.cos.size(-1)
if max_seq_length is None:
max_seq_length = self.max_seq_length
# initialize the kv cache for all blocks
for block in self.transformer.h:
block.attn.kv_cache = block.attn.build_kv_cache(
batch_size,
max_seq_length,
rope_cache_length,
device,
dtype,
)
if self.mask_cache is None or self.mask_cache.size(3) != max_seq_length:
# passing `attn_mask` to SDPA disables the flash implementation. since we only need the mask
# for the kv-cache support (only during inference), we only create it in that situation
self.mask_cache = build_mask_cache(max_seq_length, device)
def clear_kv_cache(self) -> None:
self.mask_cache = None
for block in self.transformer.h:
block.attn.kv_cache = None
class Block(nn.Module):
def __init__(
self,
config: Config,
block_idx: int,
) -> None:
super().__init__()
if not config.parallel_residual and config.shared_attention_norm:
raise NotImplementedError(
"No checkpoint amongst the ones we support uses this configuration"
" (non-parallel residual and shared attention norm)."
)
self.norm_1 = config.norm_class(config.n_embd, eps=config.norm_eps)
self.attn = CausalSelfAttention(config, block_idx)
self.post_attention_norm = (
config.norm_class(config.n_embd, eps=config.norm_eps) if config.post_attention_norm else nn.Identity()
)
self.norm_2 = None if config.shared_attention_norm else config.norm_class(config.n_embd, eps=config.norm_eps)
self.mlp = config.mlp_class(config)
self.post_mlp_norm = (
config.norm_class(config.n_embd, eps=config.norm_eps) if config.post_mlp_norm else nn.Identity()
)
self.config = config
def forward(
self,
x: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
mask: Optional[torch.Tensor] = None,
input_pos: Optional[torch.Tensor] = None,
input_pos_maxp1: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""
Non-parallel residual Parallel residual
┌─ x ┌─ x ──────────────────┐ Note: if `shared_attention_norm` is True,
│ ↓ │ ↓ ↓ the output from `norm_1` is reused
│ norm_1 │ norm_1 ───────► norm_2
│ ↓ │ ↓ ↓
│ attn │ attn MLP
│ ↓ │ ↓ ↓
| post_attn_norm | post_attn_norm post_mlp_norm
| ↓ | ↓ ↓
┌─ └► + └► + ◄─────────────────┘
| ↓
│ norm_2
│ ↓
│ MLP
│ ↓
| post_mlp_norm
| ↓
└───► +
"""
x_normed = self.norm_1(x)
attention_output = self.attn(
x_normed, cos, sin, mask, input_pos, input_pos_maxp1
)
attention_output = self.post_attention_norm(attention_output)
if self.config.parallel_residual:
if not self.config.shared_attention_norm:
x_normed = self.norm_2(x)
x = attention_output + x
else:
x = attention_output + x
x_normed = self.norm_2(x)
return self.post_mlp_norm(self.mlp(x_normed)) + x
class CausalSelfAttention(nn.Module):
def __init__(self, config: Config, block_idx: int) -> None:
super().__init__()
# key, query and value projections for all heads, but in a batch
self.qkv = nn.Linear(
config.n_embd,
(config.n_head + 2 * config.n_query_groups) * config.head_size, # support for grouped/multi queries
bias=config.bias or config.attn_bias,
)
# output projection
self.proj = nn.Linear(
config.head_size * config.n_head, config.n_embd, bias=config.bias
)
# disabled by default
self.kv_cache: Optional[KVCache] = None
self.apply_sliding_window_attention = (
config.sliding_window_size is not None and
block_idx % config.sliding_window_layer_stride == 0
)
if config.norm_qk:
self.norm_q = config.norm_class(config.head_size * config.n_head, eps=config.norm_eps)
self.norm_k = config.norm_class(config.head_size * config.n_query_groups, eps=config.norm_eps)
else:
self.norm_q = self.norm_k = None
self.config = config
self.block_idx = block_idx
# Attention capture flags (for analysis/visualization, disabled by default)
self.capture_attn: bool = False
self.captured_attn_weights: Optional[torch.Tensor] = None # shape: (B, n_head, T_q, T_k)
def forward(
self,
x: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
mask: Optional[torch.Tensor] = None,
input_pos: Optional[torch.Tensor] = None,
input_pos_maxp1: Optional[torch.Tensor] = None,
) -> torch.Tensor:
# Notation:
# - B | batch size
# - T | time-step (sequence length)
# - C | model's embeddings size (n_embd)
# - C* | attentions's embeddings size
# - nh_(q,k,v) | number of heads for query, key and value
# - hs | head size
head_size = self.config.head_size
n_head = self.config.n_head
n_query_groups = self.config.n_query_groups
rope_n_elem = self.config.rope_n_elem
B, T, C = x.size() # batch size, sequence length, embedding dimensionality (n_embd)
# Perform a single multiplication operation using a combined QKV matrix to calculate `query`, `key`, and `value`
# instead of individually multiplying the input `x` with the respective weight matrices.
qkv = self.qkv(x) # (B, T, 3xC*)
# Define query, key and value sizes.
# If grouped/multi query is enabled, these sizes are not equal (see the diagram in `lit_gpt/config.py::Config`).
query_size = n_head * head_size
key_size = value_size = n_query_groups * head_size
# Split qkv into query, key and value matrices.
q, k, v = qkv.split((query_size, key_size, value_size), dim=-1) # 3x(B, T, C*)
if self.config.norm_qk:
q = self.norm_q(q)
k = self.norm_k(k)
# To place the num_heads (nh) dimension right after the batch (B) dimension, the first step is to decouple the
# embedding size (C) into num_heads (nh) and head_size (hs).
q = q.view(B, T, n_head, head_size) # (B, T, nh_q, hs)
k = k.view(B, T, n_query_groups, head_size) # (B, T, nh_k, hs)
v = v.view(B, T, n_query_groups, head_size) # (B, T, nh_v, hs)
# The tensors `query`, `key`, and `value` are now accurately structured: within each batch element (B), there are
# multiple heads (nh), and within each head, there is a sequence of elements (T), each represented by a vector
# of size `hs`.
q = q.transpose(1, 2) # (B, nh_q, T, hs)
k = k.transpose(1, 2) # (B, nh_k, T, hs)
v = v.transpose(1, 2) # (B, nh_v, T, hs)
# Unlike standard positional embeddings rotary embeddings must be applied at every layer.
q_roped = apply_rope(q[..., : rope_n_elem], cos, sin)
k_roped = apply_rope(k[..., : rope_n_elem], cos, sin)
q = torch.cat((q_roped, q[..., rope_n_elem :]), dim=-1) # (B, nh_q, T, hs)
k = torch.cat((k_roped, k[..., rope_n_elem :]), dim=-1) # (B, nh_k, T, hs)
# Apply kv-cache during inference.
if input_pos is not None:
if not isinstance(self.kv_cache, KVCache):
raise TypeError("You need to call `gpt.set_kv_cache()`")
k, v = self.kv_cache(input_pos, k, v)
if input_pos_maxp1 is not None:
# Subselect along sequence dimension
k = k[..., :input_pos_maxp1, :]
v = v[..., :input_pos_maxp1, :]
# k, v: (B, nh_k, input_pos_maxp1, hs)
# If input_pos_maxp1 is None -> max_seq_length
use_flash = (getattr(self.config, "use_flash_attention", True)
and mask is None
and n_query_groups == n_head
)
if use_flash:
# FlashAttention: B H T D -> B T H D
q = q.transpose(1, 2).contiguous() # (B, T, nh_q, hs)
k = k.transpose(1, 2).contiguous() # (B, T, nh_k, hs)
v = v.transpose(1, 2).contiguous() # (B, T, nh_v, hs)
from flash_attn.flash_attn_interface import flash_attn_func
y = flash_attn_func(q, k, v, dropout_p=0.0, causal=True)
y = y.transpose(1, 2) # back to B H T D
else:
# Grouped queries: balance the number of heads across all three matrices.
# NOTE: flash attention requires it in training mode.
# Multi-query: this step can be skipped since there is only 1 head, allowing us to use broadcasting.
if n_query_groups != n_head and (input_pos is None or n_query_groups != 1):
q_per_kv = n_head // n_query_groups
k = k.repeat_interleave(q_per_kv, dim=1) # (B, nh_q, T, hs)
v = v.repeat_interleave(q_per_kv, dim=1) # (B, nh_q, T, hs)
if self.apply_sliding_window_attention:
"""
Global Window Sliding window Sliding window
attention mask + bias = attention mask
┌────────────────────────┐ ┌───────────────────────┐ ┌─────────────────────────┐
│ True False False False │ │ True True True True │ │ True False False False │
│ True True False False │ │ True True True True │ │ True True False False │
│ True True True False │ │ False True True True │ │ False True True False │
│ True True True True │ │ False False True True │ │ False False True True │
└────────────────────────┘ └───────────────────────┘ └─────────────────────────┘
"""
if mask is None:
mask = torch.ones(T, T, dtype=q.dtype, device=q.device).triu(diagonal=1)
mask.masked_fill_(mask.bool(), float("-inf"))
mask = mask.view(1, 1, *mask.shape)
sliding_window_bias = torch.ones_like(mask).tril(diagonal=-self.config.sliding_window_size)
sliding_window_bias.masked_fill_(sliding_window_bias.bool(), float("-inf"))
mask += sliding_window_bias
# Efficient attention using Flash Attention CUDA kernels.
# NOTE: efficient implementation is disabled if `mask` is not None or softcapping is enabled.
# ↓ (B, nh, T, hs) @ (B, nh, T, hs).mT --> (B, nh, T, T) @ (B, nh, T, hs) --> (B, nh, T, hs)
y = self.scaled_dot_product_attention(q, k, v, mask)
# Re-assemble all head outputs side by side.
y = y.reshape(B, T, head_size * n_head)
# Output projection.
return self.proj(y) # (B, T, C)
def scaled_dot_product_attention(
self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, mask: Optional[torch.Tensor] = None
) -> torch.Tensor:
scale = 1.0 / math.sqrt(self.config.attention_scores_scalar or self.config.head_size)
# with softcapping we cannot use SDPA
if self.config.attention_logit_softcapping is not None:
scores = q @ k.mT * scale
scores = do_softcapping(scores, self.config.attention_logit_softcapping)
if mask is None:
mask = torch.ones(q.size(2), q.size(2), dtype=q.dtype, device=q.device).triu(diagonal=1)
mask.masked_fill_(mask.bool(), torch.finfo(q.dtype).min)
scores = scores + mask
scores = F.softmax(scores, dim=-1, dtype=torch.float).to(dtype=q.dtype)
if self.capture_attn:
self.captured_attn_weights = scores.detach()
y = scores @ v
elif self.capture_attn:
# Manual attention computation to capture weights (bypasses fused SDPA kernel)
# q: (B, n_head, T_q, hs), k: (B, n_head, T_k, hs)
scores = torch.matmul(q.float(), k.float().transpose(-2, -1)) * scale
if mask is not None:
if mask.dtype == torch.bool:
scores = scores.masked_fill(~mask, float('-inf'))
else:
scores = scores + mask.float()
else:
# Apply causal mask manually when no mask is provided (training mode)
T_q, T_k = q.size(-2), k.size(-2)
causal = torch.ones(T_q, T_k, device=q.device, dtype=torch.bool).tril(diagonal=T_k - T_q)
scores = scores.masked_fill(~causal, float('-inf'))
attn_weights = F.softmax(scores, dim=-1)
self.captured_attn_weights = attn_weights.detach()
y = torch.matmul(attn_weights.to(dtype=v.dtype), v)
return y.transpose(1, 2)
else:
y = F.scaled_dot_product_attention(
q, k, v, attn_mask=mask, dropout_p=0.0, scale=scale, is_causal=mask is None
)
return y.transpose(1, 2)
def build_kv_cache(
self,
batch_size: int,
max_seq_length: int,
rope_cache_length: Optional[int] = None,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
) -> "KVCache":
v_shape = (batch_size, self.config.n_query_groups, max_seq_length, self.config.head_size)
if rope_cache_length is None:
if self.config.rotary_percentage != 1.0:
raise TypeError("Please pass the `rope_cache_length=gpt.cos.size(-1)` value")
k_shape = v_shape
else:
k_shape = (
batch_size,
self.config.n_query_groups,
max_seq_length,
rope_cache_length + self.config.head_size - self.config.rope_n_elem,
)
return KVCache(k_shape, v_shape, device=device, dtype=dtype)
def _load_from_state_dict(self, state_dict: dict, prefix: str, *args: Any, **kwargs: Any) -> None:
"""For compatibility with legacy checkpoints."""
for attr in ("weight", "bias"):
legacy_key = f"{prefix}attn.{attr}"
current_key = f"{prefix}qkv.{attr}"
if legacy_key in state_dict:
state_dict[current_key] = qkv_reassemble(state_dict.pop(legacy_key), self.config)
super()._load_from_state_dict(state_dict, prefix, *args, **kwargs)
class GptNeoxMLP(nn.Module):
def __init__(self, config: Config) -> None:
super().__init__()
self.fc = nn.Linear(
config.n_embd, config.intermediate_size, bias=config.bias
)
self.proj = nn.Linear(
config.intermediate_size, config.n_embd, bias=config.bias
)
self.config = config
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.fc(x)
x = F.gelu(x, approximate=self.config.gelu_approximate)
return self.proj(x)
class LLaMAMLP(nn.Module):
def __init__(self, config: Config) -> None:
super().__init__()
self.fc_1 = nn.Linear(
config.n_embd, config.intermediate_size, bias=config.bias
)
self.fc_2 = nn.Linear(
config.n_embd, config.intermediate_size, bias=config.bias
)
self.proj = nn.Linear(
config.intermediate_size, config.n_embd, bias=config.bias
)
self.config = config
def forward(self, x: torch.Tensor) -> torch.Tensor:
x_fc_1 = self.fc_1(x)
x_fc_2 = self.fc_2(x)
x = F.silu(x_fc_1) * x_fc_2
return self.proj(x)
class GemmaMLP(LLaMAMLP):
def forward(self, x: torch.Tensor) -> torch.Tensor:
x_fc_1 = self.fc_1(x)
x_fc_2 = self.fc_2(x)
x = F.gelu(x_fc_1, approximate=self.config.gelu_approximate) * x_fc_2
return self.proj(x)
class LLaMAMoE(nn.Module):
def __init__(self, config: Config) -> None:
super().__init__()
self.gate = nn.Linear(config.n_embd, config.n_expert, bias=False)
self.experts = nn.ModuleList(LLaMAMLP(config) for _ in range(config.n_expert))
self.config = config
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Derived from: https://github.com/mistralai/mistral-src/blob/b46d6/moe_one_file_ref.py#L203-L219
See also figure 1 in https://arxiv.org/abs/2211.15841
"""
B, T, C = x.size() # batch size, sequence length, embedding dimensionality (n_embd)
x = x.view(-1, C) # (B*T, C)
router = self.gate(x) # (B*T, n_expert)
probs, indices = torch.topk(router, self.config.n_expert_per_token) # (B*T, n_expert_per_token)
probs = probs.softmax(dim=1, dtype=torch.float).to(dtype=x.dtype)
masks = indices.unsqueeze(-1) == torch.arange(self.config.n_expert, device=x.device)
masks = masks.permute(2, 0, 1) # (n_expert, B*T, n_expert_per_token)
y = torch.zeros_like(x) # (B*T, C)
for mask, expert in zip(masks, self.experts):
token_idx, expert_idx = torch.where(mask)
y[token_idx] += probs[token_idx, expert_idx, None] * expert(x[token_idx])
return y.view(B, T, C)
def build_rope_cache(
seq_len: int,
n_elem: int,
device: Optional[torch.device] = None,
base: int = 10000,
condense_ratio: int = 1,
extra_config: Optional[dict] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Enhanced Transformer with Rotary Position Embedding.
Args:
seq_len (int): Sequence length.
n_elem (int): Number of elements (head dimension).
device (torch.device, optional): Device for tensor allocations.
base (int, optional): Base for computing inverse frequencies.
condense_ratio (int, optional): Ratio to condense the position indices.
extra_config (dict, optional): Configuration parameters for frequency adjustments (used by Llama 3.1 and 3.2)
Returns:
Tuple[torch.Tensor, torch.Tensor]: Cosine and sine caches for RoPE.
Shapes are `(seq_len, n_elem)`.
"""
# Compute the inverse frequencies theta
theta = 1.0 / (base ** (torch.arange(0, n_elem, 2, device=device).float() / n_elem))
if extra_config is not None:
orig_context_len = extra_config["original_max_seq_len"]
factor = extra_config["factor"]
low_freq_factor = extra_config["low_freq_factor"]
high_freq_factor = extra_config["high_freq_factor"]
wavelen = 2 * torch.pi / theta
ratio = orig_context_len / wavelen
smooth_factor = (ratio - low_freq_factor) / (high_freq_factor - low_freq_factor)
smooth_factor = torch.clamp(smooth_factor, min=0.0, max=1.0)
# Compute adjusted_theta without masked indexing
adjusted_theta = (1 - smooth_factor) * (theta / factor) + smooth_factor * theta
theta = adjusted_theta
# Create position indices `[0, 1, ..., seq_len - 1]`
### Zhifei fix bug 1:.
seq_idx = torch.arange(seq_len, device=device, dtype=torch.float16) / float(condense_ratio)
# seq_idx = torch.arange(seq_len, device=device) / condense_ratio
# Calculate the product of position index and $\theta_i$
idx_theta = torch.outer(seq_idx, theta).repeat(1, 2)
# If `n_elem` is odd, the final dimension of `idx_theta` has size
# `n_elem + 1`, so need to cut something off.
# Due to a current bug in Hugging Face, in the case `n_elem == 1`, we leave
# `idx_theta`, `cos`, `sin` as is. Things work out in `apply_rope` due to
# broadcasting. If we shorten `idx_theta`, unit tests comparing to
# Hugging Face fail.
# https://github.com/huggingface/transformers/issues/35233
if idx_theta.shape[-1] > n_elem > 1:
idx_theta = idx_theta[..., :n_elem]
return torch.cos(idx_theta), torch.sin(idx_theta)
def batched_index_select(t, dim, idx):
"""index_select for batched index and unbatched t"""
if idx.dim() == 1:
return torch.index_select(t, dim, idx)
*batch_shape, idx_size = idx.shape
res = torch.index_select(t, dim, idx.reshape(-1)) # flat index
# split out single batch idx
res = res.view(*t.shape[:dim], -1, idx_size, *t.shape[dim + 1 :])
if dim > 0:
# move batch dim to front, this is np.rollaxis(res, dim, 0) for tensors
dims = [dim] + list(range(res.dim()))
del dims[dim + 1]
res = res.permute(dims)
# unflatten batch dims
res = res.view(*batch_shape, *res.shape[1:])
return res
def batched_index_copy_(t, dim, idx, val):
"""Index copy for batched t, idx, val"""
if t.device.type == "mps":
# Normalize negative dimensions
if dim < 0:
dim = t.dim() + dim
if idx.dim() == 1:
idx_shape = [1] * val.dim()
idx_shape[dim] = -1
idx_expanded = idx.view(*idx_shape)
idx_expanded = idx_expanded.expand_as(val)
t.scatter_(dim, idx_expanded, val)
return t
elif idx.dim() == 2:
assert dim != 0, "Cannot index the batch dimension"
batch_size = idx.size(0)
idx_size = idx.size(1)
assert batch_size == t.size(0) == val.size(0)
idx_shape = [batch_size] + [1] * (val.dim() - 1)
idx_shape[dim] = idx_size
idx_expanded = idx.view(*idx_shape)
idx_expanded = idx_expanded.expand_as(val)
t.scatter_(dim, idx_expanded, val)
return t
else:
raise NotImplementedError(f"idx.dim() == {idx.dim()} not supported")
else:
if idx.dim() == 1:
return t.index_copy_(dim, idx, val)
assert idx.dim() == 2, f"multiple batch dims not yet {idx.shape=}"
assert dim != 0, f"cannot index batch dim {dim=}"
batch_size, idx_size = idx.shape
assert batch_size == t.size(0)
assert batch_size == val.size(0)
# if we can view the batch and indexed dimensions together, we could
# do index trickery. This is, sadly, not the case for kvcache so we
# fall back to for loop
for i in range(batch_size):
unbatched_dim = dim if dim < 0 else dim - 1
t[i].index_copy_(unbatched_dim, idx[i], val[i])
return t
def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
"""
Applies RoPE transform to `x`. Note that `cos`, `sin` need to have a batch
dimension.
Args:
x: Input tensor, `(B, ..., T, head_size)`
cos: Cached cosines, `(B, T, head_size)` or `(1, T, head_size)`
sin: Cached sines, `(B, T, head_size)` or `(1, T, head_size)`
Returns:
Encoded tensor, `(B, ..., T, head_size)`
"""
if cos.dim() != 3:
raise ValueError(f"cos must be three-dimensional, but shape is {cos.shape}")
if cos.shape != sin.shape:
raise ValueError(f"cos, sin must have same shape, but cos.shape={cos.shape}, sin.shape={sin.shape}")
head_size_half = x.size(-1) // 2
x1 = x[..., : head_size_half] # (B, ..., T, head_size/2)
x2 = x[..., head_size_half :] # (B, ..., T, head_size/2)
rotated = torch.cat((-x2, x1), dim=-1) # (B, ..., T, head_size)
dims_diff = x.dim() - cos.dim()
if dims_diff > 0:
# Ensure that shapes of `x`, `cos`, `sin` align
new_shape = cos.shape[0:1] + (1,) * dims_diff + cos.shape[1:]
cos = cos.view(*new_shape)
sin = sin.view(*new_shape)
roped = (x * cos) + (rotated * sin)
return roped.to(dtype=x.dtype)
def do_softcapping(x: torch.Tensor, thresh: float) -> torch.Tensor:
return torch.tanh(x / thresh) * thresh
class KVCache(nn.Module):
"""
Buffers `k`, `v` have shape
`(batch_size, n_query_groups, max_seq_length, head_size)`.
"""
def __init__(
self,
k_shape: Tuple[int, int, int, int],
v_shape: Tuple[int, int, int, int],
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
) -> None:
super().__init__()
self.register_buffer("k", torch.zeros(k_shape, device=device, dtype=dtype), persistent=False)
self.register_buffer("v", torch.zeros(v_shape, device=device, dtype=dtype), persistent=False)
def forward(self, input_pos: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Writes new values `k` and `v` into the cache at the positions specified
by `input_pos` along the sequence dimension (`max_seq_length`). The batch
size of `k` and `v` (`bs`) must be smaller or equal to `KVCache` batch
size. Returns the full buffers, adjusted to the batch size `bs`.
Args:
input_pos: Position index, `(bs, T)` or `(T,)`
k: New values, `(bs, n_query_groups, T, head_size)`
v: New values, `(bs, n_query_groups, T, head_size)`
Returns:
k_full, v_full, `(bs, n_query_groups, max_seq_length, head_size)`
"""
# move the buffer to the activation dtype for when AMP is used
self.k = self.k.to(k.dtype)
self.v = self.v.to(v.dtype)
# update the cache
bs = k.size(0)
k = batched_index_copy_(self.k[:bs, ...], -2, input_pos, k)
v = batched_index_copy_(self.v[:bs, ...], -2, input_pos, v)
return k, v
def reset_parameters(self) -> None:
torch.nn.init.zeros_(self.k)
torch.nn.init.zeros_(self.v)
def build_mask_cache(max_seq_length: int, device: Optional[torch.device] = None) -> torch.Tensor:
ones = torch.ones((max_seq_length, max_seq_length), device=device, dtype=torch.bool)
return torch.tril(ones).unsqueeze(0).unsqueeze(0)
class RMSNorm(torch.nn.Module):
"""Root Mean Square Layer Normalization.
Derived from https://github.com/bzhangGo/rmsnorm/blob/master/rmsnorm_torch.py. BSD 3-Clause License:
https://github.com/bzhangGo/rmsnorm/blob/master/LICENSE.
"""
def __init__(self, size: int, dim: int = -1, eps: float = 1e-6, add_unit_offset: bool = False) -> None:
super().__init__()
self.weight = torch.nn.Parameter(torch.ones(size))
self.eps = eps
self.dim = dim
self.add_unit_offset = add_unit_offset
def forward(self, x: torch.Tensor) -> torch.Tensor:
dtype = x.dtype
x = x.float()
# NOTE: the original RMSNorm paper implementation is not equivalent
norm_x = torch.mean(x * x, dim=self.dim, keepdim=True)
x_normed = x * torch.rsqrt(norm_x + self.eps)
weight = (1 + self.weight) if self.add_unit_offset else self.weight
return (x_normed * weight.float()).to(dtype=dtype)
def reset_parameters(self) -> None:
torch.nn.init.ones_(self.weight)