File size: 13,777 Bytes
7ed86c3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
"""Treebank -> model batches.

Each syntactic word's FORM is encoded independently through Stoicheia's
normalize_record (guaranteeing exact word<->char-span alignment), sentences are the
concatenation of their encodable words, and whole sentences are greedily packed into
fixed-length rows with per-sentence seg_ids (block-diagonal attention, exactly like
pretraining's document packing). All input planes carry their true values — chars,
boundary (word/sentence ends), dia, punct — since all of them are known from raw text
at inference time.
"""
from __future__ import annotations

from dataclasses import dataclass, field

import numpy as np
import torch

from tagger.backbone import Stats, normalize_record
from tagger.edits import compute_script, form_key

# punctuation class LUT from the pretraining normalizer (comma/high-dot/colon/period/question)
from data.normalize import _PUNCT as PUNCT_LUT  # noqa: E402

PAD_ID = 26


def encode_word(form: str):
    """(chars, dia, cap) uint8 arrays for one FORM, or None if it has no Greek letters."""
    r = normalize_record(form, Stats(), with_punct=True)
    if r is None:
        return None
    chars, _boundary, dia, cap, _punct = r
    return chars, dia, cap


def punct_class(form: str) -> int:
    """Punctuation class a non-Greek token contributes to the preceding word."""
    return max((int(PUNCT_LUT[ord(c)]) for c in form if ord(c) < len(PUNCT_LUT)), default=0)


@dataclass
class SentEnc:
    chars: np.ndarray
    boundary: np.ndarray
    dia: np.ndarray
    punct: np.ndarray
    cap: np.ndarray
    spans: list          # per token: (start, end) char span or None (unencodable)
    y_xpos: np.ndarray   # (n_enc_words, 9) int64, -100 = unseen-in-train
    y_script: np.ndarray # (n_enc_words,)
    y_upos: np.ndarray   # (n_enc_words,)
    y_tag: np.ndarray    # (n_enc_words,) full-XPOS-tag id

    def __len__(self):
        return len(self.chars)


def encode_sentence(sent, vocab=None) -> SentEnc | None:
    """vocab=None -> encode inputs only (labels filled with -100)."""
    parts, spans = [], []
    n = 0
    for t in sent.tokens:
        enc = encode_word(t.form)
        if enc is None:
            spans.append(None)
            # non-Greek token: contribute its punctuation class to the previous word
            if parts:
                pc = punct_class(t.form)
                if pc:
                    parts[-1]["punct"][-1] = max(parts[-1]["punct"][-1], pc)
            continue
        chars, dia, cap = enc
        parts.append(dict(chars=chars, dia=dia, cap=cap,
                          boundary=np.zeros(len(chars), dtype=np.uint8),
                          punct=np.zeros(len(chars), dtype=np.uint8), tok=t))
        parts[-1]["boundary"][-1] = 1
        spans.append((n, n + len(chars)))
        n += len(chars)
    if not parts:
        return None
    parts[-1]["boundary"][-1] = 2   # sentence end

    labs = np.full((len(parts), 12), -100, dtype=np.int64)
    if vocab is not None:
        for i, p in enumerate(parts):
            t = p["tok"]
            labs[i, :9] = vocab.xpos_ids(t.xpos)
            labs[i, 9] = vocab.script_id(compute_script(form_key(t.form), t.lemma))
            labs[i, 10] = vocab.upos_id(t.upos)
            labs[i, 11] = vocab.tag_id(t.xpos)

    return SentEnc(
        chars=np.concatenate([p["chars"] for p in parts]),
        boundary=np.concatenate([p["boundary"] for p in parts]),
        dia=np.concatenate([p["dia"] for p in parts]),
        punct=np.concatenate([p["punct"] for p in parts]),
        cap=np.concatenate([p["cap"] for p in parts]),
        spans=spans,
        y_xpos=labs[:, :9], y_script=labs[:, 9], y_upos=labs[:, 10], y_tag=labs[:, 11],
    )


@dataclass
class Row:
    """One packed model row plus everything needed to map predictions back."""
    sents: list = field(default_factory=list)   # (sent_index, SentEnc)


def pack_rows(encs, T=2048, W=384, order=None):
    """Greedy packing of whole sentences (in `order`) into rows of <=T chars, <=W words.
    Oversize sentences are truncated to T at a word boundary (span-less tail words fall
    back to the lexicon rule at decode time); truncation count is returned for logging."""
    order = range(len(encs)) if order is None else order
    rows, truncated = [], 0
    cur, cur_c, cur_w = Row(), 0, 0
    for si in order:
        e = encs[si]
        if e is None:
            continue
        nc, nw = len(e), len(e.y_script)
        if nc > T or nw > W:
            truncated += 1
            continue   # pathological; handled by rule fallback at decode time
        if cur_c + nc > T or cur_w + nw > W:
            rows.append(cur)
            cur, cur_c, cur_w = Row(), 0, 0
        cur.sents.append((si, e))
        cur_c += nc
        cur_w += nw
    if cur.sents:
        rows.append(cur)
    return rows, truncated


