Abhaykoul commited on
Commit
104e2d9
Β·
verified Β·
1 Parent(s): 628cdc8

Add transformers configuration_vortex.py + modeling_vortex.py and auto_map

Browse files

Registers the architecture with AutoConfig/AutoModelForCausalLM so the model loads via trust_remote_code. Adds the KV cache (generate() was previously impossible), a bottom-right-aligned attention mask for cached decoding, and real ModelOutput types. Weights are unchanged - model.safetensors is byte-identical and not re-uploaded.

config.json CHANGED
@@ -2,7 +2,13 @@
2
  "architectures": [
3
  "VortexForCausalLM"
4
  ],
 
 
 
 
 
5
  "dtype": "float32",
 
6
  "hidden_size": 512,
7
  "initializer_range": 0.02,
8
  "intermediate_size": 1072,
@@ -11,6 +17,7 @@
11
  "num_attention_heads": 8,
12
  "num_hidden_layers": 18,
13
  "num_key_value_heads": 2,
 
14
  "rms_norm_eps": 1e-06,
15
  "rope_interleaved": true,
16
  "rope_theta": 10000.0,
 
2
  "architectures": [
3
  "VortexForCausalLM"
4
  ],
5
+ "auto_map": {
6
+ "AutoConfig": "configuration_vortex.VortexConfig",
7
+ "AutoModelForCausalLM": "modeling_vortex.VortexForCausalLM"
8
+ },
9
+ "bos_token_id": 1,
10
  "dtype": "float32",
11
+ "eos_token_id": 2,
12
  "hidden_size": 512,
13
  "initializer_range": 0.02,
14
  "intermediate_size": 1072,
 
17
  "num_attention_heads": 8,
18
  "num_hidden_layers": 18,
19
  "num_key_value_heads": 2,
20
+ "pad_token_id": 0,
21
  "rms_norm_eps": 1e-06,
22
  "rope_interleaved": true,
23
  "rope_theta": 10000.0,
