ainouche-abderahmane commited on
Commit
1c795a3
·
verified ·
1 Parent(s): 635a726

Fix tokenizer class, fp16 numerics and repo layout

Browse files

- tokenizer_config: DebertaV2Tokenizer, not LlamaTokenizer. The Llama
wrapper builds a BPE fast tokenizer from this SentencePiece Unigram
vocabulary, which re-segmented Latin text ('wesh rak khoya' as 5
pieces instead of 3), added no [CLS]/[SEP], and padded left - under
which the sequence-classification head pools a pad token.
- RMSNorm computes its statistic in float32. In float16 the square
overflowed (activations reach 316; 316^2 = 99856 > 65504), so the
encoder returned an all-zero hidden state. fp32 outputs are
bit-identical to before (max abs diff exactly 0.0).
- Weights unchanged: model.safetensors sha256 is the same file.
- Removed the duplicated hub/ tree, __pycache__ and stray tokenizer
copies; the flattened modeling_dzair.py is the code that ships.
- Card rewritten with every number traced to export_report.json.

README.md CHANGED
@@ -22,7 +22,7 @@ model-index:
22
  - name: DZAIR
23
  results:
24
  - task:
25
- type: feature-extraction
26
  name: Sentiment analysis (Latin Arabizi)
27
  dataset:
28
  type: narabizi-sentiment
@@ -35,7 +35,7 @@ model-index:
35
  value: 0.5961
36
  name: Macro F1 (10-seed mean)
37
  - task:
38
- type: feature-extraction
39
  name: Sentiment analysis (forum Arabic)
40
  dataset:
41
  type: ranim-sentiment
@@ -121,7 +121,9 @@ texts = [
121
  "ya3tik saha khoya, bon courage f projet", # Arabizi + French code-switch
122
  ]
123
 
124
- # Lowercase Latin input first for optimal tokenization
 
 
125
  inputs = tokenizer([t.lower() for t in texts], padding=True, return_tensors="pt")
126
 
127
  with torch.inference_mode():
@@ -237,15 +239,20 @@ micro-batches in a single run with no divergence.
237
 
238
  | file | size | contents |
239
  |---|---|---|
240
- | `model.safetensors` | about 421 MB | averaged discriminator weights |
241
  | `config.json` | 1 KB | architecture plus `auto_map` for `trust_remote_code` |
242
- | `modeling_dzair.py` | 58 KB | the architecture in code |
243
- | `tokenizer.model`, `tokenizer_config.json` | about 1.0 MB | 48k Unigram via `LlamaTokenizer`; `[PAD]`/`[UNK]`/`[CLS]`/`[SEP]`/`[MASK]` at ids 0–4 |
244
- | `dzair-tok-48k.model`, `dzair-tok-48k.vocab`, `dzair-tok-48k.metadata.json`, `export.proof.json` | about 2.0 MB | provenance copies of the tokenizer export |
245
- | `hub/` | 68 KB | verbatim source package for audit |
246
-
247
- Weight variants beside fp32: fp16 at 212MB, ONNX fp32 at 425MB, ONNX int8 at
248
- 109MB.
 
 
 
 
 
249
 
250
  ## Reproduction
251
 
 
22
  - name: DZAIR
23
  results:
24
  - task:
25
+ type: text-classification
26
  name: Sentiment analysis (Latin Arabizi)
27
  dataset:
28
  type: narabizi-sentiment
 
35
  value: 0.5961
36
  name: Macro F1 (10-seed mean)
37
  - task:
38
+ type: text-classification
39
  name: Sentiment analysis (forum Arabic)
40
  dataset:
41
  type: ranim-sentiment
 
121
  "ya3tik saha khoya, bon courage f projet", # Arabizi + French code-switch
122
  ]
123
 
124
+ # Lowercase Latin input first. The tokenizer wraps each row as
125
+ # [CLS] ... [SEP], pads on the right, and segments identically to the
126
+ # SentencePiece model the encoder was pretrained with.
127
  inputs = tokenizer([t.lower() for t in texts], padding=True, return_tensors="pt")
128
 
129
  with torch.inference_mode():
 
239
 
240
  | file | size | contents |
241
  |---|---|---|
242
+ | `model.safetensors` | 421.2 MB | folded discriminator backbone (E_G + delta) |
243
  | `config.json` | 1 KB | architecture plus `auto_map` for `trust_remote_code` |
244
+ | `modeling_dzair.py` | 59 KB | the architecture in one self-contained file |
245
+ | `tokenizer.model`, `tokenizer_config.json` | about 1.0 MB | 48k SentencePiece Unigram via `DebertaV2Tokenizer`; `[PAD]`/`[UNK]`/`[CLS]`/`[SEP]`/`[MASK]` at ids 0–4, `[CLS] … [SEP]` wrapping, right padding |
246
+ | `tokenizer_rules.yaml` | 2 KB | the versioned normalisation rules the tokenizer was trained under |
247
+ | `export_report.json` | 1 KB | measured sizes, SHA-256 per file, parameter counts, reload parity |
248
+
249
+ `modeling_dzair.py` is the hub package flattened into one file by the export;
250
+ it is the code that ships, and the export proves the staged directory reloads
251
+ to bit-identical weights before writing anything.
252
+
253
+ Weight variants beside fp32: `DZAIR-FP16` at 210.6 MB (0.999999 cosine),
254
+ `DZAIR-ONNX` at 424.0 MB (1.000000), `DZAIR-ONNX-INT8` at 108.4 MB
255
+ (0.999579).
256
 
257
  ## Reproduction
258
 
