Stoicheia-code / tagger /backbone.py
anonymous-stoicheia's picture
Upload folder using huggingface_hub
7ed86c3 verified
Raw
History Blame Contribute Delete
5.35 kB
"""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}")