Commit ·
1b7bd7b
0
Parent(s):
Squash clean branch history
Browse files- .gitattributes +35 -0
- README.md +3 -0
- attention_config.py +5 -0
- config.json +67 -0
- configuration_dat.py +198 -0
- dat_config.py +283 -0
- dat_core.py +328 -0
- dat_lm.py +475 -0
- dat_symbols.py +263 -0
- masks.py +61 -0
- mlm.py +95 -0
- model.safetensors +3 -0
- modeling_dat.py +290 -0
- special_tokens_map.json +1 -0
- tokenizer.json +0 -0
- tokenizer_config.json +10 -0
- transformer_components.py +282 -0
- transformer_core.py +327 -0
.gitattributes
ADDED
|
@@ -0,0 +1,35 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
*.7z filter=lfs diff=lfs merge=lfs -text
|
| 2 |
+
*.arrow filter=lfs diff=lfs merge=lfs -text
|
| 3 |
+
*.bin filter=lfs diff=lfs merge=lfs -text
|
| 4 |
+
*.bz2 filter=lfs diff=lfs merge=lfs -text
|
| 5 |
+
*.ckpt filter=lfs diff=lfs merge=lfs -text
|
| 6 |
+
*.ftz filter=lfs diff=lfs merge=lfs -text
|
| 7 |
+
*.gz filter=lfs diff=lfs merge=lfs -text
|
| 8 |
+
*.h5 filter=lfs diff=lfs merge=lfs -text
|
| 9 |
+
*.joblib filter=lfs diff=lfs merge=lfs -text
|
| 10 |
+
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
| 11 |
+
*.mlmodel filter=lfs diff=lfs merge=lfs -text
|
| 12 |
+
*.model filter=lfs diff=lfs merge=lfs -text
|
| 13 |
+
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
| 14 |
+
*.npy filter=lfs diff=lfs merge=lfs -text
|
| 15 |
+
*.npz filter=lfs diff=lfs merge=lfs -text
|
| 16 |
+
*.onnx filter=lfs diff=lfs merge=lfs -text
|
| 17 |
+
*.ot filter=lfs diff=lfs merge=lfs -text
|
| 18 |
+
*.parquet filter=lfs diff=lfs merge=lfs -text
|
| 19 |
+
*.pb filter=lfs diff=lfs merge=lfs -text
|
| 20 |
+
*.pickle filter=lfs diff=lfs merge=lfs -text
|
| 21 |
+
*.pkl filter=lfs diff=lfs merge=lfs -text
|
| 22 |
+
*.pt filter=lfs diff=lfs merge=lfs -text
|
| 23 |
+
*.pth filter=lfs diff=lfs merge=lfs -text
|
| 24 |
+
*.rar filter=lfs diff=lfs merge=lfs -text
|
| 25 |
+
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
| 26 |
+
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
| 27 |
+
*.tar.* filter=lfs diff=lfs merge=lfs -text
|
| 28 |
+
*.tar filter=lfs diff=lfs merge=lfs -text
|
| 29 |
+
*.tflite filter=lfs diff=lfs merge=lfs -text
|
| 30 |
+
*.tgz filter=lfs diff=lfs merge=lfs -text
|
| 31 |
+
*.wasm filter=lfs diff=lfs merge=lfs -text
|
| 32 |
+
*.xz filter=lfs diff=lfs merge=lfs -text
|
| 33 |
+
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
+
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
+
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
README.md
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: apache-2.0
|
| 3 |
+
---
|
attention_config.py
ADDED
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Shared attention configuration constants."""
|
| 2 |
+
|
| 3 |
+
SUPPORTED_SEQUENCE_BOUNDARY_POLICIES = frozenset(
|
| 4 |
+
{"eos_document", "none", "segment_document"}
|
| 5 |
+
)
|
config.json
ADDED
|
@@ -0,0 +1,67 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architectures": [
|
| 3 |
+
"DatForCausalLM"
|
| 4 |
+
],
|
| 5 |
+
"auto_map": {
|
| 6 |
+
"AutoConfig": "configuration_dat.DatConfig",
|
| 7 |
+
"AutoModel": "modeling_dat.DatModel",
|
| 8 |
+
"AutoModelForCausalLM": "modeling_dat.DatForCausalLM",
|
| 9 |
+
"AutoModelForMaskedLM": "modeling_dat.DatForMaskedLM",
|
| 10 |
+
"AutoModelForSequenceClassification": "modeling_dat.DatForSequenceClassification"
|
| 11 |
+
},
|
| 12 |
+
"bos_token_id": 1,
|
| 13 |
+
"dff_factor": 4,
|
| 14 |
+
"dropout": 0.1,
|
| 15 |
+
"dtype": "float32",
|
| 16 |
+
"eos_token_id": 2,
|
| 17 |
+
"ffn_activation": "swiglu",
|
| 18 |
+
"ffn_hidden_dim_mode": "dff_factor",
|
| 19 |
+
"hidden_dim": 768,
|
| 20 |
+
"init_range": 0.15,
|
| 21 |
+
"init_scheme": "normal_0_02_scaled_projection",
|
| 22 |
+
"max_position_embeddings": 513,
|
| 23 |
+
"max_rel_pos": 512,
|
| 24 |
+
"max_seq_len": 513,
|
| 25 |
+
"mlm_head_enabled": false,
|
| 26 |
+
"model_type": "dat",
|
| 27 |
+
"n_heads_ra": 3,
|
| 28 |
+
"n_heads_sa": 9,
|
| 29 |
+
"n_layers": 12,
|
| 30 |
+
"n_symbols": null,
|
| 31 |
+
"norm_first": true,
|
| 32 |
+
"norm_type": "layernorm",
|
| 33 |
+
"pad_token_id": 3,
|
| 34 |
+
"pe_type": "rope",
|
| 35 |
+
"positional_symbols_sinusoidal": false,
|
| 36 |
+
"ra_n_relations": null,
|
| 37 |
+
"ra_rel_activation": "identity",
|
| 38 |
+
"ra_symmetric_rels": false,
|
| 39 |
+
"ra_type": "rca",
|
| 40 |
+
"relative_symbols_rope": true,
|
| 41 |
+
"relsymbolic_dropout": 0.0,
|
| 42 |
+
"relsymbolic_include_self": false,
|
| 43 |
+
"relsymbolic_neighborhood_size": 2,
|
| 44 |
+
"relsymbolic_normalize_rels": true,
|
| 45 |
+
"relsymbolic_rel_n_heads": 4,
|
| 46 |
+
"relsymbolic_rel_scale": null,
|
| 47 |
+
"relsymbolic_symbolic_attn_n_heads": 4,
|
| 48 |
+
"relsymbolic_symbolic_attn_scale": null,
|
| 49 |
+
"relsymbolic_trainable_symbols": true,
|
| 50 |
+
"relsymbolic_use_bias": false,
|
| 51 |
+
"rope_theta": 10000.0,
|
| 52 |
+
"segment_boundary_token_id": null,
|
| 53 |
+
"sequence_boundary_policy": "eos_document",
|
| 54 |
+
"share_attn_params": false,
|
| 55 |
+
"shared_symbol_retriever": true,
|
| 56 |
+
"symbol_dim": null,
|
| 57 |
+
"symbol_retrieval": "relative",
|
| 58 |
+
"symbolic_attn_n_heads": null,
|
| 59 |
+
"symbolic_use_bias": false,
|
| 60 |
+
"tie_lm_head": true,
|
| 61 |
+
"tie_word_embeddings": true,
|
| 62 |
+
"transformers_version": "5.12.0",
|
| 63 |
+
"use_bias_ffn": true,
|
| 64 |
+
"use_bias_out": false,
|
| 65 |
+
"use_bias_qkv": false,
|
| 66 |
+
"vocab_size": 16384
|
| 67 |
+
}
|
configuration_dat.py
ADDED
|
@@ -0,0 +1,198 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""HuggingFace PretrainedConfig for the dual-attention (DAT) decoder LM.
|
| 2 |
+
|
| 3 |
+
This is the source-of-truth copy. scripts/convert_dat_to_hf.py copies it into a
|
| 4 |
+
generated HF repository (alongside modeling_dat.py and the flattened model
|
| 5 |
+
source) so the model can be loaded with
|
| 6 |
+
``AutoModelForCausalLM.from_pretrained(path, trust_remote_code=True)``.
|
| 7 |
+
|
| 8 |
+
The config simply carries every field of models.attention.dat.config.DatLMConfig
|
| 9 |
+
so that modeling_dat.py can rebuild the exact DatLMConfig at load time.
|
| 10 |
+
"""
|
| 11 |
+
|
| 12 |
+
from transformers import PretrainedConfig
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
# Every field needed to rebuild DatLMConfig. New exported fields are appended
|
| 16 |
+
# after existing constructor parameters to preserve positional compatibility.
|
| 17 |
+
# modeling_dat.py rebuilds DatLMConfig as
|
| 18 |
+
# DatLMConfig(**{f: getattr(config, f) for f in DAT_LM_FIELDS}).
|
| 19 |
+
DAT_LM_FIELDS = (
|
| 20 |
+
"vocab_size",
|
| 21 |
+
"max_seq_len",
|
| 22 |
+
"pe_type",
|
| 23 |
+
"hidden_dim",
|
| 24 |
+
"n_heads_sa",
|
| 25 |
+
"n_heads_ra",
|
| 26 |
+
"n_layers",
|
| 27 |
+
"dropout",
|
| 28 |
+
"dff_factor",
|
| 29 |
+
"ffn_hidden_dim_mode",
|
| 30 |
+
"ffn_activation",
|
| 31 |
+
"rope_theta",
|
| 32 |
+
"max_rel_pos",
|
| 33 |
+
"init_range",
|
| 34 |
+
"init_scheme",
|
| 35 |
+
"norm_type",
|
| 36 |
+
"norm_first",
|
| 37 |
+
"use_bias_qkv",
|
| 38 |
+
"use_bias_out",
|
| 39 |
+
"use_bias_ffn",
|
| 40 |
+
"tie_lm_head",
|
| 41 |
+
"symbol_dim",
|
| 42 |
+
"n_symbols",
|
| 43 |
+
"symbolic_attn_n_heads",
|
| 44 |
+
"symbol_retrieval",
|
| 45 |
+
"symbolic_use_bias",
|
| 46 |
+
"shared_symbol_retriever",
|
| 47 |
+
"share_attn_params",
|
| 48 |
+
"positional_symbols_sinusoidal",
|
| 49 |
+
"relative_symbols_rope",
|
| 50 |
+
"relsymbolic_rel_n_heads",
|
| 51 |
+
"relsymbolic_symbolic_attn_n_heads",
|
| 52 |
+
"relsymbolic_neighborhood_size",
|
| 53 |
+
"relsymbolic_include_self",
|
| 54 |
+
"relsymbolic_normalize_rels",
|
| 55 |
+
"relsymbolic_trainable_symbols",
|
| 56 |
+
"relsymbolic_dropout",
|
| 57 |
+
"relsymbolic_rel_scale",
|
| 58 |
+
"relsymbolic_symbolic_attn_scale",
|
| 59 |
+
"relsymbolic_use_bias",
|
| 60 |
+
"ra_type",
|
| 61 |
+
"ra_n_relations",
|
| 62 |
+
"ra_rel_activation",
|
| 63 |
+
"ra_symmetric_rels",
|
| 64 |
+
"sequence_boundary_policy",
|
| 65 |
+
"segment_boundary_token_id",
|
| 66 |
+
"pad_token_id",
|
| 67 |
+
"bos_token_id",
|
| 68 |
+
"eos_token_id",
|
| 69 |
+
"mlm_head_enabled",
|
| 70 |
+
)
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
class DatConfig(PretrainedConfig):
|
| 74 |
+
model_type = "dat"
|
| 75 |
+
# Expose the model width under HF's conventional name so generic eval
|
| 76 |
+
# harnesses (e.g. babylm-eval's finetuning head, which reads
|
| 77 |
+
# config.hidden_size) can size their classification head. The model's own
|
| 78 |
+
# field is hidden_dim; this maps the alias onto it for both get and set.
|
| 79 |
+
attribute_map = {"hidden_size": "hidden_dim"}
|
| 80 |
+
|
| 81 |
+
def __init__(
|
| 82 |
+
self,
|
| 83 |
+
vocab_size: int = 16384,
|
| 84 |
+
max_seq_len: int = 513,
|
| 85 |
+
pe_type: str = "rope",
|
| 86 |
+
hidden_dim: int = 256,
|
| 87 |
+
n_heads_sa: int = 2,
|
| 88 |
+
n_heads_ra: int = 2,
|
| 89 |
+
n_layers: int = 4,
|
| 90 |
+
dropout: float = 0.0,
|
| 91 |
+
dff_factor: int = 4,
|
| 92 |
+
ffn_hidden_dim_mode: str = "dff_factor",
|
| 93 |
+
ffn_activation: str = "gelu",
|
| 94 |
+
rope_theta: float = 10000.0,
|
| 95 |
+
max_rel_pos: int | None = None,
|
| 96 |
+
init_range: float = 0.15,
|
| 97 |
+
init_scheme: str = "xavier_uniform",
|
| 98 |
+
norm_type: str = "rmsnorm",
|
| 99 |
+
norm_first: bool = True,
|
| 100 |
+
use_bias_qkv: bool = False,
|
| 101 |
+
use_bias_out: bool = True,
|
| 102 |
+
use_bias_ffn: bool = True,
|
| 103 |
+
tie_lm_head: bool = True,
|
| 104 |
+
symbol_dim: int | None = None,
|
| 105 |
+
n_symbols: int | None = None,
|
| 106 |
+
symbolic_attn_n_heads: int | None = None,
|
| 107 |
+
symbol_retrieval: str = "symbolic",
|
| 108 |
+
symbolic_use_bias: bool = False,
|
| 109 |
+
shared_symbol_retriever: bool = True,
|
| 110 |
+
share_attn_params: bool = False,
|
| 111 |
+
positional_symbols_sinusoidal: bool = False,
|
| 112 |
+
relative_symbols_rope: bool = False,
|
| 113 |
+
relsymbolic_rel_n_heads: int = 4,
|
| 114 |
+
relsymbolic_symbolic_attn_n_heads: int = 4,
|
| 115 |
+
relsymbolic_neighborhood_size: int = 2,
|
| 116 |
+
relsymbolic_include_self: bool = False,
|
| 117 |
+
relsymbolic_normalize_rels: bool = True,
|
| 118 |
+
relsymbolic_trainable_symbols: bool = True,
|
| 119 |
+
relsymbolic_dropout: float = 0.0,
|
| 120 |
+
relsymbolic_rel_scale: float | None = None,
|
| 121 |
+
relsymbolic_symbolic_attn_scale: float | None = None,
|
| 122 |
+
relsymbolic_use_bias: bool = False,
|
| 123 |
+
ra_type: str = "ra",
|
| 124 |
+
ra_n_relations: int | None = None,
|
| 125 |
+
ra_rel_activation: str = "identity",
|
| 126 |
+
ra_symmetric_rels: bool = False,
|
| 127 |
+
sequence_boundary_policy: str = "eos_document",
|
| 128 |
+
segment_boundary_token_id: int | None = None,
|
| 129 |
+
pad_token_id: int = 0,
|
| 130 |
+
bos_token_id: int = 1,
|
| 131 |
+
eos_token_id: int = 2,
|
| 132 |
+
mlm_head_enabled: bool = False,
|
| 133 |
+
**kwargs,
|
| 134 |
+
) -> None:
|
| 135 |
+
self.vocab_size = vocab_size
|
| 136 |
+
self.max_seq_len = max_seq_len
|
| 137 |
+
self.pe_type = pe_type
|
| 138 |
+
self.hidden_dim = hidden_dim
|
| 139 |
+
self.n_heads_sa = n_heads_sa
|
| 140 |
+
self.n_heads_ra = n_heads_ra
|
| 141 |
+
self.n_layers = n_layers
|
| 142 |
+
self.dropout = dropout
|
| 143 |
+
self.dff_factor = dff_factor
|
| 144 |
+
self.ffn_hidden_dim_mode = ffn_hidden_dim_mode
|
| 145 |
+
self.ffn_activation = ffn_activation
|
| 146 |
+
self.rope_theta = rope_theta
|
| 147 |
+
self.max_rel_pos = max_rel_pos
|
| 148 |
+
self.init_range = init_range
|
| 149 |
+
self.init_scheme = init_scheme
|
| 150 |
+
self.norm_type = norm_type
|
| 151 |
+
self.norm_first = norm_first
|
| 152 |
+
self.use_bias_qkv = use_bias_qkv
|
| 153 |
+
self.use_bias_out = use_bias_out
|
| 154 |
+
self.use_bias_ffn = use_bias_ffn
|
| 155 |
+
self.tie_lm_head = tie_lm_head
|
| 156 |
+
self.symbol_dim = symbol_dim
|
| 157 |
+
self.n_symbols = n_symbols
|
| 158 |
+
self.symbolic_attn_n_heads = symbolic_attn_n_heads
|
| 159 |
+
self.symbol_retrieval = symbol_retrieval
|
| 160 |
+
self.symbolic_use_bias = symbolic_use_bias
|
| 161 |
+
self.shared_symbol_retriever = shared_symbol_retriever
|
| 162 |
+
self.share_attn_params = share_attn_params
|
| 163 |
+
self.positional_symbols_sinusoidal = positional_symbols_sinusoidal
|
| 164 |
+
self.relative_symbols_rope = relative_symbols_rope
|
| 165 |
+
self.relsymbolic_rel_n_heads = relsymbolic_rel_n_heads
|
| 166 |
+
self.relsymbolic_symbolic_attn_n_heads = relsymbolic_symbolic_attn_n_heads
|
| 167 |
+
self.relsymbolic_neighborhood_size = relsymbolic_neighborhood_size
|
| 168 |
+
self.relsymbolic_include_self = relsymbolic_include_self
|
| 169 |
+
self.relsymbolic_normalize_rels = relsymbolic_normalize_rels
|
| 170 |
+
self.relsymbolic_trainable_symbols = relsymbolic_trainable_symbols
|
| 171 |
+
self.relsymbolic_dropout = relsymbolic_dropout
|
| 172 |
+
self.relsymbolic_rel_scale = relsymbolic_rel_scale
|
| 173 |
+
self.relsymbolic_symbolic_attn_scale = relsymbolic_symbolic_attn_scale
|
| 174 |
+
self.relsymbolic_use_bias = relsymbolic_use_bias
|
| 175 |
+
self.ra_type = ra_type
|
| 176 |
+
self.ra_n_relations = ra_n_relations
|
| 177 |
+
self.ra_rel_activation = ra_rel_activation
|
| 178 |
+
self.ra_symmetric_rels = ra_symmetric_rels
|
| 179 |
+
self.sequence_boundary_policy = sequence_boundary_policy
|
| 180 |
+
self.segment_boundary_token_id = segment_boundary_token_id
|
| 181 |
+
self.mlm_head_enabled = mlm_head_enabled
|
| 182 |
+
# HF base also exposes max_position_embeddings for generation utilities.
|
| 183 |
+
self.max_position_embeddings = max_seq_len
|
| 184 |
+
# tie_lm_head is the source of truth for weight tying; keep HF's
|
| 185 |
+
# tie_word_embeddings in lockstep so from_pretrained never re-ties an
|
| 186 |
+
# untied head (or fails to tie a tied one). Drop any value coming in via
|
| 187 |
+
# kwargs (e.g. a serialized config.json) so the two cannot disagree.
|
| 188 |
+
# NOTE: tie_lm_head controls only lm_head <-> token_embeddings tying.
|
| 189 |
+
# The mlm_head.linear_out <-> lm_head tying is handled separately by
|
| 190 |
+
# the get_expanded_tied_weights_keys override in modeling_dat.py.
|
| 191 |
+
kwargs.pop("tie_word_embeddings", None)
|
| 192 |
+
super().__init__(
|
| 193 |
+
pad_token_id=pad_token_id,
|
| 194 |
+
bos_token_id=bos_token_id,
|
| 195 |
+
eos_token_id=eos_token_id,
|
| 196 |
+
tie_word_embeddings=tie_lm_head,
|
| 197 |
+
**kwargs,
|
| 198 |
+
)
|
dat_config.py
ADDED
|
@@ -0,0 +1,283 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Configuration for the owned dual-attention decoder LM."""
|
| 2 |
+
|
| 3 |
+
from dataclasses import dataclass
|
| 4 |
+
|
| 5 |
+
from .attention_config import SUPPORTED_SEQUENCE_BOUNDARY_POLICIES
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
SUPPORTED_PE_TYPES = {"sinusoidal", "learned", "relative", "rope", "none"}
|
| 9 |
+
SUPPORTED_SYMBOL_RETRIEVAL = {"symbolic", "positional", "relative", "relsymbolic"}
|
| 10 |
+
SUPPORTED_RA_TYPES = {"ra", "rca", "disrca"}
|
| 11 |
+
SUPPORTED_RA_ACTIVATIONS = {"softmax", "identity", "relu", "tanh", "sigmoid", "gelu"}
|
| 12 |
+
SUPPORTED_FFN_ACTIVATIONS = {"gelu", "relu", "swiglu", "identity"}
|
| 13 |
+
SUPPORTED_FFN_HIDDEN_DIM_MODES = {"dff_factor", "swiglu_parameter_matched"}
|
| 14 |
+
SUPPORTED_INIT_SCHEMES = {"xavier_uniform", "normal_0_02_scaled_projection"}
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
@dataclass(frozen=True)
|
| 18 |
+
class DatLMConfig:
|
| 19 |
+
vocab_size: int
|
| 20 |
+
max_seq_len: int
|
| 21 |
+
pe_type: str = "rope"
|
| 22 |
+
hidden_dim: int = 256
|
| 23 |
+
n_heads_sa: int = 2
|
| 24 |
+
n_heads_ra: int = 2
|
| 25 |
+
n_layers: int = 4
|
| 26 |
+
dropout: float = 0.0
|
| 27 |
+
dff_factor: int = 4
|
| 28 |
+
ffn_hidden_dim_mode: str = "dff_factor"
|
| 29 |
+
ffn_activation: str = "gelu"
|
| 30 |
+
rope_theta: float = 10000.0
|
| 31 |
+
max_rel_pos: int | None = None
|
| 32 |
+
sequence_boundary_policy: str = "eos_document"
|
| 33 |
+
segment_boundary_token_id: int | None = None
|
| 34 |
+
init_range: float = 0.15
|
| 35 |
+
init_scheme: str = "xavier_uniform"
|
| 36 |
+
norm_type: str = "rmsnorm"
|
| 37 |
+
norm_first: bool = True
|
| 38 |
+
use_bias_qkv: bool = False
|
| 39 |
+
use_bias_out: bool = True
|
| 40 |
+
use_bias_ffn: bool = True
|
| 41 |
+
tie_lm_head: bool = True
|
| 42 |
+
symbol_dim: int | None = None
|
| 43 |
+
n_symbols: int | None = None
|
| 44 |
+
symbolic_attn_n_heads: int | None = None
|
| 45 |
+
symbol_retrieval: str = "symbolic"
|
| 46 |
+
symbolic_use_bias: bool = False
|
| 47 |
+
shared_symbol_retriever: bool = True
|
| 48 |
+
share_attn_params: bool = False
|
| 49 |
+
positional_symbols_sinusoidal: bool = False
|
| 50 |
+
relative_symbols_rope: bool = False
|
| 51 |
+
relsymbolic_rel_n_heads: int = 4
|
| 52 |
+
relsymbolic_symbolic_attn_n_heads: int = 4
|
| 53 |
+
relsymbolic_neighborhood_size: int = 2
|
| 54 |
+
relsymbolic_include_self: bool = False
|
| 55 |
+
relsymbolic_normalize_rels: bool = True
|
| 56 |
+
relsymbolic_trainable_symbols: bool = True
|
| 57 |
+
relsymbolic_dropout: float = 0.0
|
| 58 |
+
relsymbolic_rel_scale: float | None = None
|
| 59 |
+
relsymbolic_symbolic_attn_scale: float | None = None
|
| 60 |
+
relsymbolic_use_bias: bool = False
|
| 61 |
+
ra_type: str = "ra"
|
| 62 |
+
ra_n_relations: int | None = None
|
| 63 |
+
ra_rel_activation: str = "identity"
|
| 64 |
+
ra_symmetric_rels: bool = False
|
| 65 |
+
pad_token_id: int = 0
|
| 66 |
+
bos_token_id: int = 1
|
| 67 |
+
eos_token_id: int = 2
|
| 68 |
+
mlm_head_enabled: bool = False
|
| 69 |
+
|
| 70 |
+
def __post_init__(self) -> None:
|
| 71 |
+
if self.vocab_size <= 0:
|
| 72 |
+
raise ValueError(f"vocab_size must be positive, got {self.vocab_size}")
|
| 73 |
+
if self.max_seq_len <= 0:
|
| 74 |
+
raise ValueError(f"max_seq_len must be positive, got {self.max_seq_len}")
|
| 75 |
+
if self.pe_type not in SUPPORTED_PE_TYPES:
|
| 76 |
+
raise ValueError(f"Unsupported pe_type: {self.pe_type}")
|
| 77 |
+
if self.hidden_dim <= 0:
|
| 78 |
+
raise ValueError(f"hidden_dim must be positive, got {self.hidden_dim}")
|
| 79 |
+
if self.n_heads_sa <= 0:
|
| 80 |
+
raise ValueError(f"n_heads_sa must be positive for DAT, got {self.n_heads_sa}")
|
| 81 |
+
if self.n_heads_ra <= 0:
|
| 82 |
+
raise ValueError(f"n_heads_ra must be positive for DAT, got {self.n_heads_ra}")
|
| 83 |
+
total_heads = self.total_n_heads
|
| 84 |
+
if self.hidden_dim % total_heads != 0:
|
| 85 |
+
raise ValueError(
|
| 86 |
+
f"hidden_dim ({self.hidden_dim}) must be divisible by total DAT heads "
|
| 87 |
+
f"({total_heads} = {self.n_heads_sa} SA + {self.n_heads_ra} RA)"
|
| 88 |
+
)
|
| 89 |
+
if self.n_layers <= 0:
|
| 90 |
+
raise ValueError(f"n_layers must be positive, got {self.n_layers}")
|
| 91 |
+
if not 0.0 <= self.dropout < 1.0:
|
| 92 |
+
raise ValueError(f"dropout must be in [0.0, 1.0), got {self.dropout}")
|
| 93 |
+
if self.dff_factor <= 0:
|
| 94 |
+
raise ValueError(f"dff_factor must be positive, got {self.dff_factor}")
|
| 95 |
+
if self.ffn_hidden_dim_mode not in SUPPORTED_FFN_HIDDEN_DIM_MODES:
|
| 96 |
+
raise ValueError(f"Unsupported ffn_hidden_dim_mode for DAT: {self.ffn_hidden_dim_mode}")
|
| 97 |
+
if self.ffn_activation not in SUPPORTED_FFN_ACTIVATIONS:
|
| 98 |
+
raise ValueError(f"Unsupported ffn_activation for DAT: {self.ffn_activation}")
|
| 99 |
+
if self.ffn_hidden_dim_mode == "swiglu_parameter_matched" and self.ffn_activation != "swiglu":
|
| 100 |
+
raise ValueError(
|
| 101 |
+
"ffn_hidden_dim_mode='swiglu_parameter_matched' requires "
|
| 102 |
+
f"ffn_activation='swiglu', got {self.ffn_activation}"
|
| 103 |
+
)
|
| 104 |
+
if self.rope_theta <= 0.0:
|
| 105 |
+
raise ValueError(f"rope_theta must be positive, got {self.rope_theta}")
|
| 106 |
+
if self.max_rel_pos is not None and self.max_rel_pos <= 0:
|
| 107 |
+
raise ValueError(f"max_rel_pos must be positive when provided, got {self.max_rel_pos}")
|
| 108 |
+
if self.sequence_boundary_policy not in SUPPORTED_SEQUENCE_BOUNDARY_POLICIES:
|
| 109 |
+
raise ValueError(
|
| 110 |
+
"Unsupported sequence_boundary_policy for DAT: "
|
| 111 |
+
f"{self.sequence_boundary_policy}"
|
| 112 |
+
)
|
| 113 |
+
if self.sequence_boundary_policy == "segment_document":
|
| 114 |
+
if self.segment_boundary_token_id is None:
|
| 115 |
+
raise ValueError(
|
| 116 |
+
"segment_boundary_token_id is required when "
|
| 117 |
+
"sequence_boundary_policy='segment_document'"
|
| 118 |
+
)
|
| 119 |
+
if self.init_range <= 0.0:
|
| 120 |
+
raise ValueError(f"init_range must be positive, got {self.init_range}")
|
| 121 |
+
if self.init_scheme not in SUPPORTED_INIT_SCHEMES:
|
| 122 |
+
raise ValueError(f"Unsupported init_scheme for DAT: {self.init_scheme}")
|
| 123 |
+
if self.norm_type not in {"layernorm", "rmsnorm"}:
|
| 124 |
+
raise ValueError(
|
| 125 |
+
f"norm_type must be 'layernorm' or 'rmsnorm', got {self.norm_type}"
|
| 126 |
+
)
|
| 127 |
+
if self.pe_type == "rope" and self.head_dim % 2 != 0:
|
| 128 |
+
raise ValueError(
|
| 129 |
+
"RoPE requires even DAT head_dim, got "
|
| 130 |
+
f"{self.head_dim} from hidden_dim={self.hidden_dim}, total_heads={total_heads}"
|
| 131 |
+
)
|
| 132 |
+
if self.pe_type == "sinusoidal" and self.hidden_dim % 2 != 0:
|
| 133 |
+
raise ValueError(f"Sinusoidal encoding requires even hidden_dim, got {self.hidden_dim}")
|
| 134 |
+
if self.share_attn_params and self.n_heads_sa != self.n_heads_ra:
|
| 135 |
+
raise ValueError(
|
| 136 |
+
"share_attn_params=True requires n_heads_sa == n_heads_ra, "
|
| 137 |
+
f"got {self.n_heads_sa} and {self.n_heads_ra}"
|
| 138 |
+
)
|
| 139 |
+
if self.symbol_dim is not None and self.symbol_dim <= 0:
|
| 140 |
+
raise ValueError(f"symbol_dim must be positive when provided, got {self.symbol_dim}")
|
| 141 |
+
if self.symbolic_attn_n_heads is not None and self.symbolic_attn_n_heads <= 0:
|
| 142 |
+
raise ValueError(
|
| 143 |
+
"symbolic_attn_n_heads must be positive when provided, "
|
| 144 |
+
f"got {self.symbolic_attn_n_heads}"
|
| 145 |
+
)
|
| 146 |
+
if self.symbol_retrieval not in SUPPORTED_SYMBOL_RETRIEVAL:
|
| 147 |
+
raise ValueError(f"Unsupported symbol_retrieval for DAT: {self.symbol_retrieval}")
|
| 148 |
+
if self.symbol_retrieval == "symbolic":
|
| 149 |
+
symbolic_heads = self.resolved_symbolic_attn_n_heads
|
| 150 |
+
if self.hidden_dim % symbolic_heads != 0:
|
| 151 |
+
raise ValueError(
|
| 152 |
+
f"hidden_dim ({self.hidden_dim}) must be divisible by symbolic_attn_n_heads "
|
| 153 |
+
f"({symbolic_heads}) for symbolic retrieval"
|
| 154 |
+
)
|
| 155 |
+
if self.resolved_symbol_dim % symbolic_heads != 0:
|
| 156 |
+
raise ValueError(
|
| 157 |
+
f"symbol_dim ({self.resolved_symbol_dim}) must be divisible by symbolic_attn_n_heads "
|
| 158 |
+
f"({symbolic_heads}) for symbolic retrieval"
|
| 159 |
+
)
|
| 160 |
+
if (
|
| 161 |
+
self.symbol_retrieval == "positional"
|
| 162 |
+
and self.positional_symbols_sinusoidal
|
| 163 |
+
and self.resolved_symbol_dim % 2 != 0
|
| 164 |
+
):
|
| 165 |
+
raise ValueError(
|
| 166 |
+
"Sinusoidal positional symbols require even symbol_dim, "
|
| 167 |
+
f"got {self.resolved_symbol_dim}"
|
| 168 |
+
)
|
| 169 |
+
if self.relative_symbols_rope:
|
| 170 |
+
if self.symbol_retrieval != "relative":
|
| 171 |
+
raise ValueError(
|
| 172 |
+
"relative_symbols_rope=True requires symbol_retrieval='relative', "
|
| 173 |
+
f"got {self.symbol_retrieval!r}"
|
| 174 |
+
)
|
| 175 |
+
if self.resolved_symbol_dim % 2 != 0:
|
| 176 |
+
raise ValueError(
|
| 177 |
+
"RoPE relative symbols require even symbol_dim, "
|
| 178 |
+
f"got {self.resolved_symbol_dim}"
|
| 179 |
+
)
|
| 180 |
+
if self.positional_symbols_sinusoidal and self.symbol_retrieval != "positional":
|
| 181 |
+
raise ValueError(
|
| 182 |
+
"positional_symbols_sinusoidal=True requires symbol_retrieval='positional', "
|
| 183 |
+
f"got {self.symbol_retrieval!r}"
|
| 184 |
+
)
|
| 185 |
+
if self.resolved_n_symbols <= 0:
|
| 186 |
+
raise ValueError(f"resolved_n_symbols must be positive, got {self.resolved_n_symbols}")
|
| 187 |
+
if self.symbol_retrieval == "relsymbolic":
|
| 188 |
+
if self.relsymbolic_rel_n_heads <= 0:
|
| 189 |
+
raise ValueError(
|
| 190 |
+
f"relsymbolic_rel_n_heads must be positive, got {self.relsymbolic_rel_n_heads}"
|
| 191 |
+
)
|
| 192 |
+
if self.hidden_dim % self.relsymbolic_rel_n_heads != 0:
|
| 193 |
+
raise ValueError(
|
| 194 |
+
f"hidden_dim ({self.hidden_dim}) must be divisible by "
|
| 195 |
+
f"relsymbolic_rel_n_heads ({self.relsymbolic_rel_n_heads})"
|
| 196 |
+
)
|
| 197 |
+
if self.relsymbolic_symbolic_attn_n_heads <= 0:
|
| 198 |
+
raise ValueError(
|
| 199 |
+
"relsymbolic_symbolic_attn_n_heads must be positive, "
|
| 200 |
+
f"got {self.relsymbolic_symbolic_attn_n_heads}"
|
| 201 |
+
)
|
| 202 |
+
if self.hidden_dim % self.relsymbolic_symbolic_attn_n_heads != 0:
|
| 203 |
+
raise ValueError(
|
| 204 |
+
f"hidden_dim ({self.hidden_dim}) must be divisible by "
|
| 205 |
+
f"relsymbolic_symbolic_attn_n_heads ({self.relsymbolic_symbolic_attn_n_heads})"
|
| 206 |
+
)
|
| 207 |
+
if self.resolved_symbol_dim % self.relsymbolic_symbolic_attn_n_heads != 0:
|
| 208 |
+
raise ValueError(
|
| 209 |
+
f"symbol_dim ({self.resolved_symbol_dim}) must be divisible by "
|
| 210 |
+
f"relsymbolic_symbolic_attn_n_heads ({self.relsymbolic_symbolic_attn_n_heads})"
|
| 211 |
+
)
|
| 212 |
+
if self.relsymbolic_neighborhood_size <= 0:
|
| 213 |
+
raise ValueError(
|
| 214 |
+
"relsymbolic_neighborhood_size must be positive, "
|
| 215 |
+
f"got {self.relsymbolic_neighborhood_size}"
|
| 216 |
+
)
|
| 217 |
+
if not 0.0 <= self.relsymbolic_dropout < 1.0:
|
| 218 |
+
raise ValueError(
|
| 219 |
+
f"relsymbolic_dropout must be in [0.0, 1.0), got {self.relsymbolic_dropout}"
|
| 220 |
+
)
|
| 221 |
+
if self.relsymbolic_rel_scale is not None and self.relsymbolic_rel_scale <= 0.0:
|
| 222 |
+
raise ValueError(
|
| 223 |
+
"relsymbolic_rel_scale must be positive when provided, "
|
| 224 |
+
f"got {self.relsymbolic_rel_scale}"
|
| 225 |
+
)
|
| 226 |
+
if (
|
| 227 |
+
self.relsymbolic_symbolic_attn_scale is not None
|
| 228 |
+
and self.relsymbolic_symbolic_attn_scale <= 0.0
|
| 229 |
+
):
|
| 230 |
+
raise ValueError(
|
| 231 |
+
"relsymbolic_symbolic_attn_scale must be positive when provided, "
|
| 232 |
+
f"got {self.relsymbolic_symbolic_attn_scale}"
|
| 233 |
+
)
|
| 234 |
+
if self.ra_type not in SUPPORTED_RA_TYPES:
|
| 235 |
+
raise ValueError(f"Unsupported ra_type for DAT: {self.ra_type}")
|
| 236 |
+
if self.ra_n_relations is not None and self.ra_n_relations <= 0:
|
| 237 |
+
raise ValueError(
|
| 238 |
+
f"ra_n_relations must be positive when provided, got {self.ra_n_relations}"
|
| 239 |
+
)
|
| 240 |
+
if self.ra_type != "ra" and self.ra_n_relations is not None:
|
| 241 |
+
raise ValueError(f"ra_n_relations applies only to ra_type='ra', got {self.ra_type}")
|
| 242 |
+
if self.ra_type != "ra" and self.ra_symmetric_rels:
|
| 243 |
+
raise ValueError(f"ra_symmetric_rels applies only to ra_type='ra', got {self.ra_type}")
|
| 244 |
+
n_relations = self.resolved_ra_n_relations
|
| 245 |
+
if self.ra_type == "ra" and (self.head_dim * self.n_heads_ra) % n_relations != 0:
|
| 246 |
+
raise ValueError(
|
| 247 |
+
f"head_dim * n_heads_ra ({self.head_dim * self.n_heads_ra}) must be "
|
| 248 |
+
f"divisible by ra_n_relations ({n_relations})"
|
| 249 |
+
)
|
| 250 |
+
if self.ra_rel_activation not in SUPPORTED_RA_ACTIVATIONS:
|
| 251 |
+
raise ValueError(
|
| 252 |
+
f"Unsupported ra_rel_activation for DAT: {self.ra_rel_activation}"
|
| 253 |
+
)
|
| 254 |
+
|
| 255 |
+
@property
|
| 256 |
+
def total_n_heads(self) -> int:
|
| 257 |
+
return self.n_heads_sa + self.n_heads_ra
|
| 258 |
+
|
| 259 |
+
@property
|
| 260 |
+
def head_dim(self) -> int:
|
| 261 |
+
return self.hidden_dim // self.total_n_heads
|
| 262 |
+
|
| 263 |
+
@property
|
| 264 |
+
def resolved_symbol_dim(self) -> int:
|
| 265 |
+
return self.hidden_dim if self.symbol_dim is None else self.symbol_dim
|
| 266 |
+
|
| 267 |
+
@property
|
| 268 |
+
def resolved_symbolic_attn_n_heads(self) -> int:
|
| 269 |
+
return self.total_n_heads if self.symbolic_attn_n_heads is None else self.symbolic_attn_n_heads
|
| 270 |
+
|
| 271 |
+
@property
|
| 272 |
+
def resolved_ffn_hidden_dim(self) -> int:
|
| 273 |
+
if self.ffn_hidden_dim_mode == "swiglu_parameter_matched":
|
| 274 |
+
return int(8 / 3 * self.hidden_dim)
|
| 275 |
+
return self.hidden_dim * self.dff_factor
|
| 276 |
+
|
| 277 |
+
@property
|
| 278 |
+
def resolved_n_symbols(self) -> int:
|
| 279 |
+
return self.max_seq_len if self.n_symbols is None else self.n_symbols
|
| 280 |
+
|
| 281 |
+
@property
|
| 282 |
+
def resolved_ra_n_relations(self) -> int:
|
| 283 |
+
return self.n_heads_ra if self.ra_n_relations is None else self.ra_n_relations
|
dat_core.py
ADDED
|
@@ -0,0 +1,328 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Dual-attention primitives for decoder language models."""
|
| 2 |
+
|
| 3 |
+
import math
|
| 4 |
+
from typing import Optional
|
| 5 |
+
|
| 6 |
+
import torch
|
| 7 |
+
import torch.nn as nn
|
| 8 |
+
|
| 9 |
+
from .transformer_components import PositionalInfo
|
| 10 |
+
from .transformer_core import MultiHeadAttentionBase
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
def _activate_scores(scores: torch.Tensor, activation: str) -> torch.Tensor:
|
| 14 |
+
if activation == "softmax":
|
| 15 |
+
return torch.softmax(scores, dim=-1)
|
| 16 |
+
if activation == "identity":
|
| 17 |
+
return scores
|
| 18 |
+
if activation == "relu":
|
| 19 |
+
return torch.relu(scores)
|
| 20 |
+
if activation == "tanh":
|
| 21 |
+
return torch.tanh(scores)
|
| 22 |
+
if activation == "sigmoid":
|
| 23 |
+
return torch.sigmoid(scores)
|
| 24 |
+
if activation == "gelu":
|
| 25 |
+
return torch.nn.functional.gelu(scores)
|
| 26 |
+
raise ValueError(f"Unsupported attention activation: {activation}")
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
class RelationalAttentionBase(MultiHeadAttentionBase):
|
| 30 |
+
def __init__(
|
| 31 |
+
self,
|
| 32 |
+
hidden_dim: int,
|
| 33 |
+
symbol_dim: int,
|
| 34 |
+
n_heads: int,
|
| 35 |
+
total_n_heads: int,
|
| 36 |
+
dropout: float = 0.0,
|
| 37 |
+
n_relations: Optional[int] = None,
|
| 38 |
+
rel_activation: str = "identity",
|
| 39 |
+
symmetric_rels: bool = False,
|
| 40 |
+
use_relative_positional_symbols: bool = False,
|
| 41 |
+
use_bias_qkv: bool = False,
|
| 42 |
+
use_bias_out: bool = True,
|
| 43 |
+
):
|
| 44 |
+
head_dim = hidden_dim // total_n_heads
|
| 45 |
+
output_dim = n_heads * head_dim
|
| 46 |
+
super().__init__(
|
| 47 |
+
query_dim=hidden_dim,
|
| 48 |
+
output_dim=output_dim,
|
| 49 |
+
key_dim=hidden_dim,
|
| 50 |
+
value_dim=symbol_dim,
|
| 51 |
+
n_heads=n_heads,
|
| 52 |
+
hidden_dim=hidden_dim,
|
| 53 |
+
dropout=dropout,
|
| 54 |
+
total_n_heads=total_n_heads,
|
| 55 |
+
activation="softmax",
|
| 56 |
+
use_bias_qkv=use_bias_qkv,
|
| 57 |
+
use_bias_out=use_bias_out,
|
| 58 |
+
)
|
| 59 |
+
self.symbol_dim = symbol_dim
|
| 60 |
+
self.rel_activation = rel_activation
|
| 61 |
+
self.symmetric_rels = symmetric_rels
|
| 62 |
+
self.use_relative_positional_symbols = use_relative_positional_symbols
|
| 63 |
+
self.n_relations = n_heads if n_relations is None else n_relations
|
| 64 |
+
total_rel_dim = self.head_dim * n_heads
|
| 65 |
+
if total_rel_dim % self.n_relations != 0:
|
| 66 |
+
raise ValueError(
|
| 67 |
+
f"head_dim * n_heads ({total_rel_dim}) must be divisible by n_relations "
|
| 68 |
+
f"({self.n_relations})"
|
| 69 |
+
)
|
| 70 |
+
self.rel_proj_dim = total_rel_dim // self.n_relations
|
| 71 |
+
self.rel_scale = 1.0 / math.sqrt(self.rel_proj_dim)
|
| 72 |
+
rel_total_dim = self.n_relations * self.rel_proj_dim
|
| 73 |
+
self.wq_rel = nn.Linear(hidden_dim, rel_total_dim, bias=False)
|
| 74 |
+
self.wk_rel = self.wq_rel if symmetric_rels else nn.Linear(
|
| 75 |
+
hidden_dim,
|
| 76 |
+
rel_total_dim,
|
| 77 |
+
bias=False,
|
| 78 |
+
)
|
| 79 |
+
nn.init.xavier_uniform_(self.wq_rel.weight)
|
| 80 |
+
if self.wk_rel is not self.wq_rel:
|
| 81 |
+
nn.init.xavier_uniform_(self.wk_rel.weight)
|
| 82 |
+
|
| 83 |
+
def _compute_base_attention(
|
| 84 |
+
self,
|
| 85 |
+
x: torch.Tensor,
|
| 86 |
+
mask: Optional[torch.Tensor],
|
| 87 |
+
pos_info: Optional[PositionalInfo],
|
| 88 |
+
) -> torch.Tensor:
|
| 89 |
+
batch_size, seq_len, _ = x.shape
|
| 90 |
+
q = self._reshape_for_multihead(self.q_proj(x), batch_size, seq_len)
|
| 91 |
+
k = self._reshape_for_multihead(self.k_proj(x), batch_size, seq_len)
|
| 92 |
+
scores = self._compute_attn_scores(q, k, pos_info)
|
| 93 |
+
return self._apply_activation_and_mask(scores, mask)
|
| 94 |
+
|
| 95 |
+
def compute_relational_scores(
|
| 96 |
+
self,
|
| 97 |
+
x: torch.Tensor,
|
| 98 |
+
mask: Optional[torch.Tensor],
|
| 99 |
+
return_as: str,
|
| 100 |
+
) -> torch.Tensor:
|
| 101 |
+
batch_size, seq_len, _ = x.shape
|
| 102 |
+
q_rel = self.wq_rel(x).view(batch_size, seq_len, self.n_relations, self.rel_proj_dim)
|
| 103 |
+
k_rel = self.wk_rel(x).view(batch_size, seq_len, self.n_relations, self.rel_proj_dim)
|
| 104 |
+
q_rel = q_rel.transpose(1, 2)
|
| 105 |
+
k_rel = k_rel.transpose(1, 2)
|
| 106 |
+
relations = torch.matmul(q_rel, k_rel.transpose(-2, -1)) * self.rel_scale
|
| 107 |
+
processed_mask = self._process_mask(mask)
|
| 108 |
+
if self.rel_activation == "softmax" and processed_mask is not None:
|
| 109 |
+
relations.masked_fill_(~processed_mask, torch.finfo(relations.dtype).min)
|
| 110 |
+
relations = _activate_scores(relations, self.rel_activation)
|
| 111 |
+
if processed_mask is not None:
|
| 112 |
+
relations = relations.masked_fill(~processed_mask, 0.0)
|
| 113 |
+
if return_as == "vectors":
|
| 114 |
+
return relations.permute(0, 2, 3, 1)
|
| 115 |
+
if return_as == "scores":
|
| 116 |
+
return relations
|
| 117 |
+
raise ValueError(f"return_as must be 'vectors' or 'scores', got {return_as}")
|
| 118 |
+
|
| 119 |
+
def combine_attention_and_relations(
|
| 120 |
+
self,
|
| 121 |
+
attn_weights: torch.Tensor,
|
| 122 |
+
relation_scores: torch.Tensor,
|
| 123 |
+
) -> torch.Tensor:
|
| 124 |
+
return attn_weights * relation_scores
|
| 125 |
+
|
| 126 |
+
def _process_symbols_with_attention(
|
| 127 |
+
self,
|
| 128 |
+
symbols: torch.Tensor,
|
| 129 |
+
attn_weights: torch.Tensor,
|
| 130 |
+
) -> torch.Tensor:
|
| 131 |
+
batch_size, _, seq_len_q, seq_len_k = attn_weights.shape
|
| 132 |
+
values = self.v_proj(symbols)
|
| 133 |
+
if self.use_relative_positional_symbols:
|
| 134 |
+
values = values.view(seq_len_q, seq_len_k, self.n_heads, self.head_dim)
|
| 135 |
+
return torch.einsum("bhij,ijhd->bihd", attn_weights, values)
|
| 136 |
+
|
| 137 |
+
values = values.view(batch_size, seq_len_k, self.n_heads, self.head_dim)
|
| 138 |
+
values = values.transpose(1, 2)
|
| 139 |
+
output = torch.matmul(attn_weights, values)
|
| 140 |
+
return output.transpose(1, 2)
|
| 141 |
+
|
| 142 |
+
def _apply_output_projection(self, output: torch.Tensor) -> torch.Tensor:
|
| 143 |
+
batch_size, seq_len = output.shape[:2]
|
| 144 |
+
output = output.contiguous().view(batch_size, seq_len, self.output_dim)
|
| 145 |
+
output = self.o_proj(output)
|
| 146 |
+
output = self.dropout(output)
|
| 147 |
+
return output
|
| 148 |
+
|
| 149 |
+
def _validate_symbols(self, x: torch.Tensor, symbols: torch.Tensor) -> None:
|
| 150 |
+
if self.use_relative_positional_symbols:
|
| 151 |
+
seq_len = x.shape[1]
|
| 152 |
+
expected_shape = (seq_len, seq_len, self.symbol_dim)
|
| 153 |
+
if tuple(symbols.shape) != expected_shape:
|
| 154 |
+
raise ValueError(
|
| 155 |
+
f"Relative symbols must have shape {expected_shape}, got {tuple(symbols.shape)}"
|
| 156 |
+
)
|
| 157 |
+
|
| 158 |
+
|
| 159 |
+
class RelationalAttention(RelationalAttentionBase):
|
| 160 |
+
def __init__(
|
| 161 |
+
self,
|
| 162 |
+
hidden_dim: int,
|
| 163 |
+
symbol_dim: int,
|
| 164 |
+
n_heads: int,
|
| 165 |
+
total_n_heads: int,
|
| 166 |
+
n_relations: int,
|
| 167 |
+
dropout: float = 0.0,
|
| 168 |
+
rel_activation: str = "identity",
|
| 169 |
+
symmetric_rels: bool = False,
|
| 170 |
+
use_relative_positional_symbols: bool = False,
|
| 171 |
+
use_bias_qkv: bool = False,
|
| 172 |
+
use_bias_out: bool = True,
|
| 173 |
+
):
|
| 174 |
+
super().__init__(
|
| 175 |
+
hidden_dim=hidden_dim,
|
| 176 |
+
symbol_dim=symbol_dim,
|
| 177 |
+
n_heads=n_heads,
|
| 178 |
+
total_n_heads=total_n_heads,
|
| 179 |
+
dropout=dropout,
|
| 180 |
+
n_relations=n_relations,
|
| 181 |
+
rel_activation=rel_activation,
|
| 182 |
+
symmetric_rels=symmetric_rels,
|
| 183 |
+
use_relative_positional_symbols=use_relative_positional_symbols,
|
| 184 |
+
use_bias_qkv=use_bias_qkv,
|
| 185 |
+
use_bias_out=use_bias_out,
|
| 186 |
+
)
|
| 187 |
+
self.wr_proj = nn.Parameter(torch.empty(n_heads, self.head_dim, n_relations))
|
| 188 |
+
nn.init.xavier_uniform_(self.wr_proj)
|
| 189 |
+
|
| 190 |
+
def forward(
|
| 191 |
+
self,
|
| 192 |
+
x: torch.Tensor,
|
| 193 |
+
symbols: torch.Tensor,
|
| 194 |
+
mask: Optional[torch.Tensor],
|
| 195 |
+
pos_info: Optional[PositionalInfo],
|
| 196 |
+
) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
|
| 197 |
+
self._validate_symbols(x, symbols)
|
| 198 |
+
attn_weights = self._compute_base_attention(x, mask, pos_info)
|
| 199 |
+
attn_weights = self.attn_dropout(attn_weights)
|
| 200 |
+
relation_vectors = self.compute_relational_scores(x, mask, return_as="vectors")
|
| 201 |
+
attended_symbols = self._process_symbols_with_attention(symbols, attn_weights)
|
| 202 |
+
projected_relations = torch.einsum(
|
| 203 |
+
"bhij,bijr,hdr->bihd",
|
| 204 |
+
attn_weights,
|
| 205 |
+
relation_vectors,
|
| 206 |
+
self.wr_proj,
|
| 207 |
+
)
|
| 208 |
+
output = self._apply_output_projection(attended_symbols + projected_relations)
|
| 209 |
+
self.last_attn_weights = attn_weights.detach()
|
| 210 |
+
return output, {"attention": attn_weights, "relations": relation_vectors}
|
| 211 |
+
|
| 212 |
+
|
| 213 |
+
class RelationalCrossAttention(MultiHeadAttentionBase):
|
| 214 |
+
def __init__(
|
| 215 |
+
self,
|
| 216 |
+
hidden_dim: int,
|
| 217 |
+
symbol_dim: int,
|
| 218 |
+
n_heads: int,
|
| 219 |
+
total_n_heads: int,
|
| 220 |
+
dropout: float = 0.0,
|
| 221 |
+
activation: str = "identity",
|
| 222 |
+
use_relative_positional_symbols: bool = False,
|
| 223 |
+
use_bias_qkv: bool = False,
|
| 224 |
+
use_bias_out: bool = True,
|
| 225 |
+
):
|
| 226 |
+
head_dim = hidden_dim // total_n_heads
|
| 227 |
+
super().__init__(
|
| 228 |
+
query_dim=hidden_dim,
|
| 229 |
+
output_dim=n_heads * head_dim,
|
| 230 |
+
key_dim=hidden_dim,
|
| 231 |
+
value_dim=symbol_dim,
|
| 232 |
+
n_heads=n_heads,
|
| 233 |
+
hidden_dim=hidden_dim,
|
| 234 |
+
dropout=dropout,
|
| 235 |
+
total_n_heads=total_n_heads,
|
| 236 |
+
activation=activation,
|
| 237 |
+
use_bias_qkv=use_bias_qkv,
|
| 238 |
+
use_bias_out=use_bias_out,
|
| 239 |
+
)
|
| 240 |
+
self.symbol_dim = symbol_dim
|
| 241 |
+
self.use_relative_positional_symbols = use_relative_positional_symbols
|
| 242 |
+
|
| 243 |
+
def forward(
|
| 244 |
+
self,
|
| 245 |
+
x: torch.Tensor,
|
| 246 |
+
symbols: torch.Tensor,
|
| 247 |
+
mask: Optional[torch.Tensor],
|
| 248 |
+
pos_info: Optional[PositionalInfo],
|
| 249 |
+
) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
|
| 250 |
+
if self.use_relative_positional_symbols:
|
| 251 |
+
batch_size, seq_len, _ = x.shape
|
| 252 |
+
expected_shape = (seq_len, seq_len, self.symbol_dim)
|
| 253 |
+
if tuple(symbols.shape) != expected_shape:
|
| 254 |
+
raise ValueError(
|
| 255 |
+
f"Relative symbols must have shape {expected_shape}, got {tuple(symbols.shape)}"
|
| 256 |
+
)
|
| 257 |
+
|
| 258 |
+
q = self._reshape_for_multihead(self.q_proj(x), batch_size, seq_len)
|
| 259 |
+
k = self._reshape_for_multihead(self.k_proj(x), batch_size, seq_len)
|
| 260 |
+
values = self.v_proj(symbols).view(seq_len, seq_len, self.n_heads, self.head_dim)
|
| 261 |
+
scores = self._compute_attn_scores(q, k, pos_info)
|
| 262 |
+
weights = self._apply_activation_and_mask(scores, mask)
|
| 263 |
+
weights = self.attn_dropout(weights)
|
| 264 |
+
output = torch.einsum("bhij,ijhd->bihd", weights, values)
|
| 265 |
+
output = output.contiguous().view(batch_size, seq_len, self.output_dim)
|
| 266 |
+
output = self.o_proj(output)
|
| 267 |
+
output = self.dropout(output)
|
| 268 |
+
self.last_attn_weights = weights.detach()
|
| 269 |
+
return output, {"attention": weights}
|
| 270 |
+
|
| 271 |
+
output, weights = super().forward(
|
| 272 |
+
query=x,
|
| 273 |
+
key=x,
|
| 274 |
+
value=symbols,
|
| 275 |
+
mask=mask,
|
| 276 |
+
pos_info=pos_info,
|
| 277 |
+
)
|
| 278 |
+
return output, {"attention": weights}
|
| 279 |
+
|
| 280 |
+
|
| 281 |
+
class DisentangledRelationalCrossAttention(RelationalAttentionBase):
|
| 282 |
+
def __init__(
|
| 283 |
+
self,
|
| 284 |
+
hidden_dim: int,
|
| 285 |
+
symbol_dim: int,
|
| 286 |
+
n_heads: int,
|
| 287 |
+
total_n_heads: int,
|
| 288 |
+
dropout: float = 0.0,
|
| 289 |
+
rel_activation: str = "identity",
|
| 290 |
+
use_relative_positional_symbols: bool = False,
|
| 291 |
+
use_bias_qkv: bool = False,
|
| 292 |
+
use_bias_out: bool = True,
|
| 293 |
+
):
|
| 294 |
+
super().__init__(
|
| 295 |
+
hidden_dim=hidden_dim,
|
| 296 |
+
symbol_dim=symbol_dim,
|
| 297 |
+
n_heads=n_heads,
|
| 298 |
+
total_n_heads=total_n_heads,
|
| 299 |
+
dropout=dropout,
|
| 300 |
+
n_relations=None,
|
| 301 |
+
rel_activation=rel_activation,
|
| 302 |
+
symmetric_rels=False,
|
| 303 |
+
use_relative_positional_symbols=use_relative_positional_symbols,
|
| 304 |
+
use_bias_qkv=use_bias_qkv,
|
| 305 |
+
use_bias_out=use_bias_out,
|
| 306 |
+
)
|
| 307 |
+
|
| 308 |
+
def forward(
|
| 309 |
+
self,
|
| 310 |
+
x: torch.Tensor,
|
| 311 |
+
symbols: torch.Tensor,
|
| 312 |
+
mask: Optional[torch.Tensor],
|
| 313 |
+
pos_info: Optional[PositionalInfo],
|
| 314 |
+
) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
|
| 315 |
+
self._validate_symbols(x, symbols)
|
| 316 |
+
attn_weights = self._compute_base_attention(x, mask, pos_info)
|
| 317 |
+
relation_scores = self.compute_relational_scores(x, mask, return_as="scores")
|
| 318 |
+
combined_weights = self.attn_dropout(
|
| 319 |
+
self.combine_attention_and_relations(attn_weights, relation_scores)
|
| 320 |
+
)
|
| 321 |
+
output = self._process_symbols_with_attention(symbols, combined_weights)
|
| 322 |
+
output = self._apply_output_projection(output)
|
| 323 |
+
self.last_attn_weights = attn_weights.detach()
|
| 324 |
+
return output, {
|
| 325 |
+
"attention": attn_weights,
|
| 326 |
+
"relations": relation_scores,
|
| 327 |
+
"combined": combined_weights,
|
| 328 |
+
}
|
dat_lm.py
ADDED
|
@@ -0,0 +1,475 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Decoder-only LM built from dual-attention blocks."""
|
| 2 |
+
|
| 3 |
+
import math
|
| 4 |
+
from typing import Optional
|
| 5 |
+
|
| 6 |
+
import torch
|
| 7 |
+
import torch.nn as nn
|
| 8 |
+
import torch.nn.functional as F
|
| 9 |
+
from torch.nn.utils.rnn import pad_sequence
|
| 10 |
+
|
| 11 |
+
from .dat_config import DatLMConfig
|
| 12 |
+
from .dat_core import (
|
| 13 |
+
DisentangledRelationalCrossAttention,
|
| 14 |
+
RelationalAttention,
|
| 15 |
+
RelationalCrossAttention,
|
| 16 |
+
)
|
| 17 |
+
from .transformer_core import SelfAttention
|
| 18 |
+
from .dat_symbols import (
|
| 19 |
+
PositionalSymbolRetriever,
|
| 20 |
+
RelationalSymbolicAttentionRetriever,
|
| 21 |
+
RelativePositionalSymbolRetriever,
|
| 22 |
+
SymbolicAttentionRetriever,
|
| 23 |
+
)
|
| 24 |
+
from .masks import build_decoder_attention_mask
|
| 25 |
+
from .mlm import MaskClassifier
|
| 26 |
+
from .transformer_components import FeedForward, PositionalEncoding, PositionalInfo, RMSNorm
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
class DatDecoderBlock(nn.Module):
|
| 30 |
+
def __init__(self, config: DatLMConfig):
|
| 31 |
+
super().__init__()
|
| 32 |
+
self.norm_first = config.norm_first
|
| 33 |
+
# input_dim is the block state width; hidden_dim is the total DAT width
|
| 34 |
+
# used with total_n_heads to derive the shared SA/RA head_dim.
|
| 35 |
+
self.sensory_attention = SelfAttention(
|
| 36 |
+
input_dim=config.hidden_dim,
|
| 37 |
+
n_heads=config.n_heads_sa,
|
| 38 |
+
hidden_dim=config.hidden_dim,
|
| 39 |
+
total_n_heads=config.total_n_heads,
|
| 40 |
+
dropout=config.dropout,
|
| 41 |
+
supports_relative=config.pe_type == "relative",
|
| 42 |
+
use_bias_qkv=config.use_bias_qkv,
|
| 43 |
+
use_bias_out=config.use_bias_out,
|
| 44 |
+
)
|
| 45 |
+
self.relational_attention = _build_relational_attention(config)
|
| 46 |
+
if config.share_attn_params:
|
| 47 |
+
self.sensory_attention.q_proj = self.relational_attention.q_proj
|
| 48 |
+
self.sensory_attention.k_proj = self.relational_attention.k_proj
|
| 49 |
+
self.feed_forward = FeedForward(
|
| 50 |
+
input_dim=config.hidden_dim,
|
| 51 |
+
hidden_dim=config.resolved_ffn_hidden_dim,
|
| 52 |
+
dropout=config.dropout,
|
| 53 |
+
activation=config.ffn_activation,
|
| 54 |
+
use_bias=config.use_bias_ffn,
|
| 55 |
+
)
|
| 56 |
+
self.norm1 = _create_norm(config.norm_type, config.hidden_dim)
|
| 57 |
+
self.norm2 = _create_norm(config.norm_type, config.hidden_dim)
|
| 58 |
+
self.dropout = nn.Dropout(config.dropout)
|
| 59 |
+
|
| 60 |
+
def forward(
|
| 61 |
+
self,
|
| 62 |
+
x: torch.Tensor,
|
| 63 |
+
symbol_retriever: nn.Module,
|
| 64 |
+
mask: torch.Tensor,
|
| 65 |
+
pos_info: Optional[PositionalInfo],
|
| 66 |
+
) -> torch.Tensor:
|
| 67 |
+
if self.norm_first:
|
| 68 |
+
normed_x = self.norm1(x)
|
| 69 |
+
# Retrieve symbols at the block boundary so pre-norm attention and
|
| 70 |
+
# input-dependent symbol retrieval use the same representation.
|
| 71 |
+
symbols = symbol_retriever(normed_x)
|
| 72 |
+
sensory_output, _ = self.sensory_attention(normed_x, mask=mask, pos_info=pos_info)
|
| 73 |
+
relational_output, _ = self.relational_attention(
|
| 74 |
+
normed_x,
|
| 75 |
+
symbols,
|
| 76 |
+
mask=mask,
|
| 77 |
+
pos_info=pos_info,
|
| 78 |
+
)
|
| 79 |
+
x = x + self.dropout(torch.cat((sensory_output, relational_output), dim=-1))
|
| 80 |
+
x = x + self.dropout(self.feed_forward(self.norm2(x)))
|
| 81 |
+
return x
|
| 82 |
+
|
| 83 |
+
symbols = symbol_retriever(x)
|
| 84 |
+
sensory_output, _ = self.sensory_attention(x, mask=mask, pos_info=pos_info)
|
| 85 |
+
relational_output, _ = self.relational_attention(
|
| 86 |
+
x,
|
| 87 |
+
symbols,
|
| 88 |
+
mask=mask,
|
| 89 |
+
pos_info=pos_info,
|
| 90 |
+
)
|
| 91 |
+
x = self.norm1(x + self.dropout(torch.cat((sensory_output, relational_output), dim=-1)))
|
| 92 |
+
x = self.norm2(x + self.dropout(self.feed_forward(x)))
|
| 93 |
+
return x
|
| 94 |
+
|
| 95 |
+
|
| 96 |
+
class DatDecoderLM(nn.Module):
|
| 97 |
+
def __init__(self, config: DatLMConfig):
|
| 98 |
+
super().__init__()
|
| 99 |
+
self.config = config
|
| 100 |
+
self.token_embeddings = nn.Embedding(config.vocab_size, config.hidden_dim)
|
| 101 |
+
self.embedding_dropout = nn.Dropout(config.dropout)
|
| 102 |
+
position_dim = (
|
| 103 |
+
config.hidden_dim if config.pe_type in {"sinusoidal", "learned", "none"}
|
| 104 |
+
else config.head_dim
|
| 105 |
+
)
|
| 106 |
+
self.position_encoder = PositionalEncoding(
|
| 107 |
+
embedding_dim=position_dim,
|
| 108 |
+
pe_type=config.pe_type,
|
| 109 |
+
max_len=config.max_seq_len,
|
| 110 |
+
theta=config.rope_theta,
|
| 111 |
+
max_rel_pos=config.max_rel_pos,
|
| 112 |
+
init_range=config.init_range,
|
| 113 |
+
)
|
| 114 |
+
self.symbol_retrievers = _build_symbol_retrievers(config)
|
| 115 |
+
self.layers = nn.ModuleList(DatDecoderBlock(config) for _ in range(config.n_layers))
|
| 116 |
+
self.final_norm = _create_norm(config.norm_type, config.hidden_dim)
|
| 117 |
+
self.lm_head = nn.Linear(config.hidden_dim, config.vocab_size, bias=False)
|
| 118 |
+
if config.mlm_head_enabled:
|
| 119 |
+
self.mlm_head = MaskClassifier(
|
| 120 |
+
hidden_dim=config.hidden_dim,
|
| 121 |
+
vocab_size=config.vocab_size,
|
| 122 |
+
norm_type=config.norm_type,
|
| 123 |
+
ffn_activation=config.ffn_activation,
|
| 124 |
+
use_bias_ffn=config.use_bias_ffn,
|
| 125 |
+
dropout=config.dropout,
|
| 126 |
+
word_embedding=self.lm_head.weight if not config.tie_lm_head else self.token_embeddings.weight,
|
| 127 |
+
)
|
| 128 |
+
self._init_weights()
|
| 129 |
+
|
| 130 |
+
@property
|
| 131 |
+
def device(self) -> torch.device:
|
| 132 |
+
return self.token_embeddings.weight.device
|
| 133 |
+
|
| 134 |
+
def _init_weights(self) -> None:
|
| 135 |
+
if self.config.init_scheme == "xavier_uniform":
|
| 136 |
+
nn.init.xavier_uniform_(
|
| 137 |
+
self.token_embeddings.weight,
|
| 138 |
+
gain=nn.init.calculate_gain("linear"),
|
| 139 |
+
)
|
| 140 |
+
if self.config.tie_lm_head:
|
| 141 |
+
self.lm_head.weight = self.token_embeddings.weight
|
| 142 |
+
else:
|
| 143 |
+
nn.init.xavier_uniform_(
|
| 144 |
+
self.lm_head.weight,
|
| 145 |
+
gain=nn.init.calculate_gain("linear"),
|
| 146 |
+
)
|
| 147 |
+
return
|
| 148 |
+
|
| 149 |
+
if self.config.init_scheme == "normal_0_02_scaled_projection":
|
| 150 |
+
self._init_normal_0_02_scaled_projection()
|
| 151 |
+
if self.config.tie_lm_head:
|
| 152 |
+
self.lm_head.weight = self.token_embeddings.weight
|
| 153 |
+
return
|
| 154 |
+
|
| 155 |
+
raise ValueError(f"Unsupported init_scheme: {self.config.init_scheme}")
|
| 156 |
+
|
| 157 |
+
def _init_normal_0_02_scaled_projection(self) -> None:
|
| 158 |
+
for module in self.modules():
|
| 159 |
+
if isinstance(module, nn.Linear):
|
| 160 |
+
nn.init.normal_(module.weight, mean=0.0, std=0.02)
|
| 161 |
+
if module.bias is not None:
|
| 162 |
+
nn.init.zeros_(module.bias)
|
| 163 |
+
elif isinstance(module, nn.Embedding):
|
| 164 |
+
nn.init.normal_(module.weight, mean=0.0, std=0.02)
|
| 165 |
+
|
| 166 |
+
for module in self.modules():
|
| 167 |
+
if isinstance(module, RelationalAttention):
|
| 168 |
+
nn.init.normal_(module.wr_proj, mean=0.0, std=0.02)
|
| 169 |
+
elif isinstance(module, SymbolicAttentionRetriever):
|
| 170 |
+
nn.init.normal_(module.template_features, mean=0.0, std=1.0)
|
| 171 |
+
nn.init.normal_(module.symbol_library, mean=0.0, std=1.0)
|
| 172 |
+
elif isinstance(module, RelativePositionalSymbolRetriever) and not module.rope:
|
| 173 |
+
nn.init.xavier_uniform_(module.position_encoder.rel_pos_embeddings_table.weight)
|
| 174 |
+
|
| 175 |
+
scaled_std = 0.02 / math.sqrt(2 * self.config.n_layers)
|
| 176 |
+
ffn_scaled_suffix = (
|
| 177 |
+
"feed_forward.w_up.weight"
|
| 178 |
+
if self.config.ffn_activation == "swiglu"
|
| 179 |
+
else "feed_forward.linear2.weight"
|
| 180 |
+
)
|
| 181 |
+
for name, parameter in self.named_parameters():
|
| 182 |
+
# Exclude the MLM head from depth-scaled init: its feed_forward
|
| 183 |
+
# down-projection shares the same suffix as decoder layers but
|
| 184 |
+
# should use the regular 0.02 initialization.
|
| 185 |
+
if name.startswith("mlm_head"):
|
| 186 |
+
continue
|
| 187 |
+
if name.endswith("o_proj.weight") or name.endswith(ffn_scaled_suffix):
|
| 188 |
+
nn.init.normal_(parameter, mean=0.0, std=scaled_std)
|
| 189 |
+
|
| 190 |
+
def _build_attention_mask(
|
| 191 |
+
self,
|
| 192 |
+
input_ids: torch.Tensor,
|
| 193 |
+
attention_mask: Optional[torch.Tensor] = None,
|
| 194 |
+
bidirectional: bool = False,
|
| 195 |
+
) -> torch.Tensor:
|
| 196 |
+
return build_decoder_attention_mask(
|
| 197 |
+
input_ids=input_ids,
|
| 198 |
+
pad_token_id=self.config.pad_token_id,
|
| 199 |
+
eos_token_id=self.config.eos_token_id,
|
| 200 |
+
sequence_boundary_policy=self.config.sequence_boundary_policy,
|
| 201 |
+
attention_mask=attention_mask,
|
| 202 |
+
segment_boundary_token_id=self.config.segment_boundary_token_id,
|
| 203 |
+
bidirectional=bidirectional,
|
| 204 |
+
)
|
| 205 |
+
|
| 206 |
+
def encode_for_objective(
|
| 207 |
+
self,
|
| 208 |
+
input_ids: torch.Tensor,
|
| 209 |
+
attention_mask: Optional[torch.Tensor] = None,
|
| 210 |
+
bidirectional: bool = False,
|
| 211 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 212 |
+
if input_ids.dim() != 2:
|
| 213 |
+
raise ValueError(f"input_ids must be rank-2 [batch, seq], got shape {tuple(input_ids.shape)}")
|
| 214 |
+
if input_ids.shape[1] > self.config.max_seq_len:
|
| 215 |
+
raise ValueError(
|
| 216 |
+
f"Sequence length {input_ids.shape[1]} exceeds max_seq_len {self.config.max_seq_len}"
|
| 217 |
+
)
|
| 218 |
+
|
| 219 |
+
|
| 220 |
+
mask = self._build_attention_mask(input_ids, attention_mask, bidirectional=bidirectional)
|
| 221 |
+
token_embeddings = self.token_embeddings(input_ids)
|
| 222 |
+
hidden_states = token_embeddings
|
| 223 |
+
pos_info = self.position_encoder.get_positional_info(input_ids.shape[1], input_ids.device)
|
| 224 |
+
if pos_info.apply_to_embeddings:
|
| 225 |
+
if pos_info.embeddings is None:
|
| 226 |
+
raise ValueError(f"Embedding-level pe_type {pos_info.pe_type} did not provide embeddings.")
|
| 227 |
+
hidden_states = hidden_states + pos_info.embeddings.unsqueeze(0)
|
| 228 |
+
hidden_states = self.embedding_dropout(hidden_states)
|
| 229 |
+
|
| 230 |
+
for symbol_retriever, layer in zip(self.symbol_retrievers, self.layers):
|
| 231 |
+
hidden_states = layer(
|
| 232 |
+
hidden_states,
|
| 233 |
+
symbol_retriever=symbol_retriever,
|
| 234 |
+
mask=mask,
|
| 235 |
+
pos_info=pos_info,
|
| 236 |
+
)
|
| 237 |
+
|
| 238 |
+
hidden_states = self.final_norm(hidden_states)
|
| 239 |
+
return token_embeddings, hidden_states
|
| 240 |
+
|
| 241 |
+
def forward(
|
| 242 |
+
self,
|
| 243 |
+
input_ids: torch.Tensor,
|
| 244 |
+
targets: Optional[torch.Tensor] = None,
|
| 245 |
+
attention_mask: Optional[torch.Tensor] = None,
|
| 246 |
+
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
|
| 247 |
+
_, hidden_states = self.encode_for_objective(input_ids, attention_mask=attention_mask)
|
| 248 |
+
logits = self.lm_head(hidden_states)
|
| 249 |
+
|
| 250 |
+
loss = None
|
| 251 |
+
if targets is not None:
|
| 252 |
+
if targets.shape != input_ids.shape:
|
| 253 |
+
raise ValueError(
|
| 254 |
+
f"targets shape must match input_ids shape, got {tuple(targets.shape)} "
|
| 255 |
+
f"vs {tuple(input_ids.shape)}"
|
| 256 |
+
)
|
| 257 |
+
loss = F.cross_entropy(
|
| 258 |
+
logits.reshape(-1, logits.size(-1)),
|
| 259 |
+
targets.reshape(-1),
|
| 260 |
+
ignore_index=-1,
|
| 261 |
+
)
|
| 262 |
+
|
| 263 |
+
return logits, loss
|
| 264 |
+
|
| 265 |
+
def forward_mlm(
|
| 266 |
+
self,
|
| 267 |
+
input_ids: torch.Tensor,
|
| 268 |
+
attention_mask: Optional[torch.Tensor] = None,
|
| 269 |
+
) -> torch.Tensor:
|
| 270 |
+
"""Encode bidirectionally and apply the MLM head.
|
| 271 |
+
|
| 272 |
+
Args:
|
| 273 |
+
input_ids: Masked input token ids [batch, seq].
|
| 274 |
+
attention_mask: Validity mask [batch, seq].
|
| 275 |
+
|
| 276 |
+
Returns:
|
| 277 |
+
MLM logits at every position [batch, seq, vocab].
|
| 278 |
+
"""
|
| 279 |
+
if not self.config.mlm_head_enabled:
|
| 280 |
+
raise ValueError(
|
| 281 |
+
"forward_mlm requires mlm_head_enabled=True; "
|
| 282 |
+
"the MLM head is not instantiated."
|
| 283 |
+
)
|
| 284 |
+
_, hidden_states = self.encode_for_objective(
|
| 285 |
+
input_ids, attention_mask=attention_mask, bidirectional=True
|
| 286 |
+
)
|
| 287 |
+
return self.mlm_head(hidden_states)
|
| 288 |
+
|
| 289 |
+
@torch.no_grad()
|
| 290 |
+
def generate(
|
| 291 |
+
self,
|
| 292 |
+
input_ids: torch.Tensor,
|
| 293 |
+
attention_mask: Optional[torch.Tensor] = None,
|
| 294 |
+
max_new_tokens: int = 100,
|
| 295 |
+
temperature: float = 1.0,
|
| 296 |
+
top_k: Optional[int] = None,
|
| 297 |
+
do_sample: bool = True,
|
| 298 |
+
eos_token_id: Optional[int] = None,
|
| 299 |
+
) -> torch.Tensor:
|
| 300 |
+
if temperature <= 0.0:
|
| 301 |
+
raise ValueError(f"temperature must be positive, got {temperature}")
|
| 302 |
+
if max_new_tokens < 0:
|
| 303 |
+
raise ValueError(f"max_new_tokens must be non-negative, got {max_new_tokens}")
|
| 304 |
+
if top_k is not None and top_k <= 0:
|
| 305 |
+
raise ValueError(f"top_k must be positive when provided, got {top_k}")
|
| 306 |
+
if eos_token_id is None:
|
| 307 |
+
eos_token_id = self.config.eos_token_id
|
| 308 |
+
|
| 309 |
+
if attention_mask is None:
|
| 310 |
+
generated_sequences = [row.clone() for row in input_ids]
|
| 311 |
+
else:
|
| 312 |
+
current_attention_mask = attention_mask.bool()
|
| 313 |
+
generated_sequences = []
|
| 314 |
+
for row, row_mask in zip(input_ids, current_attention_mask):
|
| 315 |
+
generated_sequence = row[row_mask]
|
| 316 |
+
if generated_sequence.numel() == 0:
|
| 317 |
+
raise ValueError("Each input row must contain at least one unmasked token for generation.")
|
| 318 |
+
generated_sequences.append(generated_sequence)
|
| 319 |
+
|
| 320 |
+
finished = torch.tensor(
|
| 321 |
+
[sequence[-1].item() == eos_token_id for sequence in generated_sequences],
|
| 322 |
+
dtype=torch.bool,
|
| 323 |
+
device=input_ids.device,
|
| 324 |
+
)
|
| 325 |
+
|
| 326 |
+
for _ in range(max_new_tokens):
|
| 327 |
+
context_sequences = [sequence[-self.config.max_seq_len :] for sequence in generated_sequences]
|
| 328 |
+
context_ids = pad_sequence(
|
| 329 |
+
context_sequences,
|
| 330 |
+
batch_first=True,
|
| 331 |
+
padding_value=self.config.pad_token_id,
|
| 332 |
+
)
|
| 333 |
+
context_mask = pad_sequence(
|
| 334 |
+
[
|
| 335 |
+
torch.ones(sequence.shape[0], dtype=torch.bool, device=input_ids.device)
|
| 336 |
+
for sequence in context_sequences
|
| 337 |
+
],
|
| 338 |
+
batch_first=True,
|
| 339 |
+
padding_value=False,
|
| 340 |
+
)
|
| 341 |
+
logits, _ = self(context_ids, attention_mask=context_mask)
|
| 342 |
+
last_positions = context_mask.long().sum(dim=1) - 1
|
| 343 |
+
batch_indices = torch.arange(logits.shape[0], device=logits.device)
|
| 344 |
+
next_token_logits = logits[batch_indices, last_positions, :] / temperature
|
| 345 |
+
|
| 346 |
+
if top_k is not None:
|
| 347 |
+
k = min(top_k, next_token_logits.size(-1))
|
| 348 |
+
top_values, _ = torch.topk(next_token_logits, k=k)
|
| 349 |
+
cutoff = top_values[:, -1].unsqueeze(-1)
|
| 350 |
+
next_token_logits = next_token_logits.masked_fill(next_token_logits < cutoff, float("-inf"))
|
| 351 |
+
|
| 352 |
+
if do_sample:
|
| 353 |
+
probs = F.softmax(next_token_logits, dim=-1)
|
| 354 |
+
next_token = torch.multinomial(probs, num_samples=1)
|
| 355 |
+
else:
|
| 356 |
+
next_token = torch.argmax(next_token_logits, dim=-1, keepdim=True)
|
| 357 |
+
|
| 358 |
+
for row_index in range(next_token.shape[0]):
|
| 359 |
+
if finished[row_index]:
|
| 360 |
+
continue
|
| 361 |
+
generated_sequences[row_index] = torch.cat(
|
| 362 |
+
[generated_sequences[row_index], next_token[row_index]],
|
| 363 |
+
)
|
| 364 |
+
if next_token[row_index, 0].item() == eos_token_id:
|
| 365 |
+
finished[row_index] = True
|
| 366 |
+
|
| 367 |
+
if torch.all(finished):
|
| 368 |
+
break
|
| 369 |
+
|
| 370 |
+
return pad_sequence(
|
| 371 |
+
generated_sequences,
|
| 372 |
+
batch_first=True,
|
| 373 |
+
padding_value=self.config.pad_token_id,
|
| 374 |
+
)
|
| 375 |
+
|
| 376 |
+
|
| 377 |
+
def _build_symbol_retrievers(config: DatLMConfig) -> nn.ModuleList:
|
| 378 |
+
retriever = _build_symbol_retriever(config)
|
| 379 |
+
if config.shared_symbol_retriever:
|
| 380 |
+
return nn.ModuleList([retriever] * config.n_layers)
|
| 381 |
+
return nn.ModuleList(_build_symbol_retriever(config) for _ in range(config.n_layers))
|
| 382 |
+
|
| 383 |
+
|
| 384 |
+
def _build_symbol_retriever(config: DatLMConfig) -> nn.Module:
|
| 385 |
+
if config.symbol_retrieval == "symbolic":
|
| 386 |
+
return SymbolicAttentionRetriever(
|
| 387 |
+
hidden_dim=config.hidden_dim,
|
| 388 |
+
symbol_dim=config.resolved_symbol_dim,
|
| 389 |
+
n_symbols=config.resolved_n_symbols,
|
| 390 |
+
n_heads=config.resolved_symbolic_attn_n_heads,
|
| 391 |
+
dropout=config.dropout,
|
| 392 |
+
use_bias=config.symbolic_use_bias,
|
| 393 |
+
)
|
| 394 |
+
if config.symbol_retrieval == "positional":
|
| 395 |
+
return PositionalSymbolRetriever(
|
| 396 |
+
symbol_dim=config.resolved_symbol_dim,
|
| 397 |
+
max_len=config.max_seq_len,
|
| 398 |
+
sinusoidal=config.positional_symbols_sinusoidal,
|
| 399 |
+
)
|
| 400 |
+
if config.symbol_retrieval == "relative":
|
| 401 |
+
max_rel_distance = config.max_rel_pos if config.max_rel_pos is not None else max(1, config.max_seq_len // 2)
|
| 402 |
+
return RelativePositionalSymbolRetriever(
|
| 403 |
+
symbol_dim=config.resolved_symbol_dim,
|
| 404 |
+
max_rel_distance=max_rel_distance,
|
| 405 |
+
rope=config.relative_symbols_rope,
|
| 406 |
+
theta=config.rope_theta,
|
| 407 |
+
)
|
| 408 |
+
if config.symbol_retrieval == "relsymbolic":
|
| 409 |
+
return RelationalSymbolicAttentionRetriever(
|
| 410 |
+
hidden_dim=config.hidden_dim,
|
| 411 |
+
symbol_dim=config.resolved_symbol_dim,
|
| 412 |
+
rel_n_heads=config.relsymbolic_rel_n_heads,
|
| 413 |
+
symbolic_attn_n_heads=config.relsymbolic_symbolic_attn_n_heads,
|
| 414 |
+
n_symbols=config.resolved_n_symbols,
|
| 415 |
+
neighborhood_size=config.relsymbolic_neighborhood_size,
|
| 416 |
+
include_self=config.relsymbolic_include_self,
|
| 417 |
+
normalize_rels=config.relsymbolic_normalize_rels,
|
| 418 |
+
dropout=config.relsymbolic_dropout,
|
| 419 |
+
trainable_symbols=config.relsymbolic_trainable_symbols,
|
| 420 |
+
rel_scale=config.relsymbolic_rel_scale,
|
| 421 |
+
symbolic_attn_scale=config.relsymbolic_symbolic_attn_scale,
|
| 422 |
+
use_bias=config.relsymbolic_use_bias,
|
| 423 |
+
)
|
| 424 |
+
raise ValueError(f"Unsupported symbol_retrieval: {config.symbol_retrieval}")
|
| 425 |
+
|
| 426 |
+
|
| 427 |
+
def _build_relational_attention(config: DatLMConfig) -> nn.Module:
|
| 428 |
+
use_relative_symbols = config.symbol_retrieval == "relative"
|
| 429 |
+
if config.ra_type == "ra":
|
| 430 |
+
return RelationalAttention(
|
| 431 |
+
hidden_dim=config.hidden_dim,
|
| 432 |
+
symbol_dim=config.resolved_symbol_dim,
|
| 433 |
+
n_heads=config.n_heads_ra,
|
| 434 |
+
total_n_heads=config.total_n_heads,
|
| 435 |
+
n_relations=config.resolved_ra_n_relations,
|
| 436 |
+
dropout=config.dropout,
|
| 437 |
+
rel_activation=config.ra_rel_activation,
|
| 438 |
+
symmetric_rels=config.ra_symmetric_rels,
|
| 439 |
+
use_relative_positional_symbols=use_relative_symbols,
|
| 440 |
+
use_bias_qkv=config.use_bias_qkv,
|
| 441 |
+
use_bias_out=config.use_bias_out,
|
| 442 |
+
)
|
| 443 |
+
if config.ra_type == "rca":
|
| 444 |
+
return RelationalCrossAttention(
|
| 445 |
+
hidden_dim=config.hidden_dim,
|
| 446 |
+
symbol_dim=config.resolved_symbol_dim,
|
| 447 |
+
n_heads=config.n_heads_ra,
|
| 448 |
+
total_n_heads=config.total_n_heads,
|
| 449 |
+
dropout=config.dropout,
|
| 450 |
+
activation=config.ra_rel_activation,
|
| 451 |
+
use_relative_positional_symbols=use_relative_symbols,
|
| 452 |
+
use_bias_qkv=config.use_bias_qkv,
|
| 453 |
+
use_bias_out=config.use_bias_out,
|
| 454 |
+
)
|
| 455 |
+
if config.ra_type == "disrca":
|
| 456 |
+
return DisentangledRelationalCrossAttention(
|
| 457 |
+
hidden_dim=config.hidden_dim,
|
| 458 |
+
symbol_dim=config.resolved_symbol_dim,
|
| 459 |
+
n_heads=config.n_heads_ra,
|
| 460 |
+
total_n_heads=config.total_n_heads,
|
| 461 |
+
dropout=config.dropout,
|
| 462 |
+
rel_activation=config.ra_rel_activation,
|
| 463 |
+
use_relative_positional_symbols=use_relative_symbols,
|
| 464 |
+
use_bias_qkv=config.use_bias_qkv,
|
| 465 |
+
use_bias_out=config.use_bias_out,
|
| 466 |
+
)
|
| 467 |
+
raise ValueError(f"Unsupported ra_type: {config.ra_type}")
|
| 468 |
+
|
| 469 |
+
|
| 470 |
+
def _create_norm(norm_type: str, hidden_dim: int) -> nn.Module:
|
| 471 |
+
if norm_type == "layernorm":
|
| 472 |
+
return nn.LayerNorm(hidden_dim)
|
| 473 |
+
if norm_type == "rmsnorm":
|
| 474 |
+
return RMSNorm(hidden_dim)
|
| 475 |
+
raise ValueError(f"Unsupported norm_type: {norm_type}")
|
dat_symbols.py
ADDED
|
@@ -0,0 +1,263 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Symbol retrieval modules for the owned DAT backend."""
|
| 2 |
+
|
| 3 |
+
import math
|
| 4 |
+
|
| 5 |
+
import torch
|
| 6 |
+
import torch.nn as nn
|
| 7 |
+
import torch.nn.functional as F
|
| 8 |
+
|
| 9 |
+
from .transformer_components import PositionalEncoding
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class SymbolicAttentionRetriever(nn.Module):
|
| 13 |
+
def __init__(
|
| 14 |
+
self,
|
| 15 |
+
hidden_dim: int,
|
| 16 |
+
symbol_dim: int,
|
| 17 |
+
n_symbols: int,
|
| 18 |
+
n_heads: int,
|
| 19 |
+
dropout: float = 0.0,
|
| 20 |
+
trainable_symbols: bool = True,
|
| 21 |
+
scale: float | None = None,
|
| 22 |
+
use_bias: bool = False,
|
| 23 |
+
):
|
| 24 |
+
super().__init__()
|
| 25 |
+
if hidden_dim % n_heads != 0:
|
| 26 |
+
raise ValueError(f"hidden_dim ({hidden_dim}) must be divisible by n_heads ({n_heads})")
|
| 27 |
+
if symbol_dim % n_heads != 0:
|
| 28 |
+
raise ValueError(f"symbol_dim ({symbol_dim}) must be divisible by n_heads ({n_heads})")
|
| 29 |
+
if scale is not None and scale <= 0.0:
|
| 30 |
+
raise ValueError(f"scale must be positive when provided, got {scale}")
|
| 31 |
+
|
| 32 |
+
self.hidden_dim = hidden_dim
|
| 33 |
+
self.symbol_dim = symbol_dim
|
| 34 |
+
self.n_symbols = n_symbols
|
| 35 |
+
self.n_heads = n_heads
|
| 36 |
+
self.head_dim = hidden_dim // n_heads
|
| 37 |
+
self.symbol_head_dim = symbol_dim // n_heads
|
| 38 |
+
self.scale = 1.0 / math.sqrt(self.head_dim) if scale is None else scale
|
| 39 |
+
|
| 40 |
+
self.q_proj = nn.Linear(hidden_dim, hidden_dim, bias=use_bias)
|
| 41 |
+
self.template_features = nn.Parameter(torch.empty(n_symbols, hidden_dim))
|
| 42 |
+
self.symbol_library = nn.Parameter(
|
| 43 |
+
torch.empty(n_symbols, symbol_dim),
|
| 44 |
+
requires_grad=trainable_symbols,
|
| 45 |
+
)
|
| 46 |
+
self.dropout = nn.Dropout(dropout)
|
| 47 |
+
self._init_weights()
|
| 48 |
+
|
| 49 |
+
def _init_weights(self) -> None:
|
| 50 |
+
nn.init.normal_(self.template_features, mean=0.0, std=0.9)
|
| 51 |
+
nn.init.normal_(self.symbol_library, mean=0.0, std=0.9)
|
| 52 |
+
|
| 53 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 54 |
+
batch_size, seq_len, _ = x.shape
|
| 55 |
+
|
| 56 |
+
queries = self.q_proj(x)
|
| 57 |
+
queries = queries.view(batch_size, seq_len, self.n_heads, self.head_dim)
|
| 58 |
+
queries = queries.transpose(1, 2)
|
| 59 |
+
|
| 60 |
+
keys = self.template_features.view(self.n_symbols, self.n_heads, self.head_dim)
|
| 61 |
+
keys = keys.transpose(0, 1).unsqueeze(0).expand(batch_size, -1, -1, -1)
|
| 62 |
+
|
| 63 |
+
values = self.symbol_library.view(self.n_symbols, self.n_heads, self.symbol_head_dim)
|
| 64 |
+
values = values.transpose(0, 1).unsqueeze(0).expand(batch_size, -1, -1, -1)
|
| 65 |
+
|
| 66 |
+
scores = torch.matmul(queries, keys.transpose(-2, -1)) * self.scale
|
| 67 |
+
weights = F.softmax(scores, dim=-1)
|
| 68 |
+
weights = self.dropout(weights)
|
| 69 |
+
retrieved = torch.matmul(weights, values)
|
| 70 |
+
retrieved = retrieved.transpose(1, 2).contiguous()
|
| 71 |
+
return retrieved.view(batch_size, seq_len, self.symbol_dim)
|
| 72 |
+
|
| 73 |
+
|
| 74 |
+
class PositionalSymbolRetriever(nn.Module):
|
| 75 |
+
def __init__(
|
| 76 |
+
self,
|
| 77 |
+
symbol_dim: int,
|
| 78 |
+
max_len: int,
|
| 79 |
+
sinusoidal: bool = False,
|
| 80 |
+
):
|
| 81 |
+
super().__init__()
|
| 82 |
+
self.symbol_dim = symbol_dim
|
| 83 |
+
self.max_len = max_len
|
| 84 |
+
self.sinusoidal = sinusoidal
|
| 85 |
+
pe_type = "sinusoidal" if sinusoidal else "learned"
|
| 86 |
+
self.position_encoder = PositionalEncoding(
|
| 87 |
+
embedding_dim=symbol_dim,
|
| 88 |
+
pe_type=pe_type,
|
| 89 |
+
max_len=max_len,
|
| 90 |
+
)
|
| 91 |
+
|
| 92 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 93 |
+
batch_size, seq_len, _ = x.shape
|
| 94 |
+
pos_info = self.position_encoder.get_positional_info(seq_len, x.device)
|
| 95 |
+
if pos_info.embeddings is None:
|
| 96 |
+
raise ValueError(f"Positional symbol retrieval produced no embeddings for seq_len={seq_len}")
|
| 97 |
+
return pos_info.embeddings.unsqueeze(0).expand(batch_size, -1, -1)
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
class RelativePositionalSymbolRetriever(nn.Module):
|
| 101 |
+
def __init__(
|
| 102 |
+
self,
|
| 103 |
+
symbol_dim: int,
|
| 104 |
+
max_rel_distance: int,
|
| 105 |
+
rope: bool = False,
|
| 106 |
+
theta: float = 10000.0,
|
| 107 |
+
):
|
| 108 |
+
super().__init__()
|
| 109 |
+
if rope and symbol_dim % 2 != 0:
|
| 110 |
+
raise ValueError(f"RoPE relative symbols require even symbol_dim, got {symbol_dim}")
|
| 111 |
+
if theta <= 0.0:
|
| 112 |
+
raise ValueError(f"theta must be positive, got {theta}")
|
| 113 |
+
self.symbol_dim = symbol_dim
|
| 114 |
+
self.max_rel_distance = max_rel_distance
|
| 115 |
+
self.rope = rope
|
| 116 |
+
self.theta = theta
|
| 117 |
+
if rope:
|
| 118 |
+
self.register_buffer("_rope_relative_cache", torch.empty(0), persistent=False)
|
| 119 |
+
else:
|
| 120 |
+
self.position_encoder = PositionalEncoding(
|
| 121 |
+
embedding_dim=symbol_dim,
|
| 122 |
+
pe_type="relative",
|
| 123 |
+
max_len=max_rel_distance * 2,
|
| 124 |
+
max_rel_pos=max_rel_distance,
|
| 125 |
+
)
|
| 126 |
+
|
| 127 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 128 |
+
seq_len = x.shape[1]
|
| 129 |
+
if self.rope:
|
| 130 |
+
return self._rope_relative_symbols(seq_len, x.device, x.dtype)
|
| 131 |
+
|
| 132 |
+
pos_info = self.position_encoder.get_positional_info(seq_len, x.device)
|
| 133 |
+
if pos_info.rel_embeddings is None:
|
| 134 |
+
raise ValueError(f"Relative symbol retrieval produced no embeddings for seq_len={seq_len}")
|
| 135 |
+
return pos_info.rel_embeddings
|
| 136 |
+
|
| 137 |
+
def _rope_relative_symbols(
|
| 138 |
+
self,
|
| 139 |
+
seq_len: int,
|
| 140 |
+
device: torch.device,
|
| 141 |
+
dtype: torch.dtype,
|
| 142 |
+
) -> torch.Tensor:
|
| 143 |
+
cached = self._rope_relative_cache
|
| 144 |
+
if (
|
| 145 |
+
cached.shape[0] >= seq_len
|
| 146 |
+
and cached.device == device
|
| 147 |
+
and cached.dtype == dtype
|
| 148 |
+
):
|
| 149 |
+
return cached[:seq_len, :seq_len, :]
|
| 150 |
+
|
| 151 |
+
positions = torch.arange(seq_len, device=device)
|
| 152 |
+
distances = positions[None, :] - positions[:, None]
|
| 153 |
+
if self.max_rel_distance is not None:
|
| 154 |
+
distances = torch.clamp(distances, -self.max_rel_distance, self.max_rel_distance)
|
| 155 |
+
inv_freq = PositionalEncoding._rope_inv_freq(self.symbol_dim, self.theta, device=device)
|
| 156 |
+
phases = distances.to(torch.float32).unsqueeze(-1) * inv_freq
|
| 157 |
+
symbols = torch.empty(seq_len, seq_len, self.symbol_dim, device=device)
|
| 158 |
+
symbols[..., 0::2] = torch.cos(phases)
|
| 159 |
+
symbols[..., 1::2] = torch.sin(phases)
|
| 160 |
+
self._rope_relative_cache = symbols.to(dtype=dtype)
|
| 161 |
+
return self._rope_relative_cache
|
| 162 |
+
|
| 163 |
+
|
| 164 |
+
class RelationalSymbolicAttentionRetriever(nn.Module):
|
| 165 |
+
def __init__(
|
| 166 |
+
self,
|
| 167 |
+
hidden_dim: int,
|
| 168 |
+
symbol_dim: int,
|
| 169 |
+
rel_n_heads: int,
|
| 170 |
+
symbolic_attn_n_heads: int,
|
| 171 |
+
n_symbols: int,
|
| 172 |
+
neighborhood_size: int = 2,
|
| 173 |
+
include_self: bool = False,
|
| 174 |
+
normalize_rels: bool = True,
|
| 175 |
+
dropout: float = 0.0,
|
| 176 |
+
trainable_symbols: bool = True,
|
| 177 |
+
rel_scale: float | None = None,
|
| 178 |
+
symbolic_attn_scale: float | None = None,
|
| 179 |
+
use_bias: bool = False,
|
| 180 |
+
):
|
| 181 |
+
super().__init__()
|
| 182 |
+
if hidden_dim % rel_n_heads != 0:
|
| 183 |
+
raise ValueError(
|
| 184 |
+
f"hidden_dim ({hidden_dim}) must be divisible by rel_n_heads ({rel_n_heads})"
|
| 185 |
+
)
|
| 186 |
+
if rel_scale is not None and rel_scale <= 0.0:
|
| 187 |
+
raise ValueError(f"rel_scale must be positive when provided, got {rel_scale}")
|
| 188 |
+
if symbolic_attn_scale is not None and symbolic_attn_scale <= 0.0:
|
| 189 |
+
raise ValueError(
|
| 190 |
+
"symbolic_attn_scale must be positive when provided, "
|
| 191 |
+
f"got {symbolic_attn_scale}"
|
| 192 |
+
)
|
| 193 |
+
|
| 194 |
+
self.hidden_dim = hidden_dim
|
| 195 |
+
self.symbol_dim = symbol_dim
|
| 196 |
+
self.rel_n_heads = rel_n_heads
|
| 197 |
+
self.neighborhood_size = neighborhood_size
|
| 198 |
+
self.include_self = include_self
|
| 199 |
+
self.normalize_rels = normalize_rels
|
| 200 |
+
self.neighborhood_dim = neighborhood_size + (1 if include_self else 0)
|
| 201 |
+
self.rel_feature_dim = rel_n_heads * self.neighborhood_dim
|
| 202 |
+
rel_head_dim = hidden_dim // rel_n_heads
|
| 203 |
+
self.rel_scale = 1.0 / math.sqrt(rel_head_dim) if rel_scale is None else rel_scale
|
| 204 |
+
|
| 205 |
+
self.q_proj = nn.Linear(hidden_dim, hidden_dim, bias=use_bias)
|
| 206 |
+
self.k_proj = nn.Linear(hidden_dim, hidden_dim, bias=use_bias)
|
| 207 |
+
self.rel_to_hidden = nn.Linear(self.rel_feature_dim, hidden_dim, bias=True)
|
| 208 |
+
self.symbolic_attention = SymbolicAttentionRetriever(
|
| 209 |
+
hidden_dim=hidden_dim,
|
| 210 |
+
symbol_dim=symbol_dim,
|
| 211 |
+
n_symbols=n_symbols,
|
| 212 |
+
n_heads=symbolic_attn_n_heads,
|
| 213 |
+
dropout=dropout,
|
| 214 |
+
trainable_symbols=trainable_symbols,
|
| 215 |
+
scale=symbolic_attn_scale,
|
| 216 |
+
use_bias=use_bias,
|
| 217 |
+
)
|
| 218 |
+
self._init_weights()
|
| 219 |
+
|
| 220 |
+
def _init_weights(self) -> None:
|
| 221 |
+
nn.init.xavier_uniform_(self.q_proj.weight)
|
| 222 |
+
nn.init.xavier_uniform_(self.k_proj.weight)
|
| 223 |
+
if self.q_proj.bias is not None:
|
| 224 |
+
nn.init.zeros_(self.q_proj.bias)
|
| 225 |
+
if self.k_proj.bias is not None:
|
| 226 |
+
nn.init.zeros_(self.k_proj.bias)
|
| 227 |
+
nn.init.xavier_uniform_(self.rel_to_hidden.weight)
|
| 228 |
+
nn.init.zeros_(self.rel_to_hidden.bias)
|
| 229 |
+
|
| 230 |
+
def _compute_neighborhood_indices(self, seq_len: int, device: torch.device) -> torch.Tensor:
|
| 231 |
+
positions = torch.arange(seq_len, device=device).unsqueeze(1)
|
| 232 |
+
# Decoder-LM adaptation: DSSL uses bidirectional neighborhoods, but this
|
| 233 |
+
# owned backend must keep symbol retrieval causal. Keep the old causal
|
| 234 |
+
# neighborhood order: current/immediate-past positions before older ones.
|
| 235 |
+
if self.include_self:
|
| 236 |
+
offsets = torch.arange(0, self.neighborhood_size + 1, device=device).unsqueeze(0)
|
| 237 |
+
else:
|
| 238 |
+
offsets = torch.arange(1, self.neighborhood_size + 1, device=device).unsqueeze(0)
|
| 239 |
+
return (positions - offsets).clamp(0, seq_len - 1)
|
| 240 |
+
|
| 241 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 242 |
+
batch_size, seq_len, _ = x.shape
|
| 243 |
+
head_dim = self.hidden_dim // self.rel_n_heads
|
| 244 |
+
|
| 245 |
+
queries = self.q_proj(x).view(batch_size, seq_len, self.rel_n_heads, head_dim)
|
| 246 |
+
keys = self.k_proj(x).view(batch_size, seq_len, self.rel_n_heads, head_dim)
|
| 247 |
+
queries = queries.transpose(1, 2)
|
| 248 |
+
keys = keys.transpose(1, 2)
|
| 249 |
+
|
| 250 |
+
neighbor_indices = self._compute_neighborhood_indices(seq_len, x.device)
|
| 251 |
+
neighborhood_keys = keys[:, :, neighbor_indices]
|
| 252 |
+
neighborhood_relations = torch.einsum("bhid,bhijd->bhij", queries, neighborhood_keys)
|
| 253 |
+
if self.normalize_rels:
|
| 254 |
+
neighborhood_relations = F.softmax(neighborhood_relations * self.rel_scale, dim=-1)
|
| 255 |
+
|
| 256 |
+
neighborhood_relations = neighborhood_relations.permute(0, 2, 3, 1)
|
| 257 |
+
neighborhood_relations = neighborhood_relations.contiguous().view(
|
| 258 |
+
batch_size,
|
| 259 |
+
seq_len,
|
| 260 |
+
self.rel_feature_dim,
|
| 261 |
+
)
|
| 262 |
+
relational_features = self.rel_to_hidden(neighborhood_relations)
|
| 263 |
+
return self.symbolic_attention(relational_features)
|
masks.py
ADDED
|
@@ -0,0 +1,61 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Attention-mask helpers shared by repo-owned decoder LMs."""
|
| 2 |
+
|
| 3 |
+
from typing import Optional
|
| 4 |
+
|
| 5 |
+
import torch
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
def build_decoder_attention_mask(
|
| 9 |
+
input_ids: torch.Tensor,
|
| 10 |
+
pad_token_id: int,
|
| 11 |
+
eos_token_id: int,
|
| 12 |
+
sequence_boundary_policy: str,
|
| 13 |
+
attention_mask: Optional[torch.Tensor] = None,
|
| 14 |
+
segment_boundary_token_id: Optional[int] = None,
|
| 15 |
+
bidirectional: bool = False,
|
| 16 |
+
) -> torch.Tensor:
|
| 17 |
+
if input_ids.ndim != 2:
|
| 18 |
+
raise ValueError(
|
| 19 |
+
"input_ids must be rank-2 [batch, seq], "
|
| 20 |
+
f"got shape {tuple(input_ids.shape)}"
|
| 21 |
+
)
|
| 22 |
+
_, seq_len = input_ids.shape
|
| 23 |
+
if attention_mask is None:
|
| 24 |
+
valid_tokens = input_ids != pad_token_id
|
| 25 |
+
else:
|
| 26 |
+
valid_tokens = attention_mask.bool()
|
| 27 |
+
|
| 28 |
+
query_mask = valid_tokens.unsqueeze(2)
|
| 29 |
+
key_mask = valid_tokens.unsqueeze(1)
|
| 30 |
+
if bidirectional:
|
| 31 |
+
mask = query_mask & key_mask
|
| 32 |
+
else:
|
| 33 |
+
directionality = torch.tril(
|
| 34 |
+
torch.ones(seq_len, seq_len, dtype=torch.bool, device=input_ids.device)
|
| 35 |
+
).unsqueeze(0)
|
| 36 |
+
mask = directionality & query_mask & key_mask
|
| 37 |
+
|
| 38 |
+
if sequence_boundary_policy == "none":
|
| 39 |
+
return mask
|
| 40 |
+
if sequence_boundary_policy == "segment_document":
|
| 41 |
+
if segment_boundary_token_id is None:
|
| 42 |
+
raise ValueError(
|
| 43 |
+
"segment_boundary_token_id is required when "
|
| 44 |
+
"sequence_boundary_policy='segment_document'"
|
| 45 |
+
)
|
| 46 |
+
# The boundary token starts the next segment: cumsum increments on the
|
| 47 |
+
# boundary position, so the marker attends with the following tokens.
|
| 48 |
+
segment_ids = torch.cumsum(
|
| 49 |
+
input_ids == segment_boundary_token_id, dim=1
|
| 50 |
+
)
|
| 51 |
+
same_segment = segment_ids.unsqueeze(1) == segment_ids.unsqueeze(2)
|
| 52 |
+
return mask & same_segment
|
| 53 |
+
if sequence_boundary_policy != "eos_document":
|
| 54 |
+
raise ValueError(f"Unsupported sequence_boundary_policy: {sequence_boundary_policy}")
|
| 55 |
+
|
| 56 |
+
# Next-token training predicts EOS from the preceding document token.
|
| 57 |
+
# Once EOS is present as an input token, it starts the next segment so the
|
| 58 |
+
# following document is not predicted with prior-document context.
|
| 59 |
+
document_ids = torch.cumsum(input_ids == eos_token_id, dim=1)
|
| 60 |
+
same_document = document_ids.unsqueeze(1) == document_ids.unsqueeze(2)
|
| 61 |
+
return mask & same_document
|
mlm.py
ADDED
|
@@ -0,0 +1,95 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""MLM head for the hybrid MLM/CLM objective.
|
| 2 |
+
|
| 3 |
+
Ported from temp/gpt-bert-main/pretraining/model.py MaskClassifier, but
|
| 4 |
+
configurable to match the trunk's customizability:
|
| 5 |
+
- norm_type: "layernorm" or "rmsnorm" (reuses the existing field).
|
| 6 |
+
- ffn_activation: "gelu", "swiglu", "relu", or "identity" (reuses the
|
| 7 |
+
existing field). The head's nonlinearity follows this choice.
|
| 8 |
+
- use_bias_ffn: whether the linear layers in the head use bias.
|
| 9 |
+
- weight tying: the final linear is tied to the word embedding, matching
|
| 10 |
+
gpt-bert's MaskClassifier and our lm_head.
|
| 11 |
+
|
| 12 |
+
The head is controlled by model config (mlm_head_enabled) so checkpoint
|
| 13 |
+
loading and HF export remain strict and honest.
|
| 14 |
+
"""
|
| 15 |
+
|
| 16 |
+
import torch
|
| 17 |
+
import torch.nn as nn
|
| 18 |
+
|
| 19 |
+
from .transformer_components import RMSNorm, FeedForward
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def _create_norm(norm_type: str, hidden_dim: int) -> nn.Module:
|
| 23 |
+
if norm_type == "layernorm":
|
| 24 |
+
return nn.LayerNorm(hidden_dim)
|
| 25 |
+
if norm_type == "rmsnorm":
|
| 26 |
+
return RMSNorm(hidden_dim)
|
| 27 |
+
raise ValueError(f"Unsupported norm_type: {norm_type}")
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
class MaskClassifier(nn.Module):
|
| 31 |
+
"""MLM prediction head with configurable norm and activation.
|
| 32 |
+
|
| 33 |
+
Architecture (matching gpt-bert's MaskClassifier):
|
| 34 |
+
norm -> FeedForward (activation, bias configurable) -> norm -> dropout -> linear (tied)
|
| 35 |
+
|
| 36 |
+
The final linear is tied to the word embedding weight, matching
|
| 37 |
+
gpt-bert's MaskClassifier and our lm_head.
|
| 38 |
+
|
| 39 |
+
Args:
|
| 40 |
+
hidden_dim: Model hidden dimension.
|
| 41 |
+
vocab_size: Vocabulary size for the output projection.
|
| 42 |
+
norm_type: "layernorm" or "rmsnorm".
|
| 43 |
+
ffn_activation: "gelu", "swiglu", "relu", or "identity".
|
| 44 |
+
use_bias_ffn: Whether the intermediate linear layers use bias.
|
| 45 |
+
dropout: Dropout probability.
|
| 46 |
+
word_embedding: The word embedding weight to tie the final linear to.
|
| 47 |
+
"""
|
| 48 |
+
|
| 49 |
+
def __init__(
|
| 50 |
+
self,
|
| 51 |
+
hidden_dim: int,
|
| 52 |
+
vocab_size: int,
|
| 53 |
+
norm_type: str,
|
| 54 |
+
ffn_activation: str,
|
| 55 |
+
use_bias_ffn: bool,
|
| 56 |
+
dropout: float,
|
| 57 |
+
word_embedding: nn.Parameter,
|
| 58 |
+
):
|
| 59 |
+
super().__init__()
|
| 60 |
+
self.hidden_dim = hidden_dim
|
| 61 |
+
self.vocab_size = vocab_size
|
| 62 |
+
self.norm_type = norm_type
|
| 63 |
+
self.ffn_activation = ffn_activation
|
| 64 |
+
self.use_bias_ffn = use_bias_ffn
|
| 65 |
+
self.dropout_p = dropout
|
| 66 |
+
|
| 67 |
+
self.norm1 = _create_norm(norm_type, hidden_dim)
|
| 68 |
+
self.feed_forward = FeedForward(
|
| 69 |
+
input_dim=hidden_dim,
|
| 70 |
+
hidden_dim=hidden_dim,
|
| 71 |
+
dropout=dropout,
|
| 72 |
+
activation=ffn_activation,
|
| 73 |
+
use_bias=use_bias_ffn,
|
| 74 |
+
)
|
| 75 |
+
self.norm2 = _create_norm(norm_type, hidden_dim)
|
| 76 |
+
self.dropout = nn.Dropout(dropout)
|
| 77 |
+
self.linear_out = nn.Linear(hidden_dim, vocab_size, bias=False)
|
| 78 |
+
self.linear_out.weight = word_embedding
|
| 79 |
+
|
| 80 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 81 |
+
"""Apply the MLM head to contextualized embeddings.
|
| 82 |
+
|
| 83 |
+
Args:
|
| 84 |
+
x: Contextualized embeddings [batch, seq_len, hidden_dim] or
|
| 85 |
+
flattened [num_masked, hidden_dim].
|
| 86 |
+
|
| 87 |
+
Returns:
|
| 88 |
+
Logits [batch, seq_len, vocab_size] or [num_masked, vocab_size].
|
| 89 |
+
"""
|
| 90 |
+
x = self.norm1(x)
|
| 91 |
+
x = self.feed_forward(x)
|
| 92 |
+
x = self.norm2(x)
|
| 93 |
+
x = self.dropout(x)
|
| 94 |
+
x = self.linear_out(x)
|
| 95 |
+
return x
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:37693c16dbb081d2adff7f0de2db3ecca7c22f6fa710007ee1caaed7b7072f4b
|
| 3 |
+
size 493209616
|
modeling_dat.py
ADDED
|
@@ -0,0 +1,290 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""HuggingFace PreTrainedModel wrapper for the dual-attention (DAT) decoder LM.
|
| 2 |
+
|
| 3 |
+
Source-of-truth copy. scripts/convert_dat_to_hf.py copies it into a generated HF
|
| 4 |
+
repository together with configuration_dat.py and the flattened model source
|
| 5 |
+
(dat_config.py, dat_core.py, dat_symbols.py, dat_lm.py, transformer_core.py,
|
| 6 |
+
transformer_components.py).
|
| 7 |
+
|
| 8 |
+
Why every model module is imported directly here: with trust_remote_code on a
|
| 9 |
+
local directory, transformers 4.46.3 only copies the entry module's *direct*
|
| 10 |
+
relative imports into its dynamic-module cache (it does not recurse). Importing
|
| 11 |
+
all flattened modules here forces every one of them to be copied, after which
|
| 12 |
+
their own relative imports resolve because the files are co-located.
|
| 13 |
+
"""
|
| 14 |
+
|
| 15 |
+
import collections
|
| 16 |
+
|
| 17 |
+
import torch
|
| 18 |
+
from transformers import PreTrainedModel
|
| 19 |
+
from transformers.modeling_outputs import BaseModelOutput, CausalLMOutput, MaskedLMOutput, SequenceClassifierOutput
|
| 20 |
+
from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss
|
| 21 |
+
|
| 22 |
+
# Force every flattened module to be copied into the dynamic-module cache.
|
| 23 |
+
from .attention_config import SUPPORTED_SEQUENCE_BOUNDARY_POLICIES # noqa: F401
|
| 24 |
+
from .masks import build_decoder_attention_mask # noqa: F401
|
| 25 |
+
from .mlm import MaskClassifier # noqa: F401
|
| 26 |
+
from .transformer_components import FeedForward, PositionalEncoding # noqa: F401
|
| 27 |
+
from .transformer_core import SelfAttention # noqa: F401
|
| 28 |
+
from .dat_symbols import SymbolicAttentionRetriever # noqa: F401
|
| 29 |
+
from .dat_core import RelationalAttention # noqa: F401
|
| 30 |
+
from .dat_config import DatLMConfig
|
| 31 |
+
from .dat_lm import DatDecoderLM
|
| 32 |
+
from .configuration_dat import DatConfig, DAT_LM_FIELDS
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
def _refresh_rope_buffers(model: DatDecoderLM, rope_theta: float) -> None:
|
| 36 |
+
pe = model.position_encoder
|
| 37 |
+
if hasattr(pe, "pe_type") and pe.pe_type == "rope":
|
| 38 |
+
cos, sin = pe._precompute_rope_freqs(pe.embedding_dim, pe.max_len, rope_theta)
|
| 39 |
+
pe.register_buffer("rope_cos", cos, persistent=False)
|
| 40 |
+
pe.register_buffer("rope_sin", sin, persistent=False)
|
| 41 |
+
pe._rope_uninitialized = True
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def _collect_shared_weight_keys(module: torch.nn.Module) -> dict[str, str]:
|
| 45 |
+
# remove_duplicate=False is essential: by default named_parameters yields
|
| 46 |
+
# each shared tensor only once, which would hide exactly the duplicates we
|
| 47 |
+
# need to declare (and which state_dict still emits under every name).
|
| 48 |
+
names_by_storage: dict[int, list[str]] = collections.defaultdict(list)
|
| 49 |
+
for name, parameter in module.named_parameters(remove_duplicate=False):
|
| 50 |
+
names_by_storage[id(parameter)].append(name)
|
| 51 |
+
shared_keys: dict[str, str] = {}
|
| 52 |
+
for names in names_by_storage.values():
|
| 53 |
+
if len(names) > 1:
|
| 54 |
+
names = sorted(names)
|
| 55 |
+
# Prefer token_embeddings.weight as the canonical source
|
| 56 |
+
# (HF convention: output weights tied to input embeddings).
|
| 57 |
+
source = names[0]
|
| 58 |
+
for n in names:
|
| 59 |
+
if n.endswith("token_embeddings.weight"):
|
| 60 |
+
source = n
|
| 61 |
+
break
|
| 62 |
+
for target in names:
|
| 63 |
+
if target != source:
|
| 64 |
+
shared_keys[target] = source
|
| 65 |
+
return shared_keys
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
class _ExpandedTiedWeightsMixin:
|
| 69 |
+
def get_expanded_tied_weights_keys(self, all_submodels: bool = False) -> dict:
|
| 70 |
+
# HF's default returns {} when tie_word_embeddings=False, which would
|
| 71 |
+
# discard the mlm_head.linear_out <-> lm_head tie for NextLat models
|
| 72 |
+
# (tie_lm_head=False). Return the dynamically computed mapping directly
|
| 73 |
+
# so all shared weights are properly tied on load.
|
| 74 |
+
return dict(self._tied_weights_keys)
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
class DatModel(_ExpandedTiedWeightsMixin, PreTrainedModel):
|
| 78 |
+
config_class = DatConfig
|
| 79 |
+
base_model_prefix = ""
|
| 80 |
+
|
| 81 |
+
def __init__(self, config: DatConfig) -> None:
|
| 82 |
+
super().__init__(config)
|
| 83 |
+
dat_config = DatLMConfig(**{field: getattr(config, field) for field in DAT_LM_FIELDS})
|
| 84 |
+
self.model = DatDecoderLM(dat_config)
|
| 85 |
+
self._tied_weights_keys = _collect_shared_weight_keys(self)
|
| 86 |
+
self.post_init()
|
| 87 |
+
|
| 88 |
+
def post_init(self) -> None:
|
| 89 |
+
super().post_init()
|
| 90 |
+
# Work around accelerate/transformers leaving persistent=False buffers uninitialized
|
| 91 |
+
_refresh_rope_buffers(self.model, self.config.rope_theta)
|
| 92 |
+
|
| 93 |
+
def get_input_embeddings(self) -> torch.nn.Module:
|
| 94 |
+
return self.model.token_embeddings
|
| 95 |
+
|
| 96 |
+
def set_input_embeddings(self, value: torch.nn.Module) -> None:
|
| 97 |
+
self.model.token_embeddings = value
|
| 98 |
+
|
| 99 |
+
def forward(
|
| 100 |
+
self,
|
| 101 |
+
input_ids: torch.Tensor,
|
| 102 |
+
attention_mask: torch.Tensor | None = None,
|
| 103 |
+
**kwargs,
|
| 104 |
+
) -> BaseModelOutput:
|
| 105 |
+
_, hidden_states = self.model.encode_for_objective(input_ids, attention_mask=attention_mask)
|
| 106 |
+
return BaseModelOutput(last_hidden_state=hidden_states)
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
class DatForCausalLM(_ExpandedTiedWeightsMixin, PreTrainedModel):
|
| 110 |
+
config_class = DatConfig
|
| 111 |
+
base_model_prefix = "model"
|
| 112 |
+
_tied_weights_keys = {"model.lm_head.weight": "model.token_embeddings.weight"}
|
| 113 |
+
|
| 114 |
+
def __init__(self, config: DatConfig) -> None:
|
| 115 |
+
super().__init__(config)
|
| 116 |
+
dat_config = DatLMConfig(**{field: getattr(config, field) for field in DAT_LM_FIELDS})
|
| 117 |
+
self.model = DatDecoderLM(dat_config)
|
| 118 |
+
self._tied_weights_keys = _collect_shared_weight_keys(self)
|
| 119 |
+
self.post_init()
|
| 120 |
+
|
| 121 |
+
def post_init(self) -> None:
|
| 122 |
+
super().post_init()
|
| 123 |
+
# Work around accelerate/transformers leaving persistent=False buffers uninitialized
|
| 124 |
+
_refresh_rope_buffers(self.model, self.config.rope_theta)
|
| 125 |
+
|
| 126 |
+
def get_input_embeddings(self) -> torch.nn.Module:
|
| 127 |
+
return self.model.token_embeddings
|
| 128 |
+
|
| 129 |
+
def set_input_embeddings(self, value: torch.nn.Module) -> None:
|
| 130 |
+
self.model.token_embeddings = value
|
| 131 |
+
|
| 132 |
+
def get_output_embeddings(self) -> torch.nn.Module:
|
| 133 |
+
return self.model.lm_head
|
| 134 |
+
|
| 135 |
+
def set_output_embeddings(self, new_embeddings: torch.nn.Module) -> None:
|
| 136 |
+
self.model.lm_head = new_embeddings
|
| 137 |
+
|
| 138 |
+
def forward(
|
| 139 |
+
self,
|
| 140 |
+
input_ids: torch.Tensor,
|
| 141 |
+
attention_mask: torch.Tensor | None = None,
|
| 142 |
+
labels: torch.Tensor | None = None,
|
| 143 |
+
**kwargs,
|
| 144 |
+
) -> CausalLMOutput:
|
| 145 |
+
logits, _ = self.model(input_ids, attention_mask=attention_mask)
|
| 146 |
+
|
| 147 |
+
loss = None
|
| 148 |
+
if labels is not None:
|
| 149 |
+
shift_logits = logits[:, :-1, :].contiguous()
|
| 150 |
+
shift_labels = labels[:, 1:].contiguous()
|
| 151 |
+
loss = torch.nn.functional.cross_entropy(
|
| 152 |
+
shift_logits.view(-1, shift_logits.size(-1)),
|
| 153 |
+
shift_labels.view(-1),
|
| 154 |
+
ignore_index=-100,
|
| 155 |
+
)
|
| 156 |
+
|
| 157 |
+
return CausalLMOutput(loss=loss, logits=logits)
|
| 158 |
+
|
| 159 |
+
|
| 160 |
+
class DatForMaskedLM(_ExpandedTiedWeightsMixin, PreTrainedModel):
|
| 161 |
+
config_class = DatConfig
|
| 162 |
+
base_model_prefix = "model"
|
| 163 |
+
|
| 164 |
+
def __init__(self, config: DatConfig) -> None:
|
| 165 |
+
super().__init__(config)
|
| 166 |
+
dat_config = DatLMConfig(**{field: getattr(config, field) for field in DAT_LM_FIELDS})
|
| 167 |
+
self.model = DatDecoderLM(dat_config)
|
| 168 |
+
self._tied_weights_keys = _collect_shared_weight_keys(self)
|
| 169 |
+
self.post_init()
|
| 170 |
+
|
| 171 |
+
def post_init(self) -> None:
|
| 172 |
+
super().post_init()
|
| 173 |
+
_refresh_rope_buffers(self.model, self.config.rope_theta)
|
| 174 |
+
|
| 175 |
+
def get_input_embeddings(self) -> torch.nn.Module:
|
| 176 |
+
return self.model.token_embeddings
|
| 177 |
+
|
| 178 |
+
def set_input_embeddings(self, value: torch.nn.Module) -> None:
|
| 179 |
+
self.model.token_embeddings = value
|
| 180 |
+
|
| 181 |
+
def get_output_embeddings(self) -> torch.nn.Module:
|
| 182 |
+
return self.model.mlm_head.linear_out
|
| 183 |
+
|
| 184 |
+
def set_output_embeddings(self, new_embeddings: torch.nn.Module) -> None:
|
| 185 |
+
self.model.mlm_head.linear_out = new_embeddings
|
| 186 |
+
# Keep lm_head in sync with mlm_head.linear_out so that
|
| 187 |
+
# resize_token_embeddings updates both tied heads.
|
| 188 |
+
self.model.lm_head.weight = new_embeddings.weight
|
| 189 |
+
|
| 190 |
+
def forward(
|
| 191 |
+
self,
|
| 192 |
+
input_ids: torch.Tensor,
|
| 193 |
+
attention_mask: torch.Tensor | None = None,
|
| 194 |
+
labels: torch.Tensor | None = None,
|
| 195 |
+
**kwargs,
|
| 196 |
+
) -> MaskedLMOutput | tuple:
|
| 197 |
+
return_dict = kwargs.pop("return_dict", None)
|
| 198 |
+
if return_dict is None:
|
| 199 |
+
return_dict = self.config.return_dict
|
| 200 |
+
if kwargs:
|
| 201 |
+
raise TypeError(f"Unexpected keyword argument(s): {list(kwargs)}")
|
| 202 |
+
logits = self.model.forward_mlm(input_ids, attention_mask=attention_mask)
|
| 203 |
+
|
| 204 |
+
loss = None
|
| 205 |
+
if labels is not None:
|
| 206 |
+
loss = CrossEntropyLoss(ignore_index=-100)(
|
| 207 |
+
logits.view(-1, logits.size(-1)),
|
| 208 |
+
labels.view(-1),
|
| 209 |
+
)
|
| 210 |
+
|
| 211 |
+
if not return_dict:
|
| 212 |
+
output = (logits,)
|
| 213 |
+
return ((loss,) + output) if loss is not None else output
|
| 214 |
+
|
| 215 |
+
return MaskedLMOutput(loss=loss, logits=logits)
|
| 216 |
+
|
| 217 |
+
class DatForSequenceClassification(_ExpandedTiedWeightsMixin, PreTrainedModel):
|
| 218 |
+
config_class = DatConfig
|
| 219 |
+
base_model_prefix = "model"
|
| 220 |
+
|
| 221 |
+
def __init__(self, config: DatConfig) -> None:
|
| 222 |
+
super().__init__(config)
|
| 223 |
+
self.num_labels = getattr(config, "num_labels", 2)
|
| 224 |
+
dat_config = DatLMConfig(**{field: getattr(config, field) for field in DAT_LM_FIELDS})
|
| 225 |
+
self.model = DatDecoderLM(dat_config)
|
| 226 |
+
self.score = torch.nn.Linear(config.hidden_dim, self.num_labels, bias=False)
|
| 227 |
+
self._tied_weights_keys = _collect_shared_weight_keys(self)
|
| 228 |
+
self.post_init()
|
| 229 |
+
|
| 230 |
+
def post_init(self) -> None:
|
| 231 |
+
super().post_init()
|
| 232 |
+
_refresh_rope_buffers(self.model, self.config.rope_theta)
|
| 233 |
+
|
| 234 |
+
def get_input_embeddings(self) -> torch.nn.Module:
|
| 235 |
+
return self.model.token_embeddings
|
| 236 |
+
|
| 237 |
+
def set_input_embeddings(self, value: torch.nn.Module) -> None:
|
| 238 |
+
self.model.token_embeddings = value
|
| 239 |
+
|
| 240 |
+
def forward(
|
| 241 |
+
self,
|
| 242 |
+
input_ids: torch.Tensor,
|
| 243 |
+
attention_mask: torch.Tensor | None = None,
|
| 244 |
+
labels: torch.Tensor | None = None,
|
| 245 |
+
**kwargs,
|
| 246 |
+
) -> SequenceClassifierOutput:
|
| 247 |
+
_, hidden_states = self.model.encode_for_objective(input_ids, attention_mask=attention_mask)
|
| 248 |
+
logits = self.score(hidden_states)
|
| 249 |
+
|
| 250 |
+
batch_size = input_ids.shape[0]
|
| 251 |
+
|
| 252 |
+
if self.config.pad_token_id is None:
|
| 253 |
+
last_non_pad_token = -1
|
| 254 |
+
else:
|
| 255 |
+
non_pad_mask = (input_ids != self.config.pad_token_id).to(logits.device, torch.int32)
|
| 256 |
+
if (non_pad_mask.sum(-1) == 0).any().item():
|
| 257 |
+
raise ValueError("Cannot pool sequence-classification logits for all-padding sequences")
|
| 258 |
+
token_indices = torch.arange(input_ids.shape[-1], device=logits.device, dtype=torch.int32)
|
| 259 |
+
last_non_pad_token = (token_indices * non_pad_mask).argmax(-1)
|
| 260 |
+
|
| 261 |
+
pooled_logits = logits[torch.arange(batch_size, device=logits.device), last_non_pad_token]
|
| 262 |
+
|
| 263 |
+
loss = None
|
| 264 |
+
if labels is not None:
|
| 265 |
+
if getattr(self.config, "problem_type", None) is None:
|
| 266 |
+
if self.num_labels == 1:
|
| 267 |
+
self.config.problem_type = "regression"
|
| 268 |
+
elif self.num_labels > 1 and (labels.dtype == torch.long or labels.dtype == torch.int):
|
| 269 |
+
self.config.problem_type = "single_label_classification"
|
| 270 |
+
else:
|
| 271 |
+
self.config.problem_type = "multi_label_classification"
|
| 272 |
+
|
| 273 |
+
if self.config.problem_type == "regression":
|
| 274 |
+
loss_fct = MSELoss()
|
| 275 |
+
if self.num_labels == 1:
|
| 276 |
+
loss = loss_fct(pooled_logits.squeeze(), labels.squeeze())
|
| 277 |
+
else:
|
| 278 |
+
loss = loss_fct(pooled_logits, labels)
|
| 279 |
+
elif self.config.problem_type == "single_label_classification":
|
| 280 |
+
loss_fct = CrossEntropyLoss()
|
| 281 |
+
loss = loss_fct(pooled_logits.view(-1, self.num_labels), labels.view(-1))
|
| 282 |
+
elif self.config.problem_type == "multi_label_classification":
|
| 283 |
+
loss_fct = BCEWithLogitsLoss()
|
| 284 |
+
loss = loss_fct(pooled_logits, labels)
|
| 285 |
+
|
| 286 |
+
return SequenceClassifierOutput(
|
| 287 |
+
loss=loss,
|
| 288 |
+
logits=pooled_logits,
|
| 289 |
+
hidden_states=hidden_states,
|
| 290 |
+
)
|
special_tokens_map.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
{"bos_token": "<s>", "eos_token": "</s>", "unk_token": "<unk>", "sep_token": "</s>", "pad_token": "<pad>", "cls_token": "<s>", "mask_token": "<mask>"}
|
tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
tokenizer_config.json
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"tokenizer_class": "PreTrainedTokenizerFast",
|
| 3 |
+
"bos_token": "<s>",
|
| 4 |
+
"eos_token": "</s>",
|
| 5 |
+
"unk_token": "<unk>",
|
| 6 |
+
"sep_token": "</s>",
|
| 7 |
+
"pad_token": "<pad>",
|
| 8 |
+
"cls_token": "<s>",
|
| 9 |
+
"mask_token": "<mask>"
|
| 10 |
+
}
|
transformer_components.py
ADDED
|
@@ -0,0 +1,282 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Minimal transformer components for decoder LM training."""
|
| 2 |
+
|
| 3 |
+
import math
|
| 4 |
+
from dataclasses import dataclass
|
| 5 |
+
from typing import Optional, Tuple
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
import torch.nn as nn
|
| 9 |
+
from torch import Tensor
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
@dataclass(frozen=True)
|
| 13 |
+
class PositionalInfo:
|
| 14 |
+
pe_type: str
|
| 15 |
+
embeddings: Optional[Tensor] = None
|
| 16 |
+
rope_freqs: Optional[Tuple[Tensor, Tensor]] = None
|
| 17 |
+
rel_embeddings: Optional[Tensor] = None
|
| 18 |
+
apply_to_embeddings: bool = False
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
class RMSNorm(nn.Module):
|
| 22 |
+
def __init__(self, dim: int, eps: float = 1e-5):
|
| 23 |
+
super().__init__()
|
| 24 |
+
self.eps = eps
|
| 25 |
+
self.weight = nn.Parameter(torch.ones(dim))
|
| 26 |
+
|
| 27 |
+
def _norm(self, x: Tensor) -> Tensor:
|
| 28 |
+
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
|
| 29 |
+
|
| 30 |
+
def forward(self, x: Tensor) -> Tensor:
|
| 31 |
+
if x.dtype in (torch.float16, torch.bfloat16):
|
| 32 |
+
output = self._norm(x.float()).type_as(x)
|
| 33 |
+
else:
|
| 34 |
+
output = self._norm(x)
|
| 35 |
+
return output * self.weight
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
class FeedForward(nn.Module):
|
| 39 |
+
def __init__(
|
| 40 |
+
self,
|
| 41 |
+
input_dim: int,
|
| 42 |
+
hidden_dim: int,
|
| 43 |
+
dropout: float = 0.0,
|
| 44 |
+
activation: str = "gelu",
|
| 45 |
+
use_bias: bool = True,
|
| 46 |
+
):
|
| 47 |
+
super().__init__()
|
| 48 |
+
self.activation = activation
|
| 49 |
+
self.use_swiglu = activation == "swiglu"
|
| 50 |
+
self.dropout = nn.Dropout(dropout)
|
| 51 |
+
if self.use_swiglu:
|
| 52 |
+
self.w_gate = nn.Linear(input_dim, hidden_dim, bias=use_bias)
|
| 53 |
+
self.w_up = nn.Linear(input_dim, hidden_dim, bias=use_bias)
|
| 54 |
+
self.w_down = nn.Linear(hidden_dim, input_dim, bias=use_bias)
|
| 55 |
+
self.activation_fn = nn.SiLU()
|
| 56 |
+
else:
|
| 57 |
+
self.linear1 = nn.Linear(input_dim, hidden_dim, bias=use_bias)
|
| 58 |
+
self.linear2 = nn.Linear(hidden_dim, input_dim, bias=use_bias)
|
| 59 |
+
if activation == "gelu":
|
| 60 |
+
self.activation_fn = nn.GELU()
|
| 61 |
+
elif activation == "relu":
|
| 62 |
+
self.activation_fn = nn.ReLU()
|
| 63 |
+
elif activation == "identity":
|
| 64 |
+
self.activation_fn = nn.Identity()
|
| 65 |
+
else:
|
| 66 |
+
raise ValueError(f"Unsupported ffn activation: {activation}")
|
| 67 |
+
self._init_weights()
|
| 68 |
+
|
| 69 |
+
def _init_weights(self) -> None:
|
| 70 |
+
for module in self.modules():
|
| 71 |
+
if isinstance(module, nn.Linear):
|
| 72 |
+
nn.init.xavier_uniform_(module.weight, gain=1.0)
|
| 73 |
+
if module.bias is not None:
|
| 74 |
+
nn.init.zeros_(module.bias)
|
| 75 |
+
|
| 76 |
+
def forward(self, x: Tensor) -> Tensor:
|
| 77 |
+
if self.use_swiglu:
|
| 78 |
+
gate = self.activation_fn(self.w_gate(x))
|
| 79 |
+
up = self.w_up(x)
|
| 80 |
+
x = gate * up
|
| 81 |
+
x = self.dropout(x)
|
| 82 |
+
return self.w_down(x)
|
| 83 |
+
|
| 84 |
+
x = self.linear1(x)
|
| 85 |
+
x = self.activation_fn(x)
|
| 86 |
+
x = self.dropout(x)
|
| 87 |
+
x = self.linear2(x)
|
| 88 |
+
return x
|
| 89 |
+
|
| 90 |
+
|
| 91 |
+
class PositionalEncoding(nn.Module):
|
| 92 |
+
def __init__(
|
| 93 |
+
self,
|
| 94 |
+
embedding_dim: int,
|
| 95 |
+
pe_type: str,
|
| 96 |
+
max_len: int,
|
| 97 |
+
theta: float = 10000.0,
|
| 98 |
+
max_rel_pos: Optional[int] = None,
|
| 99 |
+
init_range: float = 0.15,
|
| 100 |
+
):
|
| 101 |
+
super().__init__()
|
| 102 |
+
if embedding_dim <= 0:
|
| 103 |
+
raise ValueError(f"embedding_dim must be positive, got {embedding_dim}")
|
| 104 |
+
if max_len <= 0:
|
| 105 |
+
raise ValueError(f"max_len must be positive, got {max_len}")
|
| 106 |
+
if pe_type not in {"sinusoidal", "learned", "relative", "rope", "none"}:
|
| 107 |
+
raise ValueError(f"Unsupported pe_type: {pe_type}")
|
| 108 |
+
if theta <= 0.0:
|
| 109 |
+
raise ValueError(f"theta must be positive, got {theta}")
|
| 110 |
+
if init_range <= 0.0:
|
| 111 |
+
raise ValueError(f"init_range must be positive, got {init_range}")
|
| 112 |
+
|
| 113 |
+
self.embedding_dim = embedding_dim
|
| 114 |
+
self.pe_type = pe_type
|
| 115 |
+
self.max_len = max_len
|
| 116 |
+
|
| 117 |
+
if pe_type == "sinusoidal":
|
| 118 |
+
if embedding_dim % 2 != 0:
|
| 119 |
+
raise ValueError(f"Sinusoidal encoding requires even embedding_dim, got {embedding_dim}")
|
| 120 |
+
pe = self._precompute_sinusoidal_embeddings(
|
| 121 |
+
embedding_dim,
|
| 122 |
+
max_len,
|
| 123 |
+
device=torch.empty(0).device,
|
| 124 |
+
)
|
| 125 |
+
self.register_buffer("sinusoidal_embeddings", pe, persistent=False)
|
| 126 |
+
self._sinusoidal_uninitialized = True
|
| 127 |
+
elif pe_type == "learned":
|
| 128 |
+
self.position_embeddings = nn.Embedding(max_len, embedding_dim)
|
| 129 |
+
nn.init.uniform_(self.position_embeddings.weight, -init_range, init_range)
|
| 130 |
+
elif pe_type == "relative":
|
| 131 |
+
self.max_relative_position = max_rel_pos if max_rel_pos is not None else max_len // 2
|
| 132 |
+
if self.max_relative_position <= 0:
|
| 133 |
+
raise ValueError(
|
| 134 |
+
f"max_relative_position must be positive, got {self.max_relative_position}"
|
| 135 |
+
)
|
| 136 |
+
num_embeddings = 2 * self.max_relative_position + 1
|
| 137 |
+
self.rel_pos_embeddings_table = nn.Embedding(num_embeddings, embedding_dim)
|
| 138 |
+
nn.init.uniform_(self.rel_pos_embeddings_table.weight, -init_range, init_range)
|
| 139 |
+
elif pe_type == "rope":
|
| 140 |
+
if embedding_dim % 2 != 0:
|
| 141 |
+
raise ValueError(f"RoPE requires even embedding_dim, got {embedding_dim}")
|
| 142 |
+
self.theta = theta
|
| 143 |
+
cos, sin = self._precompute_rope_freqs(embedding_dim, max_len, theta)
|
| 144 |
+
self.register_buffer("rope_cos", cos, persistent=False)
|
| 145 |
+
self.register_buffer("rope_sin", sin, persistent=False)
|
| 146 |
+
self._rope_uninitialized = True
|
| 147 |
+
|
| 148 |
+
def get_positional_info(self, seq_len: int, device: torch.device) -> PositionalInfo:
|
| 149 |
+
if self.pe_type == "none":
|
| 150 |
+
return PositionalInfo(pe_type="none")
|
| 151 |
+
if self.pe_type == "sinusoidal":
|
| 152 |
+
if seq_len > self.sinusoidal_embeddings.size(0):
|
| 153 |
+
raise ValueError(
|
| 154 |
+
f"seq_len {seq_len} exceeds precomputed sinusoidal length {self.sinusoidal_embeddings.size(0)}"
|
| 155 |
+
)
|
| 156 |
+
if (
|
| 157 |
+
getattr(self, "_sinusoidal_uninitialized", False)
|
| 158 |
+
or self.sinusoidal_embeddings.device.type == "meta"
|
| 159 |
+
):
|
| 160 |
+
embeddings = self._precompute_sinusoidal_embeddings(
|
| 161 |
+
self.embedding_dim,
|
| 162 |
+
self.max_len,
|
| 163 |
+
device=device,
|
| 164 |
+
)
|
| 165 |
+
self.register_buffer("sinusoidal_embeddings", embeddings, persistent=False)
|
| 166 |
+
self._sinusoidal_uninitialized = False
|
| 167 |
+
return PositionalInfo(
|
| 168 |
+
pe_type="sinusoidal",
|
| 169 |
+
embeddings=self.sinusoidal_embeddings[:seq_len].to(device),
|
| 170 |
+
apply_to_embeddings=True,
|
| 171 |
+
)
|
| 172 |
+
if self.pe_type == "learned":
|
| 173 |
+
if seq_len > self.max_len:
|
| 174 |
+
raise ValueError(
|
| 175 |
+
f"seq_len {seq_len} exceeds maximum learned position length {self.max_len}"
|
| 176 |
+
)
|
| 177 |
+
positions = torch.arange(seq_len, dtype=torch.long, device=device)
|
| 178 |
+
return PositionalInfo(
|
| 179 |
+
pe_type="learned",
|
| 180 |
+
embeddings=self.position_embeddings(positions),
|
| 181 |
+
apply_to_embeddings=True,
|
| 182 |
+
)
|
| 183 |
+
if self.pe_type == "relative":
|
| 184 |
+
return PositionalInfo(
|
| 185 |
+
pe_type="relative",
|
| 186 |
+
rel_embeddings=self._generate_relative_embeddings(seq_len, device),
|
| 187 |
+
apply_to_embeddings=False,
|
| 188 |
+
)
|
| 189 |
+
if seq_len > self.rope_cos.size(0):
|
| 190 |
+
# Dynamically extend RoPE buffer (standard practice for extrapolation)
|
| 191 |
+
cos, sin = self._precompute_rope_freqs(self.embedding_dim, seq_len, self.theta)
|
| 192 |
+
self.register_buffer("rope_cos", cos, persistent=False)
|
| 193 |
+
self.register_buffer("rope_sin", sin, persistent=False)
|
| 194 |
+
self._rope_uninitialized = False
|
| 195 |
+
if getattr(self, "_rope_uninitialized", True) or torch.isnan(self.rope_cos[0, 0]):
|
| 196 |
+
cos, sin = self._precompute_rope_freqs(self.embedding_dim, self.max_len, self.theta)
|
| 197 |
+
self.rope_cos.copy_(cos)
|
| 198 |
+
self.rope_sin.copy_(sin)
|
| 199 |
+
self._rope_uninitialized = False
|
| 200 |
+
return PositionalInfo(
|
| 201 |
+
pe_type="rope",
|
| 202 |
+
rope_freqs=(
|
| 203 |
+
self.rope_cos[:seq_len].to(device),
|
| 204 |
+
self.rope_sin[:seq_len].to(device),
|
| 205 |
+
),
|
| 206 |
+
apply_to_embeddings=False,
|
| 207 |
+
)
|
| 208 |
+
|
| 209 |
+
@staticmethod
|
| 210 |
+
def _rope_inv_freq(dim: int, theta: float, device: torch.device | str) -> Tensor:
|
| 211 |
+
return 1.0 / (theta ** (torch.arange(0, dim, 2, device=device).float() / dim))
|
| 212 |
+
|
| 213 |
+
@staticmethod
|
| 214 |
+
def _precompute_sinusoidal_embeddings(
|
| 215 |
+
dim: int,
|
| 216 |
+
end: int,
|
| 217 |
+
device: torch.device | str,
|
| 218 |
+
) -> Tensor:
|
| 219 |
+
pe = torch.zeros(end, dim, device=device)
|
| 220 |
+
position = torch.arange(0, end, dtype=torch.float, device=device).unsqueeze(1)
|
| 221 |
+
div_term = PositionalEncoding._rope_inv_freq(dim, 10000.0, device=device)
|
| 222 |
+
pe[:, 0::2] = torch.sin(position * div_term)
|
| 223 |
+
pe[:, 1::2] = torch.cos(position * div_term)
|
| 224 |
+
return pe
|
| 225 |
+
|
| 226 |
+
@staticmethod
|
| 227 |
+
def _precompute_rope_freqs(dim: int, end: int, theta: float) -> Tuple[Tensor, Tensor]:
|
| 228 |
+
freqs = PositionalEncoding._rope_inv_freq(dim, theta, device="cpu")
|
| 229 |
+
positions = torch.arange(end, device="cpu")
|
| 230 |
+
freqs = torch.outer(positions, freqs).float()
|
| 231 |
+
return torch.cos(freqs), torch.sin(freqs)
|
| 232 |
+
|
| 233 |
+
def _generate_relative_embeddings(self, seq_len: int, device: torch.device) -> Tensor:
|
| 234 |
+
range_q = torch.arange(seq_len, device=device)
|
| 235 |
+
range_k = torch.arange(seq_len, device=device)
|
| 236 |
+
distance_mat = range_k[None, :] - range_q[:, None]
|
| 237 |
+
distance_mat_clipped = torch.clamp(
|
| 238 |
+
distance_mat,
|
| 239 |
+
-self.max_relative_position,
|
| 240 |
+
self.max_relative_position,
|
| 241 |
+
)
|
| 242 |
+
final_indices = distance_mat_clipped + self.max_relative_position
|
| 243 |
+
return self.rel_pos_embeddings_table(final_indices.long())
|
| 244 |
+
|
| 245 |
+
|
| 246 |
+
def reshape_for_broadcast(freqs_cis: Tensor, x: Tensor) -> Tensor:
|
| 247 |
+
ndim = x.ndim
|
| 248 |
+
if ndim < 2:
|
| 249 |
+
raise ValueError(f"Input tensor x must have at least 2 dimensions, got {ndim}")
|
| 250 |
+
if x.shape[-2] != freqs_cis.shape[0] or x.shape[-1] != freqs_cis.shape[-1]:
|
| 251 |
+
raise ValueError(
|
| 252 |
+
f"Shape mismatch for RoPE broadcasting: freqs_cis {freqs_cis.shape}, x {x.shape}"
|
| 253 |
+
)
|
| 254 |
+
shape = [1] * ndim
|
| 255 |
+
shape[-2] = x.shape[-2]
|
| 256 |
+
shape[-1] = x.shape[-1]
|
| 257 |
+
return freqs_cis.view(shape)
|
| 258 |
+
|
| 259 |
+
|
| 260 |
+
def apply_rotary_emb(xq: Tensor, xk: Tensor, freqs_cos: Tensor, freqs_sin: Tensor) -> Tuple[Tensor, Tensor]:
|
| 261 |
+
if xq.shape[-1] % 2 != 0:
|
| 262 |
+
raise ValueError(f"Query feature dimension must be even for RoPE, got {xq.shape[-1]}")
|
| 263 |
+
if xk.shape[-1] % 2 != 0:
|
| 264 |
+
raise ValueError(f"Key feature dimension must be even for RoPE, got {xk.shape[-1]}")
|
| 265 |
+
|
| 266 |
+
xq_r, xq_i = xq.float().reshape(xq.shape[:-1] + (-1, 2)).unbind(-1)
|
| 267 |
+
xk_r, xk_i = xk.float().reshape(xk.shape[:-1] + (-1, 2)).unbind(-1)
|
| 268 |
+
|
| 269 |
+
freqs_cos_q = reshape_for_broadcast(freqs_cos[: xq.shape[-2]], xq_r)
|
| 270 |
+
freqs_sin_q = reshape_for_broadcast(freqs_sin[: xq.shape[-2]], xq_r)
|
| 271 |
+
freqs_cos_k = reshape_for_broadcast(freqs_cos[: xk.shape[-2]], xk_r)
|
| 272 |
+
freqs_sin_k = reshape_for_broadcast(freqs_sin[: xk.shape[-2]], xk_r)
|
| 273 |
+
|
| 274 |
+
xq_out_r = xq_r * freqs_cos_q - xq_i * freqs_sin_q
|
| 275 |
+
xq_out_i = xq_r * freqs_sin_q + xq_i * freqs_cos_q
|
| 276 |
+
xk_out_r = xk_r * freqs_cos_k - xk_i * freqs_sin_k
|
| 277 |
+
xk_out_i = xk_r * freqs_sin_k + xk_i * freqs_cos_k
|
| 278 |
+
|
| 279 |
+
xq_out = torch.stack([xq_out_r, xq_out_i], dim=-1).flatten(-2)
|
| 280 |
+
xk_out = torch.stack([xk_out_r, xk_out_i], dim=-1).flatten(-2)
|
| 281 |
+
|
| 282 |
+
return xq_out.type_as(xq), xk_out.type_as(xk)
|
transformer_core.py
ADDED
|
@@ -0,0 +1,327 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Minimal self-attention core for decoder language models."""
|
| 2 |
+
|
| 3 |
+
import math
|
| 4 |
+
from typing import Optional, Tuple
|
| 5 |
+
|
| 6 |
+
import torch
|
| 7 |
+
import torch.nn as nn
|
| 8 |
+
|
| 9 |
+
from .transformer_components import PositionalInfo, apply_rotary_emb
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def _activate_scores(scores: torch.Tensor, activation: str) -> torch.Tensor:
|
| 13 |
+
if activation == "softmax":
|
| 14 |
+
return torch.softmax(scores, dim=-1)
|
| 15 |
+
if activation == "identity":
|
| 16 |
+
return scores
|
| 17 |
+
if activation == "relu":
|
| 18 |
+
return torch.relu(scores)
|
| 19 |
+
if activation == "tanh":
|
| 20 |
+
return torch.tanh(scores)
|
| 21 |
+
if activation == "sigmoid":
|
| 22 |
+
return torch.sigmoid(scores)
|
| 23 |
+
if activation == "gelu":
|
| 24 |
+
return torch.nn.functional.gelu(scores)
|
| 25 |
+
raise ValueError(f"Unsupported attention activation: {activation}")
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
class MultiHeadAttentionBase(nn.Module):
|
| 29 |
+
def __init__(
|
| 30 |
+
self,
|
| 31 |
+
query_dim: int,
|
| 32 |
+
output_dim: int,
|
| 33 |
+
key_dim: Optional[int] = None,
|
| 34 |
+
value_dim: Optional[int] = None,
|
| 35 |
+
n_heads: int = 4,
|
| 36 |
+
hidden_dim: int = 64,
|
| 37 |
+
dropout: float = 0.0,
|
| 38 |
+
total_n_heads: Optional[int] = None,
|
| 39 |
+
activation: str = "softmax",
|
| 40 |
+
use_bias_qkv: bool = False,
|
| 41 |
+
use_bias_out: bool = True,
|
| 42 |
+
):
|
| 43 |
+
super().__init__()
|
| 44 |
+
self.query_dim = query_dim
|
| 45 |
+
self.key_dim = query_dim if key_dim is None else key_dim
|
| 46 |
+
self.value_dim = self.key_dim if value_dim is None else value_dim
|
| 47 |
+
self.output_dim = output_dim
|
| 48 |
+
self.n_heads = n_heads
|
| 49 |
+
self.hidden_dim = hidden_dim
|
| 50 |
+
self.total_n_heads = n_heads if total_n_heads is None else total_n_heads
|
| 51 |
+
self.activation = activation
|
| 52 |
+
self.use_bias_qkv = use_bias_qkv
|
| 53 |
+
self.use_bias_out = use_bias_out
|
| 54 |
+
|
| 55 |
+
if self.total_n_heads <= 0:
|
| 56 |
+
raise ValueError(f"total_n_heads must be positive, got {self.total_n_heads}")
|
| 57 |
+
if hidden_dim % self.total_n_heads != 0:
|
| 58 |
+
raise ValueError(
|
| 59 |
+
f"hidden_dim ({hidden_dim}) must be divisible by total_n_heads ({self.total_n_heads})"
|
| 60 |
+
)
|
| 61 |
+
|
| 62 |
+
self.head_dim = hidden_dim // self.total_n_heads
|
| 63 |
+
self.scale = 1.0 / math.sqrt(self.head_dim)
|
| 64 |
+
projected_dim = n_heads * self.head_dim
|
| 65 |
+
self.q_proj = nn.Linear(self.query_dim, projected_dim, bias=use_bias_qkv)
|
| 66 |
+
self.k_proj = nn.Linear(self.key_dim, projected_dim, bias=use_bias_qkv)
|
| 67 |
+
self.v_proj = nn.Linear(self.value_dim, projected_dim, bias=use_bias_qkv)
|
| 68 |
+
self.o_proj = nn.Linear(projected_dim, output_dim, bias=use_bias_out)
|
| 69 |
+
self.dropout = nn.Dropout(dropout)
|
| 70 |
+
self.attn_dropout = nn.Dropout(dropout)
|
| 71 |
+
self.last_attn_weights = None
|
| 72 |
+
|
| 73 |
+
self._init_weights()
|
| 74 |
+
|
| 75 |
+
def _init_weights(self) -> None:
|
| 76 |
+
gain = 1.0 / math.sqrt(2.0)
|
| 77 |
+
for projection in (self.q_proj, self.k_proj, self.v_proj, self.o_proj):
|
| 78 |
+
nn.init.xavier_uniform_(projection.weight, gain=gain)
|
| 79 |
+
if projection.bias is not None:
|
| 80 |
+
nn.init.zeros_(projection.bias)
|
| 81 |
+
|
| 82 |
+
def _reshape_for_multihead(self, x: torch.Tensor, batch_size: int, seq_len: int) -> torch.Tensor:
|
| 83 |
+
x = x.view(batch_size, seq_len, self.n_heads, self.head_dim)
|
| 84 |
+
return x.transpose(1, 2)
|
| 85 |
+
|
| 86 |
+
def _process_mask(self, mask: Optional[torch.Tensor]) -> Optional[torch.Tensor]:
|
| 87 |
+
if mask is None:
|
| 88 |
+
return None
|
| 89 |
+
if mask.dim() == 2:
|
| 90 |
+
if mask.shape[0] != mask.shape[1]:
|
| 91 |
+
raise ValueError(
|
| 92 |
+
"2D decoder attention masks must be square [seq, seq]; "
|
| 93 |
+
f"got {tuple(mask.shape)}. Expand padding masks at the LM boundary."
|
| 94 |
+
)
|
| 95 |
+
return mask.bool().unsqueeze(0).unsqueeze(0)
|
| 96 |
+
if mask.dim() == 3:
|
| 97 |
+
return mask.bool().unsqueeze(1)
|
| 98 |
+
raise ValueError(f"Mask must be 2D or 3D, got {mask.dim()}D")
|
| 99 |
+
|
| 100 |
+
def _compute_attn_scores(
|
| 101 |
+
self,
|
| 102 |
+
q: torch.Tensor,
|
| 103 |
+
k: torch.Tensor,
|
| 104 |
+
pos_info: Optional[PositionalInfo] = None,
|
| 105 |
+
) -> torch.Tensor:
|
| 106 |
+
if pos_info is not None and pos_info.rope_freqs is not None:
|
| 107 |
+
freqs_cos, freqs_sin = pos_info.rope_freqs
|
| 108 |
+
q, k = apply_rotary_emb(q, k, freqs_cos, freqs_sin)
|
| 109 |
+
return torch.matmul(q, k.transpose(-2, -1)) * self.scale
|
| 110 |
+
|
| 111 |
+
def _apply_activation_and_mask(self, scores: torch.Tensor, mask: Optional[torch.Tensor]) -> torch.Tensor:
|
| 112 |
+
processed_mask = self._process_mask(mask)
|
| 113 |
+
if processed_mask is not None and self.activation in {"softmax", "sigmoid", "tanh"}:
|
| 114 |
+
scores.masked_fill_(~processed_mask, torch.finfo(scores.dtype).min)
|
| 115 |
+
weights = _activate_scores(scores, self.activation)
|
| 116 |
+
if processed_mask is not None:
|
| 117 |
+
weights = weights.masked_fill(~processed_mask, 0.0)
|
| 118 |
+
return weights
|
| 119 |
+
|
| 120 |
+
def _apply_mask(self, scores: torch.Tensor, mask: Optional[torch.Tensor]) -> torch.Tensor:
|
| 121 |
+
return self._apply_activation_and_mask(scores, mask)
|
| 122 |
+
|
| 123 |
+
def forward(
|
| 124 |
+
self,
|
| 125 |
+
query: torch.Tensor,
|
| 126 |
+
key: Optional[torch.Tensor] = None,
|
| 127 |
+
value: Optional[torch.Tensor] = None,
|
| 128 |
+
mask: Optional[torch.Tensor] = None,
|
| 129 |
+
pos_info: Optional[PositionalInfo] = None,
|
| 130 |
+
) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 131 |
+
batch_size, seq_len, _ = query.shape
|
| 132 |
+
key = query if key is None else key
|
| 133 |
+
value = key if value is None else value
|
| 134 |
+
key_len = key.shape[1]
|
| 135 |
+
|
| 136 |
+
if key.shape[0] != batch_size or value.shape[0] != batch_size:
|
| 137 |
+
raise ValueError(
|
| 138 |
+
f"Batch size mismatch: query={batch_size}, key={key.shape[0]}, value={value.shape[0]}"
|
| 139 |
+
)
|
| 140 |
+
if key.shape[1] != value.shape[1]:
|
| 141 |
+
raise ValueError(
|
| 142 |
+
f"Key and value sequence length mismatch: {key.shape[1]} vs {value.shape[1]}"
|
| 143 |
+
)
|
| 144 |
+
|
| 145 |
+
q = self.q_proj(query)
|
| 146 |
+
k = self.k_proj(key)
|
| 147 |
+
v = self.v_proj(value)
|
| 148 |
+
|
| 149 |
+
q = self._reshape_for_multihead(q, batch_size, seq_len)
|
| 150 |
+
k = self._reshape_for_multihead(k, batch_size, key_len)
|
| 151 |
+
v = self._reshape_for_multihead(v, batch_size, key_len)
|
| 152 |
+
|
| 153 |
+
attn_scores = self._compute_attn_scores(q, k, pos_info)
|
| 154 |
+
attn_weights = self._apply_activation_and_mask(attn_scores, mask)
|
| 155 |
+
attn_weights = self.attn_dropout(attn_weights)
|
| 156 |
+
|
| 157 |
+
attn_output = torch.matmul(attn_weights, v)
|
| 158 |
+
attn_output = attn_output.transpose(1, 2).contiguous()
|
| 159 |
+
attn_output = attn_output.view(batch_size, seq_len, self.n_heads * self.head_dim)
|
| 160 |
+
attn_output = self.o_proj(attn_output)
|
| 161 |
+
attn_output = self.dropout(attn_output)
|
| 162 |
+
|
| 163 |
+
self.last_attn_weights = attn_weights.detach()
|
| 164 |
+
return attn_output, attn_weights
|
| 165 |
+
|
| 166 |
+
|
| 167 |
+
class SelfAttention(MultiHeadAttentionBase):
|
| 168 |
+
def __init__(
|
| 169 |
+
self,
|
| 170 |
+
input_dim: int,
|
| 171 |
+
n_heads: int = 4,
|
| 172 |
+
hidden_dim: int = 64,
|
| 173 |
+
dropout: float = 0.0,
|
| 174 |
+
supports_relative: bool = False,
|
| 175 |
+
total_n_heads: Optional[int] = None,
|
| 176 |
+
use_bias_qkv: bool = False,
|
| 177 |
+
use_bias_out: bool = True,
|
| 178 |
+
):
|
| 179 |
+
head_dim = hidden_dim // (n_heads if total_n_heads is None else total_n_heads)
|
| 180 |
+
super().__init__(
|
| 181 |
+
query_dim=input_dim,
|
| 182 |
+
output_dim=n_heads * head_dim,
|
| 183 |
+
key_dim=input_dim,
|
| 184 |
+
value_dim=input_dim,
|
| 185 |
+
n_heads=n_heads,
|
| 186 |
+
hidden_dim=hidden_dim,
|
| 187 |
+
dropout=dropout,
|
| 188 |
+
total_n_heads=total_n_heads,
|
| 189 |
+
activation="softmax",
|
| 190 |
+
use_bias_qkv=use_bias_qkv,
|
| 191 |
+
use_bias_out=use_bias_out,
|
| 192 |
+
)
|
| 193 |
+
self.supports_relative = supports_relative
|
| 194 |
+
if self.supports_relative:
|
| 195 |
+
self.rel_k_proj = nn.Linear(self.head_dim, self.head_dim, bias=False)
|
| 196 |
+
self.rel_v_proj = nn.Linear(self.head_dim, self.head_dim, bias=False)
|
| 197 |
+
nn.init.xavier_uniform_(self.rel_k_proj.weight)
|
| 198 |
+
nn.init.xavier_uniform_(self.rel_v_proj.weight)
|
| 199 |
+
|
| 200 |
+
def forward(
|
| 201 |
+
self,
|
| 202 |
+
x: torch.Tensor,
|
| 203 |
+
mask: Optional[torch.Tensor] = None,
|
| 204 |
+
pos_info: Optional[PositionalInfo] = None,
|
| 205 |
+
) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 206 |
+
if self.supports_relative and pos_info is not None and pos_info.rel_embeddings is not None:
|
| 207 |
+
batch_size, seq_len, _ = x.shape
|
| 208 |
+
rel_embeddings = pos_info.rel_embeddings
|
| 209 |
+
if rel_embeddings.shape[:2] != (seq_len, seq_len):
|
| 210 |
+
raise ValueError(
|
| 211 |
+
f"Relative embeddings shape {rel_embeddings.shape} does not match sequence length {seq_len}"
|
| 212 |
+
)
|
| 213 |
+
if rel_embeddings.shape[2] != self.head_dim:
|
| 214 |
+
raise ValueError(
|
| 215 |
+
f"Relative embeddings dim {rel_embeddings.shape[2]} does not match head_dim {self.head_dim}"
|
| 216 |
+
)
|
| 217 |
+
|
| 218 |
+
rel_k = self.rel_k_proj(rel_embeddings)
|
| 219 |
+
rel_v = self.rel_v_proj(rel_embeddings)
|
| 220 |
+
|
| 221 |
+
q = self.q_proj(x)
|
| 222 |
+
k = self.k_proj(x)
|
| 223 |
+
v = self.v_proj(x)
|
| 224 |
+
|
| 225 |
+
q = self._reshape_for_multihead(q, batch_size, seq_len)
|
| 226 |
+
k = self._reshape_for_multihead(k, batch_size, seq_len)
|
| 227 |
+
v = self._reshape_for_multihead(v, batch_size, seq_len)
|
| 228 |
+
|
| 229 |
+
attn_scores = self._compute_attn_scores(q, k, pos_info)
|
| 230 |
+
rel_scores = torch.einsum("bhid,ijd->bhij", q, rel_k) * self.scale
|
| 231 |
+
attn_scores = attn_scores + rel_scores
|
| 232 |
+
|
| 233 |
+
attn_weights = self._apply_mask(attn_scores, mask)
|
| 234 |
+
attn_weights = self.attn_dropout(attn_weights)
|
| 235 |
+
|
| 236 |
+
attn_output = torch.matmul(attn_weights, v)
|
| 237 |
+
rel_output = torch.einsum("bhij,ijd->bhid", attn_weights, rel_v)
|
| 238 |
+
attn_output = attn_output + rel_output
|
| 239 |
+
|
| 240 |
+
attn_output = attn_output.transpose(1, 2).contiguous()
|
| 241 |
+
attn_output = attn_output.view(batch_size, seq_len, self.n_heads * self.head_dim)
|
| 242 |
+
attn_output = self.o_proj(attn_output)
|
| 243 |
+
attn_output = self.dropout(attn_output)
|
| 244 |
+
self.last_attn_weights = attn_weights.detach()
|
| 245 |
+
return attn_output, attn_weights
|
| 246 |
+
|
| 247 |
+
return super().forward(query=x, mask=mask, pos_info=pos_info)
|
| 248 |
+
|
| 249 |
+
|
| 250 |
+
class CrossAttention(MultiHeadAttentionBase):
|
| 251 |
+
def __init__(
|
| 252 |
+
self,
|
| 253 |
+
input_dim: int,
|
| 254 |
+
output_dim: int,
|
| 255 |
+
context_dim: Optional[int] = None,
|
| 256 |
+
n_heads: int = 8,
|
| 257 |
+
hidden_dim: int = 64,
|
| 258 |
+
dropout: float = 0.0,
|
| 259 |
+
supports_relative: bool = False,
|
| 260 |
+
use_bias_qkv: bool = False,
|
| 261 |
+
use_bias_out: bool = True,
|
| 262 |
+
):
|
| 263 |
+
resolved_context_dim = input_dim if context_dim is None else context_dim
|
| 264 |
+
super().__init__(
|
| 265 |
+
query_dim=input_dim,
|
| 266 |
+
key_dim=resolved_context_dim,
|
| 267 |
+
value_dim=resolved_context_dim,
|
| 268 |
+
output_dim=output_dim,
|
| 269 |
+
n_heads=n_heads,
|
| 270 |
+
hidden_dim=hidden_dim,
|
| 271 |
+
dropout=dropout,
|
| 272 |
+
activation="softmax",
|
| 273 |
+
use_bias_qkv=use_bias_qkv,
|
| 274 |
+
use_bias_out=use_bias_out,
|
| 275 |
+
)
|
| 276 |
+
self.supports_relative = supports_relative
|
| 277 |
+
if supports_relative:
|
| 278 |
+
self.rel_k_proj = nn.Linear(self.head_dim, self.head_dim, bias=False)
|
| 279 |
+
self.rel_v_proj = nn.Linear(self.head_dim, self.head_dim, bias=False)
|
| 280 |
+
nn.init.xavier_uniform_(self.rel_k_proj.weight)
|
| 281 |
+
nn.init.xavier_uniform_(self.rel_v_proj.weight)
|
| 282 |
+
|
| 283 |
+
def forward(
|
| 284 |
+
self,
|
| 285 |
+
inputs: torch.Tensor,
|
| 286 |
+
context: torch.Tensor,
|
| 287 |
+
mask: Optional[torch.Tensor] = None,
|
| 288 |
+
pos_info: Optional[PositionalInfo] = None,
|
| 289 |
+
) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 290 |
+
if self.supports_relative and pos_info is not None and pos_info.rel_embeddings is not None:
|
| 291 |
+
batch_size, seq_len_q, _ = inputs.shape
|
| 292 |
+
_, seq_len_k, _ = context.shape
|
| 293 |
+
rel_embeddings = pos_info.rel_embeddings
|
| 294 |
+
if rel_embeddings.shape[:2] != (seq_len_q, seq_len_k):
|
| 295 |
+
raise ValueError(
|
| 296 |
+
f"Relative embeddings shape {rel_embeddings.shape} does not match "
|
| 297 |
+
f"sequence lengths {(seq_len_q, seq_len_k)}"
|
| 298 |
+
)
|
| 299 |
+
if rel_embeddings.shape[2] != self.head_dim:
|
| 300 |
+
raise ValueError(
|
| 301 |
+
f"Relative embeddings dim {rel_embeddings.shape[2]} does not match head_dim {self.head_dim}"
|
| 302 |
+
)
|
| 303 |
+
|
| 304 |
+
rel_k = self.rel_k_proj(rel_embeddings)
|
| 305 |
+
rel_v = self.rel_v_proj(rel_embeddings)
|
| 306 |
+
q = self._reshape_for_multihead(self.q_proj(inputs), batch_size, seq_len_q)
|
| 307 |
+
k = self._reshape_for_multihead(self.k_proj(context), batch_size, seq_len_k)
|
| 308 |
+
v = self._reshape_for_multihead(self.v_proj(context), batch_size, seq_len_k)
|
| 309 |
+
attn_scores = self._compute_attn_scores(q, k, pos_info)
|
| 310 |
+
attn_scores = attn_scores + torch.einsum("bhid,ijd->bhij", q, rel_k) * self.scale
|
| 311 |
+
attn_weights = self._apply_activation_and_mask(attn_scores, mask)
|
| 312 |
+
attn_weights = self.attn_dropout(attn_weights)
|
| 313 |
+
self.last_attn_weights = attn_weights.detach()
|
| 314 |
+
attn_output = torch.matmul(attn_weights, v)
|
| 315 |
+
attn_output = attn_output + torch.einsum("bhij,ijd->bhid", attn_weights, rel_v)
|
| 316 |
+
attn_output = attn_output.transpose(1, 2).contiguous().view(batch_size, seq_len_q, -1)
|
| 317 |
+
attn_output = self.o_proj(attn_output)
|
| 318 |
+
attn_output = self.dropout(attn_output)
|
| 319 |
+
return attn_output, attn_weights
|
| 320 |
+
|
| 321 |
+
return super().forward(
|
| 322 |
+
query=inputs,
|
| 323 |
+
key=context,
|
| 324 |
+
value=context,
|
| 325 |
+
mask=mask,
|
| 326 |
+
pos_info=pos_info,
|
| 327 |
+
)
|