File size: 5,345 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
"""Bridge to the Stoicheia (Stoicheia) backbone — the ONLY module that touches it.

The pretraining repo is imported live via $STOICHEIA_ROOT (see env.sh); nothing there is
modified. At release time this shim is the single place to swap in the published
Stoicheia package or a vendored copy of model/ + data/normalize.py.
"""
from __future__ import annotations

import os
import sys

_GCB = os.path.expandvars(os.environ.get("STOICHEIA_ROOT", "$STOICHEIA_DATA"))
if _GCB not in sys.path:
    sys.path.insert(0, _GCB)

import torch

from data.normalize import (  # noqa: F401  (re-exported for the rest of the package)
    ALPHABET, DIA_STATES, N_PUNCT, Stats, normalize_record, restore_polytonic, unpack_dia,
)
from model.char_bert import CharBertConfig, CharBertEncoder

UNK_BND, UNK_DIA, UNK_PUNCT = 3, 48, 6
PAD_ID = 26


class CharBertWithHidden(CharBertEncoder):
    """CharBertEncoder whose forward also returns the final hidden state.

    forward() is a verbatim copy of the parent's with `hidden=x` added to the output —
    the parent discards x after the output heads. State dict is identical, so
    pretraining checkpoints load strictly.
    """

    def forward(self, batch):
        from model.layers import build_attn_mask, build_block_mask

        cfg = self.cfg
        ids = batch["input_ids"]
        B, T = ids.shape
        pos = torch.arange(T, device=ids.device)
        seg = batch["seg_id"]

        x = (self.e_char(ids) + self.e_bnd(batch["boundary"]) + self.e_dia(batch["dia"])
             + self.e_punct(batch["punct"]))
        # optional fine-tune-only capitalization channel (pretraining treats cap as
        # output-only); zero-init so loading a pretraining checkpoint is a no-op
        cap_emb = getattr(self, "cap_emb", None)
        if cap_emb is not None and "cap" in batch:
            x = x + cap_emb(batch["cap"])

        if cfg.attn_impl == "flex":
            char_mask = build_block_mask(seg, cfg.char_window, ids.device)
            glob_mask = build_block_mask(seg, 0, ids.device)
        else:
            char_mask = build_attn_mask(seg, cfg.char_window, ids.device, x.dtype)
            glob_mask = build_attn_mask(seg, 0, ids.device, x.dtype)

        collect = getattr(self, "return_layers", False)
        layers = []
        for blk in self.blocks:
            m = glob_mask if blk.window == 0 else char_mask
            x = blk(x, pos, m)
            if collect:
                layers.append(x)

        x = self.norm_out(x)
        return dict(
            layers=layers,
            hidden=x,
            char=self.head_char(x),
            boundary=self.head_bnd(x),
            dia=self.head_dia(x),
            cap=self.head_cap(x),
            punct=self.head_punct(x),
        )


def load_backbone(ckpt_path, device, attn_impl="sdpa"):
    """Load a Stoicheia pretraining checkpoint into a hidden-exposing encoder.

    Mirrors eval/intrinsic.py::load_model; returns (model, pretrain_cfg_dict).
    """
    # Pretraining ablation: "random:<ckpt>" builds the SAME architecture as <ckpt> but
    # leaves the weights at their random init, so capacity/tokenisation/finetune recipe are
    # identical and the metric delta is exactly what pretraining contributed. Mirrors
    # meter/backbone.py::load_backbone.
    random_init = False
    ckpt_path = str(ckpt_path)
    if ckpt_path.startswith("random:"):
        random_init = True
        ckpt_path = ckpt_path.split(":", 1)[1]
    sd = torch.load(os.path.expandvars(ckpt_path), map_location="cpu")
    c = sd["cfg"]
    mcfg = CharBertConfig(attn_impl=attn_impl, d_model=c["d_model"],
                          n_heads=c["d_model"] // 64, depth=c["depth"],
                          char_window=c["char_window"], qk_norm=c.get("qk_norm", True))
    model = CharBertWithHidden(mcfg)
    if random_init:
        print(f"RANDOM-INIT backbone (architecture from {ckpt_path}, weights NOT loaded)",
              flush=True)
    else:
        model.load_state_dict(sd["model"])
    return model.to(device), c


def load_backbone_auto(cfg, device):
    """Dispatch on cfg['backbone']: 'charbert' (default; existing torch.save checkpoint via
    load_backbone above) or 'hf' (a HuggingFace hub id or local directory, given as
    cfg['hf_model_name_or_path'], loaded through tagger.hf_backbone.load_hf_backbone).

    Used by both tagger/train.py and parser/joint_train.py so a training config can request
    either encoder family without either script special-casing the two paths beyond this call.

    Returns (encoder, pretrain_cfg_dict, tokenizer_or_None) -- tokenizer is only non-None for
    the 'hf' path; pass it straight through to TaggerDataset(..., tokenizer=tokenizer) to switch
    the dataset onto the HF subword batching path.
    """
    kind = cfg.get("backbone", "charbert")
    if kind == "charbert":
        encoder, pcfg = load_backbone(cfg["ckpt"], device, attn_impl=cfg.get("attn", "sdpa"))
        return encoder, pcfg, None
    if kind == "hf":
        from tagger.hf_backbone import load_hf_backbone
        encoder, pcfg, tokenizer = load_hf_backbone(
            cfg["hf_model_name_or_path"], device,
            add_prefix_space=cfg.get("hf_add_prefix_space", True))
        return encoder, pcfg, tokenizer
    raise ValueError(f"tagger.backbone.load_backbone_auto: unknown backbone kind {kind!r}")