Antreas commited on
Commit
a19c6b5
·
verified ·
1 Parent(s): 2572623

Enable AutoModel loading

Browse files
Files changed (1) hide show
  1. configuration_ogma.py +144 -0
configuration_ogma.py ADDED
@@ -0,0 +1,144 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Hugging Face AutoConfig support for Ogma models."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from enum import StrEnum
6
+ from typing import Any
7
+
8
+ from transformers import PretrainedConfig
9
+
10
+ __all__ = ["OgmaConfig", "VariantType", "PoolingType", "TaskToken"]
11
+
12
+
13
+ class VariantType(StrEnum):
14
+ """Architecture variant identifiers."""
15
+
16
+ TRANSFORMER = "transformer"
17
+ DEEP_NARROW = "deep_narrow"
18
+ CONV = "conv"
19
+ LINEAR_ATTENTION = "linear_attention"
20
+ MLP_MIXER = "mlp_mixer"
21
+ TRANSFORMER_RESA = "transformer_resa"
22
+ GLA = "gla"
23
+
24
+
25
+ class PoolingType(StrEnum):
26
+ """Pooling strategy identifiers."""
27
+
28
+ TASK_TOKEN = "task_token"
29
+ LATENT_ATTENTION = "latent_attention"
30
+ MEAN = "mean"
31
+
32
+
33
+ class TaskToken(StrEnum):
34
+ """Task token identifiers for asymmetric encoding."""
35
+
36
+ QRY = "QRY"
37
+ DOC = "DOC"
38
+ SYM = "SYM"
39
+
40
+
41
+ class OgmaConfig(PretrainedConfig):
42
+ """Configuration for Ogma embedding models."""
43
+
44
+ model_type = "ogma"
45
+
46
+ def __init__(
47
+ self,
48
+ variant: str | VariantType = VariantType.TRANSFORMER,
49
+ d_embed: int = 128,
50
+ d_model: int = 256,
51
+ n_layers: int = 1,
52
+ n_heads: int = 4,
53
+ vocab_size: int = 30_000,
54
+ max_seq_len: int = 512,
55
+ matryoshka_dims: list[int] | None = None,
56
+ pooling: str | PoolingType = PoolingType.TASK_TOKEN,
57
+ d_output: int = 256,
58
+ ffn_mult: float = 8 / 3,
59
+ conv_kernel_size: int = 7,
60
+ spatial_rank: int = 32,
61
+ n_random_features: int = 128,
62
+ dropout: float = 0.0,
63
+ scorer_type: str = "dot",
64
+ scorer_alpha_init: float = 0.1,
65
+ scorer_hidden: int = 0,
66
+ gla_expand_k: float = 0.5,
67
+ gla_expand_v: float = 1.0,
68
+ gla_gate_low_rank_dim: int = 16,
69
+ gla_gate_logit_normalizer: int = 16,
70
+ gla_use_short_conv: bool = True,
71
+ gla_conv_size: int = 4,
72
+ pad_id: int = 0,
73
+ unk_id: int = 1,
74
+ bos_id: int = 2,
75
+ eos_id: int = 3,
76
+ qry_id: int = 4,
77
+ doc_id: int = 5,
78
+ sym_id: int = 6,
79
+ n_special_tokens: int = 7,
80
+ **kwargs: Any,
81
+ ) -> None:
82
+ kwargs.setdefault("pad_token_id", pad_id)
83
+ kwargs.setdefault("bos_token_id", bos_id)
84
+ kwargs.setdefault("eos_token_id", eos_id)
85
+ super().__init__(**kwargs)
86
+ self.variant = VariantType(variant)
87
+ self.d_embed = d_embed
88
+ self.d_model = d_model
89
+ self.n_layers = n_layers
90
+ self.n_heads = n_heads
91
+ self.vocab_size = vocab_size
92
+ self.max_seq_len = max_seq_len
93
+ self.matryoshka_dims = matryoshka_dims or [32, 64, 128, 256]
94
+ self.pooling = PoolingType(pooling)
95
+ self.d_output = d_output
96
+ self.ffn_mult = ffn_mult
97
+ self.conv_kernel_size = conv_kernel_size
98
+ self.spatial_rank = spatial_rank
99
+ self.n_random_features = n_random_features
100
+ self.dropout = dropout
101
+ self.scorer_type = scorer_type
102
+ self.scorer_alpha_init = scorer_alpha_init
103
+ self.scorer_hidden = scorer_hidden
104
+ self.gla_expand_k = gla_expand_k
105
+ self.gla_expand_v = gla_expand_v
106
+ self.gla_gate_low_rank_dim = gla_gate_low_rank_dim
107
+ self.gla_gate_logit_normalizer = gla_gate_logit_normalizer
108
+ self.gla_use_short_conv = gla_use_short_conv
109
+ self.gla_conv_size = gla_conv_size
110
+ self.pad_id = pad_id
111
+ self.unk_id = unk_id
112
+ self.bos_id = bos_id
113
+ self.eos_id = eos_id
114
+ self.qry_id = qry_id
115
+ self.doc_id = doc_id
116
+ self.sym_id = sym_id
117
+ self.n_special_tokens = n_special_tokens
118
+
119
+ @property
120
+ def d_head(self) -> int:
121
+ """Per-head dimension."""
122
+ return self.d_model // self.n_heads
123
+
124
+ @property
125
+ def ffn_hidden(self) -> int:
126
+ """SwiGLU FFN hidden dimension."""
127
+ return int(self.d_model * self.ffn_mult)
128
+
129
+ def task_token_id(self, task: TaskToken | str) -> int:
130
+ """Return token ID for a task token."""
131
+ task = TaskToken(task)
132
+ return {
133
+ TaskToken.QRY: self.qry_id,
134
+ TaskToken.DOC: self.doc_id,
135
+ TaskToken.SYM: self.sym_id,
136
+ }[task]
137
+
138
+ def to_dict(self) -> dict[str, Any]:
139
+ """Serialize config to a JSON-compatible dictionary."""
140
+ output = super().to_dict()
141
+ output["variant"] = self.variant.value
142
+ output["pooling"] = self.pooling.value
143
+ return output
144
+