nur-dev's picture
Add files using upload-large-folder tool
7c5e40e verified
Raw
History Blame Contribute Delete
2.23 kB
"""Canonical Universal Dependencies label vocabularies (basic UD, v1).
These fixed id maps make the model's graph heads supervised targets stable and
reproducible. Two label spaces:
* Node types <- UPOS (17 universal POS tags) -> ``node_type_logits``.
* Relations <- universal deprel (subtype after ':' dropped) -> arc/rel head.
A model trained with UD supervision must have ``node_type_vocab_size`` >=
``NUM_UPOS`` and ``graph_relation_types`` >= ``NUM_DEPREL``.
"""
from __future__ import annotations
# 17 universal POS tags (UD v2).
UPOS_TAGS: tuple[str, ...] = (
"ADJ", "ADP", "ADV", "AUX", "CCONJ", "DET", "INTJ", "NOUN", "NUM",
"PART", "PRON", "PROPN", "PUNCT", "SCONJ", "SYM", "VERB", "X",
)
UPOS_TO_ID: dict[str, int] = {tag: i for i, tag in enumerate(UPOS_TAGS)}
NUM_UPOS = len(UPOS_TAGS)
PUNCT_UPOS_ID = UPOS_TO_ID["PUNCT"]
# 37 universal syntactic relations (UD v2), including ``root``.
UD_DEPRELS: tuple[str, ...] = (
"acl", "advcl", "advmod", "amod", "appos", "aux", "case", "cc", "ccomp",
"clf", "compound", "conj", "cop", "csubj", "dep", "det", "discourse",
"dislocated", "expl", "fixed", "flat", "goeswith", "iobj", "list", "mark",
"nmod", "nsubj", "nummod", "obj", "obl", "orphan", "parataxis", "punct",
"reparandum", "root", "vocative", "xcomp",
)
DEPREL_TO_ID: dict[str, int] = {rel: i for i, rel in enumerate(UD_DEPRELS)}
NUM_DEPREL = len(UD_DEPRELS)
PUNCT_DEPREL_ID = DEPREL_TO_ID["punct"]
ROOT_DEPREL_ID = DEPREL_TO_ID["root"]
# Ignore index shared with torch cross-entropy for masked/unsupervised positions.
IGNORE_INDEX = -100
def universal_deprel(deprel: str) -> str:
"""Drop the language-specific subtype (``nsubj:pass`` -> ``nsubj``)."""
return deprel.split(":", 1)[0].lower()
def deprel_to_id(deprel: str) -> int:
"""Map a (possibly subtyped) deprel to a universal relation id.
Unknown relations fall back to the generic ``dep`` bucket rather than
failing, so a new treebank cannot crash training.
"""
return DEPREL_TO_ID.get(universal_deprel(deprel), DEPREL_TO_ID["dep"])
def upos_to_id(upos: str) -> int | None:
"""Map a UPOS tag to a node-type id, or ``None`` if unrecognised."""
return UPOS_TO_ID.get(upos.upper())