File size: 43,038 Bytes
7b592f7 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 897 898 899 900 901 902 903 904 905 906 907 908 909 910 911 912 913 914 915 916 917 918 919 920 921 922 923 924 925 926 927 928 929 930 931 932 933 934 935 936 937 938 939 940 941 942 943 944 945 946 947 948 949 950 951 952 953 954 955 956 957 958 959 960 961 962 963 964 965 966 967 968 969 970 971 972 973 | # 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)
|