-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmodel.py
More file actions
3308 lines (3029 loc) · 135 KB
/
Copy pathmodel.py
File metadata and controls
3308 lines (3029 loc) · 135 KB
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
974
975
976
977
978
979
980
981
982
983
984
985
986
987
988
989
990
991
992
993
994
995
996
997
998
999
1000
# model.py
import torch
from torch import Tensor
import torch.nn as nn
from torch.nn import functional as F
from dataclasses import dataclass
from functools import partial
import math
from fast_attnres import (
FAST_ATTNRES_DTYPES,
FAST_ATTNRES_MAX_SOURCES,
FAST_ATTNRES_MAX_WIDTH,
FAST_ATTNRES_RELEASE_COMMIT,
FAST_ATTNRES_SOURCE_SHA256,
FAST_ATTNRES_VERSION,
FAST_ATTNRES_WHEEL_SHA256,
FastAttnResDecision,
fast_attnres_config_decision,
fast_attnres_package_provenance,
load_fast_attnres,
)
from attnres_ops import (
attention_residual_phase1_from_logits,
attention_residual_phase2,
attention_residual_phase2_torch,
attention_residual_phase2_from_logit,
attention_residual_read,
lrid_attention_residual_phase2,
lrid_attention_residual_phase2_torch,
lrid_attention_residual_read,
)
def rms_norm_eps(x: torch.Tensor, eps: float = None) -> float:
if eps is not None:
return eps
return torch.finfo(x.dtype).eps
def norm(x: Tensor, eps: float = None):
if hasattr(F, "rms_norm"):
return F.rms_norm(x, (x.size(-1),), eps=eps)
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + rms_norm_eps(x, eps))
def _norm_lrid_key(x: Tensor, num_heads: int):
x_shape = x.shape
x = x.reshape(*x_shape[:-1], num_heads, x_shape[-1] // num_heads)
return norm(x).reshape(*x_shape)
class TokenEmbedding(nn.Module):
def __init__(self, num_embeddings, embedding_dim):
super().__init__()
self.num_embeddings = num_embeddings
self.embedding_dim = embedding_dim
self.weight = nn.Parameter(torch.empty(num_embeddings, embedding_dim))
def forward(self, idx):
return F.embedding(idx, self.weight)
@dataclass
class ModelConfig:
block_size: int = 1024
vocab_size: int = 100277
n_layer: int = 12
n_head: int = 12
n_embd: int = 768
mlp_hidden_dim: int = None
mlp_ratio: float = 4.0
weight_tying: bool = False
rope_theta: float = 500000.0
norm_pos: str = "after"
qk_norm: bool = True
clip_qkv: float = None
flash_attention: bool = False
init_std: float = 0.02
init_cutoff_factor: float = None
attnres_type: str = None
use_attnres: bool = False
use_fused_attnres: bool = False
attnres_num_blocks: int = 8
attnres_block_average: bool = True
attnres_block_average_mode: str = "count"
attnres_block_count_prior: bool = True
attnres_block_alpha: object = "legacy"
attnres_block_beta: object = "legacy"
attnres_block_alpha_learned: bool = False
attnres_block_beta_learned: bool = False
attnres_block_alpha_scope: str = "shared"
attnres_block_beta_scope: str = "shared"
attnres_block_split_sublayers: bool = False
attnres_block_learned_scale: bool = False
attnres_block_learned_scale_init: str = "count"
attnres_block_value_norm: bool = False
attnres_key_norm: bool = True
attn_res_query_norm: bool = False
attn_res_query_init: str = "zero"
attnres_training_cache_phase1: bool = True
attnres_training_torch_phase2: bool = True
attnres_fuse_read_norm: bool = True
use_lrid: bool = False
lrid_rank: int = 64
lrid_projection_rank: int = None
lrid_num_heads: int = 1
lrid_input_dependent_query: bool = False
lrid_static_embedding_key: bool = False
lrid_add_static_embedding_key: bool = False
lrid_add_static_source_key: bool = False
lrid_key_from_value: bool = False
lrid_key_from_value_shared: bool = False
lrid_key_from_output_tail: bool = False
lrid_key_value_norm: bool = True
lrid_query_from_value: bool = False
lrid_query_from_value_shared: bool = False
lrid_use_logit_scale: bool = True
lrid_logit_scale: float = None
# ``auto`` selects Fast for standard (R=D) and sliced output-tail LRID
# (R<=D), while leaving projected/experimental LRID on the legacy path.
# Keep the dataclass default legacy for old direct-library callers and
# checkpoints; train.py opts into auto by default.
attnres_backend: str = "legacy"
def __post_init__(self):
requested_attnres_backend = str(self.attnres_backend or "legacy").lower()
if requested_attnres_backend not in {"auto", "legacy", "fast"}:
raise ValueError("attnres_backend must be one of: auto, legacy, fast")
self.attnres_type = (self.attnres_type or "block")
self.attnres_type = self.attnres_type.lower()
self.attnres_block_average_mode = (self.attnres_block_average_mode or "count").lower()
if self.attnres_block_average_mode not in {"count", "sqrt"}:
raise ValueError("attnres_block_average_mode must be one of: count, sqrt")
self.attnres_block_alpha_scope = self._normalize_block_power_scope(
self.attnres_block_alpha_scope,
"attnres_block_alpha_scope",
)
self.attnres_block_beta_scope = self._normalize_block_power_scope(
self.attnres_block_beta_scope,
"attnres_block_beta_scope",
)
self.attnres_block_alpha = self._normalize_block_power_value(
self.attnres_block_alpha,
"attnres_block_alpha",
)
self.attnres_block_beta = self._normalize_block_power_value(
self.attnres_block_beta,
"attnres_block_beta",
)
self._validate_block_power_value_length(
self.attnres_block_alpha,
self.attnres_block_alpha_scope,
"attnres_block_alpha",
)
self._validate_block_power_value_length(
self.attnres_block_beta,
self.attnres_block_beta_scope,
"attnres_block_beta",
)
self.attnres_block_learned_scale_init = self._normalize_block_scale_init(
self.attnres_block_learned_scale_init
)
if self.attnres_block_learned_scale and self.attnres_type != "block":
raise ValueError("attnres_block_learned_scale requires attnres_type='block'")
if self.attnres_block_value_norm and self.attnres_type != "block":
raise ValueError("attnres_block_value_norm requires attnres_type='block'")
if self.attnres_block_value_norm and self.attnres_block_learned_scale:
raise ValueError("attnres_block_value_norm and attnres_block_learned_scale are mutually exclusive")
explicit_alpha_formula = (
self.attnres_block_alpha != "legacy"
or self.attnres_block_alpha_learned
)
if explicit_alpha_formula and (self.attnres_block_learned_scale or self.attnres_block_value_norm):
raise ValueError(
"attnres_block_alpha formula mode is mutually exclusive with "
"attnres_block_learned_scale and attnres_block_value_norm"
)
if (
self.attnres_type == "block"
and self.attnres_block_count_prior
and (
self.attnres_block_learned_scale
or self.attnres_block_value_norm
)
):
raise ValueError(
"attnres_block_count_prior requires alpha-formula block summaries: "
"attnres_block_learned_scale=False, and attnres_block_value_norm=False"
)
self.attn_res_query_init = (self.attn_res_query_init or "zero").lower()
if self.attn_res_query_init not in {"zero", "normal", "trunc_normal"}:
raise ValueError("attn_res_query_init must be one of: zero, normal, trunc_normal")
if self.lrid_static_embedding_key and self.lrid_add_static_embedding_key:
raise ValueError("lrid_static_embedding_key and lrid_add_static_embedding_key are mutually exclusive")
if self.lrid_key_from_value_shared:
self.lrid_key_from_value = True
if self.lrid_key_from_output_tail:
if self.lrid_key_from_value:
raise ValueError("lrid_key_from_output_tail and lrid_key_from_value are mutually exclusive")
if self.lrid_static_embedding_key or self.lrid_add_static_embedding_key or self.lrid_add_static_source_key:
raise ValueError(
"lrid_key_from_output_tail cannot be combined with static LRID key additions"
)
if self.lrid_query_from_value_shared:
self.lrid_query_from_value = True
if self.lrid_rank < 1:
raise ValueError("lrid_rank must be >= 1")
if self.lrid_key_from_output_tail and self.lrid_rank > self.n_embd:
raise ValueError("lrid_rank must be <= n_embd when lrid_key_from_output_tail=True")
if self.lrid_projection_rank is None:
self.lrid_projection_rank = self.lrid_rank
if self.lrid_projection_rank < self.lrid_rank:
raise ValueError("lrid_projection_rank must be >= lrid_rank")
if self.lrid_num_heads < 1:
raise ValueError("lrid_num_heads must be >= 1")
if self.lrid_rank % self.lrid_num_heads != 0:
raise ValueError("lrid_rank must be divisible by lrid_num_heads")
if self.n_embd % self.lrid_num_heads != 0:
raise ValueError("n_embd must be divisible by lrid_num_heads")
if not self.lrid_use_logit_scale:
self.lrid_logit_scale = 1.0
elif self.lrid_logit_scale is None:
self.lrid_logit_scale = 1.0 / math.sqrt(self.lrid_rank // self.lrid_num_heads)
elif self.lrid_logit_scale <= 0.0:
raise ValueError("lrid_logit_scale must be positive")
if self.use_lrid:
self.use_attnres = True
self._attnres_backend_requested = requested_attnres_backend
if requested_attnres_backend == "auto":
self.attnres_backend = (
"fast"
if self.use_attnres and (not self.use_lrid or self.lrid_key_from_output_tail)
else "legacy"
)
else:
self.attnres_backend = requested_attnres_backend
@staticmethod
def _normalize_block_scale_init(value):
value = str(value or "count").lower().replace(" ", "")
if value in {"count", "1/c", "inv_count", "inverse_count"}:
return "count"
if value in {"sqrt", "1/sqrtc", "1/sqrt(c)", "inv_sqrt", "inverse_sqrt"}:
return "sqrt"
if value in {"one", "1"}:
return "one"
raise ValueError(
"attnres_block_learned_scale_init must be one of: count, sqrt, one"
)
@staticmethod
def _normalize_block_power_scope(value, name):
value = str(value or "shared").lower().replace("-", "_")
aliases = {
"all": "shared",
"global": "shared",
"layer": "per_residual",
"layers": "per_residual",
"per_layer": "per_residual",
"residual": "per_residual",
"block": "per_block",
"blocks": "per_block",
}
value = aliases.get(value, value)
if value not in {"shared", "per_residual", "per_block"}:
raise ValueError(f"{name} must be one of: shared, per_residual, per_block")
return value
@staticmethod
def _normalize_block_power_value(value, name):
if isinstance(value, str):
text = value.strip().lower()
if text == "legacy":
return "legacy"
parts = [part.strip() for part in text.split(",")]
try:
if len(parts) == 1:
return float(parts[0])
return [float(part) for part in parts]
except ValueError as exc:
raise ValueError(f"{name} must be 'legacy', a float, or comma-separated floats") from exc
if isinstance(value, (int, float)):
return float(value)
try:
return [float(part) for part in value]
except TypeError as exc:
raise ValueError(f"{name} must be 'legacy', a float, or comma-separated floats") from exc
def _block_power_scope_length(self, scope):
if scope == "shared":
return 1
if scope == "per_residual":
return 2 * self.n_layer
return int(self.attnres_num_blocks)
def _validate_block_power_value_length(self, value, scope, name):
if value == "legacy" or isinstance(value, float):
return
expected = self._block_power_scope_length(scope)
if len(value) != expected:
raise ValueError(f"{name} list length must be {expected} for {scope} scope")
class LRIDStaticKey(nn.Module):
def __init__(self, config):
super().__init__()
self.key = nn.Parameter(torch.empty(config.lrid_rank))
def reset_parameters(self, std=0.02, init_cutoff_factor=None):
if init_cutoff_factor is not None:
cutoff = init_cutoff_factor * std
nn.init.trunc_normal_(self.key, mean=0.0, std=std, a=-cutoff, b=cutoff)
else:
nn.init.normal_(self.key, mean=0.0, std=std)
def forward(self, reference):
return self.key.to(reference.dtype).view(1, 1, -1).expand(reference.size(0), reference.size(1), -1)
class RotaryEmbedding(nn.Module):
def __init__(self, config):
super().__init__()
dim = config.n_embd // config.n_head
max_seq_len = config.block_size
base = config.rope_theta
inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2) / dim))
freq = torch.outer(torch.arange(max_seq_len), inv_freq)
self.register_buffer("sin", freq.sin()[None, None])
self.register_buffer("cos", freq.cos()[None, None])
def _forward_single(self, x, offset=0):
T = x.size(-2)
sin = self.sin[:, :, offset:offset + T]
cos = self.cos[:, :, offset:offset + T]
x1, x2 = x[..., 0::2], x[..., 1::2]
return torch.stack([cos * x1 - sin * x2, sin * x1 + cos * x2], dim=-1).flatten(-2)
def forward(self, q, k=None, offset=0):
if k is None:
return self._forward_single(q, offset=offset)
return self._forward_single(q, offset=offset), self._forward_single(k, offset=offset)
class LRIDFusedProjection(nn.Module):
def __init__(self, config, input_dim, output_dim):
super().__init__()
self.output_dim = output_dim
self.rank = config.lrid_rank
self.projection_rank = config.lrid_projection_rank
self.use_query = config.lrid_input_dependent_query
self.use_output_tail_key = config.lrid_key_from_output_tail
self.use_value_key = config.lrid_key_from_value
self.use_shared_value_key = config.lrid_key_from_value_shared
self.use_local_value_key = self.use_value_key and not self.use_shared_value_key
self.use_key = not (self.use_value_key or self.use_output_tail_key)
self.use_value_query = self.use_query and config.lrid_query_from_value
self.use_shared_value_query = self.use_query and config.lrid_query_from_value_shared
self.use_local_value_query = self.use_value_query and not self.use_shared_value_query
self.use_fused_query = self.use_query and not self.use_value_query
self.key_offset = output_dim
self.query_offset = output_dim + (self.projection_rank if self.use_key else 0)
extra_dim = (self.projection_rank if self.use_key else 0) + (self.rank if self.use_fused_query else 0)
self.proj = nn.Linear(input_dim, output_dim + extra_dim, bias=False)
self.use_key_norm = self.use_key and config.attnres_key_norm
self.use_value_norm = (self.use_local_value_key or self.use_local_value_query) and config.lrid_key_value_norm
self.num_heads = config.lrid_num_heads
if self.use_local_value_key:
self.value_key_proj = nn.Linear(output_dim, config.lrid_rank, bias=False)
else:
self.value_key_proj = None
if self.use_local_value_query:
self.value_query_proj = nn.Linear(output_dim, config.lrid_rank, bias=False)
else:
self.value_query_proj = None
def _prepare_value_projection_input(self, value):
if self.use_value_norm:
return norm(value)
return value
def project_key_from_value(self, value):
if not self.use_local_value_key:
raise RuntimeError("Local value-key projection is only available in unshared lrid_key_from_value mode")
return self.value_key_proj(self._prepare_value_projection_input(value))
def project_query_from_value(self, value):
if not self.use_local_value_query:
raise RuntimeError("Local value-query projection is only available in unshared lrid_query_from_value mode")
return self.value_query_proj(self._prepare_value_projection_input(value))
def forward(self, x, emit_lrid_key=True):
if not emit_lrid_key:
output = F.linear(x, self.proj.weight[:self.output_dim], None)
if self.use_query:
return output, None, None
return output, None
projected = self.proj(x)
output = projected[..., :self.output_dim].contiguous()
key = None
query = None
if self.use_key:
key = projected[..., self.key_offset:self.key_offset + self.rank]
if self.use_key_norm:
key = _norm_lrid_key(key, self.num_heads)
elif self.use_output_tail_key:
key = output[..., -self.rank:].contiguous()
elif self.use_local_value_key:
key = self.project_key_from_value(output)
if self.use_fused_query:
query = projected[..., self.query_offset:self.query_offset + self.rank]
elif self.use_local_value_query:
query = self.project_query_from_value(output)
if self.use_query:
return output, key, query
return output, key
class LRIDSourceKeyProjection(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
self.rank = config.lrid_rank
self.num_heads = config.lrid_num_heads
self.use_output_tail_key = config.lrid_key_from_output_tail
uses_value_projection = (
config.lrid_key_from_value
or (config.lrid_input_dependent_query and config.lrid_query_from_value_shared)
)
self.use_value_norm = uses_value_projection and config.lrid_key_value_norm
self.proj = None if self.use_output_tail_key else nn.Linear(config.n_embd, config.lrid_rank, bias=False)
if config.lrid_input_dependent_query and config.lrid_query_from_value_shared:
self.query_proj = nn.Linear(config.n_embd, config.lrid_rank, bias=False)
else:
self.query_proj = None
self.use_key_norm = config.attnres_key_norm and not config.lrid_key_from_value
def _prepare_value_projection_input(self, x):
if self.use_value_norm:
return norm(x)
return x
def forward(self, x):
if self.use_output_tail_key:
return x[..., -self.rank:].contiguous()
if self.config.lrid_key_from_value:
x = self._prepare_value_projection_input(x)
key = self.proj(x)
if self.use_key_norm:
key = _norm_lrid_key(key, self.num_heads)
return key
def project_query_from_value(self, x):
if self.query_proj is None:
raise RuntimeError("Shared value-query projection is only available when lrid_query_from_value_shared=True")
return self.query_proj(self._prepare_value_projection_input(x))
class MultiHeadAttention(nn.Module):
flash_attn_func = None
flash_attn_varlen_func = None
flash_tried = False
def __init__(self, config, layer_idx=0):
super().__init__()
assert config.n_embd % config.n_head == 0
self.n_head = config.n_head
self.n_embd = config.n_embd
self.head_dim = config.n_embd // config.n_head
self.rope = RotaryEmbedding(config)
self.layer_idx = layer_idx
self.config = config
self.c_attn = nn.Linear(config.n_embd, 3 * config.n_embd, bias=False)
if config.use_lrid:
self.c_proj = LRIDFusedProjection(config, config.n_embd, config.n_embd)
else:
self.c_proj = nn.Linear(config.n_embd, config.n_embd, bias=False)
self.use_qk_norm = config.qk_norm
self.clip_qkv = config.clip_qkv
if config.flash_attention and not MultiHeadAttention.flash_tried:
try:
from flash_attn import flash_attn_func, flash_attn_varlen_func
MultiHeadAttention.flash_attn_func = flash_attn_func
MultiHeadAttention.flash_attn_varlen_func = flash_attn_varlen_func
MultiHeadAttention.flash_tried = True
except Exception as e:
print(f"Error with flash-attn {e}.")
MultiHeadAttention.flash_tried = True
def _scaled_dot_product_attention(self, q, k, v, attn_mask=None, is_causal=True,
cu_doc_len=None, max_doc_len=None):
B, H, T, D = q.size()
if cu_doc_len is not None and max_doc_len is not None and MultiHeadAttention.flash_attn_varlen_func is not None:
q_flat = q.transpose(1, 2).reshape(B * T, H, D)
k_flat = k.transpose(1, 2).reshape(B * T, H, D)
v_flat = v.transpose(1, 2).reshape(B * T, H, D)
cu_doc_len = cu_doc_len.to(device=q.device, dtype=torch.int32)
x = MultiHeadAttention.flash_attn_varlen_func(
q_flat, k_flat, v_flat,
cu_seqlens_q=cu_doc_len,
cu_seqlens_k=cu_doc_len,
max_seqlen_q=max_doc_len,
max_seqlen_k=max_doc_len,
causal=is_causal,
)
return x.view(B, T, H, D).contiguous().view(B, T, self.n_embd)
elif cu_doc_len is not None or max_doc_len is not None:
raise RuntimeError(
"Document masking requires flash-attn varlen support. "
"Install flash-attn or disable use_doc_masking."
)
elif MultiHeadAttention.flash_attn_func is not None and attn_mask is None:
x = MultiHeadAttention.flash_attn_func(
q.transpose(1, 2),
k.transpose(1, 2),
v.transpose(1, 2),
causal=is_causal,
)
return x.contiguous().view(B, T, self.n_embd)
else:
x = F.scaled_dot_product_attention(
q, k, v,
attn_mask=attn_mask,
is_causal=is_causal,
)
return x.transpose(1, 2).contiguous().view(B, T, self.n_embd)
def forward(self, x, past_kv=None, use_cache=False, cu_doc_len=None, max_doc_len=None, emit_lrid_key=True):
B, T, C = x.size()
q, k, v = self.c_attn(x).split(self.n_embd, dim=2)
if self.clip_qkv is not None:
q.clamp_(min=-self.clip_qkv, max=self.clip_qkv)
k.clamp_(min=-self.clip_qkv, max=self.clip_qkv)
v.clamp_(min=-self.clip_qkv, max=self.clip_qkv)
q = q.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
k = k.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
v = v.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
if self.use_qk_norm:
q = norm(q)
k = norm(k)
if past_kv is not None:
past_k, past_v = past_kv
pos_offset = past_k.size(-2)
else:
pos_offset = 0
q, k = self.rope(q, k, offset=pos_offset)
if past_kv is not None:
k = torch.cat([past_k, k], dim=2)
v = torch.cat([past_v, v], dim=2)
is_causal = past_kv is None
attention_output = self._scaled_dot_product_attention(
q, k, v,
is_causal=is_causal,
cu_doc_len=cu_doc_len,
max_doc_len=max_doc_len,
)
if self.config.use_lrid:
projected = self.c_proj(attention_output, emit_lrid_key=emit_lrid_key)
else:
projected = self.c_proj(attention_output)
if self.config.use_lrid:
if self.config.lrid_input_dependent_query:
x, lrid_key, lrid_query = projected
else:
x, lrid_key = projected
else:
x = projected
if use_cache:
if self.config.use_lrid:
if self.config.lrid_input_dependent_query:
return x, (k, v), lrid_key, lrid_query
return x, (k, v), lrid_key
return x, (k, v)
if self.config.use_lrid:
if self.config.lrid_input_dependent_query:
return x, lrid_key, lrid_query
return x, lrid_key
return x
class MLP(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
self.hidden_dim = config.mlp_hidden_dim if config.mlp_hidden_dim is not None else int(config.n_embd * config.mlp_ratio)
self.fc1 = nn.Linear(config.n_embd, self.hidden_dim * 2, bias=False)
if config.use_lrid:
self.fc2 = LRIDFusedProjection(config, self.hidden_dim, config.n_embd)
else:
self.fc2 = nn.Linear(self.hidden_dim, config.n_embd, bias=False)
def forward(self, x, emit_lrid_key=True):
x = self.fc1(x)
x, gate = x.chunk(2, dim=-1)
x = F.silu(gate) * x
if self.config.use_lrid:
return self.fc2(x, emit_lrid_key=emit_lrid_key)
return self.fc2(x)
class AttentionResidual(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
self.use_key_norm = config.attnres_key_norm
self.query = nn.Parameter(torch.empty(config.n_embd))
def _query(self, dtype):
query = self.query
if self.config.attn_res_query_norm:
query = norm(query.float())
return query.to(dtype)
def forward(self, values, source_counts=None, source_logit_biases=None):
keys = norm(values) if self.use_key_norm else values
logits = torch.einsum("d,sbtd->sbt", self._query(keys.dtype), keys)
if source_counts is not None and source_logit_biases is not None:
raise RuntimeError("source_counts and source_logit_biases are mutually exclusive")
if source_counts is not None:
log_counts = torch.as_tensor(source_counts, device=logits.device, dtype=torch.float32).log()
logits = logits + log_counts.view(-1, 1, 1)
if source_logit_biases is not None:
bias_values = [
bias.to(device=logits.device, dtype=torch.float32).reshape(())
if torch.is_tensor(bias)
else torch.tensor(float(bias), device=logits.device, dtype=torch.float32)
for bias in source_logit_biases
]
logit_bias = torch.stack(bias_values)
logits = logits + logit_bias.view(-1, 1, 1)
weights = F.softmax(logits.float(), dim=0).to(values.dtype)
return torch.einsum("sbt,sbtd->btd", weights, values)
class Block(nn.Module):
def __init__(self, config, layer_idx=0):
super().__init__()
self.norm_pos = config.norm_pos
self.attn = MultiHeadAttention(config, layer_idx=layer_idx)
self.mlp = MLP(config)
self.layer_idx = layer_idx
self.config = config
def forward_attention(self, x, past_kv=None, use_cache=False, cu_doc_len=None, max_doc_len=None, x_is_normalized=False, emit_lrid_key=True):
if self.norm_pos in {"before", "both"} and not x_is_normalized:
x = norm(x)
attn_out = self.attn(
x,
past_kv=past_kv,
use_cache=use_cache,
cu_doc_len=cu_doc_len,
max_doc_len=max_doc_len,
emit_lrid_key=emit_lrid_key,
)
if self.config.use_lrid:
if self.config.lrid_input_dependent_query:
if use_cache:
x, new_kv, lrid_key, lrid_query = attn_out
else:
x, lrid_key, lrid_query = attn_out
new_kv = None
else:
if use_cache:
x, new_kv, lrid_key = attn_out
else:
x, lrid_key = attn_out
new_kv = None
elif use_cache:
x, new_kv = attn_out
else:
x = attn_out
new_kv = None
if self.norm_pos in {"after", "both"}:
x = norm(x)
if use_cache:
if self.config.use_lrid:
if self.config.lrid_input_dependent_query:
return x, new_kv, lrid_key, lrid_query
return x, new_kv, lrid_key
return x, new_kv
if self.config.use_lrid:
if self.config.lrid_input_dependent_query:
return x, lrid_key, lrid_query
return x, lrid_key
return x
def forward_mlp(self, x, x_is_normalized=False, emit_lrid_key=True):
if self.norm_pos in {"before", "both"} and not x_is_normalized:
x = norm(x)
mlp_out = self.mlp(x, emit_lrid_key=emit_lrid_key)
if self.config.use_lrid:
if self.config.lrid_input_dependent_query:
x, lrid_key, lrid_query = mlp_out
else:
x, lrid_key = mlp_out
else:
x = mlp_out
if self.norm_pos in {"after", "both"}:
x = norm(x)
if self.config.use_lrid:
if self.config.lrid_input_dependent_query:
return x, lrid_key, lrid_query
return x, lrid_key
return x
def forward(self, x, past_kv=None, use_cache=False, cu_doc_len=None, max_doc_len=None):
residual = x
attn_out = self.forward_attention(
x,
past_kv=past_kv,
use_cache=use_cache,
cu_doc_len=cu_doc_len,
max_doc_len=max_doc_len,
)
if use_cache:
x, new_kv = attn_out
else:
x = attn_out
new_kv = None
x = residual + x
residual = x
x = self.forward_mlp(x)
x = residual + x
if use_cache:
return x, new_kv
else:
return x
class OBPM(nn.Module):
def __init__(self, config: ModelConfig):
super().__init__()
self.config = config
self.use_attnres = config.use_attnres
self.use_fused_attnres = config.use_fused_attnres
self.attnres_type = config.attnres_type
self.use_lrid = config.use_lrid
self.attnres_block_ends = self._make_attnres_block_ends()
torch.backends.cuda.enable_flash_sdp(True)
torch.backends.cuda.enable_mem_efficient_sdp(False)
transformer_modules = dict(
wte=TokenEmbedding(config.vocab_size, config.n_embd),
layers=nn.ModuleList([Block(config, layer_idx=i) for i in range(config.n_layer)])
)
if self.use_attnres:
if self.attnres_type not in {"full", "block"}:
raise ValueError("attnres_type must be 'full' or 'block'")
if self.attnres_type == "block" and config.attnres_num_blocks < 1:
raise ValueError("attnres_num_blocks must be >= 1 when using block AttnRes")
if self.use_lrid:
needs_lrid_source_projection = (
(not config.lrid_static_embedding_key and not config.lrid_key_from_output_tail)
or config.lrid_key_from_value_shared
or (config.lrid_input_dependent_query and config.lrid_query_from_value_shared)
)
if needs_lrid_source_projection:
transformer_modules["lrid_embedding_key"] = LRIDSourceKeyProjection(config)
if config.lrid_static_embedding_key or config.lrid_add_static_embedding_key:
transformer_modules["lrid_static_embedding_key"] = LRIDStaticKey(config)
if config.lrid_add_static_source_key:
transformer_modules["lrid_static_source_key"] = LRIDStaticKey(config)
key_head_dim = config.lrid_rank // config.lrid_num_heads
transformer_modules["lrid_queries"] = nn.ParameterList(
[
nn.Parameter(torch.empty(config.lrid_num_heads, key_head_dim))
for _ in range(2 * config.n_layer)
]
)
if config.lrid_input_dependent_query:
transformer_modules["lrid_query_gates"] = nn.ParameterList(
[
nn.Parameter(torch.zeros(config.lrid_num_heads))
for _ in range(2 * config.n_layer)
]
)
else:
transformer_modules["attn_residuals"] = nn.ModuleList(
[AttentionResidual(config) for _ in range(2 * config.n_layer)]
)
self.transformer = nn.ModuleDict(transformer_modules)
if self.use_attnres and self.attnres_type == "block" and config.attnres_block_learned_scale:
self.transformer.register_parameter(
"attnres_block_scales",
nn.Parameter(self._make_attnres_block_scale_init()),
)
if self.use_attnres and self.attnres_type == "block" and config.attnres_block_alpha_learned:
self.transformer.register_parameter(
"attnres_block_alphas",
nn.Parameter(self._make_attnres_block_power_init("alpha")),
)
if self.use_attnres and self.attnres_type == "block" and config.attnres_block_beta_learned:
self.transformer.register_parameter(
"attnres_block_betas",
nn.Parameter(self._make_attnres_block_power_init("beta")),
)
if not config.weight_tying:
self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)
self.apply(partial(self._init_weights, std=config.init_std, init_cutoff_factor=config.init_cutoff_factor))
if self.use_lrid:
for query in self.transformer.lrid_queries:
self._init_attnres_query(query, config.init_std, config.init_cutoff_factor)
if self.use_lrid and config.lrid_input_dependent_query:
for module in self.modules():
if isinstance(module, LRIDFusedProjection):
self._init_lrid_dynamic_query_projection(module, config.init_std, config.init_cutoff_factor)
self._fast_attnres_op = None
# Fast-AttnRes v2 is CUDA/BF16-only. Construction happens on CPU, so
# the production route is armed only after the caller moves the full
# model to its final CUDA/BF16 runtime and calls require_fast_attnres().
self._fast_attnres_enabled = False
def to_mixed_precision(self, dtype=torch.bfloat16):
# These learned block exponents intentionally remain fp32. Preserve
# their exact values instead of casting fp32 -> bf16 -> fp32 along with
# the rest of the module, which would silently perturb checkpoints on
# resume/evaluation.
preserved_fp32 = {}
for name in ("attnres_block_alphas", "attnres_block_betas"):
if hasattr(self.transformer, name):
preserved_fp32[name] = getattr(self.transformer, name).detach().clone()
self.to(dtype=dtype)
for name, value in preserved_fp32.items():
param = getattr(self.transformer, name)
param.data = value.to(device=param.device, dtype=torch.float32)
return self
def get_num_params(self):
return sum(p.numel() for p in self.parameters())
def fast_attnres_startup_report(self, validate_package: bool = False):
"""Describe the statically selected Fast route without mutable counters."""
decision = (
self._fast_attnres_config_decision()
if self.config.attnres_backend == "fast"
else FastAttnResDecision("legacy", "backend_legacy", "attnres_backend is legacy")
)
total_reads = 2 * self.config.n_layer if self.use_attnres else 0
parameter = next(self.parameters(), None)
operator_width = self.config.n_embd + (
self.config.lrid_rank if self._fast_lrid_requires_key_payload() else 0
)
max_sources = (
2 * self.config.n_layer + 1
if self.attnres_type == "full"
else (
min(2 * self.config.n_layer, self.config.attnres_num_blocks)
+ (2 if self.config.attnres_block_split_sublayers else 1)
)
)
if decision.eligible and (
parameter is None
or parameter.dtype not in FAST_ATTNRES_DTYPES
or parameter.device.type != "cuda"
or operator_width > FAST_ATTNRES_MAX_WIDTH
or max_sources > FAST_ATTNRES_MAX_SOURCES
):
decision = FastAttnResDecision(
"legacy",
"unsupported_runtime_envelope",
"model must be CUDA BF16 and inside the Fast-AttnRes v2.0.1 shape envelope",
)
if decision.eligible and self.attnres_type == "block":
if self._use_attnres_block_count_prior():
decision = FastAttnResDecision(
"legacy",
"source_count_prior",
"Fast-AttnRes v2.0.1 has no source-logit-prior API",
)
active_reads = total_reads if decision.eligible else 0
package = {"version": None, "source_hashes": {}}
if validate_package and active_reads:
self._fast_attnres_op = load_fast_attnres()
package = fast_attnres_package_provenance()
reasons = (
{}
if active_reads
else ({decision.reason or "unspecified": total_reads} if total_reads else {})
)
return {
**decision.as_dict(),
"backend": self.config.attnres_backend,
"requested_backend": getattr(
self.config, "_attnres_backend_requested", self.config.attnres_backend
),
"resolved_backend": "fast-attnres" if active_reads else "legacy",
"active_reads": active_reads,
"total_reads": total_reads,
"legacy_fallback_reads": total_reads - active_reads,
"fast_reads": active_reads,
"legacy_reads": total_reads - active_reads,
"fallback_reasons": reasons,
"value_width": self.config.n_embd,
"operator_width": operator_width,
"uses_exact_key_payload": self._fast_lrid_requires_key_payload(),
**package,
}
def require_fast_attnres(self, validate_package: bool = True):
"""Arm Fast only if every multi-source routed read is guaranteed Fast."""
report = self.fast_attnres_startup_report(validate_package=validate_package)
total_reads = int(report["total_reads"])
failures = []
if self.config.attnres_backend != "fast":
failures.append("backend did not resolve to fast")
if total_reads <= 0:
failures.append("model has no routed reads")
if int(report["active_reads"]) != total_reads:
failures.append("not every routed read is active")
if int(report["fast_reads"]) != total_reads:
failures.append("not every routed read resolves to Fast-AttnRes")
if int(report["legacy_reads"]) or int(report["legacy_fallback_reads"]):
failures.append("a legacy fallback remains")
if report["resolved_backend"] != "fast-attnres":
failures.append("resolved backend is not Fast-AttnRes")
if validate_package:
if report.get("version") != FAST_ATTNRES_VERSION:
failures.append(f"Fast-AttnRes {FAST_ATTNRES_VERSION} provenance is missing")
if report.get("release_commit") != FAST_ATTNRES_RELEASE_COMMIT:
failures.append("Fast-AttnRes release commit does not match the qualified release")
if report.get("wheel_sha256") != FAST_ATTNRES_WHEEL_SHA256:
failures.append("Fast-AttnRes wheel hash does not match the qualified release")
if report.get("distribution_source_sha256") != FAST_ATTNRES_SOURCE_SHA256:
failures.append("installed Fast-AttnRes source hash is not the qualified release")
if not report.get("source_hashes"):
failures.append("installed Fast-AttnRes source provenance is incomplete")
if failures:
self._fast_attnres_enabled = False
reason = report.get("reason") or "unknown"
detail = report.get("detail") or ""
raise RuntimeError(
"Fast-AttnRes-only contract failed: "
+ "; ".join(failures)
+ f" (reason={reason}: {detail})"
)
self._fast_attnres_enabled = True
return report
def _assert_fast_runtime_read(
self,
sources,
*,
source_counts=None,
source_logit_biases=None,
):
if not self._fast_attnres_enabled:
raise RuntimeError(
"Fast-AttnRes model is not armed; call require_fast_attnres() "
"after moving it to CUDA BF16"
)
if source_counts is not None or source_logit_biases is not None:
raise RuntimeError("Fast-AttnRes cannot use source priors or logit biases")
first = sources[0]
if first.device.type != "cuda" or first.dtype not in FAST_ATTNRES_DTYPES:
raise RuntimeError("Fast-AttnRes residual reads require CUDA BF16 sources")
def fast_attnres_route_report(self):
report = self.fast_attnres_startup_report(validate_package=False)
report["will_execute"] = bool(report["active_reads"])