def batch_rows(rows, T=2048, W=384, device=None):
    """Stack a list of Rows into model tensors + label tensors + slot metadata.

    Returns dict with input_ids/boundary/dia/punct/seg_id (B,T), word_id (B,T) in
    [-1,W), y_xpos (B,W,9), y_script (B,W), y_upos (B,W), and slots: per row, a list
    of (sent_index, token_index) per word slot (for mapping predictions back).
    """
    B = len(rows)
    ids = np.full((B, T), PAD_ID, dtype=np.int64)
    bnd = np.zeros((B, T), dtype=np.int64)
    dia = np.zeros((B, T), dtype=np.int64)
    pct = np.zeros((B, T), dtype=np.int64)
    cp = np.zeros((B, T), dtype=np.int64)
    seg = np.zeros((B, T), dtype=np.int64)
    wid = np.full((B, T), -1, dtype=np.int64)
    y = np.full((B, W, 12), -100, dtype=np.int64)
    slots = []
    for b, row in enumerate(rows):
        c = w = 0
        rs = []
        for k, (si, e) in enumerate(row.sents):
            n = len(e)
            ids[b, c:c + n] = e.chars
            bnd[b, c:c + n] = e.boundary
            dia[b, c:c + n] = e.dia
            pct[b, c:c + n] = e.punct
            cp[b, c:c + n] = e.cap
            seg[b, c:c + n] = k + 1
            j = 0
            for ti, span in enumerate(e.spans):
                if span is None:
                    continue
                s0, s1 = span
                wid[b, c + s0:c + s1] = w
                y[b, w, :9] = e.y_xpos[j]
                y[b, w, 9] = e.y_script[j]
                y[b, w, 10] = e.y_upos[j]
                y[b, w, 11] = e.y_tag[j]
                rs.append((si, ti))
                w += 1
                j += 1
            c += n
        slots.append(rs)
    t = lambda a: torch.from_numpy(a) if device is None else torch.from_numpy(a).to(device)
    return dict(input_ids=t(ids), boundary=t(bnd), dia=t(dia), punct=t(pct), cap=t(cp),
                seg_id=t(seg),
                word_id=t(wid), y_xpos=t(y[:, :, :9]), y_script=t(y[:, :, 9]),
                y_upos=t(y[:, :, 10]), y_tag=t(y[:, :, 11]), slots=slots)


@dataclass
class HFSentEnc:
    """One sentence's HF subword encoding: real tokenizer ids for the WHOLE sentence text
    (Greek and non-Greek tokens alike -- a subword LM was pretrained on running text and should
    see punctuation etc. as context), plus a word_id-style alignment and the same label arrays
    encode_sentence produces, in the same order (only "encodable" = has-Greek-letters tokens,
    per encode_word, get a pooled word slot / a label row -- exactly the CharBERT convention, so
    XPOS/script/UPOS/lemma-edit-script targets and parser.model.build_gold's gold-arc indexing
    line up 1:1 across both backbones)."""
    input_ids: list
    word_id: list           # length == len(input_ids); slot in [0, n_enc) or -1 (incl. specials
                             # and non-Greek tokens, which get real subwords but no word slot)
    enc_orig_idx: list       # original sent.tokens index for each of the n_enc word slots, in
                             # slot order -- mirrors the char path's (sent_index, token_index)
                             # bookkeeping in `slots` for build_gold / JointModel._regroup
    y_xpos: np.ndarray
    y_script: np.ndarray
    y_upos: np.ndarray
    y_tag: np.ndarray

    def __len__(self):
        return len(self.input_ids)


def encode_sentence_hf(sent, tokenizer, vocab=None, max_len=512):
    """HF subword tokenization + word alignment for one sentence, or None if it has no
    encodable (Greek) tokens, or if the untruncated sequence exceeds max_len subword positions
    (dropped whole, like pack_rows' oversize-sentence rule for the char path -- no partial/
    misaligned sentences)."""
    words = [t.form for t in sent.tokens]
    if not words:
        return None
    enc_idx = [i for i, t in enumerate(sent.tokens) if encode_word(t.form) is not None]
    if not enc_idx:
        return None
    slot_of = {orig: k for k, orig in enumerate(enc_idx)}

    labs = np.full((len(enc_idx), 12), -100, dtype=np.int64)
    if vocab is not None:
        for k, i in enumerate(enc_idx):
            t = sent.tokens[i]
            labs[k, :9] = vocab.xpos_ids(t.xpos)
            labs[k, 9] = vocab.script_id(compute_script(form_key(t.form), t.lemma))
            labs[k, 10] = vocab.upos_id(t.upos)
            labs[k, 11] = vocab.tag_id(t.xpos)

    tok_out = tokenizer(words, is_split_into_words=True)
    ids = tok_out["input_ids"]
    if len(ids) > max_len:
        return None
    wraw = tok_out.word_ids()
    wid = [slot_of.get(w, -1) if w is not None else -1 for w in wraw]

    return HFSentEnc(input_ids=ids, word_id=wid, enc_orig_idx=enc_idx,
                      y_xpos=labs[:, :9], y_script=labs[:, 9], y_upos=labs[:, 10],
                      y_tag=labs[:, 11])