configuration_vortex.py ADDED
@@ -0,0 +1,205 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Vortex configuration β€” Hugging Face `PretrainedConfig` subclass.
2
+
3
+ Self-contained on purpose. When `trust_remote_code=True` is used,
4
+ `transformers` copies `configuration_vortex.py` and `modeling_vortex.py` into
5
+ `~/.cache/huggingface/modules/transformers_modules/<repo>/` and imports them as
6
+ *top-level* modules. Any import of a sibling file in this repository (e.g.
7
+ `from config import VortexArch`) would fail at that point, so this file may only
8
+ depend on the standard library and `transformers`.
9
+
10
+ Registering with the auto classes is what makes the checkpoint loadable with a
11
+ plain `AutoModelForCausalLM.from_pretrained(...)`:
12
+
13
+ AutoConfig.register("vortex", VortexConfig)
14
+ AutoModelForCausalLM.register(VortexConfig, VortexForCausalLM)
15
+
16
+ `src/export_hf.py` writes the equivalent `auto_map` block into `config.json`,
17
+ which is the serialised form of those two calls.
18
+ """
19
+
20
+ from __future__ import annotations
21
+
22
+ from transformers.configuration_utils import PretrainedConfig
23
+ from transformers.utils import logging
24
+
25
+ logger = logging.get_logger(__name__)
26
+
27
+
28
+ class VortexConfig(PretrainedConfig):
29
+ """Configuration for the Vortex decoder-only Transformer.
30
+
31
+ The defaults are the `vortex-50m-16k` preset: a 512d x 18L model with a
32
+ 16,384-token tied embedding table and 8Q/2KV grouped-query attention.
33
+
34
+ Args:
35
+ vocab_size (`int`, *optional*, defaults to 16384):
36
+ Size of the token embedding table. With `tie_word_embeddings=True`
37
+ this is also the size of the output head, and it is the single
38
+ biggest lever on the parameter budget at this scale.
39
+ hidden_size (`int`, *optional*, defaults to 512):
40
+ Model dimension. Must be divisible by `num_attention_heads`.
41
+ num_hidden_layers (`int`, *optional*, defaults to 18):
42
+ Number of decoder blocks.
43
+ num_attention_heads (`int`, *optional*, defaults to 8):
44
+ Number of query heads. `hidden_size // num_attention_heads` is the
45
+ head dimension and must be even for RoPE.
46
+ num_key_value_heads (`int`, *optional*, defaults to 2):
47
+ Number of key/value heads. Fewer than `num_attention_heads` selects
48
+ grouped-query attention (GQA); must divide `num_attention_heads`.
49
+ intermediate_size (`int`, *optional*, defaults to 1072):
50
+ SwiGLU feed-forward width, ~2.09x `hidden_size`.
51
+ rms_norm_eps (`float`, *optional*, defaults to 1e-6):
52
+ Epsilon inside every RMSNorm.
53
+ rope_theta (`float`, *optional*, defaults to 10000.0):
54
+ RoPE base. Higher values stretch the wavelength of the
55
+ high-frequency rotary components.
56
+ max_position_embeddings (`int`, *optional*, defaults to 2048):
57
+ Maximum context length. The RoPE tables are built to this size and
58
+ grow on demand if a longer sequence is actually seen.
59
+ use_qk_norm (`bool`, *optional*, defaults to `True`):
60
+ Per-head RMSNorm on queries and keys before the attention matmul.
61
+ The main defence against attention entropy collapse in small
62
+ models; costs 2 * head_dim parameters per layer.
63
+ tie_word_embeddings (`bool`, *optional*, defaults to `True`):
64
+ Share the `lm_head` weight with the input embedding. Halves the
65
+ vocabulary-sized parameter cost.
66
+ zero_init_residual (`bool`, *optional*, defaults to `True`):
67
+ Initialise `o_proj` and `down_proj` to exactly zero so every block is
68
+ an identity at step 0. Only affects fresh initialisation β€” it has no
69
+ effect on loading trained weights.
70
+ initializer_range (`float`, *optional*, defaults to 0.02):
71
+ Standard deviation of the normal init for linear and embedding
72
+ weights.
73
+ use_cache (`bool`, *optional*, defaults to `True`):
74
+ Return a key/value `Cache` from `forward` so `generate` runs in
75
+ O(1) per token instead of re-running the full prefix.
76
+ scale_residual (`bool`, *optional*, defaults to `False`):
77
+ Scale residual branch outputs by `1/sqrt(2 * num_hidden_layers)`
78
+ (GPT-2 style). Redundant next to zero-init residuals, so off.
79
+ rope_interleaved (`bool`, *optional*, defaults to `True`):
80
+ `True` uses the GPT-NeoX split-half pairing (`x1, x2 = x.chunk(2)`);
81
+ `False` uses the interleaved-even/odd pairing. Recorded for
82
+ provenance; the split-half layout is what the released weights were
83
+ trained with.
84
+ """
85
+
86
+ model_type = "vortex"
87
+ keys_to_ignore_at_inference = ["past_key_values"]
88
+
89
+ # Defaults mirror the `vortex-50m-16k` preset (src/config.py::VortexArch).
90
+ # They are duplicated rather than imported so this file stays standalone.
91
+ def __init__(
92
+ self,
93
+ vocab_size: int = 16_384,
94
+ hidden_size: int = 512,
95
+ num_hidden_layers: int = 18,
96
+ num_attention_heads: int = 8,
97
+ num_key_value_heads: int = 2,
98
+ intermediate_size: int = 1_072,
99
+ rms_norm_eps: float = 1e-6,
100
+ rope_theta: float = 10_000.0,
101
+ max_position_embeddings: int = 2_048,
102
+ use_qk_norm: bool = True,
103
+ tie_word_embeddings: bool = True,
104
+ zero_init_residual: bool = True,
105
+ initializer_range: float = 0.02,
106
+ use_cache: bool = True,
107
+ scale_residual: bool = False,
108
+ rope_interleaved: bool = True,
109
+ bos_token_id: int = 1,
110
+ eos_token_id: int = 2,
111
+ pad_token_id: int = 0,
112
+ **kwargs,
113
+ ):
114
+ self.vocab_size = int(vocab_size)
115
+ self.hidden_size = int(hidden_size)
116
+ self.num_hidden_layers = int(num_hidden_layers)
117
+ self.num_attention_heads = int(num_attention_heads)
118
+ self.num_key_value_heads = int(num_key_value_heads)
119
+ self.intermediate_size = int(intermediate_size)
120
+ self.rms_norm_eps = float(rms_norm_eps)
121
+ self.rope_theta = float(rope_theta)
122
+ self.max_position_embeddings = int(max_position_embeddings)
123
+ self.use_qk_norm = bool(use_qk_norm)
124
+ self.zero_init_residual = bool(zero_init_residual)
125
+ self.initializer_range = float(initializer_range)
126
+ self.scale_residual = bool(scale_residual)
127
+ self.rope_interleaved = bool(rope_interleaved)
128
+ self.name_or_path = kwargs.pop("name_or_path", "")
129
+
130
+ super().__init__(
131
+ bos_token_id=bos_token_id,
132
+ eos_token_id=eos_token_id,
133
+ pad_token_id=pad_token_id,
134
+ tie_word_embeddings=bool(tie_word_embeddings),
135
+ **kwargs,
136
+ )
137
+
138
+ # `use_cache` is a model-level flag, not a base `PretrainedConfig`
139
+ # attribute β€” transformers 5 dropped it from the base class, so setting
140
+ # it here is what makes `config.use_cache` readable on a loaded config.
141
+ self.use_cache = bool(use_cache)
142
+
143
+ self.validate()
144
+
145
+ # ── derived ──────────────────────────────────────────────────────
146
+ # `hidden_size` and `num_attention_heads` are also the names of the two
147
+ # outermost `__init__` parameters, so these are read from the instance
148
+ # rather than the caller's arguments. A `head_dim` passed in the config JSON
149
+ # is a *derived* value: recomputing it keeps the model and its config from
150
+ # disagreeing if someone edits one and not the other.
151
+
152
+ @property
153
+ def head_dim(self) -> int:
154
+ """Query/key/value head dimension."""
155
+ return self.hidden_size // self.num_attention_heads
156
+
157
+ @property
158
+ def num_query_groups(self) -> int:
159
+ """Query heads served by each KV head under GQA."""
160
+ return self.num_attention_heads // self.num_key_value_heads
161
+
162
+ # ── validation ───────────────────────────────────────────────────
163
+ def validate(self) -> None:
164
+ """Reject an illegal shape at construction time.
165
+
166
+ Without this, a bad GQA split surfaces as an opaque SDPA error
167
+ ("heads in key and value must divide the number of heads in query")
168
+ layers deep inside a forward pass.
169
+ """
170
+ if self.hidden_size <= 0:
171
+ raise ValueError(f"hidden_size must be positive, got {self.hidden_size}")
172
+ if self.num_attention_heads <= 0:
173
+ raise ValueError(
174
+ f"num_attention_heads must be positive, got {self.num_attention_heads}"
175
+ )
176
+ if self.hidden_size % self.num_attention_heads != 0:
177
+ raise ValueError(
178
+ f"hidden_size {self.hidden_size} is not divisible by "
179
+ f"num_attention_heads {self.num_attention_heads}"
180
+ )
181
+ if self.num_key_value_heads < 1:
182
+ raise ValueError(
183
+ f"num_key_value_heads must be >= 1, got {self.num_key_value_heads}"
184
+ )
185
+ if self.num_key_value_heads > self.num_attention_heads:
186
+ raise ValueError(
187
+ f"num_key_value_heads {self.num_key_value_heads} exceeds "
188
+ f"num_attention_heads {self.num_attention_heads}"
189
+ )
190
+ if self.num_attention_heads % self.num_key_value_heads != 0:
191
+ raise ValueError(
192
+ f"num_attention_heads {self.num_attention_heads} is not divisible by "
193
+ f"num_key_value_heads {self.num_key_value_heads}; GQA needs whole "
194
+ f"query groups"
195
+ )
196
+ if self.head_dim % 2 != 0:
197
+ raise ValueError(
198
+ f"head_dim {self.head_dim} must be even for RoPE; got "
199
+ f"hidden_size {self.hidden_size} / {self.num_attention_heads} heads"
200
+ )
201
+ if self.vocab_size <= 0:
202
+ raise ValueError(f"vocab_size must be positive, got {self.vocab_size}")
203
+
204
+
205
+ __all__ = ["VortexConfig"]
generation_config.json ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ {
2
+ "_from_model_config": true,
3
+ "transformers_version": "5.17.0",
4
+ "use_cache": true
5
+ }
modeling_vortex.py ADDED
@@ -0,0 +1,970 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Vortex modeling β€” Hugging Face `PreTrainedModel` implementation.
2
+
3
+ Self-contained on purpose. With `trust_remote_code=True`, `transformers` copies
4
+ `configuration_vortex.py` and `modeling_vortex.py` into
5
+ `~/.cache/huggingface/modules/transformers_modules/<repo>/` and imports them as
6
+ a package, so nothing here may import a sibling file from this repository.
7
+ `configuration_vortex` is the only dependency and it travels with this module, so
8
+ the pair is always copied together β€” see `_load_config_class` for why the import
9
+ is written the way it is.
10
+
11
+ What this adds over a bare `nn.Module` port, and why each piece is needed for
12
+ `AutoModelForCausalLM` / `generate` / `Trainer` to work:
13
+
14
+ * **Key/value cache.** `use_cache` was a config field with no implementation β€”
15
+ every `forward` recomputed the whole prefix. `VortexAttention` now consumes a
16
+ `transformers` `Cache`, which is what makes `model.generate()` viable.
17
+ * **Position offsets under a cache.** RoPE was sliced `cos[:T]`, i.e. positions
18
+ were always 0-based. With a cache the query block starts at `past_len`; the
19
+ rotary tables are now sliced `[offset : offset + T]`. RoPE is relative, so this
20
+ leaves the pretraining fast path bit-identical.
21
+ * **A correct attention mask on the cached path.** SDPA's `is_causal=True`
22
+ assumes top-left alignment and is only right when the cache is empty. Cached
23
+ steps with left padding need an explicit bottom-right-aligned mask, which is
24
+ what `VortexModel._build_causal_mask` builds. The empty-cache/no-padding case
25
+ still takes the `is_causal=True` fast path, so training numerics and memory are
26
+ unchanged.
27
+ * **Real `ModelOutput`s.** The previous `CausalLMOutput` was a plain object, so
28
+ `output.logits` worked but nothing HF-side (generation, `Trainer`, tensor
29
+ logging) recognised it.
30
+ * **Standard input plumbing** β€” `attention_mask`, `position_ids`,
31
+ `inputs_embeds`, `num_items_in_batch`, `logits_to_keep`.
32
+
33
+ Two deliberate deviations from HF naming conventions:
34
+
35
+ * The decoder submodules keep their original names (`attn`, `ln_attn`, `ln_mlp`)
36
+ rather than `self_attn`, `input_layernorm`, `post_attention_layernorm`. HF's
37
+ `self_attn` means *cross*-attention, which this architecture does not have.
38
+ More importantly, the released checkpoints on the Hub use the current names,
39
+ and `from_pretrained` matches `state_dict` keys literally β€” renaming would
40
+ break every one of them unless a key-remapping table were threaded through
41
+ `from_pretrained`, which is a per-version API in transformers 5.x. The outer
42
+ names (`model.*`, `embed_tokens`, `lm_head`, `norm`) already match HF.
43
+ * `logits_to_keep` is not decoration. Computing `(B, T, vocab_size)` logits for a
44
+ full 2048-token batch is the largest single memory term in a training step, and
45
+ the whole point of the chunked loss path is to never materialise it. That is
46
+ why `labels=` returns `logits=None` unless logits are explicitly asked for.
47
+ """
48
+
49
+ from __future__ import annotations
50
+
51
+ import importlib.util
52
+ import math
53
+ import os
54
+ import sys
55
+ from typing import Optional, Tuple, Union
56
+
57
+ import torch
58
+ import torch.nn as nn
59
+ import torch.nn.functional as F
60
+ from torch.utils.checkpoint import checkpoint
61
+ from transformers.activations import ACT2FN
62
+ from transformers.cache_utils import Cache, DynamicCache
63
+ from transformers.generation import GenerationMixin
64
+ from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast
65
+ from transformers.modeling_utils import PreTrainedModel
66
+ from transformers.utils import logging
67
+
68
+
69
+ def _load_config_class():
70
+ """Import the sibling `configuration_vortex` module.
71
+
72
+ Two different import mechanics have to be satisfied:
73
+
74
+ * Under `trust_remote_code`, `transformers` copies both files into its
75
+ dynamic-module package and imports them as a package, so a *relative*
76
+ import is the one that resolves.
77
+ * Running the repo's own tests imports this file as a top-level module from
78
+ `src/`, where there is no package and no `__package__`.
79
+
80
+ A plain top-level `from configuration_vortex import ...` is not an option:
81
+ `dynamic_module_utils.check_imports` runs `importlib.import_module` on every
82
+ statically-detected import *before* the sibling has been copied next to this
83
+ file, so it fails with "No module named 'configuration_vortex'" and a
84
+ misleading `pip install configuration_vortex`. Loading by file path keeps the
85
+ statement out of the AST the checker inspects.
86
+ """
87
+ if __package__:
88
+ from .configuration_vortex import VortexConfig
89
+
90
+ return VortexConfig
91
+
92
+ spec = importlib.util.spec_from_file_location(
93
+ "configuration_vortex",
94
+ os.path.join(os.path.dirname(os.path.abspath(__file__)), "configuration_vortex.py"),
95
+ )
96
+ module = importlib.util.module_from_spec(spec)
97
+ # Registered before exec so the dataclass-free class object survives even if
98
+ # something inside the module re-enters this lookup.
99
+ sys.modules["configuration_vortex"] = module
100
+ spec.loader.exec_module(module)
101
+ return module.VortexConfig
102
+
103
+
104
+ VortexConfig = _load_config_class()
105
+
106
+ logger = logging.get_logger(__name__)
107
+
108
+
109
+ # ──────────────────────────────────────────────────────────────────────
110
+ # Norm
111
+ # ──────────────────────────────────────────────────────────────────────
112
+ class VortexRMSNorm(nn.Module):
113
+ """RMSNorm with the reduction and the norm-weight multiply in fp32.
114
+
115
+ Upcasting is the point: with 18 pre-norm blocks in bf16 autocast, a bf16
116
+ reduction over the residual stream loses enough precision to stall training.
117
+ The output is cast back so the residual add stays in the activation dtype.
118
+ """
119
+
120
+ def __init__(self, hidden_size: int, eps: float = 1e-6):
121
+ super().__init__()
122
+ self.eps = float(eps)
123
+ self.weight = nn.Parameter(torch.ones(hidden_size))
124
+ self.normalized_shape = (hidden_size,)
125
+
126
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
127
+ input_dtype = hidden_states.dtype
128
+ hidden_states = hidden_states.to(torch.float32)
129
+ variance = hidden_states.pow(2).mean(-1, keepdim=True)
130
+ hidden_states = hidden_states * torch.rsqrt(variance + self.eps)
131
+ return (self.weight.float() * hidden_states).to(input_dtype)
132
+
133
+ def extra_repr(self) -> str:
134
+ return f"{tuple(self.weight.shape)}, eps={self.eps}"
135
+
136
+
137
+ # ──────────────────────────────────────────────────────────────────────
138
+ # Rotary position embedding
139
+ # ──────────────────────────────────────────────────────────────────────
140
+ def build_rope_cache(
141
+ head_dim: int,
142
+ max_seq_len: int,
143
+ base: float = 10_000.0,
144
+ device=None,
145
+ dtype: torch.dtype = torch.float32,
146
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
147
+ """Build the `(max_seq_len, head_dim / 2)` cos/sin tables for RoPE."""
148
+ if head_dim % 2 != 0:
149
+ raise ValueError(f"head_dim must be even, got {head_dim}")
150
+ inv_freq = 1.0 / (base ** (torch.arange(0, head_dim, 2, device=device).float() / head_dim))
151
+ position_ids = torch.arange(max_seq_len, device=device, dtype=torch.float32)
152
+ freqs = torch.outer(position_ids, inv_freq)
153
+ return freqs.cos().to(dtype), freqs.sin().to(dtype)
154
+
155
+
156
+ def apply_rope(
157
+ x: torch.Tensor,
158
+ cos: torch.Tensor,
159
+ sin: torch.Tensor,
160
+ offset: int = 0,
161
+ ) -> torch.Tensor:
162
+ """Rotate the last dim of `x` (GPT-NeoX split-half pairing).
163
+
164
+ `x` is `(batch, heads, seq, head_dim)`. `cos`/`sin` are `(seq, head_dim / 2)`
165
+ *absolute* position tables; `offset` selects the starting position, which is
166
+ what puts a cached query block on the right rotary phase.
167
+
168
+ Pre-sliced tables with the default `offset=0` are still accepted, so the
169
+ direct-call form used by the verification suite keeps working.
170
+ """
171
+ if x.shape[-2] != cos.shape[0] or offset != 0:
172
+ T = x.shape[-2]
173
+ cos = cos[offset : offset + T]
174
+ sin = sin[offset : offset + T]
175
+ cos = cos.unsqueeze(0).unsqueeze(0)
176
+ sin = sin.unsqueeze(0).unsqueeze(0)
177
+ x1, x2 = x.chunk(2, dim=-1)
178
+ return torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], dim=-1)
179
+
180
+
181
+ class VortexRotaryEmbedding(nn.Module):
182
+ """Per-model RoPE table, built once and shared by every attention layer.
183
+
184
+ Held in non-persistent state so it never becomes a checkpoint tensor β€” it is
185
+ fully determined by `head_dim`, `rope_theta` and the current device/dtype.
186
+ """
187
+
188
+ def __init__(self, config: VortexConfig, device=None):
189
+ super().__init__()
190
+ self.config = config
191
+ self.max_seq_len_cached = config.max_position_embeddings
192
+
193
+ # Plain attributes, deliberately not buffers. `from_pretrained` builds
194
+ # the model on a meta device and materialises only the tensors it finds
195
+ # in the checkpoint, so a *non-persistent* buffer is left as
196
+ # uninitialised memory: the model loads without error and every RoPE
197
+ # application is garbage. Keeping this out of `state_dict` also means the
198
+ # key layout stays identical to the released training checkpoints, which
199
+ # is what lets `load_state_dict(strict=True)` accept them.
200
+ self._inv_freq: Optional[torch.Tensor] = None
201
+ self._inv_freq_device: Optional[torch.device] = device
202
+ self._cos: Optional[torch.Tensor] = None
203
+ self._sin: Optional[torch.Tensor] = None
204
+ self._cached_len = 0
205
+ self._cached_dtype: Optional[torch.dtype] = None
206
+
207
+ def _get_inv_freq(self, device: torch.device) -> torch.Tensor:
208
+ head_dim = self.config.head_dim
209
+ if self._inv_freq is None or self._inv_freq_device != device:
210
+ self._inv_freq = 1.0 / (
211
+ self.config.rope_theta
212
+ ** (torch.arange(0, head_dim, 2, device=device, dtype=torch.float32) / head_dim)
213
+ )
214
+ self._inv_freq_device = device
215
+ # Invalidate the cos/sin tables; they were built from the old one.
216
+ self._cos = self._sin = None
217
+ return self._inv_freq
218
+
219
+ @torch.no_grad()
220
+ def forward(self, x: torch.Tensor, seq_len: int) -> Tuple[torch.Tensor, torch.Tensor]:
221
+ """Return cos/sin tables covering at least `seq_len` positions."""
222
+ device, dtype = x.device, x.dtype
223
+ if (
224
+ self._cos is None
225
+ or self._cached_len < seq_len
226
+ or self._cos.device != device
227
+ or self._cached_dtype != dtype
228
+ ):
229
+ inv_freq = self._get_inv_freq(device)
230
+ self._cached_len = max(seq_len, self.config.max_position_embeddings)
231
+ position_ids = torch.arange(self._cached_len, device=device, dtype=torch.float32)
232
+ freqs = torch.outer(position_ids, inv_freq)
233
+ self._cos = freqs.cos().to(dtype)
234
+ self._sin = freqs.sin().to(dtype)
235
+ self._cached_dtype = dtype
236
+ return self._cos, self._sin
237
+
238
+
239
+ # ──────────────────────────────────────────────────────────────────────
240
+ # Attention
241
+ # ──────────────────────────────────────────────────────────────────────
242
+ class VortexAttention(nn.Module):
243
+ """Causal grouped-query attention with optional QK-Norm.
244
+
245
+ Goes through `F.scaled_dot_product_attention` with no hand-written softmax,
246
+ which lets PyTorch dispatch to FlashAttention-2 on Ampere and later and to
247
+ the math backend everywhere else. `enable_gqa` avoids materialising repeated
248
+ KV heads; the `repeat_interleave` branch only runs on torch < 2.5.
249
+ """
250
+
251
+ def __init__(self, config: VortexConfig, layer_idx: int = 0):
252
+ super().__init__()
253
+ self.config = config
254
+ self.layer_idx = layer_idx
255
+
256
+ self.n_heads = config.num_attention_heads
257
+ self.n_kv = config.num_key_value_heads
258
+ self.head_dim = config.head_dim
259
+ self.n_groups = self.n_heads // self.n_kv
260
+ self.rope_theta = config.rope_theta
261
+ self.use_qk_norm = bool(config.use_qk_norm)
262
+ self.scale = self.head_dim**-0.5
263
+
264
+ hidden_size = config.hidden_size
265
+ self.q_proj = nn.Linear(hidden_size, self.n_heads * self.head_dim, bias=False)
266
+ self.k_proj = nn.Linear(hidden_size, self.n_kv * self.head_dim, bias=False)
267
+ self.v_proj = nn.Linear(hidden_size, self.n_kv * self.head_dim, bias=False)
268
+ self.o_proj = nn.Linear(self.n_heads * self.head_dim, hidden_size, bias=False)
269
+
270
+ if self.use_qk_norm:
271
+ # Per-head RMS over head_dim, applied before the attention matmul.
272
+ # Without it, small models hit attention entropy collapse early: a
273
+ # few heads saturate, their softmax goes one-hot, and those heads
274
+ # are dead for the rest of the run. Costs 2 * head_dim params/layer.
275
+ self.q_norm = VortexRMSNorm(self.head_dim, eps=config.rms_norm_eps)
276
+ self.k_norm = VortexRMSNorm(self.head_dim, eps=config.rms_norm_eps)
277
+ else:
278
+ self.q_norm = self.k_norm = nn.Identity()
279
+
280
+ def forward(
281
+ self,
282
+ x: torch.Tensor,
283
+ position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
284
+ position_offset: int = 0,
285
+ attention_mask: Optional[torch.Tensor] = None,
286
+ past_key_value: Optional[Cache] = None,
287
+ ) -> torch.Tensor:
288
+ B, T, C = x.shape
289
+
290
+ q = self.q_proj(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
291
+ k = self.k_proj(x).view(B, T, self.n_kv, self.head_dim).transpose(1, 2)
292
+ v = self.v_proj(x).view(B, T, self.n_kv, self.head_dim).transpose(1, 2)
293
+
294
+ # QK-Norm: bound the pre-softmax logits before RoPE mixes them.
295
+ q = self.q_norm(q)
296
+ k = self.k_norm(k)
297
+
298
+ if position_embeddings is None:
299
+ position_embeddings = build_rope_cache(
300
+ self.head_dim, position_offset + T, self.rope_theta, x.device, x.dtype
301
+ )
302
+ cos, sin = position_embeddings
303
+ q = apply_rope(q, cos, sin, offset=position_offset)
304
+ k = apply_rope(k, cos, sin, offset=position_offset)
305
+
306
+ if past_key_value is not None:
307
+ k, v = past_key_value.update(k, v, self.layer_idx)
308
+
309
+ # `is_causal=True` is only correct when the cache is empty: SDPA assumes
310
+ # top-left alignment, and a cached block queries a suffix of the key
311
+ # sequence. `VortexModel` hands over an explicit mask whenever that is
312
+ # the case and leaves it `None` for the prefill fast path.
313
+ is_causal = attention_mask is None and T > 1
314
+
315
+ try:
316
+ out = F.scaled_dot_product_attention(
317
+ q, k, v,
318
+ attn_mask=attention_mask,
319
+ dropout_p=0.0,
320
+ is_causal=is_causal,
321
+ scale=self.scale,
322
+ enable_gqa=self.n_groups > 1,
323
+ )
324
+ except TypeError: # torch < 2.5 has no `enable_gqa`
325
+ if self.n_groups > 1:
326
+ k = k.repeat_interleave(self.n_groups, dim=1)
327
+ v = v.repeat_interleave(self.n_groups, dim=1)
328
+ out = F.scaled_dot_product_attention(
329
+ q, k, v,
330
+ attn_mask=attention_mask,
331
+ dropout_p=0.0,
332
+ is_causal=is_causal,
333
+ scale=self.scale,
334
+ )
335
+
336
+ out = out.transpose(1, 2).contiguous().view(B, T, C)
337
+ return self.o_proj(out)
338
+
339
+
340
+ # ──────────────────────────────────────────────────────────────────────
341
+ # MLP
342
+ # ──────────────────────────────────────────────────────────────────────
343
+ class VortexMLP(nn.Module):
344
+ """SwiGLU feed-forward: `down(silu(gate(x)) * up(x))`."""
345
+
346
+ def __init__(self, config: VortexConfig):
347
+ super().__init__()
348
+ intermediate_size = config.intermediate_size
349
+ self.gate_proj = nn.Linear(config.hidden_size, intermediate_size, bias=False)
350
+ self.up_proj = nn.Linear(config.hidden_size, intermediate_size, bias=False)
351
+ self.down_proj = nn.Linear(intermediate_size, config.hidden_size, bias=False)
352
+ self.act_fn = ACT2FN["silu"]
353
+
354
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
355
+ return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
356
+
357
+
358
+ # ──────────────────────────────────────────────────────────────────────
359
+ # Block
360
+ # ──────────────────────────────────────────────────────────────────────
361
+ class VortexBlock(nn.Module):
362
+ """Pre-norm block: attention and MLP each add onto the residual stream."""
363
+
364
+ def __init__(self, config: VortexConfig, layer_idx: int = 0):
365
+ super().__init__()
366
+ self.layer_idx = layer_idx
367
+ self.attn = VortexAttention(config, layer_idx=layer_idx)
368
+ self.mlp = VortexMLP(config)
369
+ self.ln_attn = VortexRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
370
+ self.ln_mlp = VortexRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
371
+
372
+ # GPT-2 style 1/sqrt(2L) branch scaling. Off by default: it is redundant
373
+ # next to zero-initialised residual outputs, which already make every
374
+ # block an exact identity at init.
375
+ self.resid_scale = (
376
+ 1.0 / math.sqrt(2.0 * config.num_hidden_layers) if config.scale_residual else 1.0
377
+ )
378
+
379
+ def forward(
380
+ self,
381
+ x: torch.Tensor,
382
+ position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
383
+ position_offset: int = 0,
384
+ attention_mask: Optional[torch.Tensor] = None,
385
+ past_key_value: Optional[Cache] = None,
386
+ ) -> torch.Tensor:
387
+ x = x + self.resid_scale * self.attn(
388
+ self.ln_attn(x),
389
+ position_embeddings=position_embeddings,
390
+ position_offset=position_offset,
391
+ attention_mask=attention_mask,
392
+ past_key_value=past_key_value,
393
+ )
394
+ x = x + self.resid_scale * self.mlp(self.ln_mlp(x))
395
+ return x
396
+
397
+
398
+ # ──────────────────────────────────────────────────────────────────────
399
+ # Base
400
+ # ──────────────────────────────────────────────────────────────────────
401
+ class VortexPreTrainedModel(PreTrainedModel):
402
+ """Weight init, tied-embedding bookkeeping and tokenizer plumbing."""
403
+
404
+ config_class = VortexConfig
405
+ base_model_prefix = "model"
406
+ supports_gradient_checkpointing = True
407
+ _no_split_modules = ["VortexBlock"]
408
+ _skip_keys_device_placement = "past_key_values"
409
+ _supports_sdpa = True
410
+ # SDPA already dispatches to FlashAttention-2 kernels on Ampere+, but the
411
+ # `attn_implementation="flash_attention_2"` HF interface is not implemented
412
+ # here. Claiming support would let `from_pretrained` pick a code path that
413
+ # does not exist for this architecture.
414
+ _supports_flash_attn = False
415
+ _supports_attention_backend = False
416
+ _supports_cache_class = True
417
+ _supports_static_cache = True
418
+ _can_record_outputs = {"hidden_states": VortexBlock, "attentions": VortexAttention}
419
+
420
+ def _init_weights(self, module: nn.Module):
421
+ std = self.config.initializer_range
422
+ if isinstance(module, nn.Linear):
423
+ nn.init.normal_(module.weight, mean=0.0, std=std)
424
+ if module.bias is not None:
425
+ nn.init.zeros_(module.bias)
426
+ elif isinstance(module, nn.Embedding):
427
+ # A small vocab (16K) is far more tolerant than a 151K one, but
428
+ # scaling down keeps initial logits O(1) rather than O(10).
429
+ nn.init.normal_(module.weight, mean=0.0, std=std)
430
+ elif isinstance(module, VortexRMSNorm):
431
+ nn.init.ones_(module.weight)
432
+ # Dispatched per-submodule by `PreTrainedModel.post_init`, which applies
433
+ # this over the whole tree. Overriding it here is what makes a freshly
434
+ # constructed model an identity passthrough without a separate traversal.
435
+ self._zero_init_residuals(module)
436
+
437
+ def _zero_init_residuals(self, module: Optional[nn.Module] = None) -> None:
438
+ """Zero `o_proj` and `down_proj` so every block starts as an identity.
439
+
440
+ With 18 stacked pre-norm blocks, default init compounds the residual
441
+ variance and saturates the stream before step 0. Zeroing the two branch
442
+ outputs makes the untrained network a clean passthrough, so the initial
443
+ loss is ln(vocab_size) = 9.7 rather than the hundreds default init gives.
444
+
445
+ Dispatched per-module by `PreTrainedModel.post_init` via
446
+ `_init_weights`; the recursion is over `self.modules()` so it also works
447
+ when called with no argument.
448
+ """
449
+ if not getattr(self.config, "zero_init_residual", True):
450
+ return
451
+ if module is not None:
452
+ if isinstance(module, VortexAttention):
453
+ nn.init.zeros_(module.o_proj.weight)
454
+ elif isinstance(module, VortexMLP):
455
+ nn.init.zeros_(module.down_proj.weight)
456
+ return
457
+ for block in self.model.layers:
458
+ nn.init.zeros_(block.attn.o_proj.weight)
459
+ nn.init.zeros_(block.mlp.down_proj.weight)
460
+
461
+ def resize_token_embeddings(
462
+ self,
463
+ new_num_tokens: Optional[int] = None,
464
+ pad_to_multiple_of: Optional[int] = None,
465
+ mean_resizing: bool = True,
466
+ ) -> nn.Embedding:
467
+ """Grow the embedding table, never shrink it.
468
+
469
+ Growing pads with fresh normal noise. Shrinking is refused rather than
470
+ silently truncating: rows that have been trained keep meaning something,
471
+ and a truncated table yields a model that evaluates fine and answers
472
+ with the wrong tokens.
473
+ """
474
+ old_embeddings = self.get_input_embeddings()
475
+ if old_embeddings is None:
476
+ raise ValueError("cannot resize embeddings on a model with no input embeddings")
477
+
478
+ old_num_tokens, embedding_dim = old_embeddings.weight.shape
479
+ if new_num_tokens is None:
480
+ new_num_tokens = old_num_tokens
481
+ if pad_to_multiple_of is not None:
482
+ new_num_tokens = math.ceil(new_num_tokens / pad_to_multiple_of) * pad_to_multiple_of
483
+ new_num_tokens = int(new_num_tokens)
484
+
485
+ if new_num_tokens < old_num_tokens:
486
+ raise ValueError(
487
+ f"cannot shrink token embeddings {old_num_tokens} -> {new_num_tokens}; "
488
+ f"the vocabulary must only be extended"
489
+ )
490
+ if new_num_tokens == old_num_tokens:
491
+ return old_embeddings
492
+
493
+ new_embeddings = nn.Embedding(
494
+ new_num_tokens, embedding_dim, device=old_embeddings.weight.device
495
+ )
496
+ with torch.no_grad():
497
+ new_embeddings.weight.normal_(mean=0.0, std=self.config.initializer_range)
498
+ new_embeddings.weight[:old_num_tokens].copy_(old_embeddings.weight)
499
+ self.set_input_embeddings(new_embeddings)
500
+ self.config.vocab_size = new_num_tokens
501
+
502
+ # Keep the head in step with a tied table.
503
+ if self.config.tie_word_embeddings and self.get_output_embeddings() is not None:
504
+ self.tie_weights()
505
+ return new_embeddings
506
+
507
+ def tie_weights(self, recompute_mapping: bool = False, missing_keys=None):
508
+ """Alias the `lm_head` weight onto the embedding table.
509
+
510
+ Overridden rather than inherited because `missing_keys` has two
511
+ incompatible shapes across transformers versions: a `set` in 4.x and a
512
+ mapping in 4.56+. Only `lm_head` is tied here, so it is dropped from the
513
+ "missing" report under either shape.
514
+ """
515
+ if getattr(self.config, "tie_word_embeddings", False):
516
+ output_embeddings = self.get_output_embeddings()
517
+ input_embeddings = self.get_input_embeddings()
518
+ if output_embeddings is not None and input_embeddings is not None:
519
+ output_embeddings.weight = input_embeddings.weight
520
+ if missing_keys is None:
521
+ return
522
+ discard = getattr(missing_keys, "discard", None)
523
+ if callable(discard):
524
+ discard("lm_head.weight")
525
+ return
526
+ if hasattr(missing_keys, "pop"):
527
+ try:
528
+ missing_keys.pop("lm_head.weight")
529
+ except TypeError: # mapping-style pop(key, default)
530
+ missing_keys.pop("lm_head.weight", None)
531
+
532
+
533
+ # ──────────────────────────────────────────────────────────────────────
534
+ # Model
535
+ # ──────────────────────────────────────────────────────────────────────
536
+ class VortexModel(VortexPreTrainedModel):
537
+ """Embedding + decoder stack + final norm."""
538
+
539
+ def __init__(self, config: VortexConfig):
540
+ super().__init__(config)
541
+ self.padding_idx = config.pad_token_id
542
+ self.vocab_size = config.vocab_size
543
+
544
+ self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
545
+ self.layers = nn.ModuleList(
546
+ [VortexBlock(config, layer_idx=i) for i in range(config.num_hidden_layers)]
547
+ )
548
+ self.norm = VortexRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
549
+ self.rotary_emb = VortexRotaryEmbedding(config)
550
+
551
+ # Set by `PreTrainedModel.gradient_checkpointing_enable`, which targets
552
+ # any submodule carrying this attribute.
553
+ self.gradient_checkpointing = False
554
+ self.post_init()
555
+
556
+ def get_input_embeddings(self) -> nn.Embedding:
557
+ return self.embed_tokens
558
+
559
+ def set_input_embeddings(self, value: nn.Embedding) -> None:
560
+ self.embed_tokens = value
561
+
562
+ # ── attention mask ───────────────────────────────────────────────
563
+ @staticmethod
564
+ def _build_causal_mask(
565
+ q_len: int,
566
+ kv_len: int,
567
+ attention_mask_2d: Optional[torch.Tensor],
568
+ device: torch.device,
569
+ ) -> torch.Tensor:
570
+ """Bottom-right-aligned boolean mask, `True` = attend.
571
+
572
+ Two things SDPA's `is_causal=True` cannot express:
573
+
574
+ 1. **Alignment.** With `past_len` cached keys, query `i` sits at absolute
575
+ position `past_len + i`, so it may attend to keys `0 .. past_len + i`.
576
+ Top-left alignment would bar the cached keys from every query.
577
+ 2. **Padding.** Left-padded batches need the pad columns removed.
578
+
579
+ The self-diagonal is force-enabled on top of the mask so no query row is
580
+ ever fully masked. A fully-masked row makes softmax return `NaN`, and
581
+ those `NaN`s then ride in the padded key/value vectors into the next
582
+ layer, where a `0 * NaN` in the weighted sum spreads them. Letting a
583
+ padded query attend to itself is harmless β€” that position is masked out
584
+ for every other query, so it cannot leak.
585
+ """
586
+ key_positions = torch.arange(kv_len, device=device)
587
+ query_positions = torch.arange(q_len, device=device) + (kv_len - q_len)
588
+ mask = (key_positions[None, :] <= query_positions[:, None])[None, None, :, :]
589
+
590
+ if attention_mask_2d is not None:
591
+ padding = attention_mask_2d.to(device=device)[:, None, None, :].bool()
592
+ mask = mask & padding
593
+
594
+ self_attends = (key_positions[None, :] == query_positions[:, None])[None, None, :, :]
595
+ return (mask | self_attends).expand(-1, 1, -1, -1).contiguous()
596
+
597
+ def forward(
598
+ self,
599
+ input_ids: Optional[torch.LongTensor] = None,
600
+ attention_mask: Optional[torch.Tensor] = None,
601
+ position_ids: Optional[torch.LongTensor] = None,
602
+ past_key_values: Optional[Cache] = None,
603
+ inputs_embeds: Optional[torch.FloatTensor] = None,
604
+ use_cache: Optional[bool] = None,
605
+ output_attentions: Optional[bool] = None,
606
+ output_hidden_states: Optional[bool] = None,
607
+ return_dict: Optional[bool] = None,
608
+ **kwargs,
609
+ ) -> Union[Tuple, BaseModelOutputWithPast]:
610
+ output_attentions = bool(output_attentions)
611
+ output_hidden_states = bool(output_hidden_states)
612
+ return_dict = True if return_dict is None else bool(return_dict)
613
+
614
+ checkpointing = bool(getattr(self, "gradient_checkpointing", False)) and self.training
615
+ use_cache = self.config.use_cache if use_cache is None else bool(use_cache)
616
+
617
+ # Checkpointing recomputes each block during the backward pass, and
618
+ # `Cache.update` mutates in place -- so a cached forward would append the
619
+ # same keys a second time and corrupt every downstream layer's mask
620
+ # (observed: the cache silently doubling from 24 to 48 entries).
621
+ # Training never needs the cache anyway, so it is dropped here. Inference
622
+ # is unaffected because `self.training` is False.
623
+ if checkpointing:
624
+ use_cache = False
625
+ past_key_values = None
626
+
627
+ if output_attentions:
628
+ raise NotImplementedError(
629
+ "`output_attentions=True` is not supported: each block returns only "
630
+ "hidden states, because attention runs fused inside SDPA."
631
+ )
632
+
633
+ if (input_ids is None) == (inputs_embeds is None):
634
+ raise ValueError("provide exactly one of `input_ids` or `inputs_embeds`")
635
+ if inputs_embeds is None:
636
+ inputs_embeds = self.embed_tokens(input_ids)
637
+ if past_key_values is None and use_cache:
638
+ past_key_values = DynamicCache(config=self.config)
639
+
640
+ batch_size, seq_len, _ = inputs_embeds.shape
641
+ past_len = past_key_values.get_seq_length() if past_key_values is not None else 0
642
+ kv_len = past_len + seq_len
643
+
644
+ if position_ids is None:
645
+ # RoPE is relative, so shifting every position by the same constant
646
+ # leaves every attention score unchanged. Deriving absolute
647
+ # positions from the cache length is therefore correct even for the
648
+ # left-padded batches `generate` builds.
649
+ position_ids = torch.arange(past_len, past_len + seq_len, device=inputs_embeds.device)
650
+ position_ids = position_ids.unsqueeze(0).expand(batch_size, -1)
651
+ elif position_ids.shape[-1] == kv_len and past_len > 0:
652
+ position_ids = position_ids[:, past_len:]
653
+
654
+ # A 2D `(batch, kv_len)` padding mask is what `generate` passes; a 4D
655
+ # mask is taken as already built. Anything else is ignored rather than
656
+ # guessed at.
657
+ padding_mask_2d = None
658
+ if attention_mask is not None and attention_mask.dim() == 2:
659
+ padding_mask_2d = attention_mask
660
+ attention_mask = None
661
+
662
+ if attention_mask is None and (past_len > 0 or padding_mask_2d is not None):
663
+ attention_mask = self._build_causal_mask(
664
+ q_len=seq_len,
665
+ kv_len=kv_len,
666
+ attention_mask_2d=padding_mask_2d,
667
+ device=inputs_embeds.device,
668
+ )
669
+
670
+ # One RoPE table for the whole stack rather than one per layer.
671
+ position_embeddings = self.rotary_emb(inputs_embeds, kv_len)
672
+
673
+ hidden_states = inputs_embeds
674
+ all_hidden_states = () if output_hidden_states else None
675
+
676
+ checkpoint_fn = getattr(self, "_gradient_checkpointing_func", None)
677
+ if checkpointing and checkpoint_fn is None:
678
+ checkpoint_fn = lambda fn, *args: checkpoint(fn, *args, use_reentrant=False)
679
+
680
+ for block in self.layers:
681
+ if output_hidden_states:
682
+ all_hidden_states += (hidden_states,)
683
+ if checkpointing:
684
+ hidden_states = checkpoint_fn(
685
+ block,
686
+ hidden_states,
687
+ position_embeddings,
688
+ past_len,
689
+ attention_mask,
690
+ past_key_values,
691
+ )
692
+ else:
693
+ hidden_states = block(
694
+ hidden_states,
695
+ position_embeddings=position_embeddings,
696
+ position_offset=past_len,
697
+ attention_mask=attention_mask,
698
+ past_key_value=past_key_values,
699
+ )
700
+
701
+ hidden_states = self.norm(hidden_states)
702
+ if output_hidden_states:
703
+ all_hidden_states += (hidden_states,)
704
+
705
+ if not return_dict:
706
+ return (hidden_states, past_key_values if use_cache else None, all_hidden_states)
707
+ return BaseModelOutputWithPast(
708
+ last_hidden_state=hidden_states,
709
+ past_key_values=past_key_values if use_cache else None,
710
+ hidden_states=all_hidden_states,
711
+ attentions=None,
712
+ )
713
+
714
+
715
+ # ──────────────────────────────────────────────────────────────────────
716
+ # Causal LM
717
+ # ──────────────────────────────────────────────────────────────────────
718
+ class VortexForCausalLM(VortexPreTrainedModel, GenerationMixin):
719
+ """Vortex with a tied language-modelling head.
720
+
721
+ `VortexModel` is the base model (`base_model_prefix = "model"`), so
722
+ `save_pretrained` writes `model.embed_tokens.weight`, `model.layers.N.*` and
723
+ `model.norm.weight` β€” the same key layout as the training checkpoints, which
724
+ is what lets this class load them unchanged.
725
+
726
+ `GenerationMixin` is inherited explicitly. From transformers v4.50 onward
727
+ `PreTrainedModel` no longer provides it, so without this second base the
728
+ model silently loses `generate`, `generate_from_model` and sampling helpers.
729
+ It must come *after* `PreTrainedModel` in the MRO.
730
+ """
731
+
732
+ _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"}
733
+ _tp_plan = {"lm_head.weight": "model.embed_tokens.weight"}
734
+ _pp_plan = {"embed_tokens": ["model.embed_tokens"], "layers": ["model.layers"]}
735
+
736
+ def __init__(self, config: VortexConfig):
737
+ super().__init__(config)
738
+ self.model = VortexModel(config)
739
+ self.vocab_size = config.vocab_size
740
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
741
+
742
+ # Without this a directly-constructed model keeps PyTorch's default init
743
+ # β€” no `initializer_range`, no zeroed residual outputs. `from_pretrained`
744
+ # calls it too, but construction has to be self-sufficient or
745
+ # `VortexForCausalLM(config).to(device)` silently trains a broken model.
746
+ self.post_init()
747
+ if config.tie_word_embeddings:
748
+ self.tie_weights()
749
+
750
+ def post_init(self) -> None:
751
+ """Initialise weights, then apply the zero-init residual scheme.
752
+
753
+ `PreTrainedModel.post_init` is what registers `all_tied_weights_keys`,
754
+ parallel plans and device-map hints, and it is also what drives
755
+ `_init_weights` over every submodule. It must be delegated to rather than
756
+ shadowed, but on its own it leaves `o_proj` and `down_proj` at their
757
+ normal init, so the second pass below is what actually makes an untrained
758
+ model a passthrough.
759
+ """
760
+ super().post_init()
761
+ self._zero_init_residuals()
762
+ if getattr(self.config, "tie_word_embeddings", False):
763
+ self.tie_weights()
764
+
765
+ def get_input_embeddings(self) -> nn.Embedding:
766
+ return self.model.embed_tokens
767
+
768
+ def set_input_embeddings(self, value: nn.Embedding) -> None:
769
+ self.model.embed_tokens = value
770
+
771
+ def get_output_embeddings(self) -> nn.Linear:
772
+ return self.lm_head
773
+
774
+ def set_output_embeddings(self, new_embeddings: nn.Module) -> None:
775
+ self.lm_head = new_embeddings
776
+
777
+ def get_decoder(self) -> VortexModel:
778
+ return self.model
779
+
780
+ # ── loss ─────────────────────────────────────────────────────────
781
+ def _chunked_cross_entropy(
782
+ self,
783
+ hidden_states: torch.Tensor,
784
+ labels: torch.Tensor,
785
+ chunk_size: int,
786
+ num_items_in_batch: Optional[torch.Tensor] = None,
787
+ ) -> torch.Tensor:
788
+ """Cross-entropy without materialising `(N, vocab_size)` logits.
789
+
790
+ This is the single largest avoidable memory term in a training step: at
791
+ batch 32 x 2048 tokens x 16384 vocab in fp32 the logits alone are 4.3GB.
792
+ Accumulating in time-chunks holds the peak at `chunk_size` rows instead.
793
+ """
794
+ n_valid = (labels != -100).sum()
795
+ if n_valid.item() == 0:
796
+ return hidden_states.sum() * 0.0 # keep the graph connected
797
+
798
+ total = hidden_states.new_zeros((), dtype=torch.float32)
799
+ for start in range(0, hidden_states.shape[0], chunk_size):
800
+ logits = self.lm_head(hidden_states[start : start + chunk_size]).float()
801
+ total = total + F.cross_entropy(
802
+ logits,
803
+ labels[start : start + chunk_size],
804
+ ignore_index=-100,
805
+ reduction="sum",
806
+ )
807
+ del logits
808
+
809
+ if num_items_in_batch is not None:
810
+ # `Trainer` normalises by a token count accumulated across
811
+ # gradient-accumulation steps. Matching it is what keeps the loss it
812
+ # reports comparable to the standalone training loop's.
813
+ return total / num_items_in_batch.to(total.device)
814
+ return total / n_valid.clamp_min(1).float()
815
+
816
+ def forward(
817
+ self,
818
+ input_ids: Optional[torch.LongTensor] = None,
819
+ attention_mask: Optional[torch.Tensor] = None,
820
+ position_ids: Optional[torch.LongTensor] = None,
821
+ past_key_values: Optional[Cache] = None,
822
+ inputs_embeds: Optional[torch.FloatTensor] = None,
823
+ labels: Optional[torch.LongTensor] = None,
824
+ use_cache: Optional[bool] = None,
825
+ output_attentions: Optional[bool] = None,
826
+ output_hidden_states: Optional[bool] = None,
827
+ return_dict: Optional[bool] = None,
828
+ logits_to_keep: Union[int, torch.Tensor] = 0,
829
+ chunk_size: int = 0,
830
+ num_items_in_batch: Optional[torch.Tensor] = None,
831
+ **kwargs,
832
+ ) -> Union[Tuple, CausalLMOutputWithPast]:
833
+ r"""Causal language modelling.
834
+
835
+ Args:
836
+ labels (`torch.LongTensor`, *optional*):
837
+ Targets for next-token prediction. When given, `logits` comes back
838
+ `None` unless `logits_to_keep` asks for it β€” the loss accumulates
839
+ in chunks precisely so the full `(batch, seq, vocab)` tensor never
840
+ has to exist.
841
+ logits_to_keep (`int`, *optional*, defaults to 0):
842
+ Return logits for only the last `n` positions. 0 means all of
843
+ them when no loss is being computed, and none when one is.
844
+ `transformers` sets this to 1 during `generate`; passing any value
845
+ alongside `labels` is how to ask for a loss *and* logits.
846
+ chunk_size (`int`, *optional*, defaults to 0):
847
+ Rows per cross-entropy chunk. 0 selects 1024.
848
+
849
+ Returns:
850
+ [`CausalLMOutputWithPast`]: `logits`, `loss`, and `past_key_values`
851
+ when `use_cache` is set.
852
+ """
853
+ return_dict = True if return_dict is None else bool(return_dict)
854
+ outputs = self.model(
855
+ input_ids=input_ids,
856
+ attention_mask=attention_mask,
857
+ position_ids=position_ids,
858
+ past_key_values=past_key_values,
859
+ inputs_embeds=inputs_embeds,
860
+ use_cache=use_cache,
861
+ output_attentions=output_attentions,
862
+ output_hidden_states=output_hidden_states,
863
+ return_dict=True,
864
+ **kwargs,
865
+ )
866
+ hidden_states = outputs.last_hidden_state
867
+ past_key_values = outputs.past_key_values
868
+
869
+ # Keep only the tail when asked. During generation this is the single new
870
+ # position, so the vocab-sized projection runs on one row instead of the
871
+ # whole sequence.
872
+ keep = int(logits_to_keep.item()) if isinstance(logits_to_keep, torch.Tensor) else int(logits_to_keep)
873
+
874
+ loss = None
875
+ if labels is not None:
876
+ # Next-token alignment: predict token t+1 from position t.
877
+ shift_hidden = hidden_states[..., :-1, :].reshape(-1, hidden_states.shape[-1])
878
+ shift_labels = labels[..., 1:].reshape(-1)
879
+ loss = self._chunked_cross_entropy(
880
+ shift_hidden, shift_labels, chunk_size or 1024, num_items_in_batch
881
+ )
882
+
883
+ if labels is not None and keep == 0:
884
+ # Keep the memory win. Ask with `logits_to_keep=1` if you need logits
885
+ # alongside a loss.
886
+ logits = None
887
+ else:
888
+ tail = hidden_states[:, -keep:, :] if keep > 0 else hidden_states
889
+ logits = self.lm_head(tail)
890
+
891
+ if not return_dict:
892
+ return (logits, loss) if loss is None else (logits, loss, past_key_values)
893
+
894
+ return CausalLMOutputWithPast(
895
+ loss=loss,
896
+ logits=logits,
897
+ past_key_values=past_key_values,
898
+ hidden_states=outputs.hidden_states,
899
+ attentions=outputs.attentions,
900
+ )
901
+
902
+ # ── generation ───────────────────────────────────────────────────
903
+ def prepare_inputs_for_generation(
904
+ self,
905
+ input_ids: torch.LongTensor,
906
+ past_key_values: Optional[Cache] = None,
907
+ attention_mask: Optional[torch.Tensor] = None,
908
+ inputs_embeds: Optional[torch.FloatTensor] = None,
909
+ position_ids: Optional[torch.LongTensor] = None,
910
+ use_cache: Optional[bool] = None,
911
+ **kwargs,
912
+ ) -> dict:
913
+ """Trim model inputs to the block `generate` is about to run.
914
+
915
+ Handled entirely by `GenerationMixin`: it slices `input_ids` down to the
916
+ tokens not yet in the cache. That slice is load-bearing rather than an
917
+ optimisation β€” resending the full prefix would recompute it and corrupt
918
+ the cache. Overridden only to keep the signature aligned with
919
+ `transformers` 5.x and to forward `position_ids`, which the base
920
+ implementation pops and re-slices.
921
+ """
922
+ return super().prepare_inputs_for_generation(
923
+ input_ids=input_ids,
924
+ past_key_values=past_key_values,
925
+ attention_mask=attention_mask,
926
+ inputs_embeds=inputs_embeds,
927
+ position_ids=position_ids,
928
+ use_cache=use_cache,
929
+ **kwargs,
930
+ )
931
+
932
+ # ── gradient checkpointing ───────────────────────────────────────
933
+ def gradient_checkpointing_enable(
934
+ self,
935
+ gradient_checkpointing_kwargs: Optional[dict] = None,
936
+ **kwargs,
937
+ ) -> None:
938
+ """Recompute decoder activations in the backward pass instead of storing them.
939
+
940
+ Trades roughly 20-30% step time for most of the activation memory, which
941
+ is what lets one 40GB card hold a large batch at 2K context. Only active
942
+ in training mode β€” `VortexModel.forward` gates on `self.training`.
943
+ """
944
+ super().gradient_checkpointing_enable(
945
+ gradient_checkpointing_kwargs=gradient_checkpointing_kwargs, **kwargs
946
+ )
947
+ self.model.gradient_checkpointing = True
948
+ if not hasattr(self.model, "_gradient_checkpointing_func"):
949
+ self.model._gradient_checkpointing_func = lambda fn, *args: checkpoint(
950
+ fn, *args, use_reentrant=False
951
+ )
952
+
953
+ def gradient_checkpointing_disable(self, **kwargs) -> None:
954
+ super().gradient_checkpointing_disable(**kwargs)
955
+ self.model.gradient_checkpointing = False
956
+
957
+
958
+ __all__ = [
959
+ "VortexConfig",
960
+ "VortexPreTrainedModel",
961
+ "VortexModel",
962
+ "VortexForCausalLM",
963
+ "VortexBlock",
964
+ "VortexAttention",
965
+ "VortexMLP",
966
+ "VortexRMSNorm",
967
+ "VortexRotaryEmbedding",
968
+ "build_rope_cache",
969
+ "apply_rope",
970
+ ]