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)
|