File size: 45,538 Bytes
6150318
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Vortex modeling β€” Hugging Face `PreTrainedModel` implementation.

Self-contained on purpose. With `trust_remote_code=True`, `transformers` copies
`configuration_vortex.py` and `modeling_vortex.py` into
`~/.cache/huggingface/modules/transformers_modules/<repo>/` and imports them as
a package, so nothing here may import a sibling file from this repository.
`configuration_vortex` is the only dependency and it travels with this module, so
the pair is always copied together β€” see `_load_config_class` for why the import
is written the way it is.

What this adds over a bare `nn.Module` port, and why each piece is needed for
`AutoModelForCausalLM` / `generate` / `Trainer` to work:

* **Key/value cache.** `use_cache` was a config field with no implementation β€”
  every `forward` recomputed the whole prefix. `VortexAttention` now consumes a
  `transformers` `Cache`, which is what makes `model.generate()` viable.
* **Position offsets under a cache.** RoPE was sliced `cos[:T]`, i.e. positions
  were always 0-based. With a cache the query block starts at `past_len`; the
  rotary tables are now sliced `[offset : offset + T]`. RoPE is relative, so this
  leaves the pretraining fast path bit-identical.
* **A correct attention mask on the cached path.** SDPA's `is_causal=True`
  assumes top-left alignment and is only right when the cache is empty. Cached
  steps with left padding need an explicit bottom-right-aligned mask, which is
  what `VortexModel._build_causal_mask` builds. The empty-cache/no-padding case
  still takes the `is_causal=True` fast path, so training numerics and memory are
  unchanged.
* **Real `ModelOutput`s.** The previous `CausalLMOutput` was a plain object, so
  `output.logits` worked but nothing HF-side (generation, `Trainer`, tensor
  logging) recognised it.
* **Standard input plumbing** β€” `attention_mask`, `position_ids`,
  `inputs_embeds`, `num_items_in_batch`, `logits_to_keep`.

Two deliberate deviations from HF naming conventions:

* The decoder submodules keep their original names (`attn`, `ln_attn`, `ln_mlp`)
  rather than `self_attn`, `input_layernorm`, `post_attention_layernorm`. HF's
  `self_attn` means *cross*-attention, which this architecture does not have.
  More importantly, the released checkpoints on the Hub use the current names,
  and `from_pretrained` matches `state_dict` keys literally β€” renaming would
  break every one of them unless a key-remapping table were threaded through
  `from_pretrained`, which is a per-version API in transformers 5.x. The outer
  names (`model.*`, `embed_tokens`, `lm_head`, `norm`) already match HF.
* `logits_to_keep` is not decoration. Computing `(B, T, vocab_size)` logits for a
  full 2048-token batch is the largest single memory term in a training step, and
  the whole point of the chunked loss path is to never materialise it. That is
  why `labels=` returns `logits=None` unless logits are explicitly asked for.