config.json CHANGED
@@ -1,15 +1,22 @@
1
  {
 
2
  "architectures": [
3
  "DzairModel"
4
  ],
5
  "attention_probs_dropout_prob": 0.1,
 
6
  "cls_token_id": 2,
 
 
7
  "dtype": "float32",
 
 
8
  "generator_hidden_size": 384,
9
  "generator_intermediate_size": 1024,
10
  "hidden_dropout_prob": 0.1,
11
  "hidden_size": 768,
12
  "intermediate_size": 1792,
 
13
  "layer_norm_eps": 1e-05,
14
  "mask_token_id": 4,
15
  "max_position_embeddings": 512,
@@ -19,11 +26,20 @@
19
  "num_hidden_layers": 12,
20
  "num_key_value_heads": 4,
21
  "pad_token_id": 0,
 
 
22
  "qk_norm": true,
23
  "rope_theta": 10000.0,
24
  "sep_token_id": 3,
25
  "share_generator_embeddings": true,
26
- "transformers_version": "4.57.6",
 
 
 
 
 
 
 
27
  "vocab_size": 48000,
28
  "auto_map": {
29
  "AutoConfig": "modeling_dzair.DzairConfig",
 
1
  {
2
+ "add_cross_attention": false,
3
  "architectures": [
4
  "DzairModel"
5
  ],
6
  "attention_probs_dropout_prob": 0.1,
7
+ "bos_token_id": null,
8
  "cls_token_id": 2,
9
+ "cross_attention_hidden_size": null,
10
+ "decoder_start_token_id": null,
11
  "dtype": "float32",
12
+ "eos_token_id": null,
13
+ "finetuning_task": null,
14
  "generator_hidden_size": 384,
15
  "generator_intermediate_size": 1024,
16
  "hidden_dropout_prob": 0.1,
17
  "hidden_size": 768,
18
  "intermediate_size": 1792,
19
+ "is_decoder": false,
20
  "layer_norm_eps": 1e-05,
21
  "mask_token_id": 4,
22
  "max_position_embeddings": 512,
 
26
  "num_hidden_layers": 12,
27
  "num_key_value_heads": 4,
28
  "pad_token_id": 0,
29
+ "prefix": null,
30
+ "pruned_heads": {},
31
  "qk_norm": true,
32
  "rope_theta": 10000.0,
33
  "sep_token_id": 3,
34
  "share_generator_embeddings": true,
35
+ "task_specific_params": null,
36
+ "tf_legacy_loss": false,
37
+ "tie_encoder_decoder": false,
38
+ "tie_word_embeddings": true,
39
+ "tokenizer_class": null,
40
+ "torchscript": false,
41
+ "transformers_version": "5.17.0",
42
+ "use_bfloat16": false,
43
  "vocab_size": 48000,
44
  "auto_map": {
45
  "AutoConfig": "modeling_dzair.DzairConfig",
dzair-tok-48k.metadata.json DELETED
@@ -1,12 +0,0 @@
1
- {
2
- "evidence_version": 1,
3
- "model_sha256": "6d50d8c119d0e9e15e548a473269a1e8e24609682e0e7243432a0a31de7897e6",
4
- "name": "dzair-tok-48k",
5
- "normalisation_version": 1,
6
- "normalization": "identity",
7
- "seed": 42,
8
- "tokenizer_version": "1.0.0",
9
- "train_sha256": "e9a6b5b2a233834bcb8486104e5e863c45deef92683294ed031656ef8a504c28",
10
- "train_txt": "train.txt",
11
- "vocab_size": 48000
12
- }
 
 
 
 
 
 
 
 
 
 
 
 
 
dzair-tok-48k.model DELETED
@@ -1,3 +0,0 @@
1
- version https://git-lfs.github.com/spec/v1
2
- oid sha256:6d50d8c119d0e9e15e548a473269a1e8e24609682e0e7243432a0a31de7897e6
3
- size 967834
 
 
 
 
dzair-tok-48k.vocab DELETED
The diff for this file is too large to render. See raw diff
 
export.proof.json DELETED
@@ -1,7 +0,0 @@
1
- {
2
- "encoding_mismatches": 0,
3
- "model_sha256": "6d50d8c119d0e9e15e548a473269a1e8e24609682e0e7243432a0a31de7897e6",
4
- "name": "dzair-tok-48k",
5
- "probe_sentences": 500,
6
- "rules": "tokenizer_rules.yaml"
7
- }
 
 
 
 
 
 
 
 
export_report.json ADDED
@@ -0,0 +1,45 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "repo": "DZAIR",
3
+ "source_checkpoint": "artifacts/train-checkpoints/dzair_v3_final.pt",
4
+ "source_sha256": "a35953b590d7d4d4b78ac7529b2801725aa5fdddc7fd88bf3b9f04c28d381e56",
5
+ "released_params": 105304320,
6
+ "backbone_params": 68440320,
7
+ "hidden_size": 768,
8
+ "num_hidden_layers": 12,
9
+ "files": {
10
+ "README.md": {
11
+ "bytes": 12103,
12
+ "sha256": "e2ccaa35a71a8591307422f43b1d18a1346751f8b40fea32df019b412ef2993d"
13
+ },
14
+ "config.json": {
15
+ "bytes": 1478,
16
+ "sha256": "86e324fff8b2ff7e93b471630e38f3bbe81f20cd24ad2f275ed91fdb899c29f9"
17
+ },
18
+ "model.safetensors": {
19
+ "bytes": 421228128,
20
+ "sha256": "b2448ca8dbc015cd90e13fce851122a15ac2855ff445c09f4af36d382709c006"
21
+ },
22
+ "modeling_dzair.py": {
23
+ "bytes": 58984,
24
+ "sha256": "4a339ae69b7f25985a4308a81bbf6c2f4ce2ba8e529adc0b1fa92d8d3f157a6a"
25
+ },
26
+ "tokenizer.model": {
27
+ "bytes": 967834,
28
+ "sha256": "6d50d8c119d0e9e15e548a473269a1e8e24609682e0e7243432a0a31de7897e6"
29
+ },
30
+ "tokenizer_config.json": {
31
+ "bytes": 563,
32
+ "sha256": "5a547e7c63af5a3183a3f929171e0d947022769a396e2ef63f3e67c6b015b251"
33
+ },
34
+ "tokenizer_rules.yaml": {
35
+ "bytes": 2058,
36
+ "sha256": "516b5a05d572fd229130f19f66648b23effb0d437836c03069858273b8e081c1"
37
+ }
38
+ },
39
+ "verification": {
40
+ "state_max_abs_diff": 0.0,
41
+ "hidden_max_abs_diff": 0.0,
42
+ "hidden_max_abs": 5.477493762969971,
43
+ "cosine": 0.9999999999999539
44
+ }
45
+ }
hub/__init__.py DELETED
@@ -1,44 +0,0 @@
1
- """DZAIR encoder hub module (base + small sizes).
2
-
3
- This module ships verbatim in releases. It imports nothing back from the package.
4
- """
5
-
6
- from dzair.hub.dzair_base.configuration import DZAIR_BASE_CONFIG, DZAIR_SMALL_CONFIG, DzairConfig
7
- from dzair.hub.dzair_base.modeling import (
8
- RTD_LOSS_WEIGHT,
9
- DzairEncoderOutput,
10
- DzairForMaskedLM,
11
- DzairForSequenceClassification,
12
- DzairForTokenClassification,
13
- DzairModel,
14
- DzairPreTrainedModel,
15
- DzairRTDOutput,
16
- MaskSpec,
17
- ResumeCompat,
18
- draw_token_mask,
19
- draw_word_mask,
20
- load_pretrain_state,
21
- load_resume_weights,
22
- set_gradient_checkpointing,
23
- )
24
-
25
- __all__ = [
26
- "DZAIR_BASE_CONFIG",
27
- "DZAIR_SMALL_CONFIG",
28
- "RTD_LOSS_WEIGHT",
29
- "DzairConfig",
30
- "DzairEncoderOutput",
31
- "DzairForMaskedLM",
32
- "DzairForSequenceClassification",
33
- "DzairForTokenClassification",
34
- "DzairModel",
35
- "DzairPreTrainedModel",
36
- "DzairRTDOutput",
37
- "MaskSpec",
38
- "ResumeCompat",
39
- "draw_token_mask",
40
- "draw_word_mask",
41
- "load_pretrain_state",
42
- "load_resume_weights",
43
- "set_gradient_checkpointing",
44
- ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
hub/__pycache__/__init__.cpython-312.pyc DELETED
Binary file (1.06 kB)
 
hub/dzair_base/__init__.py DELETED
@@ -1,32 +0,0 @@
1
- """DZAIR encoder hub module (base + small sizes).
2
-
3
- This module ships verbatim in releases. It imports nothing back from the package.
4
- """
5
-
6
- from dzair.hub.dzair_base.configuration import DZAIR_BASE_CONFIG, DZAIR_SMALL_CONFIG, DzairConfig
7
- from dzair.hub.dzair_base.modeling import (
8
- DzairEncoderOutput,
9
- DzairForMaskedLM,
10
- DzairForSequenceClassification,
11
- DzairForTokenClassification,
12
- DzairModel,
13
- DzairPreTrainedModel,
14
- ResumeCompat,
15
- load_pretrain_state,
16
- load_resume_weights,
17
- )
18
-
19
- __all__ = [
20
- "DZAIR_BASE_CONFIG",
21
- "DZAIR_SMALL_CONFIG",
22
- "DzairConfig",
23
- "DzairEncoderOutput",
24
- "DzairForMaskedLM",
25
- "DzairForSequenceClassification",
26
- "DzairForTokenClassification",
27
- "DzairModel",
28
- "DzairPreTrainedModel",
29
- "ResumeCompat",
30
- "load_pretrain_state",
31
- "load_resume_weights",
32
- ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
hub/dzair_base/__pycache__/__init__.cpython-312.pyc DELETED
Binary file (868 Bytes)
 
hub/dzair_base/__pycache__/configuration.cpython-312.pyc DELETED
Binary file (5.57 kB)
 
hub/dzair_base/__pycache__/modeling.cpython-312.pyc DELETED
Binary file (69.1 kB)
 
hub/dzair_base/configuration.py DELETED
@@ -1,146 +0,0 @@
1
- """Configuration for DZAIR encoder family (base + small)."""
2
-
3
- from __future__ import annotations
4
-
5
- from typing import Any
6
-
7
- from transformers import PretrainedConfig
8
-
9
- _BASE_HIDDEN_SIZE = 768
10
- _BASE_LAYERS = 12
11
- _SMALL_HIDDEN_SIZE = 384
12
- _SMALL_LAYERS = 6
13
-
14
-
15
- class DzairConfig(PretrainedConfig):
16
- """DZAIR encoder configuration.
17
-
18
- Every architectural choice is a declared field so ``config.json``
19
- round-trips exactly (base and small share this class, never a hidden
20
- ``arch`` object). ``**kwargs`` forwards only transformers-managed keys
21
- (e.g. ``transformers_version``) to ``PretrainedConfig``.
22
-
23
- Field names match the published `config.json`. Two sizes share this config:
24
- - base: 12Lx768 (discriminator, grouped-query 12Q/4KV) + 3Lx384 generator,
25
- shared embeddings
26
- - small: 6Lx384 (discriminator, grouped-query 6Q/2KV) + 3Lx384 generator,
27
- shared embeddings
28
- """
29
-
30
- model_type = "dzair"
31
-
32
- def __init__(
33
- self,
34
- vocab_size: int = 48000,
35
- hidden_size: int = 768,
36
- intermediate_size: int = 1792,
37
- num_attention_heads: int = 12,
38
- num_key_value_heads: int = 0,
39
- num_hidden_layers: int = 12,
40
- num_generator_layers: int = 3,
41
- generator_hidden_size: int = 384,
42
- generator_intermediate_size: int = 1024,
43
- max_position_embeddings: int = 512,
44
- rope_theta: float = 10000.0,
45
- hidden_dropout_prob: float = 0.1,
46
- attention_probs_dropout_prob: float = 0.1,
47
- layer_norm_eps: float = 1e-5,
48
- pad_token_id: int = 0,
49
- cls_token_id: int = 2,
50
- sep_token_id: int = 3,
51
- mask_token_id: int = 4,
52
- tie_word_embeddings: bool = True,
53
- share_generator_embeddings: bool = False,
54
- qk_norm: bool = False,
55
- **kwargs: Any,
56
- ) -> None:
57
- if hidden_size % num_attention_heads != 0:
58
- msg = f"hidden_size {hidden_size} must split over {num_attention_heads} heads"
59
- raise ValueError(msg)
60
- if num_key_value_heads == 0:
61
- num_key_value_heads = num_attention_heads
62
- if num_attention_heads % num_key_value_heads != 0:
63
- msg = (
64
- f"{num_attention_heads} query heads must split over "
65
- f"{num_key_value_heads} key-value heads"
66
- )
67
- raise ValueError(msg)
68
- if generator_hidden_size % 64 != 0:
69
- msg = f"generator_hidden_size {generator_hidden_size} must be a multiple of 64"
70
- raise ValueError(msg)
71
- self.vocab_size = vocab_size
72
- self.hidden_size = hidden_size
73
- self.intermediate_size = intermediate_size
74
- self.num_attention_heads = num_attention_heads
75
- self.num_key_value_heads = num_key_value_heads
76
- self.num_hidden_layers = num_hidden_layers
77
- self.num_generator_layers = num_generator_layers
78
- self.generator_hidden_size = generator_hidden_size
79
- self.generator_intermediate_size = generator_intermediate_size
80
- self.max_position_embeddings = max_position_embeddings
81
- self.rope_theta = rope_theta
82
- self.hidden_dropout_prob = hidden_dropout_prob
83
- self.attention_probs_dropout_prob = attention_probs_dropout_prob
84
- self.layer_norm_eps = layer_norm_eps
85
- self.share_generator_embeddings = share_generator_embeddings
86
- self.qk_norm = qk_norm
87
-
88
- super().__init__(
89
- pad_token_id=pad_token_id,
90
- cls_token_id=cls_token_id,
91
- sep_token_id=sep_token_id,
92
- tie_word_embeddings=tie_word_embeddings,
93
- **kwargs,
94
- )
95
- # Ensure mask_token_id and explicit IDs are preserved as ints
96
- self.pad_token_id = pad_token_id
97
- self.cls_token_id = cls_token_id
98
- self.sep_token_id = sep_token_id
99
- self.mask_token_id = mask_token_id
100
-
101
- @property
102
- def head_size(self) -> int:
103
- return self.hidden_size // self.num_attention_heads
104
-
105
- @property
106
- def generator_num_heads(self) -> int:
107
- """Generator query heads at head_dim 64 (always divides, checked above)."""
108
- return self.generator_hidden_size // 64
109
-
110
- @property
111
- def kv_dim(self) -> int:
112
- """Key/value width: key-value heads at the trunk head_dim."""
113
- return self.num_key_value_heads * self.head_size
114
-
115
- @property
116
- def is_base(self) -> bool:
117
- return self.hidden_size == _BASE_HIDDEN_SIZE and self.num_hidden_layers == _BASE_LAYERS
118
-
119
- @property
120
- def is_small(self) -> bool:
121
- return self.hidden_size == _SMALL_HIDDEN_SIZE and self.num_hidden_layers == _SMALL_LAYERS
122
-
123
-
124
- # Predefined configurations. Both sizes share the generator embedding table
125
- # with the discriminator (GDES, DeBERTaV3) and apply QK-norm; the FFN
126
- # intermediate is 128-aligned for tensor cores (1792 = 14x128, 1024 = 8x128).
127
- DZAIR_BASE_CONFIG = DzairConfig(
128
- num_key_value_heads=4,
129
- share_generator_embeddings=True,
130
- qk_norm=True,
131
- )
132
- DZAIR_SMALL_CONFIG = DzairConfig(
133
- hidden_size=384,
134
- intermediate_size=1024,
135
- num_attention_heads=6,
136
- num_key_value_heads=2,
137
- num_hidden_layers=6,
138
- num_generator_layers=3,
139
- generator_hidden_size=384,
140
- generator_intermediate_size=1024,
141
- share_generator_embeddings=True,
142
- qk_norm=True,
143
- )
144
-
145
-
146
- __all__ = ["DZAIR_BASE_CONFIG", "DZAIR_SMALL_CONFIG", "DzairConfig"]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
hub/dzair_base/modeling.py DELETED
@@ -1,1335 +0,0 @@
1
- """DZAIR encoder: RTD + GDES (DeBERTaV3 objective) with ModernBERT-speed architecture.
2
-
3
- Architecture: pre-RMSNorm, RoPE, SwiGLU, fused scaled-dot-product
4
- attention (FlashAttention-2 path when available) attending globally,
5
- single-chunk sequences.
6
- Objective: RTD on all tokens. Generator MLM corrupts; GDES detaches
7
- generator embeddings.
8
- """
9
-
10
- from __future__ import annotations
11
-
12
- import copy
13
- import hashlib
14
- import math
15
- from dataclasses import dataclass
16
- from pathlib import Path
17
- from typing import Any, ClassVar
18
-
19
- import torch
20
- from torch import Tensor, _dynamo, nn
21
- from torch.nn import functional
22
- from torch.utils import checkpoint as checkpoint_utils
23
- from transformers import PretrainedConfig, PreTrainedModel
24
- from transformers.utils.generic import ModelOutput
25
-
26
- from dzair.hub.dzair_base.configuration import DzairConfig
27
-
28
- IGNORE_INDEX = -100
29
-
30
- # ELECTRA (Clark et al., 2020, §3.3): small models weight the discriminator
31
- # loss at 50 relative to the generator MLM loss.
32
- RTD_LOSS_WEIGHT = 50.0
33
-
34
- # BERT 80/10/10 corruption splits (Devlin et al., 2019): below REPLACE the
35
- # token becomes [MASK], below REPLACE+RANDOM it becomes a random vocab id,
36
- # otherwise it is kept (but still predicted by the generator).
37
- MASK_REPLACE_CUTOFF = 0.8
38
- MASK_RANDOM_CUTOFF = 0.9
39
-
40
-
41
- def _apply_rope(x: Tensor, cos: Tensor, sin: Tensor) -> Tensor:
42
- """Apply rotary positional embeddings to half the head dim."""
43
- x1, x2 = x.chunk(2, dim=-1)
44
- return torch.cat((x1 * cos - x2 * sin, x1 * sin + x2 * cos), dim=-1)
45
-
46
-
47
- def _build_rope_cache(
48
- max_seq_len: int, head_dim: int, theta: float, device: torch.device
49
- ) -> tuple[Tensor, Tensor]:
50
- """Build RoPE cos/sin cache for sequence length up to max_seq_len."""
51
- inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=device).float() / head_dim))
52
- t = torch.arange(max_seq_len, device=device).float()
53
- freqs = torch.outer(t, inv_freq)
54
- cos = freqs.cos().to(torch.get_default_dtype())
55
- sin = freqs.sin().to(torch.get_default_dtype())
56
- return cos, sin
57
-
58
-
59
- class RMSNorm(nn.Module):
60
- """Root Mean Square Layer Normalization (affine weight, no bias)."""
61
-
62
- def __init__(self, dim: int, eps: float = 1e-5) -> None:
63
- super().__init__()
64
- self.eps = eps
65
- self.weight = nn.Parameter(torch.ones(dim))
66
-
67
- def forward(self, x: Tensor) -> Tensor:
68
- norm = x.pow(2).mean(dim=-1, keepdim=True).add(self.eps).rsqrt()
69
- return x * norm * self.weight
70
-
71
-
72
- class SwiGLU(nn.Module):
73
- """Swish-Gated Linear Unit."""
74
-
75
- def forward(self, x: Tensor) -> Tensor:
76
- x, gate = x.chunk(2, dim=-1)
77
- return x * functional.silu(gate)
78
-
79
-
80
- class FeedForward(nn.Module):
81
- """Pre-RMSNorm SwiGLU FFN with dropout."""
82
-
83
- def __init__(self, config: DzairConfig) -> None:
84
- super().__init__()
85
- self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps)
86
- self.up = nn.Linear(config.hidden_size, 2 * config.intermediate_size, bias=False)
87
- self.act = SwiGLU()
88
- self.down = nn.Linear(config.intermediate_size, config.hidden_size, bias=False)
89
- self.dropout = nn.Dropout(config.hidden_dropout_prob)
90
-
91
- def forward(self, x: Tensor) -> Tensor:
92
- residual = x
93
- x = self.norm(x)
94
- x = self.up(x)
95
- x = self.act(x)
96
- x = self.dropout(self.down(x))
97
- return residual + x
98
-
99
-
100
- class Attention(nn.Module):
101
- """Grouped-query attention with RoPE: every layer attends globally.
102
-
103
- Query heads share fewer key/value heads (``num_key_value_heads`` groups).
104
- Local-window alternation was cut 2026-09-14: at 512 tokens it saves
105
- ~4% wall-clock (measured FLOP arithmetic) while full attention is the
106
- literature default every baseline trains — the deviation bought
107
- complexity without evidence. Fused projections, bias-free, pre-RMSNorm.
108
- """
109
-
110
- def __init__(self, config: DzairConfig) -> None:
111
- super().__init__()
112
- self.config = config
113
- self.num_heads = config.num_attention_heads
114
- self.num_kv_heads = config.num_key_value_heads
115
- self.head_size = config.head_size
116
- self.scale = 1.0 / math.sqrt(self.head_size)
117
-
118
- # Separate Q and fused KV projections (bias-free for FA-2 compatibility).
119
- # KV groups repeat to the query count at forward time.
120
- self.q_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=False)
121
- self.kv_proj = nn.Linear(config.hidden_size, 2 * config.kv_dim, bias=False)
122
- self.out_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=False)
123
-
124
- self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps)
125
- # QK-norm (Gemma 2/3 practice): one shared RMSNorm over head_dim
126
- # applied to queries and keys before RoPE. Norm-then-rotate is a
127
- # fixed convention, not a commutation (rotation mixes dims, so an
128
- # affine weight does not commute with it) — the trained weights
129
- # bake in this order, so it must never change under them.
130
- self.qk_norm: RMSNorm | None = (
131
- RMSNorm(self.head_size, eps=config.layer_norm_eps) if config.qk_norm else None
132
- )
133
- self.dropout = nn.Dropout(config.attention_probs_dropout_prob)
134
-
135
- # RoPE cache (non-persistent, rebuilt on first use)
136
- self._cos: Tensor | None = None
137
- self._sin: Tensor | None = None
138
- self._cache_len = 0
139
-
140
- def _get_rope(self, seq_len: int, device: torch.device) -> tuple[Tensor, Tensor]:
141
- if self._cos is None or self._cache_len < seq_len or self._cos.device != device:
142
- self._cos, self._sin = _build_rope_cache(
143
- max(seq_len, self.config.max_position_embeddings),
144
- self.head_size,
145
- self.config.rope_theta,
146
- device,
147
- )
148
- self._cache_len = max(seq_len, self.config.max_position_embeddings)
149
- return self._cos[:seq_len], self._sin[:seq_len]
150
-
151
- def forward(
152
- self,
153
- x: Tensor,
154
- attention_mask: Tensor | None = None,
155
- is_causal: bool = False,
156
- ) -> Tensor:
157
- """x: [B, T, D], attention_mask: [B, T] (1=keep, 0=pad). Returns [B, T, D]."""
158
- batch_size, seq_len, _ = x.shape
159
-
160
- # Pre-norm
161
- x_norm = self.norm(x)
162
-
163
- # Grouped-query projections.
164
- q = self.q_proj(x_norm) # [B, T, D]
165
- kv = self.kv_proj(x_norm) # [B, T, 2 * kv_dim]
166
- k, v = kv.chunk(2, dim=-1)
167
-
168
- # Reshape for attention: Q [B, H, T, head_dim], K/V [B, KV, T, head_dim].
169
- q = q.view(batch_size, seq_len, self.num_heads, self.head_size).transpose(1, 2)
170
- k = k.view(batch_size, seq_len, self.num_kv_heads, self.head_size).transpose(1, 2)
171
- v = v.view(batch_size, seq_len, self.num_kv_heads, self.head_size).transpose(1, 2)
172
- # Repeat KV groups to the query count (exact: heads split evenly, checked).
173
- repeat = self.num_heads // self.num_kv_heads
174
- if repeat > 1:
175
- k = k.repeat_interleave(repeat, dim=1)
176
- v = v.repeat_interleave(repeat, dim=1)
177
-
178
- if self.qk_norm is not None:
179
- q = self.qk_norm(q)
180
- k = self.qk_norm(k)
181
-
182
- # RoPE (cast to the working dtype: an fp32 cache multiplied into bf16
183
- # queries upcasts them and drops out of the fused-attention fast path)
184
- cos, sin = self._get_rope(seq_len, x.device)
185
- cos = cos.unsqueeze(0).unsqueeze(0).to(x.dtype) # [1, 1, T, head_dim/2]
186
- sin = sin.unsqueeze(0).unsqueeze(0).to(x.dtype)
187
- q = _apply_rope(q, cos, sin)
188
- k = _apply_rope(k, cos, sin)
189
-
190
- # Scaled dot-product attention. A bool mask (True = attend) keeps the
191
- # fused fast path; the old additive float mask did not.
192
- attn_mask: Tensor | None = None
193
- if attention_mask is not None:
194
- attn_mask = attention_mask.to(torch.bool).view(batch_size, 1, 1, seq_len)
195
- # Guard against all-False mask rows (all-pad inputs): SDPA under
196
- # CUDA/Inductor produces NaNs when a row has zero attendable keys.
197
- positions = torch.arange(seq_len, device=x.device)
198
- has_key = attn_mask.any(dim=-1, keepdim=True)
199
- attn_mask = attn_mask | (~has_key & (positions == 0).view(1, 1, 1, seq_len))
200
-
201
- # Use PyTorch's scaled_dot_product_attention (uses FA-2 when available)
202
- attn_out = functional.scaled_dot_product_attention(
203
- q,
204
- k,
205
- v,
206
- attn_mask=attn_mask,
207
- dropout_p=self.config.attention_probs_dropout_prob if self.training else 0.0,
208
- is_causal=is_causal,
209
- scale=self.scale,
210
- )
211
-
212
- # Merge heads: [B, H, T, head_dim] -> [B, T, D]
213
- attn_out = attn_out.transpose(1, 2).contiguous().view(batch_size, seq_len, -1)
214
-
215
- # Zero out padding positions so unused positions never dominate
216
- # downstream means (the residual still carries the pad embedding;
217
- # losses and CLS pooling ignore pads by mask, which is what makes
218
- # this safe rather than the zeroing alone).
219
- if attention_mask is not None:
220
- attn_out = attn_out * attention_mask.view(batch_size, seq_len, 1).to(attn_out.dtype)
221
-
222
- # Output projection + residual
223
- out = self.out_proj(attn_out)
224
- out = self.dropout(out)
225
- return x + out
226
-
227
-
228
- class TransformerLayer(nn.Module):
229
- """Pre-RMSNorm transformer block: Attention + FFN."""
230
-
231
- def __init__(self, config: DzairConfig) -> None:
232
- super().__init__()
233
- self.attention = Attention(config)
234
- self.ffn = FeedForward(config)
235
-
236
- def forward(
237
- self,
238
- x: Tensor,
239
- attention_mask: Tensor | None = None,
240
- is_causal: bool = False,
241
- ) -> Tensor:
242
- x = self.attention(x, attention_mask, is_causal)
243
- return self.ffn(x)
244
-
245
-
246
- class Embeddings(nn.Module):
247
- """Token embeddings with RMSNorm and dropout.
248
-
249
- No positional embeddings (RoPE handles position). Single-chunk inputs
250
- only: ``[CLS] chunk [SEP]``.
251
- """
252
-
253
- def __init__(self, config: DzairConfig) -> None:
254
- super().__init__()
255
- self.word_embeddings = nn.Embedding(
256
- config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id
257
- )
258
- self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps)
259
- self.dropout = nn.Dropout(config.hidden_dropout_prob)
260
-
261
- def forward(self, input_ids: Tensor) -> Tensor:
262
- x = self.word_embeddings(input_ids)
263
- x = self.norm(x)
264
- return self.dropout(x)
265
-
266
-
267
- class Encoder(nn.Module):
268
- """Stack of transformer layers."""
269
-
270
- def __init__(self, config: DzairConfig) -> None:
271
- super().__init__()
272
- self.config = config
273
- self.layers = nn.ModuleList(
274
- TransformerLayer(config) for _ in range(config.num_hidden_layers)
275
- )
276
- self.gradient_checkpointing = False
277
-
278
- def forward(
279
- self,
280
- x: Tensor,
281
- attention_mask: Tensor | None = None,
282
- ) -> Tensor:
283
- for layer in self.layers:
284
- if self.gradient_checkpointing and self.training:
285
- x = checkpoint_utils.checkpoint(
286
- layer, x, attention_mask, False, use_reentrant=False
287
- )
288
- else:
289
- x = layer(x, attention_mask, is_causal=False) # bidirectional
290
- return x
291
-
292
- def forward_with_states(
293
- self,
294
- x: Tensor,
295
- attention_mask: Tensor | None = None,
296
- ) -> tuple[Tensor, tuple[Tensor, ...]]:
297
- """Forward pass returning the last output plus each layer's output."""
298
- states: list[Tensor] = []
299
- for layer in self.layers:
300
- if self.gradient_checkpointing and self.training:
301
- x = checkpoint_utils.checkpoint(
302
- layer, x, attention_mask, False, use_reentrant=False
303
- )
304
- else:
305
- x = layer(x, attention_mask, is_causal=False) # bidirectional
306
- states.append(x)
307
- return x, tuple(states)
308
-
309
-
310
- def set_gradient_checkpointing(model: nn.Module, value: bool) -> None:
311
- """Toggle activation checkpointing on every Encoder in a model.
312
-
313
- Plain attribute propagation, deliberately not via
314
- ``PreTrainedModel.gradient_checkpointing_enable`` whose signature drifted
315
- across transformers versions. The smoke test pins that outputs match and
316
- gradients flow with it on.
317
- """
318
- for module in model.modules():
319
- if isinstance(module, (Encoder, Generator)):
320
- module.gradient_checkpointing = value
321
-
322
-
323
- _ARCH_FIELDS: tuple[str, ...] = (
324
- "vocab_size",
325
- "hidden_size",
326
- "intermediate_size",
327
- "num_attention_heads",
328
- "num_key_value_heads",
329
- "num_hidden_layers",
330
- "num_generator_layers",
331
- "generator_hidden_size",
332
- "generator_intermediate_size",
333
- "max_position_embeddings",
334
- "rope_theta",
335
- "hidden_dropout_prob",
336
- "attention_probs_dropout_prob",
337
- "layer_norm_eps",
338
- "pad_token_id",
339
- "cls_token_id",
340
- "sep_token_id",
341
- "mask_token_id",
342
- "tie_word_embeddings",
343
- "share_generator_embeddings",
344
- "qk_norm",
345
- )
346
-
347
- _CONFIG_MISMATCH_MSG = (
348
- "explicit config disagrees with the checkpoint's stored config on {field}: "
349
- "explicit={explicit!r} stored={stored!r} — pass config=None to trust the checkpoint"
350
- )
351
-
352
- _NO_STORED_CONFIG_MSG = (
353
- "checkpoint {path} carries no stored config and none was passed — "
354
- "pass config=<DzairConfig> explicitly"
355
- )
356
-
357
- _FOLD_MISSING_MSG = (
358
- "GDES checkpoint is missing {missing} — found prefixes: {prefixes}; "
359
- "cannot fold E_G + delta into the released embedding"
360
- )
361
-
362
-
363
- _CHECKSUM_MISMATCH_MSG = "checkpoint checksum mismatch for {path}"
364
-
365
-
366
- def _normalize_stored_config(stored_dict: dict[str, Any]) -> dict[str, Any]:
367
- """Replace a pre-GQA null key-value count with the full-MHA default.
368
-
369
- Runs written before grouped-query attention store no (or null)
370
- key-value count; all of them trained full multi-head attention.
371
- """
372
- normalized = dict(stored_dict)
373
- if normalized.get("num_key_value_heads") is None:
374
- normalized.pop("num_key_value_heads", None)
375
- return normalized
376
-
377
-
378
- def _check_config_match(config: DzairConfig, stored_dict: dict[str, Any]) -> None:
379
- """Raise on any recorded field the explicit config disagrees on.
380
-
381
- Fields the checkpoint predates (absent) or left null are not compared:
382
- the explicit config decides those, so era-appropriate explicit configs
383
- (global attention, full MHA, eval-only dropout) load instead of
384
- refusing on a formatting technicality.
385
- """
386
- for field in _ARCH_FIELDS:
387
- if field not in stored_dict or stored_dict[field] is None:
388
- continue
389
- explicit_value = getattr(config, field, None)
390
- if explicit_value != stored_dict[field]:
391
- raise ValueError(
392
- _CONFIG_MISMATCH_MSG.format(
393
- field=field, explicit=explicit_value, stored=stored_dict[field]
394
- )
395
- )
396
-
397
-
398
- def _read_pretrain_checkpoint(
399
- checkpoint_path: str | Path,
400
- config: DzairConfig | None,
401
- ) -> tuple[DzairConfig, dict[str, Tensor]]:
402
- """Resolve (config, state) from a pretraining checkpoint.
403
-
404
- The checkpoint's stored config wins unless an explicit config is passed;
405
- an explicit config that disagrees with the stored one on a field the
406
- checkpoint actually records raises instead of silently misloading.
407
- Fields the checkpoint predates (absent) or left null are not compared:
408
- the explicit config decides those, so era-appropriate explicit configs
409
- (global attention, full MHA, eval-only dropout) load instead of
410
- refusing on a formatting technicality. Stored nulls/absences for the
411
- key-value count mean the run predates grouped-query attention and
412
- trained full multi-head attention, so they normalize to the default —
413
- never to a silent mismatch.
414
- """
415
- path_obj = Path(checkpoint_path)
416
- sidecar = path_obj.parent / f"{path_obj.name}.sha256"
417
- if sidecar.is_file():
418
- want = sidecar.read_text(encoding="utf-8").strip()
419
- digest = hashlib.sha256()
420
- with path_obj.open("rb") as f:
421
- for chunk in iter(lambda: f.read(1 << 20), b""):
422
- digest.update(chunk)
423
- if digest.hexdigest() != want:
424
- raise ValueError(_CHECKSUM_MISMATCH_MSG.format(path=checkpoint_path))
425
- raw = torch.load(checkpoint_path, map_location="cpu", weights_only=True)
426
- if not isinstance(raw, dict):
427
- msg = f"checkpoint payload is not a mapping: {checkpoint_path}"
428
- raise TypeError(msg)
429
- inner = raw.get("model")
430
- state: dict[str, Tensor] = inner if isinstance(inner, dict) else raw
431
- stored = raw.get("config")
432
- stored_dict = stored if isinstance(stored, dict) else None
433
- if config is not None:
434
- if stored_dict is not None:
435
- _check_config_match(config, stored_dict)
436
- return copy.deepcopy(config), state
437
- if stored_dict is None:
438
- raise ValueError(_NO_STORED_CONFIG_MSG.format(path=checkpoint_path))
439
- return DzairConfig(**_normalize_stored_config(stored_dict)), state
440
-
441
-
442
- def _fold_shared_backbone(state_dict: dict[str, Tensor]) -> dict[str, Tensor]:
443
- """Fold a GDES checkpoint's shared table into one released embedding.
444
-
445
- The released table is ``proj(E_G) + Δ`` — the generator's table through
446
- the width bridge plus the discriminator's delta — with the input norm
447
- taken from the discriminator's ``input_norm``. Same-width (or pre-bridge)
448
- checkpoints skip the projection, exactly like the forward does.
449
- """
450
- out: dict[str, Tensor] = {}
451
- gen_key = "rtd_head.generator.embeddings.word_embeddings.weight"
452
- delta_key = "rtd_head.discriminator.delta_embeddings.weight"
453
- proj_key = "rtd_head.discriminator.gen_proj.weight"
454
- gen_table = state_dict.get(gen_key)
455
- delta = state_dict.get(delta_key)
456
- if gen_table is None or delta is None:
457
- missing = [k for k in (gen_key, delta_key) if k not in state_dict]
458
- prefixes = sorted({".".join(k.split(".")[:2]) if "." in k else k for k in state_dict})
459
- raise KeyError(_FOLD_MISSING_MSG.format(missing=missing, prefixes=prefixes[:8]))
460
- proj = state_dict.get(proj_key)
461
- if proj is None or gen_table.size(-1) == delta.size(-1):
462
- folded = gen_table + delta.to(gen_table.dtype)
463
- else:
464
- folded = gen_table.to(proj.dtype) @ proj.T + delta.to(proj.dtype)
465
- out["embeddings.word_embeddings.weight"] = folded
466
- for key, value in state_dict.items():
467
- if key.startswith("rtd_head.discriminator.input_norm."):
468
- out["embeddings.norm." + key[len("rtd_head.discriminator.input_norm.") :]] = value
469
- elif key.startswith("rtd_head.discriminator.encoder.") or key.startswith(
470
- "rtd_head.discriminator.norm."
471
- ):
472
- out[key[len("rtd_head.discriminator.") :]] = value
473
- return out
474
-
475
-
476
- _FUSED_SPLIT_MSG = (
477
- "cannot map fused {key}: expected ({fused}, {hidden}), "
478
- "or the target is grouped-query ({kv} KV heads over {nq} query heads) "
479
- "which a fused full-MHA table cannot feed without lossy subsampling — "
480
- "retrain or load into a full-MHA config"
481
- )
482
-
483
- _AMBIGUOUS_PROJ_MSG = (
484
- "checkpoint mixes fused ({fused}) and split ({split}) attention projections — "
485
- "refusing instead of guessing which one owns the layer"
486
- )
487
-
488
- _FUSED_BIAS_MSG = (
489
- "cannot map fused {key}: biased projections have no split target — "
490
- "retrain or load into a matching config"
491
- )
492
-
493
-
494
- def _unfuse_in_proj(state_dict: dict[str, Tensor], config: PretrainedConfig) -> dict[str, Tensor]:
495
- """Split fused full-MHA ``in_proj`` tables into ``q_proj`` + ``kv_proj``.
496
-
497
- Checkpoints written before grouped-query attention carry one fused
498
- QKV matrix per layer; current code keeps separate query and fused
499
- key/value projections. The split is exact only into full MHA
500
- (KV heads == query heads) with the canonical Q,K,V row order —
501
- anything else raises instead of silently remapping. Passes through
502
- states without fused tables untouched.
503
- """
504
- fused_keys = [k for k in state_dict if k.endswith("attention.in_proj.weight")]
505
- if not fused_keys:
506
- return state_dict
507
- hidden = int(config.hidden_size)
508
- num_queries = int(config.num_attention_heads)
509
- num_kv = int(getattr(config, "num_key_value_heads", 0) or num_queries)
510
- split_keys = [
511
- k
512
- for k in state_dict
513
- if k.endswith("attention.q_proj.weight") or k.endswith("attention.kv_proj.weight")
514
- ]
515
- if split_keys:
516
- msg = _AMBIGUOUS_PROJ_MSG.format(fused=fused_keys[0], split=split_keys[0])
517
- raise RuntimeError(msg)
518
- biased = [k for k in state_dict if k.endswith("attention.in_proj.bias")]
519
- if biased:
520
- msg = _FUSED_BIAS_MSG.format(key=biased[0])
521
- raise RuntimeError(msg)
522
- out = dict(state_dict)
523
- for key in fused_keys:
524
- weight = state_dict[key]
525
- if tuple(weight.shape) != (3 * hidden, hidden) or num_kv != num_queries:
526
- msg = _FUSED_SPLIT_MSG.format(
527
- key=key,
528
- fused=tuple(weight.shape),
529
- hidden=hidden,
530
- kv=num_kv,
531
- nq=num_queries,
532
- )
533
- raise RuntimeError(msg)
534
- prefix = key[: -len("in_proj.weight")]
535
- query, key_p, value = weight.split([hidden, hidden, hidden], dim=0)
536
- del out[key]
537
- out[prefix + "q_proj.weight"] = query
538
- out[prefix + "kv_proj.weight"] = torch.cat([key_p, value], dim=0)
539
- return out
540
-
541
-
542
- _GEN_TABLE_KEY = "rtd_head.generator.embeddings.word_embeddings.weight"
543
-
544
-
545
- def discriminator_backbone_state(
546
- state_dict: dict[str, Tensor], config: PretrainedConfig
547
- ) -> dict[str, Tensor]:
548
- """Map a pretraining checkpoint's discriminator weights onto ``DzairModel``.
549
-
550
- GDES checkpoints (shared table): the released embedding is the fold
551
- ``proj(E_G) + Δ`` — the generator's table through the width bridge plus
552
- the discriminator's delta — with the input norm taken from the
553
- discriminator's ``input_norm``. Independent checkpoints:
554
- ``rtd_head.discriminator.*`` maps verbatim minus the RTD
555
- classifier. A payload that is already a ``DzairModel`` state dict (no
556
- ``rtd_head`` prefix) passes through; ``strict=True`` on the caller's
557
- ``load_state_dict`` catches anything malformed. Fused full-MHA
558
- ``in_proj`` tables are split exactly (see ``_unfuse_in_proj``);
559
- generator-trunk keys never enter the mapping, so a fused generator
560
- neither helps nor breaks the fold.
561
- """
562
- shared = bool(getattr(config, "share_generator_embeddings", False))
563
- if not any(k.startswith("rtd_head.") for k in state_dict):
564
- return {
565
- (key[len("dzair.") :] if key.startswith("dzair.") else key): value
566
- for key, value in state_dict.items()
567
- }
568
- relevant = {
569
- key: value
570
- for key, value in state_dict.items()
571
- if key.startswith("rtd_head.discriminator.")
572
- or key == _GEN_TABLE_KEY
573
- or key.startswith("dzair.")
574
- }
575
- state_dict = _unfuse_in_proj(relevant, config)
576
- out: dict[str, Tensor] = {}
577
- if shared:
578
- return _fold_shared_backbone(state_dict)
579
- for key, value in state_dict.items():
580
- if key.startswith("rtd_head.discriminator.") and not key.startswith(
581
- "rtd_head.discriminator.classifier"
582
- ):
583
- out[key[len("rtd_head.discriminator.") :]] = value
584
- elif key.startswith("dzair."):
585
- out[key[len("dzair.") :]] = value
586
- return out
587
-
588
-
589
- # Pretraining-only modules absent from older checkpoints: a checkpoint missing
590
- # exactly these still loads, everything else missing or unexpected still raises.
591
- _COMPAT_MISSING_SUBSTRINGS: tuple[str, ...] = ("gen_proj.",)
592
-
593
-
594
- _GENERATION_GAP_MSG = (
595
- "checkpoint uses independent discriminator embeddings "
596
- "('rtd_head.discriminator.embeddings.') but the model expects GDES "
597
- "('rtd_head.discriminator.delta_embeddings.'): no automatic migration — "
598
- "the v1 identity (E_D independent) cannot fold into E_G + delta without "
599
- "changing numerics; retrain or load into a share_generator_embeddings=False "
600
- "config"
601
- )
602
-
603
-
604
- def load_pretrain_state(model: nn.Module, state: dict[str, Tensor]) -> None:
605
- """Load a pretraining state dict across the width-bridge generation gap.
606
-
607
- Checkpoints written before the generator width bridge lack ``gen_proj``;
608
- anything else missing, misshapen, or unexpected still raises.
609
- Fused full-MHA ``in_proj`` tables are split exactly (see
610
- ``_unfuse_in_proj``).
611
- The independent-embeddings (v1) to GDES generation gap is
612
- refused loudly: silently mapping E_D onto delta would change numerics.
613
- """
614
- model_config = getattr(model, "config", None)
615
- if model_config is not None:
616
- gen_keys = {k: v for k, v in state.items() if k.startswith("rtd_head.generator.encoder.")}
617
- trunk_keys = {k: v for k, v in state.items() if k not in gen_keys}
618
- merged_state = _unfuse_in_proj(trunk_keys, model_config)
619
- if gen_keys:
620
- merged_state.update(_unfuse_in_proj(gen_keys, _generator_view(model_config)))
621
- state = merged_state
622
- own = model.state_dict()
623
- if any("rtd_head.discriminator.embeddings." in k for k in state) and any(
624
- "delta_embeddings" in k for k in own
625
- ):
626
- raise RuntimeError(_GENERATION_GAP_MSG)
627
- if any("delta_embeddings" in k for k in state) and any(
628
- "rtd_head.discriminator.embeddings." in k for k in own
629
- ):
630
- raise RuntimeError(_GENERATION_GAP_MSG)
631
- unexpected = [k for k in state if k not in own]
632
- if unexpected:
633
- msg = f"checkpoint holds unexpected keys: {unexpected[:8]}"
634
- raise RuntimeError(msg)
635
- merged: dict[str, Tensor] = {}
636
- absent: list[str] = []
637
- for key, value in own.items():
638
- if key not in state:
639
- absent.append(key)
640
- continue
641
- if value.shape != state[key].shape:
642
- msg = f"checkpoint shape mismatch for {key}: ckpt {tuple(state[key].shape)}"
643
- raise RuntimeError(msg)
644
- merged[key] = state[key]
645
- unaccounted = [k for k in absent if not any(s in k for s in _COMPAT_MISSING_SUBSTRINGS)]
646
- if unaccounted:
647
- msg = f"checkpoint lacks load-bearing keys: {unaccounted[:8]}"
648
- raise RuntimeError(msg)
649
- for key in absent:
650
- merged[key] = own[key]
651
- model.load_state_dict(merged, strict=True)
652
-
653
-
654
- # State keys from retired training objectives. A checkpoint carrying them
655
- # predates the current code: the trunk weights still load, the retired
656
- # heads do not come back. Centralized here so every loader agrees on
657
- # what "obsolete" means; anything else unexpected still raises.
658
- OBSOLETE_STATE_SUBSTRINGS: tuple[str, ...] = (
659
- "order_head.",
660
- "order_loss_ema",
661
- "token_loss_ema",
662
- "token_type_embeddings.",
663
- )
664
-
665
-
666
- @dataclass(frozen=True)
667
- class ResumeCompat:
668
- """How a checkpoint's weights mapped onto the current model."""
669
-
670
- generation: str # "same" (exact) or "legacy" (obsolete keys dropped)
671
- dropped: tuple[str, ...]
672
-
673
-
674
- def load_resume_weights(model: nn.Module, ckpt_model_state: dict[str, Tensor]) -> ResumeCompat:
675
- """Load training weights for an exact resume across code generations.
676
-
677
- Fused full-MHA tables split exactly (trunk and generator widths
678
- handled separately); retired keys drop loudly in the report. Any
679
- other missing, misshapen, or unexpected key raises — a half-mapped
680
- model never trains. The caller decides from ``generation`` whether
681
- the optimizer may be restored (``same``) or must restart fresh
682
- (``legacy``): stale momentum on a reshaped model is silent corruption.
683
- """
684
- raw = {k.removeprefix("_orig_mod."): v for k, v in ckpt_model_state.items()}
685
- model_config = getattr(model, "config", None)
686
- if model_config is not None:
687
- gen_keys = {k: v for k, v in raw.items() if k.startswith("rtd_head.generator.encoder.")}
688
- trunk_keys = {k: v for k, v in raw.items() if k not in gen_keys}
689
- raw = _unfuse_in_proj(trunk_keys, model_config)
690
- if gen_keys:
691
- raw.update(_unfuse_in_proj(gen_keys, _generator_view(model_config)))
692
- dropped = tuple(sorted({k for k in raw if any(s in k for s in OBSOLETE_STATE_SUBSTRINGS)}))
693
- kept = {k: v for k, v in raw.items() if k not in dropped}
694
- raw_model = getattr(model, "_orig_mod", model)
695
- own = raw_model.state_dict()
696
- unexpected = [k for k in kept if k not in own]
697
- if unexpected:
698
- msg = f"checkpoint holds unexpected keys: {unexpected[:8]}"
699
- raise RuntimeError(msg)
700
- missing = [k for k in own if k not in kept]
701
- if missing:
702
- msg = f"checkpoint lacks load-bearing keys: {missing[:8]}"
703
- raise RuntimeError(msg)
704
- for key, value in own.items():
705
- if value.shape != kept[key].shape:
706
- msg = f"checkpoint shape mismatch for {key}: ckpt {tuple(kept[key].shape)}"
707
- raise RuntimeError(msg)
708
- raw_model.load_state_dict(kept, strict=True)
709
- return ResumeCompat(generation="legacy" if dropped else "same", dropped=dropped)
710
-
711
-
712
- def _generator_view(config: DzairConfig) -> DzairConfig:
713
- """A config view sizing the generator trunk: narrow width, global attention.
714
-
715
- The generator keeps head_dim 64 and key-value groups proportional to the
716
- trunk; it always attends globally so corruption quality never depends on
717
- the discriminator's local window. Copies (never mutates) the trunk config.
718
- """
719
- view = copy.copy(config)
720
- view.hidden_size = config.generator_hidden_size
721
- view.intermediate_size = config.generator_intermediate_size
722
- view.num_attention_heads = config.generator_num_heads
723
- view.num_key_value_heads = max(
724
- 1, config.generator_num_heads * config.num_key_value_heads // config.num_attention_heads
725
- )
726
- if view.num_attention_heads % view.num_key_value_heads != 0:
727
- msg = (
728
- f"generator {view.num_attention_heads} query heads must split over "
729
- f"{view.num_key_value_heads} key-value heads"
730
- )
731
- raise ValueError(msg)
732
- view.num_hidden_layers = config.num_generator_layers
733
- return view
734
-
735
-
736
- class Generator(nn.Module):
737
- """Lightweight MLM generator for RTD corruption.
738
-
739
- GDES: embeddings shared, detached for discriminator. The generator trunk
740
- runs at ``generator_hidden_size`` behind a width projection only where it
741
- meets the discriminator (see ``Discriminator.gen_proj``); its own input
742
- and LM head stay in the narrow width with tied tables.
743
- """
744
-
745
- def __init__(self, config: DzairConfig) -> None:
746
- super().__init__()
747
- self.config = config
748
- view = _generator_view(config)
749
- self.embeddings = Embeddings(view)
750
- self.encoder = nn.ModuleList(
751
- TransformerLayer(view) for _ in range(config.num_generator_layers)
752
- )
753
- self.norm = RMSNorm(view.hidden_size, eps=config.layer_norm_eps)
754
- self.lm_head = nn.Linear(view.hidden_size, config.vocab_size, bias=False)
755
- # Tie output embeddings to input embeddings
756
- self.lm_head.weight = self.embeddings.word_embeddings.weight
757
- self.gradient_checkpointing = False
758
-
759
- def forward(
760
- self,
761
- input_ids: Tensor,
762
- attention_mask: Tensor | None = None,
763
- ) -> Tensor:
764
- x = self.embeddings(input_ids)
765
- for layer in self.encoder:
766
- if self.gradient_checkpointing and self.training:
767
- x = checkpoint_utils.checkpoint(
768
- layer, x, attention_mask, False, use_reentrant=False
769
- )
770
- else:
771
- x = layer(x, attention_mask, is_causal=False)
772
- x = self.norm(x)
773
- return self.lm_head(x)
774
-
775
-
776
- class Discriminator(nn.Module):
777
- """RTD discriminator: detects replaced tokens.
778
-
779
- Two embedding policies, selected by ``config.share_generator_embeddings``:
780
-
781
- - **GDES** (True): the discriminator reads ``proj(stop_grad(E_G)) + Δ``
782
- where ``E_G`` is the generator's own (narrow) table — generator MLM
783
- training shapes the table the discriminator reads — ``proj`` bridges
784
- the generator width to the trunk width, and ``Δ`` is this module's own
785
- table. Discriminator gradients flow to ``Δ`` and ``proj`` only, by
786
- construction. The released backbone folds ``proj(E_G) + Δ`` into one
787
- table at load time.
788
- - **Independent** (False, default): a private ``Embeddings`` table, as in
789
- classic ELECTRA. The pretraining head still passes the generator's
790
- table; in this mode it is unused, and the forward is a plain lookup.
791
- """
792
-
793
- def __init__(self, config: DzairConfig) -> None:
794
- super().__init__()
795
- self.config = config
796
- self.share_generator_embeddings = bool(config.share_generator_embeddings)
797
- if self.share_generator_embeddings:
798
- self.delta_embeddings = nn.Embedding(
799
- config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id
800
- )
801
- self.gen_proj = nn.Linear(config.generator_hidden_size, config.hidden_size, bias=False)
802
- self.input_norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps)
803
- self.input_dropout = nn.Dropout(config.hidden_dropout_prob)
804
- else:
805
- self.embeddings = Embeddings(config)
806
- self.encoder = Encoder(config)
807
- self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps)
808
- self.classifier = nn.Linear(config.hidden_size, 2, bias=False) # binary: original/replaced
809
-
810
- def forward(
811
- self,
812
- input_ids: Tensor,
813
- attention_mask: Tensor | None = None,
814
- generator_embeddings: Tensor | None = None,
815
- ) -> Tensor:
816
- """``generator_embeddings`` is the generator's table (required under GDES)."""
817
- if self.share_generator_embeddings:
818
- if generator_embeddings is None:
819
- msg = "share_generator_embeddings=True requires the generator table"
820
- raise ValueError(msg)
821
- base = functional.embedding(
822
- input_ids, generator_embeddings.detach(), padding_idx=self.config.pad_token_id
823
- )
824
- if base.size(-1) != self.config.hidden_size:
825
- base = self.gen_proj(base)
826
- x = base + self.delta_embeddings(input_ids)
827
- x = self.input_dropout(self.input_norm(x))
828
- else:
829
- x = self.embeddings(input_ids)
830
-
831
- x = self.encoder(x, attention_mask)
832
- x = self.norm(x)
833
- return self.classifier(x)
834
-
835
-
836
- def _dynamo_disabled[FnT](function: FnT) -> FnT:
837
- """Exclude CPU-scalar bookkeeping from the compiled graph.
838
-
839
- ``float()`` syncs inside the forward break Dynamo (measured: a
840
- ``Tensor.item()`` graph break every step) and stall the GPU for a value
841
- only needed as an eager loss scale. Falls back to a no-op when
842
- ``disable`` is unavailable; the pinned images always carry it.
843
- """
844
- disable = getattr(_dynamo, "disable", None)
845
- if disable is None:
846
- return function
847
- return disable(function)
848
-
849
-
850
- @_dynamo_disabled
851
- def draw_token_mask(candidate: Tensor, mask_prob: float | Tensor) -> Tensor:
852
- """Per-token uniform draw over candidate positions."""
853
- return candidate & (torch.rand(candidate.shape, device=candidate.device) < mask_prob)
854
-
855
-
856
- @_dynamo_disabled
857
- def draw_word_mask(candidate: Tensor, word_starts: Tensor, mask_prob: float | Tensor) -> Tensor:
858
- """One uniform draw per word; every candidate position in a chosen word masked.
859
-
860
- Positions before the first word start are never masked.
861
- """
862
- device = candidate.device
863
- length = candidate.size(-1)
864
- arange = torch.arange(length, device=device).expand_as(candidate)
865
- cur_start = torch.where(word_starts, arange, -1).cummax(dim=-1).values
866
- chosen = word_starts & candidate & (torch.rand(candidate.shape, device=device) < mask_prob)
867
- last_chosen = torch.where(chosen, arange, -1).cummax(dim=-1).values
868
- return candidate & (cur_start >= 0) & (cur_start == last_chosen)
869
-
870
-
871
- @dataclass(frozen=True)
872
- class MaskSpec:
873
- """What may be masked and how often (built per step from the schedule)."""
874
-
875
- special_ids: frozenset[int]
876
- vocab_size: int
877
- mask_token_id: int
878
- mask_prob: float | Tensor
879
-
880
-
881
- @_dynamo_disabled
882
- def _mask_inputs(
883
- input_ids: Tensor,
884
- eligible: Tensor,
885
- spec: MaskSpec,
886
- word_starts: Tensor | None = None,
887
- ) -> tuple[Tensor, Tensor]:
888
- """BERT 80/10/10 corruption. Returns (masked_input_ids, mlm_labels).
889
-
890
- Masking is whole-word when ``word_starts`` ([B, T] bool, True at
891
- word-initial pieces) is given, else per-token. Special ids (pad/cls/sep)
892
- and ineligible positions are never masked. Dynamic every step.
893
- """
894
- device = input_ids.device
895
- is_special = torch.zeros_like(input_ids, dtype=torch.bool)
896
- for sid in spec.special_ids:
897
- is_special |= input_ids == sid
898
- candidate = eligible & ~is_special
899
-
900
- if word_starts is not None:
901
- masked = draw_word_mask(candidate, word_starts, spec.mask_prob)
902
- else:
903
- masked = draw_token_mask(candidate, spec.mask_prob)
904
-
905
- rand = torch.rand(input_ids.shape, device=device)
906
- replace_mask = masked & (rand < MASK_REPLACE_CUTOFF)
907
- random_mask = masked & (rand >= MASK_REPLACE_CUTOFF) & (rand < MASK_RANDOM_CUTOFF)
908
- # keep_mask (last 10%): input unchanged, still predicted.
909
-
910
- masked_input = input_ids.clone()
911
- masked_input[replace_mask] = spec.mask_token_id
912
- rand_tokens = torch.randint_like(input_ids, 0, spec.vocab_size)
913
- masked_input = torch.where(random_mask, rand_tokens, masked_input)
914
-
915
- mlm_labels = torch.full_like(input_ids, IGNORE_INDEX)
916
- mlm_labels[masked] = input_ids[masked]
917
- return masked_input, mlm_labels
918
-
919
-
920
- @dataclass
921
- class DzairRTDOutput:
922
- """Pretraining output: ELECTRA-style joint loss."""
923
-
924
- loss: Tensor | None
925
- rtd_logits: Tensor
926
- gen_logits: Tensor
927
- generator_loss: Tensor | None
928
- discriminator_loss: Tensor | None
929
- replacement_rate: Tensor
930
-
931
-
932
- @_dynamo_disabled
933
- def _sample_generator_corruptions(
934
- gen_logits: Tensor,
935
- masked_input: Tensor,
936
- predict: Tensor,
937
- input_ids: Tensor,
938
- softmax_chunk: int = 2048,
939
- ) -> Tensor:
940
- """Sample generator tokens on masked positions in chunks outside Dynamo."""
941
- with torch.no_grad():
942
- flat_mask = predict.reshape(-1)
943
- idx = torch.where(flat_mask)[0]
944
- sampled = torch.empty_like(idx)
945
- flat_logits = gen_logits.reshape(-1, gen_logits.size(-1))
946
- for start in range(0, idx.numel(), softmax_chunk):
947
- group = idx[start : start + softmax_chunk]
948
- probs = flat_logits[group].float().softmax(dim=-1)
949
- sampled[start : start + softmax_chunk] = torch.multinomial(probs, 1).squeeze(-1)
950
- corrupted = masked_input.reshape(-1).clone()
951
- corrupted[idx] = sampled
952
- return corrupted.view_as(input_ids)
953
-
954
-
955
- class RTDHead(nn.Module):
956
- """Generator (MLM) corrupts, discriminator (RTD) detects, GDES detaches."""
957
-
958
- def __init__(self, config: DzairConfig, rtd_loss_weight: float = RTD_LOSS_WEIGHT) -> None:
959
- super().__init__()
960
- self.config = config
961
- self.rtd_loss_weight = rtd_loss_weight
962
- self.generator = Generator(config)
963
- self.discriminator = Discriminator(config)
964
-
965
- def forward(
966
- self,
967
- input_ids: Tensor,
968
- attention_mask: Tensor | None = None,
969
- mask_prob: float | Tensor = 0.15,
970
- word_starts: Tensor | None = None,
971
- ) -> DzairRTDOutput:
972
- """Returns the joint output. ``mask_prob`` follows the 30→15% schedule.
973
-
974
- Accepts a 0-dim tensor as well as a float: pass a tensor from any
975
- compiled caller — Dynamo specializes on float argument *values*,
976
- so a per-step float schedule would recompile every step until the
977
- cache limit forces the whole model back to eager.
978
- """
979
- eligible = (
980
- attention_mask.to(torch.bool)
981
- if attention_mask is not None
982
- else torch.ones_like(input_ids, dtype=torch.bool)
983
- )
984
- spec = MaskSpec(
985
- special_ids=frozenset(
986
- sid
987
- for sid in (
988
- self.config.pad_token_id,
989
- self.config.cls_token_id,
990
- self.config.sep_token_id,
991
- )
992
- if sid is not None
993
- ),
994
- vocab_size=self.config.vocab_size,
995
- mask_token_id=self.config.mask_token_id,
996
- mask_prob=mask_prob,
997
- )
998
- masked_input, mlm_labels = _mask_inputs(input_ids, eligible, spec, word_starts)
999
-
1000
- gen_logits = self.generator(masked_input, attention_mask)
1001
- gen_loss: Tensor | None = None
1002
- predict = mlm_labels != IGNORE_INDEX
1003
- if predict.any():
1004
- gen_loss = functional.cross_entropy(gen_logits[predict], input_ids[predict].detach())
1005
-
1006
- # Corrupt only the masked positions by sampling the generator.
1007
- # Softmax runs over the masked subset in bounded chunks to cap the
1008
- # peak transient allocation. Arithmetic (measured 2026-09-10):
1009
- # 8192 * 48000 * 4 bytes = 1.57 GB -- OOMs at 20.94 GB in use
1010
- # 2048 * 48000 * 4 bytes = 0.39 GB -- 4.7 GB headroom at 18.9 GB peak
1011
- # Chunked multinomial is mathematically identical to sampling all at once.
1012
- _softmax_chunk = 2048
1013
- corrupted = _sample_generator_corruptions(
1014
- gen_logits, masked_input, predict, input_ids, _softmax_chunk
1015
- )
1016
-
1017
- disc_logits = self.discriminator(
1018
- corrupted,
1019
- attention_mask,
1020
- generator_embeddings=self.generator.embeddings.word_embeddings.weight,
1021
- )
1022
-
1023
- rtd_labels = torch.where(corrupted == input_ids, 1, 0)
1024
- rtd_labels = torch.where(eligible, rtd_labels, IGNORE_INDEX)
1025
- disc_loss: Tensor | None = None
1026
- if (rtd_labels != IGNORE_INDEX).any():
1027
- disc_loss = functional.cross_entropy(
1028
- disc_logits.reshape(-1, 2), rtd_labels.reshape(-1), ignore_index=IGNORE_INDEX
1029
- )
1030
-
1031
- loss: Tensor | None = None
1032
- if gen_loss is not None and disc_loss is not None:
1033
- loss = gen_loss + self.rtd_loss_weight * disc_loss
1034
-
1035
- with torch.no_grad():
1036
- replacement_rate = (
1037
- (corrupted[predict] != input_ids[predict]).float().mean()
1038
- if predict.any()
1039
- else torch.zeros((), device=input_ids.device)
1040
- )
1041
-
1042
- return DzairRTDOutput(
1043
- loss=loss,
1044
- rtd_logits=disc_logits,
1045
- gen_logits=gen_logits,
1046
- generator_loss=gen_loss,
1047
- discriminator_loss=disc_loss,
1048
- replacement_rate=replacement_rate,
1049
- )
1050
-
1051
-
1052
- class DzairPreTrainedModel(PreTrainedModel):
1053
- config_class = DzairConfig
1054
- base_model_prefix = "dzair"
1055
- supports_gradient_checkpointing = True
1056
- _no_split_modules: ClassVar[list[str]] = ["TransformerLayer"]
1057
-
1058
- def _init_weights(self, module: nn.Module) -> None:
1059
- # Masinissa scaled init, depth-scaled trunc normal; deliberately not
1060
- # config-driven (a field that silently does nothing is worse than none).
1061
- std = math.sqrt(2.0 / (5.0 * self.config.hidden_size))
1062
- if isinstance(module, nn.Linear):
1063
- nn.init.trunc_normal_(module.weight, mean=0.0, std=std, a=-2 * std, b=2 * std)
1064
- if module.bias is not None:
1065
- nn.init.zeros_(module.bias)
1066
- elif isinstance(module, nn.Embedding):
1067
- nn.init.trunc_normal_(module.weight, mean=0.0, std=std, a=-2 * std, b=2 * std)
1068
- elif isinstance(module, RMSNorm):
1069
- nn.init.ones_(module.weight)
1070
-
1071
-
1072
- @dataclass
1073
- class DzairEncoderOutput(ModelOutput):
1074
- """Encoder output: last state plus optional per-layer states.
1075
-
1076
- A dedicated type because the framework's BaseModelOutput pins its state
1077
- fields to FloatTensor, which the checker treats as distinct from Tensor.
1078
- hidden_states[0] is the embedding output (HF convention). Per-layer
1079
- entries are pre-norm layer outputs; last_hidden_state is post-norm, so
1080
- hidden_states[-1] != last_hidden_state by design.
1081
- """
1082
-
1083
- last_hidden_state: Tensor
1084
- hidden_states: tuple[Tensor, ...] | None = None
1085
- attentions: tuple[Tensor, ...] | None = None
1086
-
1087
-
1088
- class DzairModel(DzairPreTrainedModel):
1089
- """The encoder alone (discriminator backbone).
1090
-
1091
- Returns contextualised token representations.
1092
- """
1093
-
1094
- def __init__(self, config: DzairConfig) -> None:
1095
- super().__init__(config)
1096
- self.embeddings = Embeddings(config)
1097
- self.encoder = Encoder(config)
1098
- self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps)
1099
- self.post_init()
1100
-
1101
- def get_input_embeddings(self) -> nn.Embedding:
1102
- return self.embeddings.word_embeddings
1103
-
1104
- def set_input_embeddings(self, value: nn.Embedding) -> None:
1105
- self.embeddings.word_embeddings = value
1106
-
1107
- def forward(
1108
- self,
1109
- input_ids: Tensor,
1110
- attention_mask: Tensor | None = None,
1111
- output_hidden_states: bool = False,
1112
- ) -> DzairEncoderOutput:
1113
- """Encode tokens. With output_hidden_states, hidden_states[0] is the
1114
- embedding output and [k] the k-th layer output (HF convention).
1115
- Layer states are pre-norm; last_hidden_state is post-norm.
1116
- """
1117
- embedded = self.embeddings(input_ids)
1118
- if output_hidden_states:
1119
- last, states = self.encoder.forward_with_states(embedded, attention_mask)
1120
- return DzairEncoderOutput(
1121
- last_hidden_state=self.norm(last),
1122
- hidden_states=(embedded, *states),
1123
- )
1124
- x = self.encoder(embedded, attention_mask)
1125
- return DzairEncoderOutput(last_hidden_state=self.norm(x))
1126
-
1127
-
1128
- class DzairForMaskedLM(DzairPreTrainedModel):
1129
- """Pretraining model: Generator (MLM) + Discriminator (RTD) with GDES."""
1130
-
1131
- _tied_weights_keys: ClassVar[dict[str, str]] = {
1132
- "rtd_head.generator.lm_head.weight": "rtd_head.generator.embeddings.word_embeddings.weight",
1133
- }
1134
-
1135
- def __init__(self, config: DzairConfig) -> None:
1136
- super().__init__(config)
1137
- self.rtd_head = RTDHead(config)
1138
- self.post_init()
1139
-
1140
- def forward(
1141
- self,
1142
- input_ids: Tensor,
1143
- attention_mask: Tensor | None = None,
1144
- mask_prob: float | Tensor = 0.15,
1145
- word_starts: Tensor | None = None,
1146
- ) -> DzairRTDOutput:
1147
- return self.rtd_head(input_ids, attention_mask, mask_prob, word_starts)
1148
-
1149
-
1150
- @dataclass
1151
- class DzairSequenceClassifierOutput(ModelOutput):
1152
- """Output type of DzairForSequenceClassification."""
1153
-
1154
- loss: Tensor | None = None
1155
- logits: Tensor | None = None
1156
- hidden_states: tuple[Tensor, ...] | None = None
1157
- attentions: tuple[Tensor, ...] | None = None
1158
-
1159
-
1160
- class DzairForSequenceClassification(DzairPreTrainedModel):
1161
- """Sequence classification head on top of the DZAIR encoder backbone.
1162
-
1163
- One method, the measured one: [CLS] pooling through an MLP projection
1164
- head (Dropout -> Dense -> GELU -> Dropout) into the classification
1165
- layer. The DZNLI head ablation picked cls+mlp over mean+linear; the
1166
- landmark and attention experiments never measured a win, so they do
1167
- not ship.
1168
- """
1169
-
1170
- def __init__(self, config: DzairConfig) -> None:
1171
- super().__init__(config)
1172
- self.num_labels = getattr(config, "num_labels", 2)
1173
- self.dzair = DzairModel(config)
1174
- self.dropout = nn.Dropout(config.hidden_dropout_prob)
1175
- self.dense = nn.Linear(config.hidden_size, config.hidden_size)
1176
- self.classifier = nn.Linear(config.hidden_size, self.num_labels)
1177
- self.post_init()
1178
-
1179
- def get_input_embeddings(self) -> nn.Embedding:
1180
- return self.dzair.get_input_embeddings()
1181
-
1182
- def set_input_embeddings(self, value: nn.Embedding) -> None:
1183
- self.dzair.set_input_embeddings(value)
1184
-
1185
- def forward(
1186
- self,
1187
- input_ids: Tensor,
1188
- attention_mask: Tensor | None = None,
1189
- labels: Tensor | None = None,
1190
- output_hidden_states: bool = False,
1191
- ) -> DzairSequenceClassifierOutput:
1192
- outputs = self.dzair(
1193
- input_ids, attention_mask=attention_mask, output_hidden_states=output_hidden_states
1194
- )
1195
- pooled_output = self.dropout(outputs.last_hidden_state[:, 0])
1196
- pooled_output = self.dense(pooled_output)
1197
- pooled_output = functional.gelu(pooled_output)
1198
- pooled_output = self.dropout(pooled_output)
1199
- logits = self.classifier(pooled_output)
1200
-
1201
- loss: Tensor | None = None
1202
- if labels is not None:
1203
- if self.num_labels == 1:
1204
- loss = functional.mse_loss(logits.view(-1), labels.view(-1).float())
1205
- else:
1206
- loss = functional.cross_entropy(logits.view(-1, self.num_labels), labels.view(-1))
1207
-
1208
- return DzairSequenceClassifierOutput(
1209
- loss=loss,
1210
- logits=logits,
1211
- hidden_states=outputs.hidden_states,
1212
- attentions=outputs.attentions,
1213
- )
1214
-
1215
- def load_backbone_weights(self, state_dict: dict[str, Tensor]) -> None:
1216
- """Load pretrained discriminator backbone weights into self.dzair."""
1217
- self.dzair.load_state_dict(
1218
- discriminator_backbone_state(state_dict, self.config), strict=True
1219
- )
1220
-
1221
- @classmethod
1222
- def from_pretrained_checkpoint(
1223
- cls,
1224
- checkpoint_path: str | Path,
1225
- config: DzairConfig | None = None,
1226
- num_labels: int = 2,
1227
- ) -> DzairForSequenceClassification:
1228
- """Instantiate classification model and load backbone from pretrain checkpoint.
1229
-
1230
- The checkpoint's stored config is preferred; an explicit config that
1231
- disagrees with it raises (see ``_read_pretrain_checkpoint``).
1232
- """
1233
- model_config, state = _read_pretrain_checkpoint(checkpoint_path, config)
1234
- model_config.num_labels = num_labels
1235
- model = cls(model_config)
1236
- model.load_backbone_weights(state)
1237
- return model
1238
-
1239
-
1240
- @dataclass
1241
- class DzairTokenClassifierOutput(ModelOutput):
1242
- """Output type of DzairForTokenClassification."""
1243
-
1244
- loss: Tensor | None = None
1245
- logits: Tensor | None = None
1246
- hidden_states: tuple[Tensor, ...] | None = None
1247
- attentions: tuple[Tensor, ...] | None = None
1248
-
1249
-
1250
- class DzairForTokenClassification(DzairPreTrainedModel):
1251
- """Token classification head on top of the DZAIR encoder backbone (e.g. for NER/POS)."""
1252
-
1253
- def __init__(self, config: DzairConfig) -> None:
1254
- super().__init__(config)
1255
- self.num_labels = getattr(config, "num_labels", 2)
1256
- self.dzair = DzairModel(config)
1257
- self.dropout = nn.Dropout(config.hidden_dropout_prob)
1258
- self.classifier = nn.Linear(config.hidden_size, self.num_labels)
1259
- self.post_init()
1260
-
1261
- def get_input_embeddings(self) -> nn.Embedding:
1262
- return self.dzair.get_input_embeddings()
1263
-
1264
- def set_input_embeddings(self, value: nn.Embedding) -> None:
1265
- self.dzair.set_input_embeddings(value)
1266
-
1267
- def forward(
1268
- self,
1269
- input_ids: Tensor,
1270
- attention_mask: Tensor | None = None,
1271
- labels: Tensor | None = None,
1272
- ) -> DzairTokenClassifierOutput:
1273
- outputs = self.dzair(input_ids, attention_mask=attention_mask)
1274
- sequence_output = outputs.last_hidden_state
1275
- sequence_output = self.dropout(sequence_output)
1276
- logits = self.classifier(sequence_output)
1277
-
1278
- loss: Tensor | None = None
1279
- if labels is not None:
1280
- loss = functional.cross_entropy(
1281
- logits.view(-1, self.num_labels),
1282
- labels.view(-1),
1283
- ignore_index=-100,
1284
- )
1285
-
1286
- return DzairTokenClassifierOutput(
1287
- loss=loss,
1288
- logits=logits,
1289
- hidden_states=outputs.hidden_states,
1290
- attentions=outputs.attentions,
1291
- )
1292
-
1293
- def load_backbone_weights(self, state_dict: dict[str, Tensor]) -> None:
1294
- """Load pretrained discriminator backbone weights into self.dzair."""
1295
- self.dzair.load_state_dict(
1296
- discriminator_backbone_state(state_dict, self.config), strict=True
1297
- )
1298
-
1299
- @classmethod
1300
- def from_pretrained_checkpoint(
1301
- cls,
1302
- checkpoint_path: str | Path,
1303
- config: DzairConfig | None = None,
1304
- num_labels: int = 2,
1305
- ) -> DzairForTokenClassification:
1306
- """Instantiate token classification model and load backbone from pretrain checkpoint."""
1307
- model_config, state = _read_pretrain_checkpoint(checkpoint_path, config)
1308
- model_config.num_labels = num_labels
1309
- model = cls(model_config)
1310
- model.load_backbone_weights(state)
1311
- return model
1312
-
1313
-
1314
- __all__ = [
1315
- "OBSOLETE_STATE_SUBSTRINGS",
1316
- "RTD_LOSS_WEIGHT",
1317
- "DzairConfig",
1318
- "DzairEncoderOutput",
1319
- "DzairForMaskedLM",
1320
- "DzairForSequenceClassification",
1321
- "DzairForTokenClassification",
1322
- "DzairModel",
1323
- "DzairPreTrainedModel",
1324
- "DzairRTDOutput",
1325
- "DzairSequenceClassifierOutput",
1326
- "DzairTokenClassifierOutput",
1327
- "MaskSpec",
1328
- "ResumeCompat",
1329
- "discriminator_backbone_state",
1330
- "draw_token_mask",
1331
- "draw_word_mask",
1332
- "load_pretrain_state",
1333
- "load_resume_weights",
1334
- "set_gradient_checkpointing",
1335
- ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
hub/hub/__init__.py DELETED
@@ -1,44 +0,0 @@
1
- """DZAIR encoder hub module (base + small sizes).
2
-
3
- This module ships verbatim in releases. It imports nothing back from the package.
4
- """
5
-
6
- from dzair.hub.dzair_base.configuration import DZAIR_BASE_CONFIG, DZAIR_SMALL_CONFIG, DzairConfig
7
- from dzair.hub.dzair_base.modeling import (
8
- RTD_LOSS_WEIGHT,
9
- DzairEncoderOutput,
10
- DzairForMaskedLM,
11
- DzairForSequenceClassification,
12
- DzairForTokenClassification,
13
- DzairModel,
14
- DzairPreTrainedModel,
15
- DzairRTDOutput,
16
- MaskSpec,
17
- ResumeCompat,
18
- draw_token_mask,
19
- draw_word_mask,
20
- load_pretrain_state,
21
- load_resume_weights,
22
- set_gradient_checkpointing,
23
- )
24
-
25
- __all__ = [
26
- "DZAIR_BASE_CONFIG",
27
- "DZAIR_SMALL_CONFIG",
28
- "RTD_LOSS_WEIGHT",
29
- "DzairConfig",
30
- "DzairEncoderOutput",
31
- "DzairForMaskedLM",
32
- "DzairForSequenceClassification",
33
- "DzairForTokenClassification",
34
- "DzairModel",
35
- "DzairPreTrainedModel",
36
- "DzairRTDOutput",
37
- "MaskSpec",
38
- "ResumeCompat",
39
- "draw_token_mask",
40
- "draw_word_mask",
41
- "load_pretrain_state",
42
- "load_resume_weights",
43
- "set_gradient_checkpointing",
44
- ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
hub/hub/__pycache__/__init__.cpython-312.pyc DELETED
Binary file (1.06 kB)
 
hub/hub/dzair_base/__init__.py DELETED
@@ -1,32 +0,0 @@
1
- """DZAIR encoder hub module (base + small sizes).
2
-
3
- This module ships verbatim in releases. It imports nothing back from the package.
4
- """
5
-
6
- from dzair.hub.dzair_base.configuration import DZAIR_BASE_CONFIG, DZAIR_SMALL_CONFIG, DzairConfig
7
- from dzair.hub.dzair_base.modeling import (
8
- DzairEncoderOutput,
9
- DzairForMaskedLM,
10
- DzairForSequenceClassification,
11
- DzairForTokenClassification,
12
- DzairModel,
13
- DzairPreTrainedModel,
14
- ResumeCompat,
15
- load_pretrain_state,
16
- load_resume_weights,
17
- )
18
-
19
- __all__ = [
20
- "DZAIR_BASE_CONFIG",
21
- "DZAIR_SMALL_CONFIG",
22
- "DzairConfig",
23
- "DzairEncoderOutput",
24
- "DzairForMaskedLM",
25
- "DzairForSequenceClassification",
26
- "DzairForTokenClassification",
27
- "DzairModel",
28
- "DzairPreTrainedModel",
29
- "ResumeCompat",
30
- "load_pretrain_state",
31
- "load_resume_weights",
32
- ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
hub/hub/dzair_base/__pycache__/__init__.cpython-312.pyc DELETED
Binary file (868 Bytes)
 
hub/hub/dzair_base/__pycache__/configuration.cpython-312.pyc DELETED
Binary file (5.57 kB)
 
hub/hub/dzair_base/__pycache__/modeling.cpython-312.pyc DELETED
Binary file (69.1 kB)
 
hub/hub/dzair_base/configuration.py DELETED
@@ -1,146 +0,0 @@
1
- """Configuration for DZAIR encoder family (base + small)."""
2
-
3
- from __future__ import annotations
4
-
5
- from typing import Any
6
-
7
- from transformers import PretrainedConfig
8
-
9
- _BASE_HIDDEN_SIZE = 768
10
- _BASE_LAYERS = 12
11
- _SMALL_HIDDEN_SIZE = 384
12
- _SMALL_LAYERS = 6
13
-
14
-
15
- class DzairConfig(PretrainedConfig):
16
- """DZAIR encoder configuration.
17
-
18
- Every architectural choice is a declared field so ``config.json``
19
- round-trips exactly (base and small share this class, never a hidden
20
- ``arch`` object). ``**kwargs`` forwards only transformers-managed keys
21
- (e.g. ``transformers_version``) to ``PretrainedConfig``.
22
-
23
- Field names match the published `config.json`. Two sizes share this config:
24
- - base: 12Lx768 (discriminator, grouped-query 12Q/4KV) + 3Lx384 generator,
25
- shared embeddings
26
- - small: 6Lx384 (discriminator, grouped-query 6Q/2KV) + 3Lx384 generator,
27
- shared embeddings
28
- """
29
-
30
- model_type = "dzair"
31
-
32
- def __init__(
33
- self,
34
- vocab_size: int = 48000,
35
- hidden_size: int = 768,
36
- intermediate_size: int = 1792,
37
- num_attention_heads: int = 12,
38
- num_key_value_heads: int = 0,
39
- num_hidden_layers: int = 12,
40
- num_generator_layers: int = 3,
41
- generator_hidden_size: int = 384,
42
- generator_intermediate_size: int = 1024,
43
- max_position_embeddings: int = 512,
44
- rope_theta: float = 10000.0,
45
- hidden_dropout_prob: float = 0.1,
46
- attention_probs_dropout_prob: float = 0.1,
47
- layer_norm_eps: float = 1e-5,
48
- pad_token_id: int = 0,
49
- cls_token_id: int = 2,
50
- sep_token_id: int = 3,
51
- mask_token_id: int = 4,
52
- tie_word_embeddings: bool = True,
53
- share_generator_embeddings: bool = False,
54
- qk_norm: bool = False,
55
- **kwargs: Any,
56
- ) -> None:
57
- if hidden_size % num_attention_heads != 0:
58
- msg = f"hidden_size {hidden_size} must split over {num_attention_heads} heads"
59
- raise ValueError(msg)
60
- if num_key_value_heads == 0:
61
- num_key_value_heads = num_attention_heads
62
- if num_attention_heads % num_key_value_heads != 0:
63
- msg = (
64
- f"{num_attention_heads} query heads must split over "
65
- f"{num_key_value_heads} key-value heads"
66
- )
67
- raise ValueError(msg)
68
- if generator_hidden_size % 64 != 0:
69
- msg = f"generator_hidden_size {generator_hidden_size} must be a multiple of 64"
70
- raise ValueError(msg)
71
- self.vocab_size = vocab_size
72
- self.hidden_size = hidden_size
73
- self.intermediate_size = intermediate_size
74
- self.num_attention_heads = num_attention_heads
75
- self.num_key_value_heads = num_key_value_heads
76
- self.num_hidden_layers = num_hidden_layers
77
- self.num_generator_layers = num_generator_layers
78
- self.generator_hidden_size = generator_hidden_size
79
- self.generator_intermediate_size = generator_intermediate_size
80
- self.max_position_embeddings = max_position_embeddings
81
- self.rope_theta = rope_theta
82
- self.hidden_dropout_prob = hidden_dropout_prob
83
- self.attention_probs_dropout_prob = attention_probs_dropout_prob
84
- self.layer_norm_eps = layer_norm_eps
85
- self.share_generator_embeddings = share_generator_embeddings
86
- self.qk_norm = qk_norm
87
-
88
- super().__init__(
89
- pad_token_id=pad_token_id,
90
- cls_token_id=cls_token_id,
91
- sep_token_id=sep_token_id,
92
- tie_word_embeddings=tie_word_embeddings,
93
- **kwargs,
94
- )
95
- # Ensure mask_token_id and explicit IDs are preserved as ints
96
- self.pad_token_id = pad_token_id
97
- self.cls_token_id = cls_token_id
98
- self.sep_token_id = sep_token_id
99
- self.mask_token_id = mask_token_id
100
-
101
- @property
102
- def head_size(self) -> int:
103
- return self.hidden_size // self.num_attention_heads
104
-
105
- @property
106
- def generator_num_heads(self) -> int:
107
- """Generator query heads at head_dim 64 (always divides, checked above)."""
108
- return self.generator_hidden_size // 64
109
-
110
- @property
111
- def kv_dim(self) -> int:
112
- """Key/value width: key-value heads at the trunk head_dim."""
113
- return self.num_key_value_heads * self.head_size
114
-
115
- @property
116
- def is_base(self) -> bool:
117
- return self.hidden_size == _BASE_HIDDEN_SIZE and self.num_hidden_layers == _BASE_LAYERS
118
-
119
- @property
120
- def is_small(self) -> bool:
121
- return self.hidden_size == _SMALL_HIDDEN_SIZE and self.num_hidden_layers == _SMALL_LAYERS
122
-
123
-
124
- # Predefined configurations. Both sizes share the generator embedding table
125
- # with the discriminator (GDES, DeBERTaV3) and apply QK-norm; the FFN
126
- # intermediate is 128-aligned for tensor cores (1792 = 14x128, 1024 = 8x128).
127
- DZAIR_BASE_CONFIG = DzairConfig(
128
- num_key_value_heads=4,
129
- share_generator_embeddings=True,
130
- qk_norm=True,
131
- )
132
- DZAIR_SMALL_CONFIG = DzairConfig(
133
- hidden_size=384,
134
- intermediate_size=1024,
135
- num_attention_heads=6,
136
- num_key_value_heads=2,
137
- num_hidden_layers=6,
138
- num_generator_layers=3,
139
- generator_hidden_size=384,
140
- generator_intermediate_size=1024,
141
- share_generator_embeddings=True,
142
- qk_norm=True,
143
- )
144
-
145
-
146
- __all__ = ["DZAIR_BASE_CONFIG", "DZAIR_SMALL_CONFIG", "DzairConfig"]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
hub/hub/dzair_base/modeling.py DELETED
@@ -1,1335 +0,0 @@
1
- """DZAIR encoder: RTD + GDES (DeBERTaV3 objective) with ModernBERT-speed architecture.
2
-
3
- Architecture: pre-RMSNorm, RoPE, SwiGLU, fused scaled-dot-product
4
- attention (FlashAttention-2 path when available) attending globally,
5
- single-chunk sequences.
6
- Objective: RTD on all tokens. Generator MLM corrupts; GDES detaches
7
- generator embeddings.
8
- """
9
-
10
- from __future__ import annotations
11
-
12
- import copy
13
- import hashlib
14
- import math
15
- from dataclasses import dataclass
16
- from pathlib import Path
17
- from typing import Any, ClassVar
18
-
19
- import torch
20
- from torch import Tensor, _dynamo, nn
21
- from torch.nn import functional
22
- from torch.utils import checkpoint as checkpoint_utils
23
- from transformers import PretrainedConfig, PreTrainedModel
24
- from transformers.utils.generic import ModelOutput
25
-
26
- from dzair.hub.dzair_base.configuration import DzairConfig
27
-
28
- IGNORE_INDEX = -100
29
-
30
- # ELECTRA (Clark et al., 2020, §3.3): small models weight the discriminator
31
- # loss at 50 relative to the generator MLM loss.
32
- RTD_LOSS_WEIGHT = 50.0
33
-
34
- # BERT 80/10/10 corruption splits (Devlin et al., 2019): below REPLACE the
35
- # token becomes [MASK], below REPLACE+RANDOM it becomes a random vocab id,
36
- # otherwise it is kept (but still predicted by the generator).
37
- MASK_REPLACE_CUTOFF = 0.8
38
- MASK_RANDOM_CUTOFF = 0.9
39
-
40
-
41
- def _apply_rope(x: Tensor, cos: Tensor, sin: Tensor) -> Tensor:
42
- """Apply rotary positional embeddings to half the head dim."""
43
- x1, x2 = x.chunk(2, dim=-1)
44
- return torch.cat((x1 * cos - x2 * sin, x1 * sin + x2 * cos), dim=-1)
45
-
46
-
47
- def _build_rope_cache(
48
- max_seq_len: int, head_dim: int, theta: float, device: torch.device
49
- ) -> tuple[Tensor, Tensor]:
50
- """Build RoPE cos/sin cache for sequence length up to max_seq_len."""
51
- inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=device).float() / head_dim))
52
- t = torch.arange(max_seq_len, device=device).float()
53
- freqs = torch.outer(t, inv_freq)
54
- cos = freqs.cos().to(torch.get_default_dtype())
55
- sin = freqs.sin().to(torch.get_default_dtype())
56
- return cos, sin
57
-
58
-
59
- class RMSNorm(nn.Module):
60
- """Root Mean Square Layer Normalization (affine weight, no bias)."""
61
-
62
- def __init__(self, dim: int, eps: float = 1e-5) -> None:
63
- super().__init__()
64
- self.eps = eps
65
- self.weight = nn.Parameter(torch.ones(dim))
66
-
67
- def forward(self, x: Tensor) -> Tensor:
68
- norm = x.pow(2).mean(dim=-1, keepdim=True).add(self.eps).rsqrt()
69
- return x * norm * self.weight
70
-
71
-
72
- class SwiGLU(nn.Module):
73
- """Swish-Gated Linear Unit."""
74
-
75
- def forward(self, x: Tensor) -> Tensor:
76
- x, gate = x.chunk(2, dim=-1)
77
- return x * functional.silu(gate)
78
-
79
-
80
- class FeedForward(nn.Module):
81
- """Pre-RMSNorm SwiGLU FFN with dropout."""
82
-
83
- def __init__(self, config: DzairConfig) -> None:
84
- super().__init__()
85
- self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps)
86
- self.up = nn.Linear(config.hidden_size, 2 * config.intermediate_size, bias=False)
87
- self.act = SwiGLU()
88
- self.down = nn.Linear(config.intermediate_size, config.hidden_size, bias=False)
89
- self.dropout = nn.Dropout(config.hidden_dropout_prob)
90
-
91
- def forward(self, x: Tensor) -> Tensor:
92
- residual = x
93
- x = self.norm(x)
94
- x = self.up(x)
95
- x = self.act(x)
96
- x = self.dropout(self.down(x))
97
- return residual + x
98
-
99
-
100
- class Attention(nn.Module):
101
- """Grouped-query attention with RoPE: every layer attends globally.
102
-
103
- Query heads share fewer key/value heads (``num_key_value_heads`` groups).
104
- Local-window alternation was cut 2026-09-14: at 512 tokens it saves
105
- ~4% wall-clock (measured FLOP arithmetic) while full attention is the
106
- literature default every baseline trains — the deviation bought
107
- complexity without evidence. Fused projections, bias-free, pre-RMSNorm.
108
- """
109
-
110
- def __init__(self, config: DzairConfig) -> None:
111
- super().__init__()
112
- self.config = config
113
- self.num_heads = config.num_attention_heads
114
- self.num_kv_heads = config.num_key_value_heads
115
- self.head_size = config.head_size
116
- self.scale = 1.0 / math.sqrt(self.head_size)
117
-
118
- # Separate Q and fused KV projections (bias-free for FA-2 compatibility).
119
- # KV groups repeat to the query count at forward time.
120
- self.q_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=False)
121
- self.kv_proj = nn.Linear(config.hidden_size, 2 * config.kv_dim, bias=False)
122
- self.out_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=False)
123
-
124
- self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps)
125
- # QK-norm (Gemma 2/3 practice): one shared RMSNorm over head_dim
126
- # applied to queries and keys before RoPE. Norm-then-rotate is a
127
- # fixed convention, not a commutation (rotation mixes dims, so an
128
- # affine weight does not commute with it) — the trained weights
129
- # bake in this order, so it must never change under them.
130
- self.qk_norm: RMSNorm | None = (
131
- RMSNorm(self.head_size, eps=config.layer_norm_eps) if config.qk_norm else None
132
- )
133
- self.dropout = nn.Dropout(config.attention_probs_dropout_prob)
134
-
135
- # RoPE cache (non-persistent, rebuilt on first use)
136
- self._cos: Tensor | None = None
137
- self._sin: Tensor | None = None
138
- self._cache_len = 0
139
-
140
- def _get_rope(self, seq_len: int, device: torch.device) -> tuple[Tensor, Tensor]:
141
- if self._cos is None or self._cache_len < seq_len or self._cos.device != device:
142
- self._cos, self._sin = _build_rope_cache(
143
- max(seq_len, self.config.max_position_embeddings),
144
- self.head_size,
145
- self.config.rope_theta,
146
- device,
147
- )
148
- self._cache_len = max(seq_len, self.config.max_position_embeddings)
149
- return self._cos[:seq_len], self._sin[:seq_len]
150
-
151
- def forward(
152
- self,
153
- x: Tensor,
154
- attention_mask: Tensor | None = None,
155
- is_causal: bool = False,
156
- ) -> Tensor:
157
- """x: [B, T, D], attention_mask: [B, T] (1=keep, 0=pad). Returns [B, T, D]."""
158
- batch_size, seq_len, _ = x.shape
159
-
160
- # Pre-norm
161
- x_norm = self.norm(x)
162
-
163
- # Grouped-query projections.
164
- q = self.q_proj(x_norm) # [B, T, D]
165
- kv = self.kv_proj(x_norm) # [B, T, 2 * kv_dim]
166
- k, v = kv.chunk(2, dim=-1)
167
-
168
- # Reshape for attention: Q [B, H, T, head_dim], K/V [B, KV, T, head_dim].
169
- q = q.view(batch_size, seq_len, self.num_heads, self.head_size).transpose(1, 2)
170
- k = k.view(batch_size, seq_len, self.num_kv_heads, self.head_size).transpose(1, 2)
171
- v = v.view(batch_size, seq_len, self.num_kv_heads, self.head_size).transpose(1, 2)
172
- # Repeat KV groups to the query count (exact: heads split evenly, checked).
173
- repeat = self.num_heads // self.num_kv_heads
174
- if repeat > 1:
175
- k = k.repeat_interleave(repeat, dim=1)
176
- v = v.repeat_interleave(repeat, dim=1)
177
-
178
- if self.qk_norm is not None:
179
- q = self.qk_norm(q)
180
- k = self.qk_norm(k)
181
-
182
- # RoPE (cast to the working dtype: an fp32 cache multiplied into bf16
183
- # queries upcasts them and drops out of the fused-attention fast path)
184
- cos, sin = self._get_rope(seq_len, x.device)
185
- cos = cos.unsqueeze(0).unsqueeze(0).to(x.dtype) # [1, 1, T, head_dim/2]
186
- sin = sin.unsqueeze(0).unsqueeze(0).to(x.dtype)
187
- q = _apply_rope(q, cos, sin)
188
- k = _apply_rope(k, cos, sin)
189
-
190
- # Scaled dot-product attention. A bool mask (True = attend) keeps the
191
- # fused fast path; the old additive float mask did not.
192
- attn_mask: Tensor | None = None
193
- if attention_mask is not None:
194
- attn_mask = attention_mask.to(torch.bool).view(batch_size, 1, 1, seq_len)
195
- # Guard against all-False mask rows (all-pad inputs): SDPA under
196
- # CUDA/Inductor produces NaNs when a row has zero attendable keys.
197
- positions = torch.arange(seq_len, device=x.device)
198
- has_key = attn_mask.any(dim=-1, keepdim=True)
199
- attn_mask = attn_mask | (~has_key & (positions == 0).view(1, 1, 1, seq_len))
200
-
201
- # Use PyTorch's scaled_dot_product_attention (uses FA-2 when available)
202
- attn_out = functional.scaled_dot_product_attention(
203
- q,
204
- k,
205
- v,
206
- attn_mask=attn_mask,
207
- dropout_p=self.config.attention_probs_dropout_prob if self.training else 0.0,
208
- is_causal=is_causal,
209
- scale=self.scale,
210
- )
211
-
212
- # Merge heads: [B, H, T, head_dim] -> [B, T, D]
213
- attn_out = attn_out.transpose(1, 2).contiguous().view(batch_size, seq_len, -1)
214
-
215
- # Zero out padding positions so unused positions never dominate
216
- # downstream means (the residual still carries the pad embedding;
217
- # losses and CLS pooling ignore pads by mask, which is what makes
218
- # this safe rather than the zeroing alone).
219
- if attention_mask is not None:
220
- attn_out = attn_out * attention_mask.view(batch_size, seq_len, 1).to(attn_out.dtype)
221
-
222
- # Output projection + residual
223
- out = self.out_proj(attn_out)
224
- out = self.dropout(out)
225
- return x + out
226
-
227
-
228
- class TransformerLayer(nn.Module):
229
- """Pre-RMSNorm transformer block: Attention + FFN."""
230
-
231
- def __init__(self, config: DzairConfig) -> None:
232
- super().__init__()
233
- self.attention = Attention(config)
234
- self.ffn = FeedForward(config)
235
-
236
- def forward(
237
- self,
238
- x: Tensor,
239
- attention_mask: Tensor | None = None,
240
- is_causal: bool = False,
241
- ) -> Tensor:
242
- x = self.attention(x, attention_mask, is_causal)
243
- return self.ffn(x)
244
-
245
-
246
- class Embeddings(nn.Module):
247
- """Token embeddings with RMSNorm and dropout.
248
-
249
- No positional embeddings (RoPE handles position). Single-chunk inputs
250
- only: ``[CLS] chunk [SEP]``.
251
- """
252
-
253
- def __init__(self, config: DzairConfig) -> None:
254
- super().__init__()
255
- self.word_embeddings = nn.Embedding(
256
- config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id
257
- )
258
- self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps)
259
- self.dropout = nn.Dropout(config.hidden_dropout_prob)
260
-
261
- def forward(self, input_ids: Tensor) -> Tensor:
262
- x = self.word_embeddings(input_ids)
263
- x = self.norm(x)
264
- return self.dropout(x)
265
-
266
-
267
- class Encoder(nn.Module):
268
- """Stack of transformer layers."""
269
-
270
- def __init__(self, config: DzairConfig) -> None:
271
- super().__init__()
272
- self.config = config
273
- self.layers = nn.ModuleList(
274
- TransformerLayer(config) for _ in range(config.num_hidden_layers)
275
- )
276
- self.gradient_checkpointing = False
277
-
278
- def forward(
279
- self,
280
- x: Tensor,
281
- attention_mask: Tensor | None = None,
282
- ) -> Tensor:
283
- for layer in self.layers:
284
- if self.gradient_checkpointing and self.training:
285
- x = checkpoint_utils.checkpoint(
286
- layer, x, attention_mask, False, use_reentrant=False
287
- )
288
- else:
289
- x = layer(x, attention_mask, is_causal=False) # bidirectional
290
- return x
291
-
292
- def forward_with_states(
293
- self,
294
- x: Tensor,
295
- attention_mask: Tensor | None = None,
296
- ) -> tuple[Tensor, tuple[Tensor, ...]]:
297
- """Forward pass returning the last output plus each layer's output."""
298
- states: list[Tensor] = []
299
- for layer in self.layers:
300
- if self.gradient_checkpointing and self.training:
301
- x = checkpoint_utils.checkpoint(
302
- layer, x, attention_mask, False, use_reentrant=False
303
- )
304
- else:
305
- x = layer(x, attention_mask, is_causal=False) # bidirectional
306
- states.append(x)
307
- return x, tuple(states)
308
-
309
-
310
- def set_gradient_checkpointing(model: nn.Module, value: bool) -> None:
311
- """Toggle activation checkpointing on every Encoder in a model.
312
-
313
- Plain attribute propagation, deliberately not via
314
- ``PreTrainedModel.gradient_checkpointing_enable`` whose signature drifted
315
- across transformers versions. The smoke test pins that outputs match and
316
- gradients flow with it on.
317
- """
318
- for module in model.modules():
319
- if isinstance(module, (Encoder, Generator)):
320
- module.gradient_checkpointing = value
321
-
322
-
323
- _ARCH_FIELDS: tuple[str, ...] = (
324
- "vocab_size",
325
- "hidden_size",
326
- "intermediate_size",
327
- "num_attention_heads",
328
- "num_key_value_heads",
329
- "num_hidden_layers",
330
- "num_generator_layers",
331
- "generator_hidden_size",
332
- "generator_intermediate_size",
333
- "max_position_embeddings",
334
- "rope_theta",
335
- "hidden_dropout_prob",
336
- "attention_probs_dropout_prob",
337
- "layer_norm_eps",
338
- "pad_token_id",
339
- "cls_token_id",
340
- "sep_token_id",
341
- "mask_token_id",
342
- "tie_word_embeddings",
343
- "share_generator_embeddings",
344
- "qk_norm",
345
- )
346
-
347
- _CONFIG_MISMATCH_MSG = (
348
- "explicit config disagrees with the checkpoint's stored config on {field}: "
349
- "explicit={explicit!r} stored={stored!r} — pass config=None to trust the checkpoint"
350
- )
351
-
352
- _NO_STORED_CONFIG_MSG = (
353
- "checkpoint {path} carries no stored config and none was passed — "
354
- "pass config=<DzairConfig> explicitly"
355
- )
356
-
357
- _FOLD_MISSING_MSG = (
358
- "GDES checkpoint is missing {missing} — found prefixes: {prefixes}; "
359
- "cannot fold E_G + delta into the released embedding"
360
- )
361
-
362
-
363
- _CHECKSUM_MISMATCH_MSG = "checkpoint checksum mismatch for {path}"
364
-
365
-
366
- def _normalize_stored_config(stored_dict: dict[str, Any]) -> dict[str, Any]:
367
- """Replace a pre-GQA null key-value count with the full-MHA default.
368
-
369
- Runs written before grouped-query attention store no (or null)
370
- key-value count; all of them trained full multi-head attention.
371
- """
372
- normalized = dict(stored_dict)
373
- if normalized.get("num_key_value_heads") is None:
374
- normalized.pop("num_key_value_heads", None)
375
- return normalized
376
-
377
-
378
- def _check_config_match(config: DzairConfig, stored_dict: dict[str, Any]) -> None:
379
- """Raise on any recorded field the explicit config disagrees on.
380
-
381
- Fields the checkpoint predates (absent) or left null are not compared:
382
- the explicit config decides those, so era-appropriate explicit configs
383
- (global attention, full MHA, eval-only dropout) load instead of
384
- refusing on a formatting technicality.
385
- """
386
- for field in _ARCH_FIELDS:
387
- if field not in stored_dict or stored_dict[field] is None:
388
- continue
389
- explicit_value = getattr(config, field, None)
390
- if explicit_value != stored_dict[field]:
391
- raise ValueError(
392
- _CONFIG_MISMATCH_MSG.format(
393
- field=field, explicit=explicit_value, stored=stored_dict[field]
394
- )
395
- )
396
-
397
-
398
- def _read_pretrain_checkpoint(
399
- checkpoint_path: str | Path,
400
- config: DzairConfig | None,
401
- ) -> tuple[DzairConfig, dict[str, Tensor]]:
402
- """Resolve (config, state) from a pretraining checkpoint.
403
-
404
- The checkpoint's stored config wins unless an explicit config is passed;
405
- an explicit config that disagrees with the stored one on a field the
406
- checkpoint actually records raises instead of silently misloading.
407
- Fields the checkpoint predates (absent) or left null are not compared:
408
- the explicit config decides those, so era-appropriate explicit configs
409
- (global attention, full MHA, eval-only dropout) load instead of
410
- refusing on a formatting technicality. Stored nulls/absences for the
411
- key-value count mean the run predates grouped-query attention and
412
- trained full multi-head attention, so they normalize to the default —
413
- never to a silent mismatch.
414
- """
415
- path_obj = Path(checkpoint_path)
416
- sidecar = path_obj.parent / f"{path_obj.name}.sha256"
417
- if sidecar.is_file():
418
- want = sidecar.read_text(encoding="utf-8").strip()
419
- digest = hashlib.sha256()
420
- with path_obj.open("rb") as f:
421
- for chunk in iter(lambda: f.read(1 << 20), b""):
422
- digest.update(chunk)
423
- if digest.hexdigest() != want:
424
- raise ValueError(_CHECKSUM_MISMATCH_MSG.format(path=checkpoint_path))
425
- raw = torch.load(checkpoint_path, map_location="cpu", weights_only=True)
426
- if not isinstance(raw, dict):
427
- msg = f"checkpoint payload is not a mapping: {checkpoint_path}"
428
- raise TypeError(msg)
429
- inner = raw.get("model")
430
- state: dict[str, Tensor] = inner if isinstance(inner, dict) else raw
431
- stored = raw.get("config")
432
- stored_dict = stored if isinstance(stored, dict) else None
433
- if config is not None:
434
- if stored_dict is not None:
435
- _check_config_match(config, stored_dict)
436
- return copy.deepcopy(config), state
437
- if stored_dict is None:
438
- raise ValueError(_NO_STORED_CONFIG_MSG.format(path=checkpoint_path))
439
- return DzairConfig(**_normalize_stored_config(stored_dict)), state
440
-
441
-
442
- def _fold_shared_backbone(state_dict: dict[str, Tensor]) -> dict[str, Tensor]:
443
- """Fold a GDES checkpoint's shared table into one released embedding.
444
-
445
- The released table is ``proj(E_G) + Δ`` — the generator's table through
446
- the width bridge plus the discriminator's delta — with the input norm
447
- taken from the discriminator's ``input_norm``. Same-width (or pre-bridge)
448
- checkpoints skip the projection, exactly like the forward does.
449
- """
450
- out: dict[str, Tensor] = {}
451
- gen_key = "rtd_head.generator.embeddings.word_embeddings.weight"
452
- delta_key = "rtd_head.discriminator.delta_embeddings.weight"
453
- proj_key = "rtd_head.discriminator.gen_proj.weight"
454
- gen_table = state_dict.get(gen_key)
455
- delta = state_dict.get(delta_key)
456
- if gen_table is None or delta is None:
457
- missing = [k for k in (gen_key, delta_key) if k not in state_dict]
458
- prefixes = sorted({".".join(k.split(".")[:2]) if "." in k else k for k in state_dict})
459
- raise KeyError(_FOLD_MISSING_MSG.format(missing=missing, prefixes=prefixes[:8]))
460
- proj = state_dict.get(proj_key)
461
- if proj is None or gen_table.size(-1) == delta.size(-1):
462
- folded = gen_table + delta.to(gen_table.dtype)
463
- else:
464
- folded = gen_table.to(proj.dtype) @ proj.T + delta.to(proj.dtype)
465
- out["embeddings.word_embeddings.weight"] = folded
466
- for key, value in state_dict.items():
467
- if key.startswith("rtd_head.discriminator.input_norm."):
468
- out["embeddings.norm." + key[len("rtd_head.discriminator.input_norm.") :]] = value
469
- elif key.startswith("rtd_head.discriminator.encoder.") or key.startswith(
470
- "rtd_head.discriminator.norm."
471
- ):
472
- out[key[len("rtd_head.discriminator.") :]] = value
473
- return out
474
-
475
-
476
- _FUSED_SPLIT_MSG = (
477
- "cannot map fused {key}: expected ({fused}, {hidden}), "
478
- "or the target is grouped-query ({kv} KV heads over {nq} query heads) "
479
- "which a fused full-MHA table cannot feed without lossy subsampling — "
480
- "retrain or load into a full-MHA config"
481
- )
482
-
483
- _AMBIGUOUS_PROJ_MSG = (
484
- "checkpoint mixes fused ({fused}) and split ({split}) attention projections — "
485
- "refusing instead of guessing which one owns the layer"
486
- )
487
-
488
- _FUSED_BIAS_MSG = (
489
- "cannot map fused {key}: biased projections have no split target — "
490
- "retrain or load into a matching config"
491
- )
492
-
493
-
494
- def _unfuse_in_proj(state_dict: dict[str, Tensor], config: PretrainedConfig) -> dict[str, Tensor]:
495
- """Split fused full-MHA ``in_proj`` tables into ``q_proj`` + ``kv_proj``.
496
-
497
- Checkpoints written before grouped-query attention carry one fused
498
- QKV matrix per layer; current code keeps separate query and fused
499
- key/value projections. The split is exact only into full MHA
500
- (KV heads == query heads) with the canonical Q,K,V row order —
501
- anything else raises instead of silently remapping. Passes through
502
- states without fused tables untouched.
503
- """
504
- fused_keys = [k for k in state_dict if k.endswith("attention.in_proj.weight")]
505
- if not fused_keys:
506
- return state_dict
507
- hidden = int(config.hidden_size)
508
- num_queries = int(config.num_attention_heads)
509
- num_kv = int(getattr(config, "num_key_value_heads", 0) or num_queries)
510
- split_keys = [
511
- k
512
- for k in state_dict
513
- if k.endswith("attention.q_proj.weight") or k.endswith("attention.kv_proj.weight")
514
- ]
515
- if split_keys:
516
- msg = _AMBIGUOUS_PROJ_MSG.format(fused=fused_keys[0], split=split_keys[0])
517
- raise RuntimeError(msg)
518
- biased = [k for k in state_dict if k.endswith("attention.in_proj.bias")]
519
- if biased:
520
- msg = _FUSED_BIAS_MSG.format(key=biased[0])
521
- raise RuntimeError(msg)
522
- out = dict(state_dict)
523
- for key in fused_keys:
524
- weight = state_dict[key]
525
- if tuple(weight.shape) != (3 * hidden, hidden) or num_kv != num_queries:
526
- msg = _FUSED_SPLIT_MSG.format(
527
- key=key,
528
- fused=tuple(weight.shape),
529
- hidden=hidden,
530
- kv=num_kv,
531
- nq=num_queries,
532
- )
533
- raise RuntimeError(msg)
534
- prefix = key[: -len("in_proj.weight")]
535
- query, key_p, value = weight.split([hidden, hidden, hidden], dim=0)
536
- del out[key]
537
- out[prefix + "q_proj.weight"] = query
538
- out[prefix + "kv_proj.weight"] = torch.cat([key_p, value], dim=0)
539
- return out
540
-
541
-
542
- _GEN_TABLE_KEY = "rtd_head.generator.embeddings.word_embeddings.weight"
543
-
544
-
545
- def discriminator_backbone_state(
546
- state_dict: dict[str, Tensor], config: PretrainedConfig
547
- ) -> dict[str, Tensor]:
548
- """Map a pretraining checkpoint's discriminator weights onto ``DzairModel``.
549
-
550
- GDES checkpoints (shared table): the released embedding is the fold
551
- ``proj(E_G) + Δ`` — the generator's table through the width bridge plus
552
- the discriminator's delta — with the input norm taken from the
553
- discriminator's ``input_norm``. Independent checkpoints:
554
- ``rtd_head.discriminator.*`` maps verbatim minus the RTD
555
- classifier. A payload that is already a ``DzairModel`` state dict (no
556
- ``rtd_head`` prefix) passes through; ``strict=True`` on the caller's
557
- ``load_state_dict`` catches anything malformed. Fused full-MHA
558
- ``in_proj`` tables are split exactly (see ``_unfuse_in_proj``);
559
- generator-trunk keys never enter the mapping, so a fused generator
560
- neither helps nor breaks the fold.
561
- """
562
- shared = bool(getattr(config, "share_generator_embeddings", False))
563
- if not any(k.startswith("rtd_head.") for k in state_dict):
564
- return {
565
- (key[len("dzair.") :] if key.startswith("dzair.") else key): value
566
- for key, value in state_dict.items()
567
- }
568
- relevant = {
569
- key: value
570
- for key, value in state_dict.items()
571
- if key.startswith("rtd_head.discriminator.")
572
- or key == _GEN_TABLE_KEY
573
- or key.startswith("dzair.")
574
- }
575
- state_dict = _unfuse_in_proj(relevant, config)
576
- out: dict[str, Tensor] = {}
577
- if shared:
578
- return _fold_shared_backbone(state_dict)
579
- for key, value in state_dict.items():
580
- if key.startswith("rtd_head.discriminator.") and not key.startswith(
581
- "rtd_head.discriminator.classifier"
582
- ):
583
- out[key[len("rtd_head.discriminator.") :]] = value
584
- elif key.startswith("dzair."):
585
- out[key[len("dzair.") :]] = value
586
- return out
587
-
588
-
589
- # Pretraining-only modules absent from older checkpoints: a checkpoint missing
590
- # exactly these still loads, everything else missing or unexpected still raises.
591
- _COMPAT_MISSING_SUBSTRINGS: tuple[str, ...] = ("gen_proj.",)
592
-
593
-
594
- _GENERATION_GAP_MSG = (
595
- "checkpoint uses independent discriminator embeddings "
596
- "('rtd_head.discriminator.embeddings.') but the model expects GDES "
597
- "('rtd_head.discriminator.delta_embeddings.'): no automatic migration — "
598
- "the v1 identity (E_D independent) cannot fold into E_G + delta without "
599
- "changing numerics; retrain or load into a share_generator_embeddings=False "
600
- "config"
601
- )
602
-
603
-
604
- def load_pretrain_state(model: nn.Module, state: dict[str, Tensor]) -> None:
605
- """Load a pretraining state dict across the width-bridge generation gap.
606
-
607
- Checkpoints written before the generator width bridge lack ``gen_proj``;
608
- anything else missing, misshapen, or unexpected still raises.
609
- Fused full-MHA ``in_proj`` tables are split exactly (see
610
- ``_unfuse_in_proj``).
611
- The independent-embeddings (v1) to GDES generation gap is
612
- refused loudly: silently mapping E_D onto delta would change numerics.
613
- """
614
- model_config = getattr(model, "config", None)
615
- if model_config is not None:
616
- gen_keys = {k: v for k, v in state.items() if k.startswith("rtd_head.generator.encoder.")}
617
- trunk_keys = {k: v for k, v in state.items() if k not in gen_keys}
618
- merged_state = _unfuse_in_proj(trunk_keys, model_config)
619
- if gen_keys:
620
- merged_state.update(_unfuse_in_proj(gen_keys, _generator_view(model_config)))
621
- state = merged_state
622
- own = model.state_dict()
623
- if any("rtd_head.discriminator.embeddings." in k for k in state) and any(
624
- "delta_embeddings" in k for k in own
625
- ):
626
- raise RuntimeError(_GENERATION_GAP_MSG)
627
- if any("delta_embeddings" in k for k in state) and any(
628
- "rtd_head.discriminator.embeddings." in k for k in own
629
- ):
630
- raise RuntimeError(_GENERATION_GAP_MSG)
631
- unexpected = [k for k in state if k not in own]
632
- if unexpected:
633
- msg = f"checkpoint holds unexpected keys: {unexpected[:8]}"
634
- raise RuntimeError(msg)
635
- merged: dict[str, Tensor] = {}
636
- absent: list[str] = []
637
- for key, value in own.items():
638
- if key not in state:
639
- absent.append(key)
640
- continue
641
- if value.shape != state[key].shape:
642
- msg = f"checkpoint shape mismatch for {key}: ckpt {tuple(state[key].shape)}"
643
- raise RuntimeError(msg)
644
- merged[key] = state[key]
645
- unaccounted = [k for k in absent if not any(s in k for s in _COMPAT_MISSING_SUBSTRINGS)]
646
- if unaccounted:
647
- msg = f"checkpoint lacks load-bearing keys: {unaccounted[:8]}"
648
- raise RuntimeError(msg)
649
- for key in absent:
650
- merged[key] = own[key]
651
- model.load_state_dict(merged, strict=True)
652
-
653
-
654
- # State keys from retired training objectives. A checkpoint carrying them
655
- # predates the current code: the trunk weights still load, the retired
656
- # heads do not come back. Centralized here so every loader agrees on
657
- # what "obsolete" means; anything else unexpected still raises.
658
- OBSOLETE_STATE_SUBSTRINGS: tuple[str, ...] = (
659
- "order_head.",
660
- "order_loss_ema",
661
- "token_loss_ema",
662
- "token_type_embeddings.",
663
- )
664
-
665
-
666
- @dataclass(frozen=True)
667
- class ResumeCompat:
668
- """How a checkpoint's weights mapped onto the current model."""
669
-
670
- generation: str # "same" (exact) or "legacy" (obsolete keys dropped)
671
- dropped: tuple[str, ...]
672
-
673
-
674
- def load_resume_weights(model: nn.Module, ckpt_model_state: dict[str, Tensor]) -> ResumeCompat:
675
- """Load training weights for an exact resume across code generations.
676
-
677
- Fused full-MHA tables split exactly (trunk and generator widths
678
- handled separately); retired keys drop loudly in the report. Any
679
- other missing, misshapen, or unexpected key raises — a half-mapped
680
- model never trains. The caller decides from ``generation`` whether
681
- the optimizer may be restored (``same``) or must restart fresh
682
- (``legacy``): stale momentum on a reshaped model is silent corruption.
683
- """
684
- raw = {k.removeprefix("_orig_mod."): v for k, v in ckpt_model_state.items()}
685
- model_config = getattr(model, "config", None)
686
- if model_config is not None:
687
- gen_keys = {k: v for k, v in raw.items() if k.startswith("rtd_head.generator.encoder.")}
688
- trunk_keys = {k: v for k, v in raw.items() if k not in gen_keys}
689
- raw = _unfuse_in_proj(trunk_keys, model_config)
690
- if gen_keys:
691
- raw.update(_unfuse_in_proj(gen_keys, _generator_view(model_config)))
692
- dropped = tuple(sorted({k for k in raw if any(s in k for s in OBSOLETE_STATE_SUBSTRINGS)}))
693
- kept = {k: v for k, v in raw.items() if k not in dropped}
694
- raw_model = getattr(model, "_orig_mod", model)
695
- own = raw_model.state_dict()
696
- unexpected = [k for k in kept if k not in own]
697
- if unexpected:
698
- msg = f"checkpoint holds unexpected keys: {unexpected[:8]}"
699
- raise RuntimeError(msg)
700
- missing = [k for k in own if k not in kept]
701
- if missing:
702
- msg = f"checkpoint lacks load-bearing keys: {missing[:8]}"
703
- raise RuntimeError(msg)
704
- for key, value in own.items():
705
- if value.shape != kept[key].shape:
706
- msg = f"checkpoint shape mismatch for {key}: ckpt {tuple(kept[key].shape)}"
707
- raise RuntimeError(msg)
708
- raw_model.load_state_dict(kept, strict=True)
709
- return ResumeCompat(generation="legacy" if dropped else "same", dropped=dropped)
710
-
711
-
712
- def _generator_view(config: DzairConfig) -> DzairConfig:
713
- """A config view sizing the generator trunk: narrow width, global attention.
714
-
715
- The generator keeps head_dim 64 and key-value groups proportional to the
716
- trunk; it always attends globally so corruption quality never depends on
717
- the discriminator's local window. Copies (never mutates) the trunk config.
718
- """
719
- view = copy.copy(config)
720
- view.hidden_size = config.generator_hidden_size
721
- view.intermediate_size = config.generator_intermediate_size
722
- view.num_attention_heads = config.generator_num_heads
723
- view.num_key_value_heads = max(
724
- 1, config.generator_num_heads * config.num_key_value_heads // config.num_attention_heads
725
- )
726
- if view.num_attention_heads % view.num_key_value_heads != 0:
727
- msg = (
728
- f"generator {view.num_attention_heads} query heads must split over "
729
- f"{view.num_key_value_heads} key-value heads"
730
- )
731
- raise ValueError(msg)
732
- view.num_hidden_layers = config.num_generator_layers
733
- return view
734
-
735
-
736
- class Generator(nn.Module):
737
- """Lightweight MLM generator for RTD corruption.
738
-
739
- GDES: embeddings shared, detached for discriminator. The generator trunk
740
- runs at ``generator_hidden_size`` behind a width projection only where it
741
- meets the discriminator (see ``Discriminator.gen_proj``); its own input
742
- and LM head stay in the narrow width with tied tables.
743
- """
744
-
745
- def __init__(self, config: DzairConfig) -> None:
746
- super().__init__()
747
- self.config = config
748
- view = _generator_view(config)
749
- self.embeddings = Embeddings(view)
750
- self.encoder = nn.ModuleList(
751
- TransformerLayer(view) for _ in range(config.num_generator_layers)
752
- )
753
- self.norm = RMSNorm(view.hidden_size, eps=config.layer_norm_eps)
754
- self.lm_head = nn.Linear(view.hidden_size, config.vocab_size, bias=False)
755
- # Tie output embeddings to input embeddings
756
- self.lm_head.weight = self.embeddings.word_embeddings.weight
757
- self.gradient_checkpointing = False
758
-
759
- def forward(
760
- self,
761
- input_ids: Tensor,
762
- attention_mask: Tensor | None = None,
763
- ) -> Tensor:
764
- x = self.embeddings(input_ids)
765
- for layer in self.encoder:
766
- if self.gradient_checkpointing and self.training:
767
- x = checkpoint_utils.checkpoint(
768
- layer, x, attention_mask, False, use_reentrant=False
769
- )
770
- else:
771
- x = layer(x, attention_mask, is_causal=False)
772
- x = self.norm(x)
773
- return self.lm_head(x)
774
-
775
-
776
- class Discriminator(nn.Module):
777
- """RTD discriminator: detects replaced tokens.
778
-
779
- Two embedding policies, selected by ``config.share_generator_embeddings``:
780
-
781
- - **GDES** (True): the discriminator reads ``proj(stop_grad(E_G)) + Δ``
782
- where ``E_G`` is the generator's own (narrow) table — generator MLM
783
- training shapes the table the discriminator reads — ``proj`` bridges
784
- the generator width to the trunk width, and ``Δ`` is this module's own
785
- table. Discriminator gradients flow to ``Δ`` and ``proj`` only, by
786
- construction. The released backbone folds ``proj(E_G) + Δ`` into one
787
- table at load time.
788
- - **Independent** (False, default): a private ``Embeddings`` table, as in
789
- classic ELECTRA. The pretraining head still passes the generator's
790
- table; in this mode it is unused, and the forward is a plain lookup.
791
- """
792
-
793
- def __init__(self, config: DzairConfig) -> None:
794
- super().__init__()
795
- self.config = config
796
- self.share_generator_embeddings = bool(config.share_generator_embeddings)
797
- if self.share_generator_embeddings:
798
- self.delta_embeddings = nn.Embedding(
799
- config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id
800
- )
801
- self.gen_proj = nn.Linear(config.generator_hidden_size, config.hidden_size, bias=False)
802
- self.input_norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps)
803
- self.input_dropout = nn.Dropout(config.hidden_dropout_prob)
804
- else:
805
- self.embeddings = Embeddings(config)
806
- self.encoder = Encoder(config)
807
- self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps)
808
- self.classifier = nn.Linear(config.hidden_size, 2, bias=False) # binary: original/replaced
809
-
810
- def forward(
811
- self,
812
- input_ids: Tensor,
813
- attention_mask: Tensor | None = None,
814
- generator_embeddings: Tensor | None = None,
815
- ) -> Tensor:
816
- """``generator_embeddings`` is the generator's table (required under GDES)."""
817
- if self.share_generator_embeddings:
818
- if generator_embeddings is None:
819
- msg = "share_generator_embeddings=True requires the generator table"
820
- raise ValueError(msg)
821
- base = functional.embedding(
822
- input_ids, generator_embeddings.detach(), padding_idx=self.config.pad_token_id
823
- )
824
- if base.size(-1) != self.config.hidden_size:
825
- base = self.gen_proj(base)
826
- x = base + self.delta_embeddings(input_ids)
827
- x = self.input_dropout(self.input_norm(x))
828
- else:
829
- x = self.embeddings(input_ids)
830
-
831
- x = self.encoder(x, attention_mask)
832
- x = self.norm(x)
833
- return self.classifier(x)
834
-
835
-
836
- def _dynamo_disabled[FnT](function: FnT) -> FnT:
837
- """Exclude CPU-scalar bookkeeping from the compiled graph.
838
-
839
- ``float()`` syncs inside the forward break Dynamo (measured: a
840
- ``Tensor.item()`` graph break every step) and stall the GPU for a value
841
- only needed as an eager loss scale. Falls back to a no-op when
842
- ``disable`` is unavailable; the pinned images always carry it.
843
- """
844
- disable = getattr(_dynamo, "disable", None)
845
- if disable is None:
846
- return function
847
- return disable(function)
848
-
849
-
850
- @_dynamo_disabled
851
- def draw_token_mask(candidate: Tensor, mask_prob: float | Tensor) -> Tensor:
852
- """Per-token uniform draw over candidate positions."""
853
- return candidate & (torch.rand(candidate.shape, device=candidate.device) < mask_prob)
854
-
855
-
856
- @_dynamo_disabled
857
- def draw_word_mask(candidate: Tensor, word_starts: Tensor, mask_prob: float | Tensor) -> Tensor:
858
- """One uniform draw per word; every candidate position in a chosen word masked.
859
-
860
- Positions before the first word start are never masked.
861
- """
862
- device = candidate.device
863
- length = candidate.size(-1)
864
- arange = torch.arange(length, device=device).expand_as(candidate)
865
- cur_start = torch.where(word_starts, arange, -1).cummax(dim=-1).values
866
- chosen = word_starts & candidate & (torch.rand(candidate.shape, device=device) < mask_prob)
867
- last_chosen = torch.where(chosen, arange, -1).cummax(dim=-1).values
868
- return candidate & (cur_start >= 0) & (cur_start == last_chosen)
869
-
870
-
871
- @dataclass(frozen=True)
872
- class MaskSpec:
873
- """What may be masked and how often (built per step from the schedule)."""
874
-
875
- special_ids: frozenset[int]
876
- vocab_size: int
877
- mask_token_id: int
878
- mask_prob: float | Tensor
879
-
880
-
881
- @_dynamo_disabled
882
- def _mask_inputs(
883
- input_ids: Tensor,
884
- eligible: Tensor,
885
- spec: MaskSpec,
886
- word_starts: Tensor | None = None,
887
- ) -> tuple[Tensor, Tensor]:
888
- """BERT 80/10/10 corruption. Returns (masked_input_ids, mlm_labels).
889
-
890
- Masking is whole-word when ``word_starts`` ([B, T] bool, True at
891
- word-initial pieces) is given, else per-token. Special ids (pad/cls/sep)
892
- and ineligible positions are never masked. Dynamic every step.
893
- """
894
- device = input_ids.device
895
- is_special = torch.zeros_like(input_ids, dtype=torch.bool)
896
- for sid in spec.special_ids:
897
- is_special |= input_ids == sid
898
- candidate = eligible & ~is_special
899
-
900
- if word_starts is not None:
901
- masked = draw_word_mask(candidate, word_starts, spec.mask_prob)
902
- else:
903
- masked = draw_token_mask(candidate, spec.mask_prob)
904
-
905
- rand = torch.rand(input_ids.shape, device=device)
906
- replace_mask = masked & (rand < MASK_REPLACE_CUTOFF)
907
- random_mask = masked & (rand >= MASK_REPLACE_CUTOFF) & (rand < MASK_RANDOM_CUTOFF)
908
- # keep_mask (last 10%): input unchanged, still predicted.
909
-
910
- masked_input = input_ids.clone()
911
- masked_input[replace_mask] = spec.mask_token_id
912
- rand_tokens = torch.randint_like(input_ids, 0, spec.vocab_size)
913
- masked_input = torch.where(random_mask, rand_tokens, masked_input)
914
-
915
- mlm_labels = torch.full_like(input_ids, IGNORE_INDEX)
916
- mlm_labels[masked] = input_ids[masked]
917
- return masked_input, mlm_labels
918
-
919
-
920
- @dataclass
921
- class DzairRTDOutput:
922
- """Pretraining output: ELECTRA-style joint loss."""
923
-
924
- loss: Tensor | None
925
- rtd_logits: Tensor
926
- gen_logits: Tensor
927
- generator_loss: Tensor | None
928
- discriminator_loss: Tensor | None
929
- replacement_rate: Tensor
930
-
931
-
932
- @_dynamo_disabled
933
- def _sample_generator_corruptions(
934
- gen_logits: Tensor,
935
- masked_input: Tensor,
936
- predict: Tensor,
937
- input_ids: Tensor,
938
- softmax_chunk: int = 2048,
939
- ) -> Tensor:
940
- """Sample generator tokens on masked positions in chunks outside Dynamo."""
941
- with torch.no_grad():
942
- flat_mask = predict.reshape(-1)
943
- idx = torch.where(flat_mask)[0]
944
- sampled = torch.empty_like(idx)
945
- flat_logits = gen_logits.reshape(-1, gen_logits.size(-1))
946
- for start in range(0, idx.numel(), softmax_chunk):
947
- group = idx[start : start + softmax_chunk]
948
- probs = flat_logits[group].float().softmax(dim=-1)
949
- sampled[start : start + softmax_chunk] = torch.multinomial(probs, 1).squeeze(-1)
950
- corrupted = masked_input.reshape(-1).clone()
951
- corrupted[idx] = sampled
952
- return corrupted.view_as(input_ids)
953
-
954
-
955
- class RTDHead(nn.Module):
956
- """Generator (MLM) corrupts, discriminator (RTD) detects, GDES detaches."""
957
-
958
- def __init__(self, config: DzairConfig, rtd_loss_weight: float = RTD_LOSS_WEIGHT) -> None:
959
- super().__init__()
960
- self.config = config
961
- self.rtd_loss_weight = rtd_loss_weight
962
- self.generator = Generator(config)
963
- self.discriminator = Discriminator(config)
964
-
965
- def forward(
966
- self,
967
- input_ids: Tensor,
968
- attention_mask: Tensor | None = None,
969
- mask_prob: float | Tensor = 0.15,
970
- word_starts: Tensor | None = None,
971
- ) -> DzairRTDOutput:
972
- """Returns the joint output. ``mask_prob`` follows the 30→15% schedule.
973
-
974
- Accepts a 0-dim tensor as well as a float: pass a tensor from any
975
- compiled caller — Dynamo specializes on float argument *values*,
976
- so a per-step float schedule would recompile every step until the
977
- cache limit forces the whole model back to eager.
978
- """
979
- eligible = (
980
- attention_mask.to(torch.bool)
981
- if attention_mask is not None
982
- else torch.ones_like(input_ids, dtype=torch.bool)
983
- )
984
- spec = MaskSpec(
985
- special_ids=frozenset(
986
- sid
987
- for sid in (
988
- self.config.pad_token_id,
989
- self.config.cls_token_id,
990
- self.config.sep_token_id,
991
- )
992
- if sid is not None
993
- ),
994
- vocab_size=self.config.vocab_size,
995
- mask_token_id=self.config.mask_token_id,
996
- mask_prob=mask_prob,
997
- )
998
- masked_input, mlm_labels = _mask_inputs(input_ids, eligible, spec, word_starts)
999
-
1000
- gen_logits = self.generator(masked_input, attention_mask)
1001
- gen_loss: Tensor | None = None
1002
- predict = mlm_labels != IGNORE_INDEX
1003
- if predict.any():
1004
- gen_loss = functional.cross_entropy(gen_logits[predict], input_ids[predict].detach())
1005
-
1006
- # Corrupt only the masked positions by sampling the generator.
1007
- # Softmax runs over the masked subset in bounded chunks to cap the
1008
- # peak transient allocation. Arithmetic (measured 2026-09-10):
1009
- # 8192 * 48000 * 4 bytes = 1.57 GB -- OOMs at 20.94 GB in use
1010
- # 2048 * 48000 * 4 bytes = 0.39 GB -- 4.7 GB headroom at 18.9 GB peak
1011
- # Chunked multinomial is mathematically identical to sampling all at once.
1012
- _softmax_chunk = 2048
1013
- corrupted = _sample_generator_corruptions(
1014
- gen_logits, masked_input, predict, input_ids, _softmax_chunk
1015
- )
1016
-
1017
- disc_logits = self.discriminator(
1018
- corrupted,
1019
- attention_mask,
1020
- generator_embeddings=self.generator.embeddings.word_embeddings.weight,
1021
- )
1022
-
1023
- rtd_labels = torch.where(corrupted == input_ids, 1, 0)
1024
- rtd_labels = torch.where(eligible, rtd_labels, IGNORE_INDEX)
1025
- disc_loss: Tensor | None = None
1026
- if (rtd_labels != IGNORE_INDEX).any():
1027
- disc_loss = functional.cross_entropy(
1028
- disc_logits.reshape(-1, 2), rtd_labels.reshape(-1), ignore_index=IGNORE_INDEX
1029
- )
1030
-
1031
- loss: Tensor | None = None
1032
- if gen_loss is not None and disc_loss is not None:
1033
- loss = gen_loss + self.rtd_loss_weight * disc_loss
1034
-
1035
- with torch.no_grad():
1036
- replacement_rate = (
1037
- (corrupted[predict] != input_ids[predict]).float().mean()
1038
- if predict.any()
1039
- else torch.zeros((), device=input_ids.device)
1040
- )
1041
-
1042
- return DzairRTDOutput(
1043
- loss=loss,
1044
- rtd_logits=disc_logits,
1045
- gen_logits=gen_logits,
1046
- generator_loss=gen_loss,
1047
- discriminator_loss=disc_loss,
1048
- replacement_rate=replacement_rate,
1049
- )
1050
-
1051
-
1052
- class DzairPreTrainedModel(PreTrainedModel):
1053
- config_class = DzairConfig
1054
- base_model_prefix = "dzair"
1055
- supports_gradient_checkpointing = True
1056
- _no_split_modules: ClassVar[list[str]] = ["TransformerLayer"]
1057
-
1058
- def _init_weights(self, module: nn.Module) -> None:
1059
- # Masinissa scaled init, depth-scaled trunc normal; deliberately not
1060
- # config-driven (a field that silently does nothing is worse than none).
1061
- std = math.sqrt(2.0 / (5.0 * self.config.hidden_size))
1062
- if isinstance(module, nn.Linear):
1063
- nn.init.trunc_normal_(module.weight, mean=0.0, std=std, a=-2 * std, b=2 * std)
1064
- if module.bias is not None:
1065
- nn.init.zeros_(module.bias)
1066
- elif isinstance(module, nn.Embedding):
1067
- nn.init.trunc_normal_(module.weight, mean=0.0, std=std, a=-2 * std, b=2 * std)
1068
- elif isinstance(module, RMSNorm):
1069
- nn.init.ones_(module.weight)
1070
-
1071
-
1072
- @dataclass
1073
- class DzairEncoderOutput(ModelOutput):
1074
- """Encoder output: last state plus optional per-layer states.
1075
-
1076
- A dedicated type because the framework's BaseModelOutput pins its state
1077
- fields to FloatTensor, which the checker treats as distinct from Tensor.
1078
- hidden_states[0] is the embedding output (HF convention). Per-layer
1079
- entries are pre-norm layer outputs; last_hidden_state is post-norm, so
1080
- hidden_states[-1] != last_hidden_state by design.
1081
- """
1082
-
1083
- last_hidden_state: Tensor
1084
- hidden_states: tuple[Tensor, ...] | None = None
1085
- attentions: tuple[Tensor, ...] | None = None
1086
-
1087
-
1088
- class DzairModel(DzairPreTrainedModel):
1089
- """The encoder alone (discriminator backbone).
1090
-
1091
- Returns contextualised token representations.
1092
- """
1093
-
1094
- def __init__(self, config: DzairConfig) -> None:
1095
- super().__init__(config)
1096
- self.embeddings = Embeddings(config)
1097
- self.encoder = Encoder(config)
1098
- self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps)
1099
- self.post_init()
1100
-
1101
- def get_input_embeddings(self) -> nn.Embedding:
1102
- return self.embeddings.word_embeddings
1103
-
1104
- def set_input_embeddings(self, value: nn.Embedding) -> None:
1105
- self.embeddings.word_embeddings = value
1106
-
1107
- def forward(
1108
- self,
1109
- input_ids: Tensor,
1110
- attention_mask: Tensor | None = None,
1111
- output_hidden_states: bool = False,
1112
- ) -> DzairEncoderOutput:
1113
- """Encode tokens. With output_hidden_states, hidden_states[0] is the
1114
- embedding output and [k] the k-th layer output (HF convention).
1115
- Layer states are pre-norm; last_hidden_state is post-norm.
1116
- """
1117
- embedded = self.embeddings(input_ids)
1118
- if output_hidden_states:
1119
- last, states = self.encoder.forward_with_states(embedded, attention_mask)
1120
- return DzairEncoderOutput(
1121
- last_hidden_state=self.norm(last),
1122
- hidden_states=(embedded, *states),
1123
- )
1124
- x = self.encoder(embedded, attention_mask)
1125
- return DzairEncoderOutput(last_hidden_state=self.norm(x))
1126
-
1127
-
1128
- class DzairForMaskedLM(DzairPreTrainedModel):
1129
- """Pretraining model: Generator (MLM) + Discriminator (RTD) with GDES."""
1130
-
1131
- _tied_weights_keys: ClassVar[dict[str, str]] = {
1132
- "rtd_head.generator.lm_head.weight": "rtd_head.generator.embeddings.word_embeddings.weight",
1133
- }
1134
-
1135
- def __init__(self, config: DzairConfig) -> None:
1136
- super().__init__(config)
1137
- self.rtd_head = RTDHead(config)
1138
- self.post_init()
1139
-
1140
- def forward(
1141
- self,
1142
- input_ids: Tensor,
1143
- attention_mask: Tensor | None = None,
1144
- mask_prob: float | Tensor = 0.15,
1145
- word_starts: Tensor | None = None,
1146
- ) -> DzairRTDOutput:
1147
- return self.rtd_head(input_ids, attention_mask, mask_prob, word_starts)
1148
-
1149
-
1150
- @dataclass
1151
- class DzairSequenceClassifierOutput(ModelOutput):
1152
- """Output type of DzairForSequenceClassification."""
1153
-
1154
- loss: Tensor | None = None
1155
- logits: Tensor | None = None
1156
- hidden_states: tuple[Tensor, ...] | None = None
1157
- attentions: tuple[Tensor, ...] | None = None
1158
-
1159
-
1160
- class DzairForSequenceClassification(DzairPreTrainedModel):
1161
- """Sequence classification head on top of the DZAIR encoder backbone.
1162
-
1163
- One method, the measured one: [CLS] pooling through an MLP projection
1164
- head (Dropout -> Dense -> GELU -> Dropout) into the classification
1165
- layer. The DZNLI head ablation picked cls+mlp over mean+linear; the
1166
- landmark and attention experiments never measured a win, so they do
1167
- not ship.
1168
- """
1169
-
1170
- def __init__(self, config: DzairConfig) -> None:
1171
- super().__init__(config)
1172
- self.num_labels = getattr(config, "num_labels", 2)
1173
- self.dzair = DzairModel(config)
1174
- self.dropout = nn.Dropout(config.hidden_dropout_prob)
1175
- self.dense = nn.Linear(config.hidden_size, config.hidden_size)
1176
- self.classifier = nn.Linear(config.hidden_size, self.num_labels)
1177
- self.post_init()
1178
-
1179
- def get_input_embeddings(self) -> nn.Embedding:
1180
- return self.dzair.get_input_embeddings()
1181
-
1182
- def set_input_embeddings(self, value: nn.Embedding) -> None:
1183
- self.dzair.set_input_embeddings(value)
1184
-
1185
- def forward(
1186
- self,
1187
- input_ids: Tensor,
1188
- attention_mask: Tensor | None = None,
1189
- labels: Tensor | None = None,
1190
- output_hidden_states: bool = False,
1191
- ) -> DzairSequenceClassifierOutput:
1192
- outputs = self.dzair(
1193
- input_ids, attention_mask=attention_mask, output_hidden_states=output_hidden_states
1194
- )
1195
- pooled_output = self.dropout(outputs.last_hidden_state[:, 0])
1196
- pooled_output = self.dense(pooled_output)
1197
- pooled_output = functional.gelu(pooled_output)
1198
- pooled_output = self.dropout(pooled_output)
1199
- logits = self.classifier(pooled_output)
1200
-
1201
- loss: Tensor | None = None
1202
- if labels is not None:
1203
- if self.num_labels == 1:
1204
- loss = functional.mse_loss(logits.view(-1), labels.view(-1).float())
1205
- else:
1206
- loss = functional.cross_entropy(logits.view(-1, self.num_labels), labels.view(-1))
1207
-
1208
- return DzairSequenceClassifierOutput(
1209
- loss=loss,
1210
- logits=logits,
1211
- hidden_states=outputs.hidden_states,
1212
- attentions=outputs.attentions,
1213
- )
1214
-
1215
- def load_backbone_weights(self, state_dict: dict[str, Tensor]) -> None:
1216
- """Load pretrained discriminator backbone weights into self.dzair."""
1217
- self.dzair.load_state_dict(
1218
- discriminator_backbone_state(state_dict, self.config), strict=True
1219
- )
1220
-
1221
- @classmethod
1222
- def from_pretrained_checkpoint(
1223
- cls,
1224
- checkpoint_path: str | Path,
1225
- config: DzairConfig | None = None,
1226
- num_labels: int = 2,
1227
- ) -> DzairForSequenceClassification:
1228
- """Instantiate classification model and load backbone from pretrain checkpoint.
1229
-
1230
- The checkpoint's stored config is preferred; an explicit config that
1231
- disagrees with it raises (see ``_read_pretrain_checkpoint``).
1232
- """
1233
- model_config, state = _read_pretrain_checkpoint(checkpoint_path, config)
1234
- model_config.num_labels = num_labels
1235
- model = cls(model_config)
1236
- model.load_backbone_weights(state)
1237
- return model
1238
-
1239
-
1240
- @dataclass
1241
- class DzairTokenClassifierOutput(ModelOutput):
1242
- """Output type of DzairForTokenClassification."""
1243
-
1244
- loss: Tensor | None = None
1245
- logits: Tensor | None = None
1246
- hidden_states: tuple[Tensor, ...] | None = None
1247
- attentions: tuple[Tensor, ...] | None = None
1248
-
1249
-
1250
- class DzairForTokenClassification(DzairPreTrainedModel):
1251
- """Token classification head on top of the DZAIR encoder backbone (e.g. for NER/POS)."""
1252
-
1253
- def __init__(self, config: DzairConfig) -> None:
1254
- super().__init__(config)
1255
- self.num_labels = getattr(config, "num_labels", 2)
1256
- self.dzair = DzairModel(config)
1257
- self.dropout = nn.Dropout(config.hidden_dropout_prob)
1258
- self.classifier = nn.Linear(config.hidden_size, self.num_labels)
1259
- self.post_init()
1260
-
1261
- def get_input_embeddings(self) -> nn.Embedding:
1262
- return self.dzair.get_input_embeddings()
1263
-
1264
- def set_input_embeddings(self, value: nn.Embedding) -> None:
1265
- self.dzair.set_input_embeddings(value)
1266
-
1267
- def forward(
1268
- self,
1269
- input_ids: Tensor,
1270
- attention_mask: Tensor | None = None,
1271
- labels: Tensor | None = None,
1272
- ) -> DzairTokenClassifierOutput:
1273
- outputs = self.dzair(input_ids, attention_mask=attention_mask)
1274
- sequence_output = outputs.last_hidden_state
1275
- sequence_output = self.dropout(sequence_output)
1276
- logits = self.classifier(sequence_output)
1277
-
1278
- loss: Tensor | None = None
1279
- if labels is not None:
1280
- loss = functional.cross_entropy(
1281
- logits.view(-1, self.num_labels),
1282
- labels.view(-1),
1283
- ignore_index=-100,
1284
- )
1285
-
1286
- return DzairTokenClassifierOutput(
1287
- loss=loss,
1288
- logits=logits,
1289
- hidden_states=outputs.hidden_states,
1290
- attentions=outputs.attentions,
1291
- )
1292
-
1293
- def load_backbone_weights(self, state_dict: dict[str, Tensor]) -> None:
1294
- """Load pretrained discriminator backbone weights into self.dzair."""
1295
- self.dzair.load_state_dict(
1296
- discriminator_backbone_state(state_dict, self.config), strict=True
1297
- )
1298
-
1299
- @classmethod
1300
- def from_pretrained_checkpoint(
1301
- cls,
1302
- checkpoint_path: str | Path,
1303
- config: DzairConfig | None = None,
1304
- num_labels: int = 2,
1305
- ) -> DzairForTokenClassification:
1306
- """Instantiate token classification model and load backbone from pretrain checkpoint."""
1307
- model_config, state = _read_pretrain_checkpoint(checkpoint_path, config)
1308
- model_config.num_labels = num_labels
1309
- model = cls(model_config)
1310
- model.load_backbone_weights(state)
1311
- return model
1312
-
1313
-
1314
- __all__ = [
1315
- "OBSOLETE_STATE_SUBSTRINGS",
1316
- "RTD_LOSS_WEIGHT",
1317
- "DzairConfig",
1318
- "DzairEncoderOutput",
1319
- "DzairForMaskedLM",
1320
- "DzairForSequenceClassification",
1321
- "DzairForTokenClassification",
1322
- "DzairModel",
1323
- "DzairPreTrainedModel",
1324
- "DzairRTDOutput",
1325
- "DzairSequenceClassifierOutput",
1326
- "DzairTokenClassifierOutput",
1327
- "MaskSpec",
1328
- "ResumeCompat",
1329
- "discriminator_backbone_state",
1330
- "draw_token_mask",
1331
- "draw_word_mask",
1332
- "load_pretrain_state",
1333
- "load_resume_weights",
1334
- "set_gradient_checkpointing",
1335
- ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
modeling_dzair.py CHANGED
@@ -203,7 +203,16 @@ def _build_rope_cache(
203
 
204
 
205
  class RMSNorm(nn.Module):
206
- """Root Mean Square Layer Normalization (affine weight, no bias)."""
 
 
 
 
 
 
 
 
 
207
 
208
  def __init__(self, dim: int, eps: float = 1e-5) -> None:
209
  super().__init__()
@@ -211,8 +220,9 @@ class RMSNorm(nn.Module):
211
  self.weight = nn.Parameter(torch.ones(dim))
212
 
213
  def forward(self, x: Tensor) -> Tensor:
214
- norm = x.pow(2).mean(dim=-1, keepdim=True).add(self.eps).rsqrt()
215
- return x * norm * self.weight
 
216
 
217
 
218
  class SwiGLU(nn.Module):
@@ -253,6 +263,11 @@ class Attention(nn.Module):
253
  complexity without evidence. Fused projections, bias-free, pre-RMSNorm.
254
  """
255
 
 
 
 
 
 
256
  def __init__(self, config: DzairConfig) -> None:
257
  super().__init__()
258
  self.config = config
@@ -278,20 +293,20 @@ class Attention(nn.Module):
278
  )
279
  self.dropout = nn.Dropout(config.attention_probs_dropout_prob)
280
 
281
- # RoPE cache (non-persistent, rebuilt on first use)
282
- self._cos: Tensor | None = None
283
- self._sin: Tensor | None = None
284
- self._cache_len = 0
 
285
 
286
  def _get_rope(self, seq_len: int, device: torch.device) -> tuple[Tensor, Tensor]:
287
- if self._cos is None or self._cache_len < seq_len or self._cos.device != device:
288
  self._cos, self._sin = _build_rope_cache(
289
  max(seq_len, self.config.max_position_embeddings),
290
  self.head_size,
291
  self.config.rope_theta,
292
  device,
293
  )
294
- self._cache_len = max(seq_len, self.config.max_position_embeddings)
295
  return self._cos[:seq_len], self._sin[:seq_len]
296
 
297
  def forward(
 
203
 
204
 
205
  class RMSNorm(nn.Module):
206
+ """Root Mean Square Layer Normalization (affine weight, no bias).
207
+
208
+ The statistic is computed in float32 and the result cast back. In half
209
+ precision the square overflows: trunk activations reach 316, and
210
+ 316 squared is 99,856 against a float16 maximum of 65,504, so the mean
211
+ becomes inf, its reciprocal square root becomes 0, and the encoder
212
+ returns an all-zero hidden state (measured 2026-09-19 on the released
213
+ fp16 build). float32 inputs are unaffected — the upcast is a no-op and
214
+ outputs stay bit-identical.
215
+ """
216
 
217
  def __init__(self, dim: int, eps: float = 1e-5) -> None:
218
  super().__init__()
 
220
  self.weight = nn.Parameter(torch.ones(dim))
221
 
222
  def forward(self, x: Tensor) -> Tensor:
223
+ working = x.float()
224
+ norm = working.pow(2).mean(dim=-1, keepdim=True).add(self.eps).rsqrt()
225
+ return (working * norm * self.weight.float()).to(x.dtype)
226
 
227
 
228
  class SwiGLU(nn.Module):
 
263
  complexity without evidence. Fused projections, bias-free, pre-RMSNorm.
264
  """
265
 
266
+ # Declared so the registered buffers carry a type; register_buffer alone
267
+ # leaves them untyped for the checker.
268
+ _cos: Tensor
269
+ _sin: Tensor
270
+
271
  def __init__(self, config: DzairConfig) -> None:
272
  super().__init__()
273
  self.config = config
 
293
  )
294
  self.dropout = nn.Dropout(config.attention_probs_dropout_prob)
295
 
296
+ # RoPE cache: non-persistent buffers, so they stay out of the state
297
+ # dict but are visible to the ONNX exporter, which warns about plain
298
+ # attributes assigned during a traced forward.
299
+ self.register_buffer("_cos", torch.empty(0), persistent=False)
300
+ self.register_buffer("_sin", torch.empty(0), persistent=False)
301
 
302
  def _get_rope(self, seq_len: int, device: torch.device) -> tuple[Tensor, Tensor]:
303
+ if self._cos.numel() == 0 or self._cos.size(0) < seq_len or self._cos.device != device:
304
  self._cos, self._sin = _build_rope_cache(
305
  max(seq_len, self.config.max_position_embeddings),
306
  self.head_size,
307
  self.config.rope_theta,
308
  device,
309
  )
 
310
  return self._cos[:seq_len], self._sin[:seq_len]
311
 
312
  def forward(
tokenizer_config.json CHANGED
@@ -1,15 +1,20 @@
1
  {
2
- "tokenizer_class": "LlamaTokenizer",
 
 
 
 
3
  "unk_token": "[UNK]",
4
  "pad_token": "[PAD]",
5
- "additional_special_tokens": [
6
- "[CLS]",
7
- "[SEP]",
8
- "[MASK]"
9
- ],
10
- "add_bos_token": false,
11
- "add_eos_token": false,
 
12
  "model_max_length": 512,
13
  "latin_lowercase_before_encode": true,
14
- "note": "Wrap single chunks manually as [CLS] chunk [SEP]; see the model card."
15
  }
 
1
  {
2
+ "tokenizer_class": "DebertaV2Tokenizer",
3
+ "model_input_names": [
4
+ "input_ids",
5
+ "attention_mask"
6
+ ],
7
  "unk_token": "[UNK]",
8
  "pad_token": "[PAD]",
9
+ "cls_token": "[CLS]",
10
+ "sep_token": "[SEP]",
11
+ "mask_token": "[MASK]",
12
+ "do_lower_case": false,
13
+ "keep_accents": true,
14
+ "split_by_punct": false,
15
+ "padding_side": "right",
16
+ "truncation_side": "right",
17
  "model_max_length": 512,
18
  "latin_lowercase_before_encode": true,
19
+ "note": "Lowercase Latin spans before encoding; the tokenizer then wraps input as [CLS] chunk [SEP]. See the model card."
20
  }
tokenizer_rules.yaml ADDED
@@ -0,0 +1,48 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # DZAIR tokenizer rules. Version lives inside this file, never in the filename.
2
+ # Rule changes are version bumps with entries in tokenizer_rules.history.md.
3
+ # This file freezes task 2.5's selection; the sweep that produced it is
4
+ # recorded in data/processed/tokenizer/sweep.stats.json.
5
+ version: 1
6
+ algorithm: unigram
7
+ implementation: sentencepiece
8
+ normalization_rule_name: identity
9
+ byte_fallback: true
10
+ split_digits: true
11
+ character_coverage: 0.9995
12
+ max_sentencepiece_length: 16
13
+ specials:
14
+ pad: {piece: "[PAD]", id: 0}
15
+ unk: {piece: "[UNK]", id: 1}
16
+ cls: {piece: "[CLS]", id: 2}
17
+ sep: {piece: "[SEP]", id: 3}
18
+ mask: {piece: "[MASK]", user_defined: true}
19
+ required_chars:
20
+ arabic_block: "U+0600-U+06FF alpha, minus unified alefs U+0622/U+0623/U+0625, tatweel U+0640, Arabic-Indic digits, combining marks"
21
+ latin: "a-z only (normalisation v1 lowercases Latin; uppercase slots would never train)"
22
+ french_accents: "U+00E0 U+00E2 U+00E6 U+00E7 U+00E9 U+00E8 U+00EA U+00EB U+00EE U+00EF U+00F4 U+0153 U+00F9 U+00FB U+00FC U+00FF"
23
+ digits: "0-9 ASCII"
24
+ apostrophes: "U+0027 U+2019 (guaranteed, never pre-split)"
25
+ hyphen: "U+002D (splits Arabizi/French compounds)"
26
+ training:
27
+ cut: full
28
+ input_lines: 4000000
29
+ seed: 42
30
+ input_sentence_size: 4000000
31
+ shuffle_input_sentence: true
32
+ sweep: [16000, 24000, 32000, 48000, 64000]
33
+ nfkc_arm: dead
34
+ selection:
35
+ name: dzair-tok-48k
36
+ vocab_size: 48000
37
+ normalization: identity
38
+ criterion: min fertility at same-or-smaller vocab vs DziriBERT 1.4370 with zero UNK-class failures
39
+ fertility_overall: 1.4503
40
+ reopen: Phase 3 downstream rank overturns, or flat downstream selects 64k on compression; fertility alone never reopens
41
+ inference_preprocessing:
42
+ - NFC (training text is normalisation-v1 output, already NFC)
43
+ - lowercase Latin (training text is lowercased; raw uppercase fragments into bytes at +10% fertility, measured 2.6)
44
+ - no transliteration, no script unification (lossy by Guellil evidence)
45
+ robustness:
46
+ report: data/processed/tokenizer/robustness.stats.json
47
+ unk_hits: 0
48
+ roundtrip_failures: 0