def batch_sentences_hf(items, tokenizer, W=384, device=None):
    """items: list of (sent_index, HFSentEnc), one row per sentence -- ordinary padded batching
    (attention_mask) stands in for the char pipeline's block-diagonal packing, which existed
    only to make CharBERT's char-window local attention cheap; a standard HF encoder attends
    over the whole (padded) sentence and needs no such trick.

    Returns dict with input_ids/attention_mask (B,Tmax), word_id (B,Tmax) in [-1,W), y_xpos
    (B,W,9), y_script/y_upos/y_tag (B,W), and slots: per row, a list of (sent_index,
    token_index) per word slot -- same shape/semantics as batch_rows' `slots`.
    """
    B = len(items)
    Tmax = max(len(e) for _, e in items)
    pad_id = tokenizer.pad_token_id
    if pad_id is None:
        pad_id = tokenizer.eos_token_id if tokenizer.eos_token_id is not None else 0
    ids = np.full((B, Tmax), pad_id, dtype=np.int64)
    attn = np.zeros((B, Tmax), dtype=np.int64)
    wid = np.full((B, Tmax), -1, dtype=np.int64)
    y = np.full((B, W, 12), -100, dtype=np.int64)
    slots = []
    for b, (si, e) in enumerate(items):
        n = len(e)
        ids[b, :n] = e.input_ids
        attn[b, :n] = 1
        wid[b, :n] = e.word_id
        n_enc = e.y_xpos.shape[0]
        y[b, :n_enc, :9] = e.y_xpos
        y[b, :n_enc, 9] = e.y_script
        y[b, :n_enc, 10] = e.y_upos
        y[b, :n_enc, 11] = e.y_tag
        slots.append([(si, ti) for ti in e.enc_orig_idx])
    t = lambda a: torch.from_numpy(a) if device is None else torch.from_numpy(a).to(device)
    return dict(input_ids=t(ids), attention_mask=t(attn), word_id=t(wid),
                y_xpos=t(y[:, :, :9]), y_script=t(y[:, :, 9]), y_upos=t(y[:, :, 10]),
                y_tag=t(y[:, :, 11]), slots=slots)


def pack_dev_items(encs, W, tokenizer=None, T=2048, order=None):
    """Row/item list for evaluation (unsharded; caller shards across ranks) or for one
    training epoch's shuffled pass. CharBERT path -> pack_rows' packed Rows (T-limited,
    block-diagonal); HF path -> a flat (sent_index, HFSentEnc) list, one row per sentence.
    Returns (rows_or_items, truncated_count)."""
    if tokenizer is None:
        return pack_rows(encs, T, W, order)
    order = range(len(encs)) if order is None else order
    items = [(i, encs[i]) for i in order if encs[i] is not None]
    return items, 0


def batch_chunk(chunk, T, W, tokenizer=None, device=None):
    """Stack a chunk of pack_dev_items' output into model tensors; dispatches on backbone kind
    exactly like pack_dev_items does."""
    if tokenizer is None:
        return batch_rows(chunk, T, W, device=device)
    return batch_sentences_hf(chunk, tokenizer, W, device=device)


class TaggerDataset:
    """Encodes a .conllu once; repacks (shuffled) per epoch.

    tokenizer=None (default) -> CharBERT char-plane pipeline (encode_sentence / pack_rows /
    batch_rows), unchanged. tokenizer=<a HF fast tokenizer> -> HF subword pipeline
    (encode_sentence_hf + batch_sentences_hf, one sentence per row, no T-limited packing)."""

    def __init__(self, sentences, vocab, T=2048, W=384, tokenizer=None, hf_max_len=512):
        self.T, self.W = T, W
        self.sentences = sentences
        self.tokenizer = tokenizer
        if tokenizer is None:
            self.encs = [encode_sentence(s, vocab) for s in sentences]
        else:
            encs = [encode_sentence_hf(s, tokenizer, vocab, max_len=hf_max_len)
                    for s in sentences]
            # mirror pack_rows' oversize-sentence rule: drop (not truncate) sentences whose
            # encodable-word count can't fit a row of width W
            self.encs = [e if (e is not None and e.y_xpos.shape[0] <= W) else None for e in encs]
        self.n_enc = sum(e is not None for e in self.encs)

    def batches(self, micro, seed=None, shuffle=True):
        order = np.arange(len(self.encs))
        if shuffle:
            np.random.default_rng(seed).shuffle(order)
        rows, _ = pack_dev_items(self.encs, self.W, self.tokenizer, self.T, order)
        for i in range(0, len(rows), micro):
            yield batch_chunk(rows[i:i + micro], self.T, self.W, self.tokenizer)