"""

from __future__ import annotations

import importlib.util
import math
import os
import sys
from typing import Optional, Tuple, Union

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
from transformers.activations import ACT2FN
from transformers.cache_utils import Cache, DynamicCache
from transformers.generation import GenerationMixin
from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast
from transformers.modeling_utils import PreTrainedModel
from transformers.utils import logging


def _load_config_class():
    """Import the sibling `configuration_vortex` module.

    Two different import mechanics have to be satisfied:

    * Under `trust_remote_code`, `transformers` copies both files into its
      dynamic-module package and imports them as a package, so a *relative*
      import is the one that resolves.
    * Running the repo's own tests imports this file as a top-level module from
      `src/`, where there is no package and no `__package__`.

    A plain top-level `from configuration_vortex import ...` is not an option:
    `dynamic_module_utils.check_imports` runs `importlib.import_module` on every
    statically-detected import *before* the sibling has been copied next to this
    file, so it fails with "No module named 'configuration_vortex'" and a
    misleading `pip install configuration_vortex`. Loading by file path keeps the
    statement out of the AST the checker inspects.
    """
    if __package__:
        from .configuration_vortex import VortexConfig

        return VortexConfig

    spec = importlib.util.spec_from_file_location(
        "configuration_vortex",
        os.path.join(os.path.dirname(os.path.abspath(__file__)), "configuration_vortex.py"),
    )
    module = importlib.util.module_from_spec(spec)
    # Registered before exec so the dataclass-free class object survives even if
    # something inside the module re-enters this lookup.
    sys.modules["configuration_vortex"] = module
    spec.loader.exec_module(module)
    return module.VortexConfig


VortexConfig = _load_config_class()

logger = logging.get_logger(__name__)


# ──────────────────────────────────────────────────────────────────────
# Norm
# ──────────────────────────────────────────────────────────────────────
class VortexRMSNorm(nn.Module):
    """RMSNorm with the reduction and the norm-weight multiply in fp32.

    Upcasting is the point: with 18 pre-norm blocks in bf16 autocast, a bf16
    reduction over the residual stream loses enough precision to stall training.
    The output is cast back so the residual add stays in the activation dtype.
    """

    def __init__(self, hidden_size: int, eps: float = 1e-6):
        super().__init__()
        self.eps = float(eps)
        self.weight = nn.Parameter(torch.ones(hidden_size))
        self.normalized_shape = (hidden_size,)

    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
        input_dtype = hidden_states.dtype
        hidden_states = hidden_states.to(torch.float32)
        variance = hidden_states.pow(2).mean(-1, keepdim=True)
        hidden_states = hidden_states * torch.rsqrt(variance + self.eps)
        return (self.weight.float() * hidden_states).to(input_dtype)

    def extra_repr(self) -> str:
        return f"{tuple(self.weight.shape)}, eps={self.eps}"


# ──────────────────────────────────────────────────────────────────────
# Rotary position embedding
# ──────────────────────────────────────────────────────────────────────
def build_rope_cache(
    head_dim: int,
    max_seq_len: int,
    base: float = 10_000.0,
    device=None,
    dtype: torch.dtype = torch.float32,
) -> Tuple[torch.Tensor, torch.Tensor]:
    """Build the `(max_seq_len, head_dim / 2)` cos/sin tables for RoPE."""
    if head_dim % 2 != 0:
        raise ValueError(f"head_dim must be even, got {head_dim}")
    inv_freq = 1.0 / (base ** (torch.arange(0, head_dim, 2, device=device).float() / head_dim))
    position_ids = torch.arange(max_seq_len, device=device, dtype=torch.float32)
    freqs = torch.outer(position_ids, inv_freq)
    return freqs.cos().to(dtype), freqs.sin().to(dtype)


def apply_rope(
    x: torch.Tensor,
    cos: torch.Tensor,
    sin: torch.Tensor,
    offset: int = 0,
) -> torch.Tensor:
    """Rotate the last dim of `x` (GPT-NeoX split-half pairing).

    `x` is `(batch, heads, seq, head_dim)`. `cos`/`sin` are `(seq, head_dim / 2)`
    *absolute* position tables; `offset` selects the starting position, which is
    what puts a cached query block on the right rotary phase.

    Pre-sliced tables with the default `offset=0` are still accepted, so the
    direct-call form used by the verification suite keeps working.
    """
    if x.shape[-2] != cos.shape[0] or offset != 0:
        T = x.shape[-2]
        cos = cos[offset : offset + T]
        sin = sin[offset : offset + T]
    cos = cos.unsqueeze(0).unsqueeze(0)
    sin = sin.unsqueeze(0).unsqueeze(0)
    x1, x2 = x.chunk(2, dim=-1)
    return torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], dim=-1)


class VortexRotaryEmbedding(nn.Module):
    """Per-model RoPE table, built once and shared by every attention layer.

    Held in non-persistent state so it never becomes a checkpoint tensor β€” it is
    fully determined by `head_dim`, `rope_theta` and the current device/dtype.
    """

    def __init__(self, config: VortexConfig, device=None):
        super().__init__()
        self.config = config
        self.max_seq_len_cached = config.max_position_embeddings

        # Plain attributes, deliberately not buffers. `from_pretrained` builds
        # the model on a meta device and materialises only the tensors it finds
        # in the checkpoint, so a *non-persistent* buffer is left as
        # uninitialised memory: the model loads without error and every RoPE
        # application is garbage. Keeping this out of `state_dict` also means the
        # key layout stays identical to the released training checkpoints, which
        # is what lets `load_state_dict(strict=True)` accept them.
        self._inv_freq: Optional[torch.Tensor] = None
        self._inv_freq_device: Optional[torch.device] = device
        self._cos: Optional[torch.Tensor] = None
        self._sin: Optional[torch.Tensor] = None
        self._cached_len = 0
        self._cached_dtype: Optional[torch.dtype] = None

    def _get_inv_freq(self, device: torch.device) -> torch.Tensor:
        head_dim = self.config.head_dim
        if self._inv_freq is None or self._inv_freq_device != device:
            self._inv_freq = 1.0 / (
                self.config.rope_theta
                ** (torch.arange(0, head_dim, 2, device=device, dtype=torch.float32) / head_dim)
            )
            self._inv_freq_device = device
            # Invalidate the cos/sin tables; they were built from the old one.
            self._cos = self._sin = None
        return self._inv_freq

    @torch.no_grad()
    def forward(self, x: torch.Tensor, seq_len: int) -> Tuple[torch.Tensor, torch.Tensor]:
        """Return cos/sin tables covering at least `seq_len` positions."""
        device, dtype = x.device, x.dtype
        if (
            self._cos is None
            or self._cached_len < seq_len
            or self._cos.device != device
            or self._cached_dtype != dtype
        ):
            inv_freq = self._get_inv_freq(device)
            self._cached_len = max(seq_len, self.config.max_position_embeddings)
            position_ids = torch.arange(self._cached_len, device=device, dtype=torch.float32)
            freqs = torch.outer(position_ids, inv_freq)
            self._cos = freqs.cos().to(dtype)
            self._sin = freqs.sin().to(dtype)
            self._cached_dtype = dtype
        return self._cos, self._sin


# ──────────────────────────────────────────────────────────────────────
# Attention
# ──────────────────────────────────────────────────────────────────────
class VortexAttention(nn.Module):
    """Causal grouped-query attention with optional QK-Norm.

    Goes through `F.scaled_dot_product_attention` with no hand-written softmax,
    which lets PyTorch dispatch to FlashAttention-2 on Ampere and later and to
    the math backend everywhere else. `enable_gqa` avoids materialising repeated
    KV heads; the `repeat_interleave` branch only runs on torch < 2.5.
    """

    def __init__(self, config: VortexConfig, layer_idx: int = 0):
        super().__init__()
        self.config = config
        self.layer_idx = layer_idx

        self.n_heads = config.num_attention_heads
        self.n_kv = config.num_key_value_heads
        self.head_dim = config.head_dim
        self.n_groups = self.n_heads // self.n_kv
        self.rope_theta = config.rope_theta
        self.use_qk_norm = bool(config.use_qk_norm)
        self.scale = self.head_dim**-0.5

        hidden_size = config.hidden_size
        self.q_proj = nn.Linear(hidden_size, self.n_heads * self.head_dim, bias=False)
        self.k_proj = nn.Linear(hidden_size, self.n_kv * self.head_dim, bias=False)
        self.v_proj = nn.Linear(hidden_size, self.n_kv * self.head_dim, bias=False)
        self.o_proj = nn.Linear(self.n_heads * self.head_dim, hidden_size, bias=False)

        if self.use_qk_norm:
            # Per-head RMS over head_dim, applied before the attention matmul.
            # Without it, small models hit attention entropy collapse early: a
            # few heads saturate, their softmax goes one-hot, and those heads
            # are dead for the rest of the run. Costs 2 * head_dim params/layer.
            self.q_norm = VortexRMSNorm(self.head_dim, eps=config.rms_norm_eps)
            self.k_norm = VortexRMSNorm(self.head_dim, eps=config.rms_norm_eps)
        else:
            self.q_norm = self.k_norm = nn.Identity()

    def forward(
        self,
        x: torch.Tensor,
        position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
        position_offset: int = 0,
        attention_mask: Optional[torch.Tensor] = None,
        past_key_value: Optional[Cache] = None,
    ) -> torch.Tensor:
        B, T, C = x.shape

        q = self.q_proj(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
        k = self.k_proj(x).view(B, T, self.n_kv, self.head_dim).transpose(1, 2)
        v = self.v_proj(x).view(B, T, self.n_kv, self.head_dim).transpose(1, 2)

        # QK-Norm: bound the pre-softmax logits before RoPE mixes them.
        q = self.q_norm(q)
        k = self.k_norm(k)

        if position_embeddings is None:
            position_embeddings = build_rope_cache(
                self.head_dim, position_offset + T, self.rope_theta, x.device, x.dtype
            )
        cos, sin = position_embeddings
        q = apply_rope(q, cos, sin, offset=position_offset)
        k = apply_rope(k, cos, sin, offset=position_offset)

        if past_key_value is not None:
            k, v = past_key_value.update(k, v, self.layer_idx)

        # `is_causal=True` is only correct when the cache is empty: SDPA assumes
        # top-left alignment, and a cached block queries a suffix of the key
        # sequence. `VortexModel` hands over an explicit mask whenever that is
        # the case and leaves it `None` for the prefill fast path.
        is_causal = attention_mask is None and T > 1

        try:
            out = F.scaled_dot_product_attention(
                q, k, v,
                attn_mask=attention_mask,
                dropout_p=0.0,
                is_causal=is_causal,
                scale=self.scale,
                enable_gqa=self.n_groups > 1,
            )
        except TypeError:  # torch < 2.5 has no `enable_gqa`
            if self.n_groups > 1:
                k = k.repeat_interleave(self.n_groups, dim=1)
                v = v.repeat_interleave(self.n_groups, dim=1)
            out = F.scaled_dot_product_attention(
                q, k, v,
                attn_mask=attention_mask,
                dropout_p=0.0,
                is_causal=is_causal,
                scale=self.scale,
            )

        out = out.transpose(1, 2).contiguous().view(B, T, C)
        return self.o_proj(out)


# ──────────────────────────────────────────────────────────────────────
# MLP
# ──────────────────────────────────────────────────────────────────────
class VortexMLP(nn.Module):
    """SwiGLU feed-forward: `down(silu(gate(x)) * up(x))`."""

    def __init__(self, config: VortexConfig):
        super().__init__()
        intermediate_size = config.intermediate_size
        self.gate_proj = nn.Linear(config.hidden_size, intermediate_size, bias=False)
        self.up_proj = nn.Linear(config.hidden_size, intermediate_size, bias=False)
        self.down_proj = nn.Linear(intermediate_size, config.hidden_size, bias=False)
        self.act_fn = ACT2FN["silu"]

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))


# ──────────────────────────────────────────────────────────────────────
# Block
# ──────────────────────────────────────────────────────────────────────
class VortexBlock(nn.Module):
    """Pre-norm block: attention and MLP each add onto the residual stream."""

    def __init__(self, config: VortexConfig, layer_idx: int = 0):
        super().__init__()
        self.layer_idx = layer_idx
        self.attn = VortexAttention(config, layer_idx=layer_idx)
        self.mlp = VortexMLP(config)
        self.ln_attn = VortexRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
        self.ln_mlp = VortexRMSNorm(config.hidden_size, eps=config.rms_norm_eps)

        # GPT-2 style 1/sqrt(2L) branch scaling. Off by default: it is redundant
        # next to zero-initialised residual outputs, which already make every
        # block an exact identity at init.
        self.resid_scale = (
            1.0 / math.sqrt(2.0 * config.num_hidden_layers) if config.scale_residual else 1.0
        )

    def forward(
        self,
        x: torch.Tensor,
        position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
        position_offset: int = 0,
        attention_mask: Optional[torch.Tensor] = None,
        past_key_value: Optional[Cache] = None,
    ) -> torch.Tensor:
        x = x + self.resid_scale * self.attn(
            self.ln_attn(x),
            position_embeddings=position_embeddings,
            position_offset=position_offset,
            attention_mask=attention_mask,
            past_key_value=past_key_value,
        )
        x = x + self.resid_scale * self.mlp(self.ln_mlp(x))
        return x


# ──────────────────────────────────────────────────────────────────────
# Base
# ──────────────────────────────────────────────────────────────────────
class VortexPreTrainedModel(PreTrainedModel):
    """Weight init, tied-embedding bookkeeping and tokenizer plumbing."""

    config_class = VortexConfig
    base_model_prefix = "model"
    supports_gradient_checkpointing = True
    _no_split_modules = ["VortexBlock"]
    _skip_keys_device_placement = "past_key_values"
    _supports_sdpa = True
    # SDPA already dispatches to FlashAttention-2 kernels on Ampere+, but the
    # `attn_implementation="flash_attention_2"` HF interface is not implemented
    # here. Claiming support would let `from_pretrained` pick a code path that
    # does not exist for this architecture.
    _supports_flash_attn = False
    _supports_attention_backend = False
    _supports_cache_class = True
    _supports_static_cache = True
    _can_record_outputs = {"hidden_states": VortexBlock, "attentions": VortexAttention}

    def _init_weights(self, module: nn.Module):
        std = self.config.initializer_range
        if isinstance(module, nn.Linear):
            nn.init.normal_(module.weight, mean=0.0, std=std)
            if module.bias is not None:
                nn.init.zeros_(module.bias)
        elif isinstance(module, nn.Embedding):
            # A small vocab (16K) is far more tolerant than a 151K one, but
            # scaling down keeps initial logits O(1) rather than O(10).
            nn.init.normal_(module.weight, mean=0.0, std=std)
        elif isinstance(module, VortexRMSNorm):
            nn.init.ones_(module.weight)
        # Dispatched per-submodule by `PreTrainedModel.post_init`, which applies
        # this over the whole tree. Overriding it here is what makes a freshly
        # constructed model an identity passthrough without a separate traversal.
        self._zero_init_residuals(module)

    def _zero_init_residuals(self, module: Optional[nn.Module] = None) -> None:
        """Zero `o_proj` and `down_proj` so every block starts as an identity.

        With 18 stacked pre-norm blocks, default init compounds the residual
        variance and saturates the stream before step 0. Zeroing the two branch
        outputs makes the untrained network a clean passthrough, so the initial
        loss is ln(vocab_size) = 9.7 rather than the hundreds default init gives.

        Dispatched per-module by `PreTrainedModel.post_init` via
        `_init_weights`; the recursion is over `self.modules()` so it also works
        when called with no argument.
        """
        if not getattr(self.config, "zero_init_residual", True):
            return
        if module is not None:
            if isinstance(module, VortexAttention):
                nn.init.zeros_(module.o_proj.weight)
            elif isinstance(module, VortexMLP):
                nn.init.zeros_(module.down_proj.weight)
            return
        for block in self.model.layers:
            nn.init.zeros_(block.attn.o_proj.weight)
            nn.init.zeros_(block.mlp.down_proj.weight)

    def resize_token_embeddings(
        self,
        new_num_tokens: Optional[int] = None,
        pad_to_multiple_of: Optional[int] = None,
        mean_resizing: bool = True,
    ) -> nn.Embedding:
        """Grow the embedding table, never shrink it.

        Growing pads with fresh normal noise. Shrinking is refused rather than
        silently truncating: rows that have been trained keep meaning something,
        and a truncated table yields a model that evaluates fine and answers
        with the wrong tokens.
        """
        old_embeddings = self.get_input_embeddings()
        if old_embeddings is None:
            raise ValueError("cannot resize embeddings on a model with no input embeddings")

        old_num_tokens, embedding_dim = old_embeddings.weight.shape
        if new_num_tokens is None:
            new_num_tokens = old_num_tokens
        if pad_to_multiple_of is not None:
            new_num_tokens = math.ceil(new_num_tokens / pad_to_multiple_of) * pad_to_multiple_of
        new_num_tokens = int(new_num_tokens)

        if new_num_tokens < old_num_tokens:
            raise ValueError(
                f"cannot shrink token embeddings {old_num_tokens} -> {new_num_tokens}; "
                f"the vocabulary must only be extended"
            )
        if new_num_tokens == old_num_tokens:
            return old_embeddings

        new_embeddings = nn.Embedding(
            new_num_tokens, embedding_dim, device=old_embeddings.weight.device
        )
        with torch.no_grad():
            new_embeddings.weight.normal_(mean=0.0, std=self.config.initializer_range)
            new_embeddings.weight[:old_num_tokens].copy_(old_embeddings.weight)
        self.set_input_embeddings(new_embeddings)
        self.config.vocab_size = new_num_tokens

        # Keep the head in step with a tied table.
        if self.config.tie_word_embeddings and self.get_output_embeddings() is not None:
            self.tie_weights()
        return new_embeddings

    def tie_weights(self, recompute_mapping: bool = False, missing_keys=None):
        """Alias the `lm_head` weight onto the embedding table.

        Overridden rather than inherited because `missing_keys` has two
        incompatible shapes across transformers versions: a `set` in 4.x and a
        mapping in 4.56+. Only `lm_head` is tied here, so it is dropped from the
        "missing" report under either shape.
        """
        if getattr(self.config, "tie_word_embeddings", False):
            output_embeddings = self.get_output_embeddings()
            input_embeddings = self.get_input_embeddings()
            if output_embeddings is not None and input_embeddings is not None:
                output_embeddings.weight = input_embeddings.weight
        if missing_keys is None:
            return
        discard = getattr(missing_keys, "discard", None)
        if callable(discard):
            discard("lm_head.weight")
            return
        if hasattr(missing_keys, "pop"):
            try:
                missing_keys.pop("lm_head.weight")
            except TypeError:  # mapping-style pop(key, default)
                missing_keys.pop("lm_head.weight", None)


# ──────────────────────────────────────────────────────────────────────
# Model
# ──────────────────────────────────────────────────────────────────────
class VortexModel(VortexPreTrainedModel):
    """Embedding + decoder stack + final norm."""

    def __init__(self, config: VortexConfig):
        super().__init__(config)
        self.padding_idx = config.pad_token_id
        self.vocab_size = config.vocab_size

        self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
        self.layers = nn.ModuleList(
            [VortexBlock(config, layer_idx=i) for i in range(config.num_hidden_layers)]
        )
        self.norm = VortexRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
        self.rotary_emb = VortexRotaryEmbedding(config)

        # Set by `PreTrainedModel.gradient_checkpointing_enable`, which targets
        # any submodule carrying this attribute.
        self.gradient_checkpointing = False
        self.post_init()

    def get_input_embeddings(self) -> nn.Embedding:
        return self.embed_tokens

    def set_input_embeddings(self, value: nn.Embedding) -> None:
        self.embed_tokens = value

    # ── attention mask ───────────────────────────────────────────────
    @staticmethod
    def _build_causal_mask(
        q_len: int,
        kv_len: int,
        attention_mask_2d: Optional[torch.Tensor],
        device: torch.device,
    ) -> torch.Tensor:
        """Bottom-right-aligned boolean mask, `True` = attend.

        Two things SDPA's `is_causal=True` cannot express:

        1. **Alignment.** With `past_len` cached keys, query `i` sits at absolute
           position `past_len + i`, so it may attend to keys `0 .. past_len + i`.
           Top-left alignment would bar the cached keys from every query.
        2. **Padding.** Left-padded batches need the pad columns removed.

        The self-diagonal is force-enabled on top of the mask so no query row is
        ever fully masked. A fully-masked row makes softmax return `NaN`, and
        those `NaN`s then ride in the padded key/value vectors into the next
        layer, where a `0 * NaN` in the weighted sum spreads them. Letting a
        padded query attend to itself is harmless β€” that position is masked out
        for every other query, so it cannot leak.
        """
        key_positions = torch.arange(kv_len, device=device)
        query_positions = torch.arange(q_len, device=device) + (kv_len - q_len)
        mask = (key_positions[None, :] <= query_positions[:, None])[None, None, :, :]

        if attention_mask_2d is not None:
            padding = attention_mask_2d.to(device=device)[:, None, None, :].bool()
            mask = mask & padding

        self_attends = (key_positions[None, :] == query_positions[:, None])[None, None, :, :]
        return (mask | self_attends).expand(-1, 1, -1, -1).contiguous()

    def forward(
        self,
        input_ids: Optional[torch.LongTensor] = None,
        attention_mask: Optional[torch.Tensor] = None,
        position_ids: Optional[torch.LongTensor] = None,
        past_key_values: Optional[Cache] = None,
        inputs_embeds: Optional[torch.FloatTensor] = None,
        use_cache: Optional[bool] = None,
        output_attentions: Optional[bool] = None,
        output_hidden_states: Optional[bool] = None,
        return_dict: Optional[bool] = None,
        **kwargs,
    ) -> Union[Tuple, BaseModelOutputWithPast]:
        output_attentions = bool(output_attentions)
        output_hidden_states = bool(output_hidden_states)
        return_dict = True if return_dict is None else bool(return_dict)

        checkpointing = bool(getattr(self, "gradient_checkpointing", False)) and self.training
        use_cache = self.config.use_cache if use_cache is None else bool(use_cache)

        # Checkpointing recomputes each block during the backward pass, and
        # `Cache.update` mutates in place -- so a cached forward would append the
        # same keys a second time and corrupt every downstream layer's mask
        # (observed: the cache silently doubling from 24 to 48 entries).
        # Training never needs the cache anyway, so it is dropped here. Inference
        # is unaffected because `self.training` is False.
        if checkpointing:
            use_cache = False
            past_key_values = None

        if output_attentions:
            raise NotImplementedError(
                "`output_attentions=True` is not supported: each block returns only "
                "hidden states, because attention runs fused inside SDPA."
            )

        if (input_ids is None) == (inputs_embeds is None):
            raise ValueError("provide exactly one of `input_ids` or `inputs_embeds`")
        if inputs_embeds is None:
            inputs_embeds = self.embed_tokens(input_ids)
        if past_key_values is None and use_cache:
            past_key_values = DynamicCache(config=self.config)

        batch_size, seq_len, _ = inputs_embeds.shape
        past_len = past_key_values.get_seq_length() if past_key_values is not None else 0
        kv_len = past_len + seq_len

        if position_ids is None:
            # RoPE is relative, so shifting every position by the same constant
            # leaves every attention score unchanged. Deriving absolute
            # positions from the cache length is therefore correct even for the
            # left-padded batches `generate` builds.
            position_ids = torch.arange(past_len, past_len + seq_len, device=inputs_embeds.device)
            position_ids = position_ids.unsqueeze(0).expand(batch_size, -1)
        elif position_ids.shape[-1] == kv_len and past_len > 0:
            position_ids = position_ids[:, past_len:]

        # A 2D `(batch, kv_len)` padding mask is what `generate` passes; a 4D
        # mask is taken as already built. Anything else is ignored rather than
        # guessed at.
        padding_mask_2d = None
        if attention_mask is not None and attention_mask.dim() == 2:
            padding_mask_2d = attention_mask
            attention_mask = None

        if attention_mask is None and (past_len > 0 or padding_mask_2d is not None):
            attention_mask = self._build_causal_mask(
                q_len=seq_len,
                kv_len=kv_len,
                attention_mask_2d=padding_mask_2d,
                device=inputs_embeds.device,
            )

        # One RoPE table for the whole stack rather than one per layer.
        position_embeddings = self.rotary_emb(inputs_embeds, kv_len)

        hidden_states = inputs_embeds
        all_hidden_states = () if output_hidden_states else None

        checkpoint_fn = getattr(self, "_gradient_checkpointing_func", None)
        if checkpointing and checkpoint_fn is None:
            checkpoint_fn = lambda fn, *args: checkpoint(fn, *args, use_reentrant=False)

        for block in self.layers:
            if output_hidden_states:
                all_hidden_states += (hidden_states,)
            if checkpointing:
                hidden_states = checkpoint_fn(
                    block,
                    hidden_states,
                    position_embeddings,
                    past_len,
                    attention_mask,
                    past_key_values,
                )
            else:
                hidden_states = block(
                    hidden_states,
                    position_embeddings=position_embeddings,
                    position_offset=past_len,
                    attention_mask=attention_mask,
                    past_key_value=past_key_values,
                )

        hidden_states = self.norm(hidden_states)
        if output_hidden_states:
            all_hidden_states += (hidden_states,)

        if not return_dict:
            return (hidden_states, past_key_values if use_cache else None, all_hidden_states)
        return BaseModelOutputWithPast(
            last_hidden_state=hidden_states,
            past_key_values=past_key_values if use_cache else None,
            hidden_states=all_hidden_states,
            attentions=None,
        )


# ──────────────────────────────────────────────────────────────────────
# Causal LM
# ──────────────────────────────────────────────────────────────────────
class VortexForCausalLM(VortexPreTrainedModel, GenerationMixin):
    """Vortex with a tied language-modelling head.

    `VortexModel` is the base model (`base_model_prefix = "model"`), so
    `save_pretrained` writes `model.embed_tokens.weight`, `model.layers.N.*` and
    `model.norm.weight` β€” the same key layout as the training checkpoints, which
    is what lets this class load them unchanged.

    `GenerationMixin` is inherited explicitly. From transformers v4.50 onward
    `PreTrainedModel` no longer provides it, so without this second base the
    model silently loses `generate`, `generate_from_model` and sampling helpers.
    It must come *after* `PreTrainedModel` in the MRO.
    """

    _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"}
    _tp_plan = {"lm_head.weight": "model.embed_tokens.weight"}
    _pp_plan = {"embed_tokens": ["model.embed_tokens"], "layers": ["model.layers"]}

    def __init__(self, config: VortexConfig):
        super().__init__(config)
        self.model = VortexModel(config)
        self.vocab_size = config.vocab_size
        self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)

        # Without this a directly-constructed model keeps PyTorch's default init
        # β€” no `initializer_range`, no zeroed residual outputs. `from_pretrained`
        # calls it too, but construction has to be self-sufficient or
        # `VortexForCausalLM(config).to(device)` silently trains a broken model.
        self.post_init()
        if config.tie_word_embeddings:
            self.tie_weights()

    def post_init(self) -> None:
        """Initialise weights, then apply the zero-init residual scheme.

        `PreTrainedModel.post_init` is what registers `all_tied_weights_keys`,
        parallel plans and device-map hints, and it is also what drives
        `_init_weights` over every submodule. It must be delegated to rather than
        shadowed, but on its own it leaves `o_proj` and `down_proj` at their
        normal init, so the second pass below is what actually makes an untrained
        model a passthrough.
        """
        super().post_init()
        self._zero_init_residuals()
        if getattr(self.config, "tie_word_embeddings", False):
            self.tie_weights()

    def get_input_embeddings(self) -> nn.Embedding:
        return self.model.embed_tokens

    def set_input_embeddings(self, value: nn.Embedding) -> None:
        self.model.embed_tokens = value

    def get_output_embeddings(self) -> nn.Linear:
        return self.lm_head

    def set_output_embeddings(self, new_embeddings: nn.Module) -> None:
        self.lm_head = new_embeddings

    def get_decoder(self) -> VortexModel:
        return self.model

    # ── loss ─────────────────────────────────────────────────────────
    def _chunked_cross_entropy(
        self,
        hidden_states: torch.Tensor,
        labels: torch.Tensor,
        chunk_size: int,
        num_items_in_batch: Optional[torch.Tensor] = None,
    ) -> torch.Tensor:
        """Cross-entropy without materialising `(N, vocab_size)` logits.

        This is the single largest avoidable memory term in a training step: at
        batch 32 x 2048 tokens x 16384 vocab in fp32 the logits alone are 4.3GB.
        Accumulating in time-chunks holds the peak at `chunk_size` rows instead.
        """
        n_valid = (labels != -100).sum()
        if n_valid.item() == 0:
            return hidden_states.sum() * 0.0  # keep the graph connected

        total = hidden_states.new_zeros((), dtype=torch.float32)
        for start in range(0, hidden_states.shape[0], chunk_size):
            logits = self.lm_head(hidden_states[start : start + chunk_size]).float()
            total = total + F.cross_entropy(
                logits,
                labels[start : start + chunk_size],
                ignore_index=-100,
                reduction="sum",
            )
            del logits

        if num_items_in_batch is not None:
            # `Trainer` normalises by a token count accumulated across
            # gradient-accumulation steps. Matching it is what keeps the loss it
            # reports comparable to the standalone training loop's.
            return total / num_items_in_batch.to(total.device)
        return total / n_valid.clamp_min(1).float()

    def forward(
        self,
        input_ids: Optional[torch.LongTensor] = None,
        attention_mask: Optional[torch.Tensor] = None,
        position_ids: Optional[torch.LongTensor] = None,
        past_key_values: Optional[Cache] = None,
        inputs_embeds: Optional[torch.FloatTensor] = None,
        labels: Optional[torch.LongTensor] = None,
        use_cache: Optional[bool] = None,
        output_attentions: Optional[bool] = None,
        output_hidden_states: Optional[bool] = None,
        return_dict: Optional[bool] = None,
        logits_to_keep: Union[int, torch.Tensor] = 0,
        chunk_size: int = 0,
        num_items_in_batch: Optional[torch.Tensor] = None,
        **kwargs,
    ) -> Union[Tuple, CausalLMOutputWithPast]:
        r"""Causal language modelling.

        Args:
            labels (`torch.LongTensor`, *optional*):
                Targets for next-token prediction. When given, `logits` comes back
                `None` unless `logits_to_keep` asks for it β€” the loss accumulates
                in chunks precisely so the full `(batch, seq, vocab)` tensor never
                has to exist.
            logits_to_keep (`int`, *optional*, defaults to 0):
                Return logits for only the last `n` positions. 0 means all of
                them when no loss is being computed, and none when one is.
                `transformers` sets this to 1 during `generate`; passing any value
                alongside `labels` is how to ask for a loss *and* logits.
            chunk_size (`int`, *optional*, defaults to 0):
                Rows per cross-entropy chunk. 0 selects 1024.

        Returns:
            [`CausalLMOutputWithPast`]: `logits`, `loss`, and `past_key_values`
            when `use_cache` is set.
        """
        return_dict = True if return_dict is None else bool(return_dict)
        outputs = self.model(
            input_ids=input_ids,
            attention_mask=attention_mask,
            position_ids=position_ids,
            past_key_values=past_key_values,
            inputs_embeds=inputs_embeds,
            use_cache=use_cache,
            output_attentions=output_attentions,
            output_hidden_states=output_hidden_states,
            return_dict=True,
            **kwargs,
        )
        hidden_states = outputs.last_hidden_state
        past_key_values = outputs.past_key_values

        # Keep only the tail when asked. During generation this is the single new
        # position, so the vocab-sized projection runs on one row instead of the
        # whole sequence.
        keep = int(logits_to_keep.item()) if isinstance(logits_to_keep, torch.Tensor) else int(logits_to_keep)

        loss = None
        if labels is not None:
            # Next-token alignment: predict token t+1 from position t.
            shift_hidden = hidden_states[..., :-1, :].reshape(-1, hidden_states.shape[-1])
            shift_labels = labels[..., 1:].reshape(-1)
            loss = self._chunked_cross_entropy(
                shift_hidden, shift_labels, chunk_size or 1024, num_items_in_batch
            )

        if labels is not None and keep == 0:
            # Keep the memory win. Ask with `logits_to_keep=1` if you need logits
            # alongside a loss.
            logits = None
        else:
            tail = hidden_states[:, -keep:, :] if keep > 0 else hidden_states
            logits = self.lm_head(tail)

        if not return_dict:
            return (logits, loss) if loss is None else (logits, loss, past_key_values)

        return CausalLMOutputWithPast(
            loss=loss,
            logits=logits,
            past_key_values=past_key_values,
            hidden_states=outputs.hidden_states,
            attentions=outputs.attentions,
        )

    # ── generation ───────────────────────────────────────────────────
    def prepare_inputs_for_generation(
        self,
        input_ids: torch.LongTensor,
        past_key_values: Optional[Cache] = None,
        attention_mask: Optional[torch.Tensor] = None,
        inputs_embeds: Optional[torch.FloatTensor] = None,
        position_ids: Optional[torch.LongTensor] = None,
        use_cache: Optional[bool] = None,
        **kwargs,
    ) -> dict:
        """Trim model inputs to the block `generate` is about to run.

        Handled entirely by `GenerationMixin`: it slices `input_ids` down to the
        tokens not yet in the cache. That slice is load-bearing rather than an
        optimisation β€” resending the full prefix would recompute it and corrupt
        the cache. Overridden only to keep the signature aligned with
        `transformers` 5.x and to forward `position_ids`, which the base
        implementation pops and re-slices.
        """
        return super().prepare_inputs_for_generation(
            input_ids=input_ids,
            past_key_values=past_key_values,
            attention_mask=attention_mask,
            inputs_embeds=inputs_embeds,
            position_ids=position_ids,
            use_cache=use_cache,
            **kwargs,
        )

    # ── gradient checkpointing ───────────────────────────────────────
    def gradient_checkpointing_enable(
        self,
        gradient_checkpointing_kwargs: Optional[dict] = None,
        **kwargs,
    ) -> None:
        """Recompute decoder activations in the backward pass instead of storing them.

        Trades roughly 20-30% step time for most of the activation memory, which
        is what lets one 40GB card hold a large batch at 2K context. Only active
        in training mode β€” `VortexModel.forward` gates on `self.training`.
        """
        super().gradient_checkpointing_enable(
            gradient_checkpointing_kwargs=gradient_checkpointing_kwargs, **kwargs
        )
        self.model.gradient_checkpointing = True
        if not hasattr(self.model, "_gradient_checkpointing_func"):
            self.model._gradient_checkpointing_func = lambda fn, *args: checkpoint(
                fn, *args, use_reentrant=False
            )

    def gradient_checkpointing_disable(self, **kwargs) -> None:
        super().gradient_checkpointing_disable(**kwargs)
        self.model.gradient_checkpointing = False


__all__ = [
    "VortexConfig",
    "VortexPreTrainedModel",
    "VortexModel",
    "VortexForCausalLM",
    "VortexBlock",
    "VortexAttention",
    "VortexMLP",
    "VortexRMSNorm",
    "VortexRotaryEmbedding",
    "build_rope_cache",
    "apply_rope",
]