Ericu950 commited on
Commit
981ee48
·
verified ·
1 Parent(s): c08ac52

Publish safetensors weights, config and model card

Browse files
README.md ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ language:
4
+ - grc
5
+ library_name: transformers
6
+ tags:
7
+ - ancient-greek
8
+ - classical-philology
9
+ - character-level
10
+ - masked-diffusion
11
+ - macronization
12
+ - metrical-scansion
13
+ pipeline_tag: token-classification
14
+ ---
15
+
16
+ # Stoicheia -- macronization and metrical scansion (joint)
17
+
18
+ **Stoicheia** is a 405M-parameter character-level masked-diffusion encoder for Ancient Greek
19
+ (`d_model` 1024, depth 32, banded attention: three of every four blocks attend within a
20
+ 256-character window, the fourth globally). Its input is factored into five aligned planes --
21
+ letters, word/sentence boundaries, diacritics, capitalization, punctuation -- each of which can
22
+ be masked independently to an explicit *unknown* state at inference. That is what lets one model
23
+ read an edited text, *scriptio continua*, and a lacuna of unknown length without changing
24
+ anything but its input.
25
+
26
+ Anonymous release accompanying a paper under review.
27
+
28
+ Two per-letter heads on `Stoicheia-doc_clean`, trained jointly: **macronization** (long versus
29
+ short at ambiguous bare α/ι/υ -- Greek orthography never marks vowel length) and **scansion**
30
+ (none / heavy-end / light-end / verse-end). An optional Viterbi decoder constrains the scansion
31
+ output to valid paths through a set of metre automata.
32
+
33
+ Use this checkpoint when you want both tasks from one model; use `Stoicheia-macronizer` when you
34
+ want vowel length alone, where a dedicated head does better.
35
+
36
+ ## Usage
37
+
38
+ ```python
39
+ import torch
40
+ from transformers import AutoModel
41
+ from huggingface_hub import hf_hub_download
42
+
43
+ REPO = "Ericu950/Stoicheia-meter"
44
+ model = AutoModel.from_pretrained(REPO, trust_remote_code=True).eval()
45
+
46
+ hf_hub_download(repo_id=REPO, filename="processing_char_bert_meter.py", local_dir=".")
47
+ from processing_char_bert_meter import CharBertMeterProcessor
48
+
49
+ proc = CharBertMeterProcessor()
50
+ batch = proc("ἄνδρα μοι ἔννεπε, μοῦσα, πολύτροπον, ὃς μάλα πολλὰ")
51
+ with torch.no_grad():
52
+ out = model(**{k: v for k, v in batch.items() if not k.startswith("_")})
53
+ print(proc.decode_macronization(out, batch)) # _ long, ^ short
54
+ print(proc.decode_scansion(out, batch)) # [heavy] {light}
55
+ ```
config.json ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "n_alpha": 24,
3
+ "mask_id": 24,
4
+ "blank_id": 25,
5
+ "pad_id": 26,
6
+ "n_char_ids": 27,
7
+ "n_boundary": 4,
8
+ "n_dia": 49,
9
+ "n_punct": 7,
10
+ "d_model": 1024,
11
+ "n_heads": 16,
12
+ "depth": 32,
13
+ "char_window": 256,
14
+ "attn_impl": "sdpa",
15
+ "qk_norm": true,
16
+ "use_cap": true,
17
+ "scalar_mix": true,
18
+ "model_type": "char_bert_meter",
19
+ "auto_map": {
20
+ "AutoConfig": "configuration_char_bert_meter.CharBertMeterConfig",
21
+ "AutoModel": "modeling_char_bert_meter.CharBertMeterModel"
22
+ }
23
+ }
configuration_char_bert_meter.py ADDED
@@ -0,0 +1,50 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """HF-Hub-compatible config for Stoicheia-meter (macronization + metrical scansion).
2
+
3
+ Same backbone hyperparameters as CharBertConfig (this wraps a Stoicheia backbone
4
+ fine-tuned with two extra per-letter heads), plus the two fields that change the
5
+ model's *shape* (use_cap, scalar_mix) -- head_dropout/w_mac/w_scan/class weights are
6
+ training-only and irrelevant to inference, so they aren't part of this config.
7
+ """
8
+ from transformers import PretrainedConfig
9
+
10
+
11
+ class CharBertMeterConfig(PretrainedConfig):
12
+ model_type = "char_bert_meter"
13
+
14
+ def __init__(
15
+ self,
16
+ n_alpha: int = 24,
17
+ mask_id: int = 24,
18
+ blank_id: int = 25,
19
+ pad_id: int = 26,
20
+ n_char_ids: int = 27,
21
+ n_boundary: int = 4,
22
+ n_dia: int = 49,
23
+ n_punct: int = 7,
24
+ d_model: int = 1024,
25
+ n_heads: int = 16,
26
+ depth: int = 32,
27
+ char_window: int = 256,
28
+ attn_impl: str = "sdpa",
29
+ qk_norm: bool = True,
30
+ use_cap: bool = True,
31
+ scalar_mix: bool = True,
32
+ **kwargs,
33
+ ):
34
+ self.n_alpha = n_alpha
35
+ self.mask_id = mask_id
36
+ self.blank_id = blank_id
37
+ self.pad_id = pad_id
38
+ self.n_char_ids = n_char_ids
39
+ self.n_boundary = n_boundary
40
+ self.n_dia = n_dia
41
+ self.n_punct = n_punct
42
+ self.d_model = d_model
43
+ self.n_heads = n_heads
44
+ self.depth = depth
45
+ self.char_window = char_window
46
+ self.attn_impl = attn_impl
47
+ self.qk_norm = qk_norm
48
+ self.use_cap = use_cap
49
+ self.scalar_mix = scalar_mix
50
+ super().__init__(**kwargs)
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d17f0518e64281d3b8ae18f48922ade29aabd7581a860a3a24909f6a7149eac1
3
+ size 1620052996
modeling_char_bert_meter.py ADDED
@@ -0,0 +1,243 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """HF-Hub-compatible model for Stoicheia-meter (macronization + metrical scansion).
2
+
3
+ Self-contained: vendors the same transformer primitives as modeling_char_bert.py, plus
4
+ the fine-tune-only additions meter/model.py::MeterModel and meter/backbone.py::
5
+ CharBertWithHidden make on top of the plain backbone:
6
+ - a zero-init `cap_emb` capitalization input embedding (fine-tune-only; base
7
+ pretraining treats capitalization as output-only)
8
+ - an ELMo-style learned scalar mix over every block's output (+ the final normed
9
+ hidden state) instead of using only the last layer
10
+ - two extra per-letter heads: `head_mac` (2-way: long/short vowel quantity) and
11
+ `head_scan` (4-way: none/heavy/light/verse-final syllable weight)
12
+
13
+ The submodule layout (`self.encoder.*` for the frozen backbone, `head_mac`/
14
+ `head_scan`/`mix_w` at the top level) matches meter.model.MeterModel's real state
15
+ dict exactly -- converted checkpoints load with strict=True and no key remapping.
16
+ """
17
+ from __future__ import annotations
18
+
19
+ from dataclasses import dataclass
20
+ from typing import Optional
21
+
22
+ import torch
23
+ import torch.nn as nn
24
+ import torch.nn.functional as F
25
+ from transformers import PreTrainedModel
26
+ from transformers.modeling_outputs import ModelOutput
27
+
28
+ from .configuration_char_bert_meter import CharBertMeterConfig
29
+
30
+
31
+ class RMSNorm(nn.Module):
32
+ def __init__(self, d, eps=1e-6):
33
+ super().__init__()
34
+ self.w = nn.Parameter(torch.ones(d))
35
+ self.eps = eps
36
+
37
+ def forward(self, x):
38
+ x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
39
+ return x * self.w
40
+
41
+
42
+ class RoPE(nn.Module):
43
+ def __init__(self, dim, base=10000.0):
44
+ super().__init__()
45
+ self.dim = dim
46
+ self.base = base
47
+
48
+ def cos_sin(self, pos):
49
+ # Recomputed on every call rather than cached in a registered buffer: a
50
+ # persistent=False buffer is never covered by the checkpoint's state dict,
51
+ # so it depends entirely on __init__-time materialization -- which some
52
+ # transformers versions' meta-device/low_cpu_mem_usage loading path can
53
+ # skip, silently leaving this tensor uninitialized. Recomputing here is
54
+ # immune to that regardless of how the model was constructed/loaded.
55
+ inv = 1.0 / (self.base ** (torch.arange(0, self.dim, 2, device=pos.device).float() / self.dim))
56
+ f = torch.outer(pos.float(), inv)
57
+ emb = torch.cat([f, f], -1)
58
+ return emb.cos(), emb.sin()
59
+
60
+
61
+ def _rotate_half(x):
62
+ d = x.shape[-1] // 2
63
+ return torch.cat([-x[..., d:], x[..., :d]], -1)
64
+
65
+
66
+ def apply_rope(q, k, cos, sin):
67
+ cos = cos[None, None]
68
+ sin = sin[None, None]
69
+ return q * cos + _rotate_half(q) * sin, k * cos + _rotate_half(k) * sin
70
+
71
+
72
+ class Attention(nn.Module):
73
+ def __init__(self, d, n_heads, rope: RoPE, qk_norm=False):
74
+ super().__init__()
75
+ self.h = n_heads
76
+ self.dh = d // n_heads
77
+ self.qkv = nn.Linear(d, 3 * d, bias=False)
78
+ self.o = nn.Linear(d, d, bias=False)
79
+ self.rope = rope
80
+ self.qk_norm = qk_norm
81
+ if qk_norm:
82
+ self.q_norm = RMSNorm(self.dh)
83
+ self.k_norm = RMSNorm(self.dh)
84
+
85
+ def forward(self, x, pos, attn_mask):
86
+ B, T, D = x.shape
87
+ qkv = self.qkv(x).view(B, T, 3, self.h, self.dh).permute(2, 0, 3, 1, 4)
88
+ q, k, v = qkv[0], qkv[1], qkv[2]
89
+ if self.qk_norm:
90
+ q, k = self.q_norm(q), self.k_norm(k)
91
+ cos, sin = self.rope.cos_sin(pos)
92
+ cos, sin = cos.to(x.dtype), sin.to(x.dtype)
93
+ q, k = apply_rope(q, k, cos, sin)
94
+ out = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)
95
+ out = out.transpose(1, 2).reshape(B, T, D)
96
+ return self.o(out)
97
+
98
+
99
+ class GeGLU(nn.Module):
100
+ def __init__(self, d, mult=8 / 3):
101
+ super().__init__()
102
+ hidden = int(d * mult)
103
+ hidden = (hidden + 63) // 64 * 64
104
+ self.wi = nn.Linear(d, 2 * hidden, bias=False)
105
+ self.wo = nn.Linear(hidden, d, bias=False)
106
+
107
+ def forward(self, x):
108
+ a, b = self.wi(x).chunk(2, -1)
109
+ return self.wo(F.gelu(a) * b)
110
+
111
+
112
+ class Block(nn.Module):
113
+ def __init__(self, d, n_heads, rope, window=0, qk_norm=False):
114
+ super().__init__()
115
+ self.n1 = RMSNorm(d)
116
+ self.attn = Attention(d, n_heads, rope, qk_norm=qk_norm)
117
+ self.n2 = RMSNorm(d)
118
+ self.mlp = GeGLU(d)
119
+ self.window = window # 0 = global; >0 = local sliding window (characters)
120
+
121
+ def forward(self, x, pos, base_mask):
122
+ x = x + self.attn(self.n1(x), pos, base_mask)
123
+ x = x + self.mlp(self.n2(x))
124
+ return x
125
+
126
+
127
+ def build_attn_mask(seg_id, window, device, dtype):
128
+ """Additive mask (B,1,T,T): same-segment AND (window==0 or |i-j|<window)."""
129
+ B, T = seg_id.shape
130
+ same = seg_id[:, None, :] == seg_id[:, :, None]
131
+ if window and window > 0:
132
+ idx = torch.arange(T, device=device)
133
+ near = (idx[None, :] - idx[:, None]).abs() < window
134
+ same = same & near[None]
135
+ mask = torch.zeros(B, 1, T, T, dtype=dtype, device=device)
136
+ mask.masked_fill_(~same[:, None], float("-inf"))
137
+ return mask
138
+
139
+
140
+ class _MeterEncoder(nn.Module):
141
+ """Same submodule names/shapes as a plain CharBertEncoder (so a pretraining
142
+ backbone loads into it with no remapping), plus an optional zero-init cap_emb
143
+ and per-layer output collection for the scalar mix -- mirrors
144
+ meter.backbone.CharBertWithHidden exactly."""
145
+
146
+ def __init__(self, config: CharBertMeterConfig):
147
+ super().__init__()
148
+ self.e_char = nn.Embedding(config.n_char_ids, config.d_model)
149
+ self.e_bnd = nn.Embedding(config.n_boundary, config.d_model)
150
+ self.e_dia = nn.Embedding(config.n_dia, config.d_model)
151
+ self.e_punct = nn.Embedding(config.n_punct, config.d_model)
152
+ if config.use_cap:
153
+ self.cap_emb = nn.Embedding(2, config.d_model)
154
+ rope = RoPE(config.d_model // config.n_heads)
155
+ blocks = []
156
+ for i in range(config.depth):
157
+ win = 0 if i % 4 == 3 else config.char_window # 3 local : 1 global
158
+ blocks.append(Block(config.d_model, config.n_heads, rope, window=win, qk_norm=config.qk_norm))
159
+ self.blocks = nn.ModuleList(blocks)
160
+ self.norm_out = RMSNorm(config.d_model)
161
+ # frozen pretraining output heads: not used by the meter heads, but part of
162
+ # the backbone's real state dict (kept so a pretraining checkpoint -- or this
163
+ # converted meter checkpoint -- loads with strict=True)
164
+ self.head_char = nn.Linear(config.d_model, config.n_char_ids, bias=False)
165
+ self.head_bnd = nn.Linear(config.d_model, 3, bias=False)
166
+ self.head_dia = nn.Linear(config.d_model, 48, bias=False)
167
+ self.head_cap = nn.Linear(config.d_model, 2, bias=False)
168
+ self.head_punct = nn.Linear(config.d_model, 6, bias=False)
169
+ self.cfg = config
170
+
171
+ def forward(self, input_ids, boundary, dia, punct, cap=None, seg_id=None, collect_layers=False):
172
+ cfg = self.cfg
173
+ B, T = input_ids.shape
174
+ pos = torch.arange(T, device=input_ids.device)
175
+ seg = seg_id if seg_id is not None else torch.zeros(B, T, dtype=torch.long, device=input_ids.device)
176
+
177
+ x = self.e_char(input_ids) + self.e_bnd(boundary) + self.e_dia(dia) + self.e_punct(punct)
178
+ cap_emb = getattr(self, "cap_emb", None)
179
+ if cap_emb is not None and cap is not None:
180
+ x = x + cap_emb(cap)
181
+
182
+ attn_mask = build_attn_mask(seg, cfg.char_window, input_ids.device, x.dtype)
183
+ glob_mask = build_attn_mask(seg, 0, input_ids.device, x.dtype)
184
+
185
+ layers = []
186
+ for blk in self.blocks:
187
+ m = glob_mask if blk.window == 0 else attn_mask
188
+ x = blk(x, pos, m)
189
+ if collect_layers:
190
+ layers.append(x)
191
+ x = self.norm_out(x)
192
+ return layers, x
193
+
194
+
195
+ @dataclass
196
+ class CharBertMeterOutput(ModelOutput):
197
+ mac: torch.FloatTensor = None
198
+ scan: torch.FloatTensor = None
199
+
200
+
201
+ class CharBertMeterModel(PreTrainedModel):
202
+ config_class = CharBertMeterConfig
203
+
204
+ def __init__(self, config: CharBertMeterConfig):
205
+ super().__init__(config)
206
+ self.encoder = _MeterEncoder(config)
207
+ self.head_mac = nn.Linear(config.d_model, 2, bias=False) # 0=long, 1=short
208
+ self.head_scan = nn.Linear(config.d_model, 4, bias=False) # 0=none,1=heavy,2=light,3=verse-final
209
+ if config.scalar_mix:
210
+ self.mix_w = nn.Parameter(torch.zeros(config.depth + 1))
211
+ self.post_init()
212
+
213
+ def _init_weights(self, module):
214
+ if isinstance(module, nn.Linear):
215
+ nn.init.normal_(module.weight, std=0.02)
216
+ elif isinstance(module, nn.Embedding):
217
+ nn.init.normal_(module.weight, std=0.02)
218
+
219
+ def forward(
220
+ self,
221
+ input_ids: torch.LongTensor,
222
+ boundary: torch.LongTensor,
223
+ dia: torch.LongTensor,
224
+ punct: torch.LongTensor,
225
+ cap: Optional[torch.LongTensor] = None,
226
+ seg_id: Optional[torch.LongTensor] = None,
227
+ return_dict: bool = True,
228
+ **kwargs,
229
+ ):
230
+ collect = bool(self.config.scalar_mix)
231
+ layers, x = self.encoder(input_ids, boundary, dia, punct, cap=cap, seg_id=seg_id,
232
+ collect_layers=collect)
233
+ if self.config.scalar_mix:
234
+ h = torch.stack(layers + [x]) # (L+1, B, T, D)
235
+ mix = torch.softmax(self.mix_w, 0)
236
+ h = torch.einsum("l,lbtd->btd", mix.to(h.dtype), h)
237
+ else:
238
+ h = x
239
+ mac = self.head_mac(h)
240
+ scan = self.head_scan(h)
241
+ if not return_dict:
242
+ return (mac, scan)
243
+ return CharBertMeterOutput(mac=mac, scan=scan)
processing_char_bert_meter.py ADDED
@@ -0,0 +1,299 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """HF-Hub-compatible processor for Stoicheia-meter: text <-> the model's five
2
+ input planes (chars/boundary/dia/punct/cap -- capitalization is a real input here,
3
+ unlike the base pretrained model, where it is output-only), plus decode helpers for
4
+ the two production tasks: macronization (vowel-length marks on ambiguous dichrona)
5
+ and metrical scansion (heavy/light/verse-final syllable bracketing).
6
+
7
+ Mark conventions and the ambiguous-dichrona rule are ported verbatim from
8
+ meter/marks.py (the training-side projection code) so a converted checkpoint's
9
+ predictions decode identically to the original project's `meter.predict` CLI.
10
+ """
11
+ from __future__ import annotations
12
+
13
+ import json
14
+ import unicodedata
15
+ from dataclasses import dataclass
16
+ from pathlib import Path
17
+
18
+ import numpy as np
19
+ import torch
20
+
21
+ MASK, BLANK, PAD = 24, 25, 26
22
+ UNK_BND, UNK_DIA, UNK_PUNCT = 3, 48, 6
23
+
24
+ ALPHABET = "αβγδεζηθικλμνξοπρστυφχψω"
25
+ LETTER_IDS = {c: i for i, c in enumerate(ALPHABET)}
26
+ ID2LETTER = np.array(list(ALPHABET))
27
+
28
+ _EXTRA_BASE = {
29
+ "ς": "σ", "ϲ": "σ", "Ϲ": "σ", "ϐ": "β", "ϑ": "θ", "ϕ": "φ", "ϰ": "κ", "ϱ": "ρ", "ϖ": "π",
30
+ }
31
+ _MARK_MAP = {
32
+ 0x0301: "acute", 0x0341: "acute", 0x0300: "grave", 0x0340: "grave",
33
+ 0x0342: "circ", 0x0302: "circ", 0x0313: "smooth", 0x0343: "smooth",
34
+ 0x0314: "rough", 0x0345: "iota", 0x0308: "diaer",
35
+ }
36
+ _ACC = {"acute": 1, "grave": 2, "circ": 3}
37
+ _BR = {"smooth": 1, "rough": 2}
38
+
39
+
40
+ def _pack_dia(acc, br, iota, diaer):
41
+ return ((acc * 3 + br) * 2 + iota) * 2 + diaer
42
+
43
+
44
+ def _unpack_dia(d):
45
+ diaer = d % 2; d //= 2
46
+ iota = d % 2; d //= 2
47
+ br = d % 3; acc = d // 3
48
+ return acc, br, iota, diaer
49
+
50
+
51
+ # ---- macron/scansion label conventions (meter/marks.py, reproduced verbatim) ----
52
+ MAC_LONG, MAC_SHORT = 0, 1
53
+ SCAN_O, SCAN_HEAVY, SCAN_LIGHT, SCAN_VERSE = 0, 1, 2, 3
54
+
55
+ _A, _E, _H, _I, _O, _Y, _W = (LETTER_IDS[c] for c in "αεηιουω")
56
+ DICHRONA_IDS = np.array([_A, _I, _Y])
57
+ VOWEL_IDS = np.array([_A, _E, _H, _I, _O, _Y, _W])
58
+ DIPHTHONGS = {(_A, _I), (_A, _Y), (_E, _I), (_E, _Y), (_H, _Y),
59
+ (_O, _I), (_O, _Y), (_Y, _I), (_W, _Y)}
60
+
61
+
62
+ def ambiguous_mask(chars: np.ndarray, boundary: np.ndarray, dia: np.ndarray) -> np.ndarray:
63
+ """Which plane positions are ambiguous dichrona (the macronizer's domain)? A
64
+ position is ambiguous iff it's a base alpha/iota/upsilon without circumflex or
65
+ iota subscript, and not part of a diphthong (diaeresis on the second vowel
66
+ breaks the diphthong; pairs never span a word boundary)."""
67
+ n = len(chars)
68
+ d = np.asarray(dia, dtype=np.int64)
69
+ acc, _br, iota, diaer = _unpack_dia(d.copy())
70
+ is_dich = np.isin(chars, DICHRONA_IDS)
71
+ out = is_dich & (acc != 3) & (iota == 0)
72
+ if n > 1:
73
+ chars = np.asarray(chars)
74
+ pair = np.zeros(n - 1, dtype=bool)
75
+ for f, s in DIPHTHONGS:
76
+ pair |= (chars[:-1] == f) & (chars[1:] == s)
77
+ pair &= np.asarray(boundary[:-1]) == 0
78
+ out[1:] &= ~(pair & (diaer[1:] == 0))
79
+ out[:-1] &= ~(pair & (diaer[1:] == 0))
80
+ return out
81
+
82
+
83
+ def merge_vowelless_syllables(chars: np.ndarray, scan_labels: np.ndarray) -> np.ndarray:
84
+ """A predicted syllable span with no vowel isn't a real syllable -- it's a
85
+ boundary placed one letter early, typically at the first letter of a geminate
86
+ consonant pair (e.g. "{λε}[ν]" for what should be one closed syllable
87
+ "[λεν]"). Merge any such span into the preceding one, keeping its own weight
88
+ label: that label (usually already correct, since it's typically a closing
89
+ consonant) is normally right for the merged syllable -- only the boundary was
90
+ misplaced. A vowel-less span at the very start of the line is left as-is."""
91
+ out = np.asarray(scan_labels).copy()
92
+ is_vowel = np.isin(np.asarray(chars), VOWEL_IDS)
93
+ kept = []
94
+ start = 0
95
+ for i in range(len(out)):
96
+ if out[i] == SCAN_O:
97
+ continue
98
+ if not is_vowel[start:i + 1].any() and kept:
99
+ out[kept.pop()] = SCAN_O
100
+ kept.append(i)
101
+ start = i + 1
102
+ return out
103
+
104
+
105
+ def enforce_circumflex_heavy(dia: np.ndarray, scan_labels: np.ndarray) -> np.ndarray:
106
+ """A syllable containing a circumflexed vowel is always heavy -- a fixed rule
107
+ of Greek prosody, not something the per-letter classifier can get wrong in
108
+ principle, only in practice. Flip any SCAN_LIGHT span containing a circumflex
109
+ to SCAN_HEAVY; SCAN_VERSE is left alone (it already renders as a heavy-looking
110
+ bracket). Boundary placement itself is untouched -- this only corrects weight."""
111
+ d = np.asarray(dia, dtype=np.int64)
112
+ acc, _br, _iota, _diaer = _unpack_dia(d.copy())
113
+ has_circ = acc == 3
114
+ out = np.asarray(scan_labels).copy()
115
+ start = 0
116
+ for i in range(len(out)):
117
+ if out[i] != SCAN_O:
118
+ if out[i] == SCAN_LIGHT and has_circ[start:i + 1].any():
119
+ out[i] = SCAN_HEAVY
120
+ start = i + 1
121
+ return out
122
+
123
+
124
+ def _is_letter(ch: str) -> bool:
125
+ low = ch.lower()
126
+ if low in LETTER_IDS or low in _EXTRA_BASE:
127
+ return True
128
+ dec = unicodedata.normalize("NFD", low)
129
+ return bool(dec) and (dec[0] in LETTER_IDS or dec[0] in _EXTRA_BASE)
130
+
131
+
132
+ def insert_marks(plain: str, labels: dict) -> str:
133
+ """Write `_` (long) / `^` (short) after the letters given by
134
+ {letter_ordinal: MAC_LONG|MAC_SHORT} -- production macron output format."""
135
+ nfc = unicodedata.normalize("NFC", plain)
136
+ out, ordinal, pending = [], -1, None
137
+ for ch in nfc:
138
+ if pending is not None and not unicodedata.category(ch).startswith("M"):
139
+ out.append(pending)
140
+ pending = None
141
+ out.append(ch)
142
+ if _is_letter(ch):
143
+ ordinal += 1
144
+ if ordinal in labels:
145
+ pending = "_" if labels[ordinal] == MAC_LONG else "^"
146
+ if pending is not None:
147
+ out.append(pending)
148
+ return "".join(out)
149
+
150
+
151
+ def bracketize(plain: str, labels: dict) -> str:
152
+ """Render per-letter scan labels back into [heavy]{light} syllable spans
153
+ (verse-final span rendered as heavy, matching brevis-in-longo display)."""
154
+ nfc = unicodedata.normalize("NFC", plain)
155
+ out, cur, ordinal = [], [], -1
156
+ for ch in nfc:
157
+ cur.append(ch)
158
+ if _is_letter(ch):
159
+ ordinal += 1
160
+ lab = labels.get(ordinal, SCAN_O)
161
+ if lab != SCAN_O:
162
+ o, c = ("{", "}") if lab == SCAN_LIGHT else ("[", "]")
163
+ out.append(o + "".join(cur) + c)
164
+ cur = []
165
+ if cur:
166
+ out.append("".join(cur))
167
+ return "".join(out)
168
+
169
+
170
+ @dataclass
171
+ class _Encoded:
172
+ chars: list
173
+ boundary: list
174
+ dia: list
175
+ punct: list
176
+ cap: list
177
+
178
+
179
+ class CharBertMeterProcessor:
180
+ """`processor(text)` -> dict of batched tensors ready for `CharBertMeterModel(**batch)`."""
181
+
182
+ def __init__(self):
183
+ pass
184
+
185
+ @classmethod
186
+ def from_pretrained(cls, *_args, **_kwargs):
187
+ return cls()
188
+
189
+ def save_pretrained(self, save_directory, **_kwargs):
190
+ Path(save_directory).mkdir(parents=True, exist_ok=True)
191
+ (Path(save_directory) / "processor_config.json").write_text(
192
+ json.dumps({"processor_class": "CharBertMeterProcessor"}))
193
+
194
+ def _classify(self, text: str) -> _Encoded:
195
+ nfd = unicodedata.normalize("NFD", text)
196
+ chars, boundary, dia, punct, cap = [], [], [], [], []
197
+ acc = br = iota = diaer = 0
198
+ pending_bnd = 0
199
+ i = 0
200
+ while i < len(nfd):
201
+ ch = nfd[i]
202
+ if ch == "-":
203
+ run = 0
204
+ while i < len(nfd) and nfd[i] == "-":
205
+ run += 1
206
+ i += 1
207
+ for _ in range(run):
208
+ chars.append(MASK); boundary.append(UNK_BND)
209
+ dia.append(UNK_DIA); punct.append(UNK_PUNCT); cap.append(0)
210
+ continue
211
+ low = ch.lower()
212
+ base = low if low in LETTER_IDS else _EXTRA_BASE.get(low)
213
+ if base is not None:
214
+ if chars and pending_bnd:
215
+ boundary[-1] = pending_bnd
216
+ pending_bnd = 0
217
+ chars.append(LETTER_IDS[base])
218
+ cap.append(1 if ch != low else 0)
219
+ boundary.append(0)
220
+ dia.append(0)
221
+ punct.append(0)
222
+ acc = br = iota = diaer = 0
223
+ elif unicodedata.combining(ch) or ord(ch) in _MARK_MAP:
224
+ kind = _MARK_MAP.get(ord(ch))
225
+ if kind in _ACC:
226
+ acc = _ACC[kind]
227
+ elif kind in _BR:
228
+ br = _BR[kind]
229
+ elif kind == "iota":
230
+ iota = 1
231
+ elif kind == "diaer":
232
+ diaer = 1
233
+ if dia:
234
+ dia[-1] = _pack_dia(acc, br, iota, diaer)
235
+ elif ch.isspace():
236
+ pending_bnd = max(pending_bnd, 1)
237
+ elif ch in ".;!?":
238
+ pending_bnd = max(pending_bnd, 2)
239
+ if punct:
240
+ punct[-1] = 4 if ch == "." else 5
241
+ elif ch in ",:··":
242
+ if punct:
243
+ punct[-1] = {",": 1, "·": 2, "·": 2, ":": 3}.get(ch, 0)
244
+ i += 1
245
+ if boundary:
246
+ boundary[-1] = max(boundary[-1], 2)
247
+ return _Encoded(chars, boundary, dia, punct, cap)
248
+
249
+ def __call__(self, text: str, has_boundaries: bool = True):
250
+ """Encode `text` into model-ready tensors. Unlike the base pretraining
251
+ processor, `cap` is a real model input here (fine-tune-only channel)."""
252
+ enc = self._classify(text)
253
+ n = len(enc.chars)
254
+ chars = np.array(enc.chars, dtype=np.int64)
255
+ boundary = np.array(enc.boundary, dtype=np.int64)
256
+ dia = np.array(enc.dia, dtype=np.int64)
257
+ punct = np.array(enc.punct, dtype=np.int64)
258
+ cap = np.array(enc.cap, dtype=np.int64)
259
+ if not has_boundaries:
260
+ boundary[:] = UNK_BND
261
+
262
+ batch = dict(
263
+ input_ids=torch.from_numpy(chars)[None],
264
+ boundary=torch.from_numpy(boundary)[None],
265
+ dia=torch.from_numpy(dia)[None],
266
+ punct=torch.from_numpy(punct)[None],
267
+ cap=torch.from_numpy(cap)[None],
268
+ seg_id=torch.zeros(1, n, dtype=torch.long),
269
+ )
270
+ batch["_text"] = text # kept out-of-band for decode (marks splice into the original string)
271
+ batch["_chars"] = chars
272
+ batch["_boundary"] = boundary
273
+ batch["_dia"] = dia
274
+ return batch
275
+
276
+ # ---------------------------------------------------------------- decode
277
+
278
+ def decode_macronization(self, model_out, batch) -> str:
279
+ """Insert `_`/`^` (long/short) after every ambiguous alpha/iota/upsilon --
280
+ matches `meter.predict --macronize` exactly (only ambiguous dichrona get a
281
+ mark; unambiguous positions -- eta, omega, diphthongs, iota subscript,
282
+ circumflexed vowels -- are left bare, since their length isn't in doubt)."""
283
+ pred_mac = model_out.mac.argmax(-1)[0].numpy()
284
+ amb = ambiguous_mask(batch["_chars"], batch["_boundary"], batch["_dia"])
285
+ labels = {int(i): int(pred_mac[i]) for i in np.flatnonzero(amb)}
286
+ return insert_marks(batch["_text"], labels)
287
+
288
+ def decode_scansion(self, model_out, batch) -> str:
289
+ """Bracket every syllable the model assigns a non-trivial weight to:
290
+ [heavy], {light}, with the line-final syllable (brevis in longo) shown as
291
+ heavy -- matches `meter.predict --scan` exactly, including its two
292
+ deterministic corrections: merge_vowelless_syllables (a vowel-less
293
+ predicted span gets folded into the preceding syllable) and
294
+ enforce_circumflex_heavy (a circumflexed syllable is always heavy)."""
295
+ pred_scan = model_out.scan.argmax(-1)[0].numpy()
296
+ pred_scan = merge_vowelless_syllables(batch["_chars"], pred_scan)
297
+ pred_scan = enforce_circumflex_heavy(batch["_dia"], pred_scan)
298
+ labels = {i: int(c) for i, c in enumerate(pred_scan) if c > 0}
299
+ return bracketize(batch["_text"], labels)
training_metadata.json ADDED
@@ -0,0 +1,119 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "cfg": {
3
+ "name": "meter_joint_docclean",
4
+ "seed": 0,
5
+ "ckpt": "best.pt",
6
+ "out_dir": "meter_joint_docclean",
7
+ "encoded": "encoded",
8
+ "attn": "sdpa",
9
+ "T": 2048,
10
+ "micro_batch": 16,
11
+ "eval_micro": 16,
12
+ "lr_enc": 0.0001,
13
+ "lr_head": 0.0008,
14
+ "wd": 0.01,
15
+ "clip": 1.0,
16
+ "warmup_frac": 0.1,
17
+ "epochs": 25,
18
+ "patience": 6,
19
+ "head_dropout": 0.33,
20
+ "use_cap": true,
21
+ "scalar_mix": true,
22
+ "w_mac": 1.0,
23
+ "w_scan": 1.0,
24
+ "select": "mean",
25
+ "prose_per_epoch": 150000,
26
+ "scan_passes": 2,
27
+ "scan_max_verses": 8,
28
+ "mac_class_w": "auto",
29
+ "scan_class_w": "auto",
30
+ "log_every": 50
31
+ },
32
+ "mcfg": {
33
+ "head_dropout": 0.33,
34
+ "w_mac": 1.0,
35
+ "w_scan": 1.0,
36
+ "use_cap": true,
37
+ "scalar_mix": true,
38
+ "mac_class_w": [
39
+ 3.935332098208277,
40
+ 0.5727731680236648
41
+ ],
42
+ "scan_class_w": [
43
+ 0.4482191825355474,
44
+ 1.2170555073099973,
45
+ 1.2004924282537328,
46
+ 8.748756624191806
47
+ ]
48
+ },
49
+ "pretrain_cfg": {
50
+ "name": "gcb_doc_clean",
51
+ "seed": 0,
52
+ "d_model": 1024,
53
+ "depth": 32,
54
+ "char_window": 256,
55
+ "attn": "flex",
56
+ "compile": true,
57
+ "qk_norm": true,
58
+ "seq_len": 8192,
59
+ "micro_batch": 4,
60
+ "grad_accum": 2,
61
+ "lr": 0.0003,
62
+ "wd": 0.1,
63
+ "clip": 1.0,
64
+ "warmup_frac": 0.04,
65
+ "decay_frac": 0.4,
66
+ "tier_weights": [
67
+ 1.0,
68
+ 1.0,
69
+ 0.3
70
+ ],
71
+ "lam": 0.1,
72
+ "total_steps": 120000,
73
+ "hard_max_steps": 1000000,
74
+ "num_workers": 16,
75
+ "log_every": 20,
76
+ "ckpt_every": 2000,
77
+ "eval_every_ckpt": true,
78
+ "eval_n": 256,
79
+ "out_dir": "gcb_doc_clean",
80
+ "anneal_phases": [
81
+ [
82
+ 0.5,
83
+ [
84
+ 1.0,
85
+ 1.0,
86
+ 0.3
87
+ ]
88
+ ],
89
+ [
90
+ 0.9,
91
+ [
92
+ 1.0,
93
+ 1.0,
94
+ 0.0
95
+ ]
96
+ ],
97
+ [
98
+ 1.0,
99
+ [
100
+ 1.0,
101
+ 0.0,
102
+ 0.0
103
+ ]
104
+ ]
105
+ ]
106
+ },
107
+ "epoch": 12,
108
+ "dev": {
109
+ "mac_bal": 0.8927,
110
+ "mac_acc": 0.9195,
111
+ "mac_n": 2498,
112
+ "scan_bal": 0.9811,
113
+ "scan_acc": 0.9787,
114
+ "end_f1": 0.9791,
115
+ "scan_n": 68004
116
+ },
117
+ "T": 2048,
118
+ "_source_checkpoint": "best.pt"
119
+ }