Stoicheia-code / parser /joint_evaluate.py
Ericu950's picture
Stoicheia: training and evaluation code
5952424 verified
Raw
History Blame Contribute Delete
2.65 kB
"""Load a trained JointModel (best.pt) and report ALL CoNLL-U column scores on a split:
lemma edit-script acc, UPOS acc, factored-XPOS exact-match, and UAS/LAS — the same in-house
(encodable-token) metric the joint trainer's dev eval uses.
python -m parser.joint_evaluate --run $SYN_DATA/runs/joint_f0 --split test
"""
from __future__ import annotations
import argparse, json, os
from pathlib import Path
import torch
from tagger.backbone import load_backbone_auto
from tagger.conllu import read_conllu
from tagger.dataset import TaggerDataset, pack_dev_items
from tagger.edits import LabelVocab
from tagger.model import TaggerConfig
from parser.biaffine import ParserConfig
from parser.labels import DeprelVocab
from parser.joint_model import JointModel
from parser.joint_train import evaluate_dev
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--run", required=True)
ap.add_argument("--split", default="test")
ap.add_argument("--decode", default="greedy", choices=["greedy", "mst"])
a = ap.parse_args()
run = Path(os.path.expandvars(a.run))
sd = torch.load(run / "best.pt", map_location="cpu")
cfg = sd["cfg"]
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
vocab = LabelVocab.load(run / "vocab.json")
deprel_vocab = DeprelVocab(sd["deprel_vocab"])
T, W = sd["T"], sd["W"]
encoder, _, tokenizer = load_backbone_auto(cfg, device)
tcfg = TaggerConfig(**sd["tcfg"])
pcfg = ParserConfig(**sd["pcfg"])
core = JointModel(encoder, vocab, tcfg, pcfg, W=W).to(device)
core.load_state_dict(sd["model"])
core.eval()
kdir = Path(os.path.expandvars(cfg["kfold_dir"]))
sents = list(read_conllu(kdir / f"{a.split}.conllu"))
hf_max_len = cfg.get("hf_max_len", 512)
ds = TaggerDataset(sents, vocab, T, W, tokenizer=tokenizer, hf_max_len=hf_max_len)
rows, _ = pack_dev_items(ds.encs, W, tokenizer, T)
cnt = evaluate_dev(core, core, rows, sents, deprel_vocab, cfg.get("eval_micro", 8),
device, T, W, a.decode, tokenizer=tokenizer)
nw = max(int(cnt[3]), 1); na = max(int(cnt[6]), 1)
res = dict(split=a.split, decode=a.decode, n_words=nw, n_arc=na,
xpos_exact=round(int(cnt[0]) / nw, 4), lemma_script=round(int(cnt[1]) / nw, 4),
upos=round(int(cnt[2]) / nw, 4), uas=round(int(cnt[4]) / na, 4),
las=round(int(cnt[5]) / na, 4), dev_best=sd.get("dev"))
print("JOINT TEST " + json.dumps(res), flush=True)
with open(run / f"{a.split}_scores_{a.decode}.json", "w") as f:
json.dump(res, f, indent=2)
if __name__ == "__main__":
main()