strata-headquotient-q25 / source /src /strata /data /relation_vocab.py
nur-dev's picture
Add files using upload-large-folder tool
7c5e40e verified
Raw
History Blame Contribute Delete
2.57 kB
"""Frozen contract for the shared graph-relation vocabulary.
The model's ``edge_logits`` relation axis is a single shared id space:
* ``[0, NUM_DEPREL)`` -> UD basic dependency relations (slice 1)
* ``[SRL_BASE, RELATION_VOCAB_SIZE)`` -> SRL roles, local id 0 == NONE (slice 2)
These ids are **weight-bearing**: a trained checkpoint's ``edge_type`` projection
learns them positionally, so the ordering must not change without invalidating
checkpoints. ``relation_vocab_signature()`` pins the ordering; it is recorded in
each training run's ``run_meta.json`` so a checkpoint declares the vocabulary it
was trained against. ``require_relation_capacity()`` fails loudly if a model's
relation axis is too small for a given objective (rather than silently slicing).
"""
from __future__ import annotations
import hashlib
from strata.data.srl_labels import NUM_SRL_ROLES, RELATION_VOCAB_SIZE, SRL_BASE, SRL_ROLES
from strata.data.ud_labels import NUM_DEPREL, UD_DEPRELS
# Invariant that the whole shared-vocab scheme relies on.
assert SRL_BASE == NUM_DEPREL, "SRL_BASE must equal NUM_DEPREL for a contiguous relation vocab"
def combined_relation_labels() -> list[str]:
"""Ordered label for every relation id (index == relation id)."""
return list(UD_DEPRELS) + [f"srl:{role}" for role in SRL_ROLES]
def relation_vocab_signature() -> str:
"""Stable signature of the (id -> label) ordering; changes iff ids shift."""
labels = combined_relation_labels()
payload = "\n".join(f"{i}:{label}" for i, label in enumerate(labels))
digest = hashlib.sha256(payload.encode("utf-8")).hexdigest()[:16]
return f"relv1:{len(labels)}:{digest}"
def relation_vocab_metadata() -> dict[str, object]:
"""Compact, serialisable description of the relation-vocab contract."""
return {
"signature": relation_vocab_signature(),
"num_deprel": NUM_DEPREL,
"num_srl_roles": NUM_SRL_ROLES,
"srl_base": SRL_BASE,
"relation_vocab_size": RELATION_VOCAB_SIZE,
}
def require_relation_capacity(graph_relation_types: int, *, for_srl: bool) -> None:
"""Raise if a model's relation axis is too small for the objective."""
need = RELATION_VOCAB_SIZE if for_srl else NUM_DEPREL
what = "SRL (needs UD deprels + SRL roles)" if for_srl else "UD deprels"
if graph_relation_types < need:
raise ValueError(
f"graph_relation_types={graph_relation_types} is too small for {what}: "
f"need >= {need}. This model/checkpoint is not compatible with this objective."
)