Token Classification
Transformers
Safetensors
Ancient Greek (to 1453)
char_bert_meter
ancient-greek
classical-philology
character-level
masked-diffusion
macronization
metrical-scansion
custom_code
Instructions to use Ericu950/Stoicheia-meter with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use Ericu950/Stoicheia-meter with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("token-classification", model="Ericu950/Stoicheia-meter", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("Ericu950/Stoicheia-meter", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Publish safetensors weights, config and model card
Browse files- README.md +55 -0
- config.json +23 -0
- configuration_char_bert_meter.py +50 -0
- model.safetensors +3 -0
- modeling_char_bert_meter.py +243 -0
- processing_char_bert_meter.py +299 -0
- training_metadata.json +119 -0
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 |
+
}
|