| """Checkpoint/model registry for completed STRATA experiments. |
| |
| The registry is intentionally metadata-only. It validates artifact selection, |
| schema compatibility, and relation-vocabulary contracts without importing torch |
| or loading model weights. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import json |
| from dataclasses import dataclass |
| from pathlib import Path |
| from typing import Any |
|
|
| from strata.data.languages import CORE_LANGUAGE_CODES, require_core_languages |
| from strata.data.relation_vocab import relation_vocab_signature, require_relation_capacity |
|
|
| DEFAULT_REGISTRY_PATH = Path("configs/models/registry.json") |
| ALLOWED_STATUS = {"canonical", "specialist", "legacy", "diagnostic", "failed"} |
| REQUIRED_CHECKPOINT_FILES = ("config.json", "model.pt", "train_state.pt") |
|
|
|
|
| @dataclass(frozen=True, slots=True) |
| class ModelRegistryEntry: |
| name: str |
| status: str |
| checkpoint_path: str |
| model_config_path: str |
| tokenizer_model: str |
| corpus: str |
| languages: tuple[str, ...] |
| relation_vocab_signature: str |
| intended_use: str |
| metrics: dict[str, Any] |
| notes: tuple[str, ...] = () |
|
|
| @classmethod |
| def from_dict(cls, data: dict[str, Any]) -> "ModelRegistryEntry": |
| return cls( |
| name=str(data["name"]), |
| status=str(data["status"]), |
| checkpoint_path=str(data["checkpoint_path"]), |
| model_config_path=str(data["model_config_path"]), |
| tokenizer_model=str(data["tokenizer_model"]), |
| corpus=str(data["corpus"]), |
| languages=tuple(data.get("languages", CORE_LANGUAGE_CODES)), |
| relation_vocab_signature=str(data["relation_vocab_signature"]), |
| intended_use=str(data["intended_use"]), |
| metrics=dict(data.get("metrics", {})), |
| notes=tuple(str(note) for note in data.get("notes", ())), |
| ) |
|
|
|
|
| @dataclass(frozen=True, slots=True) |
| class ModelRegistry: |
| version: int |
| relation_vocab_signature: str |
| tokenizer_model: str |
| languages: tuple[str, ...] |
| entries: tuple[ModelRegistryEntry, ...] |
| source_docs: tuple[str, ...] = () |
|
|
| @classmethod |
| def from_dict(cls, data: dict[str, Any]) -> "ModelRegistry": |
| return cls( |
| version=int(data["version"]), |
| relation_vocab_signature=str(data["relation_vocab_signature"]), |
| tokenizer_model=str(data["tokenizer_model"]), |
| languages=tuple(data.get("languages", CORE_LANGUAGE_CODES)), |
| source_docs=tuple(str(path) for path in data.get("source_docs", ())), |
| entries=tuple(ModelRegistryEntry.from_dict(entry) for entry in data["entries"]), |
| ) |
|
|
| def by_name(self, name: str) -> ModelRegistryEntry: |
| for entry in self.entries: |
| if entry.name == name: |
| return entry |
| raise KeyError(f"unknown model registry entry {name!r}") |
|
|
| def canonical(self) -> ModelRegistryEntry: |
| canonical = [entry for entry in self.entries if entry.status == "canonical"] |
| if len(canonical) != 1: |
| raise ValueError(f"expected exactly one canonical entry, found {len(canonical)}") |
| return canonical[0] |
|
|
|
|
| def load_model_registry(path: str | Path = DEFAULT_REGISTRY_PATH) -> ModelRegistry: |
| with Path(path).open("r", encoding="utf-8") as f: |
| return ModelRegistry.from_dict(json.load(f)) |
|
|
|
|
| def _read_json(path: Path) -> dict[str, Any]: |
| with path.open("r", encoding="utf-8") as f: |
| return json.load(f) |
|
|
|
|
| def _resolve(repo_root: Path, path: str) -> Path: |
| candidate = Path(path) |
| return candidate if candidate.is_absolute() else repo_root / candidate |
|
|
|
|
| def validate_model_registry( |
| registry: ModelRegistry, |
| *, |
| repo_root: str | Path = ".", |
| check_files: bool = False, |
| ) -> list[str]: |
| """Return validation warnings; raise ``ValueError`` for hard failures.""" |
|
|
| root = Path(repo_root) |
| require_core_languages(registry.languages) |
| if registry.relation_vocab_signature != relation_vocab_signature(): |
| raise ValueError( |
| "registry relation vocab signature does not match code: " |
| f"{registry.relation_vocab_signature} != {relation_vocab_signature()}" |
| ) |
|
|
| names: set[str] = set() |
| canonical_count = 0 |
| warnings: list[str] = [] |
| for entry in registry.entries: |
| if entry.name in names: |
| raise ValueError(f"duplicate model registry entry {entry.name!r}") |
| names.add(entry.name) |
| if entry.status not in ALLOWED_STATUS: |
| raise ValueError(f"{entry.name}: unsupported status {entry.status!r}") |
| canonical_count += int(entry.status == "canonical") |
| require_core_languages(entry.languages) |
| if entry.relation_vocab_signature != registry.relation_vocab_signature: |
| raise ValueError(f"{entry.name}: relation vocab signature mismatch") |
|
|
| model_config_path = _resolve(root, entry.model_config_path) |
| if not model_config_path.exists(): |
| raise ValueError(f"{entry.name}: missing model config {model_config_path}") |
| config = _read_json(model_config_path) |
| graph_relation_types = int(config.get("graph_relation_types", 0)) |
| node_types = int(config.get("node_type_vocab_size", 0)) |
| require_relation_capacity(graph_relation_types, for_srl=True) |
| if node_types < 17: |
| raise ValueError(f"{entry.name}: node_type_vocab_size={node_types} is too small for UD") |
|
|
| checkpoint_path = _resolve(root, entry.checkpoint_path) |
| if check_files: |
| if not checkpoint_path.is_dir(): |
| raise ValueError(f"{entry.name}: missing checkpoint directory {checkpoint_path}") |
| for filename in REQUIRED_CHECKPOINT_FILES: |
| if not (checkpoint_path / filename).exists(): |
| raise ValueError(f"{entry.name}: missing checkpoint file {checkpoint_path / filename}") |
| ckpt_config = _read_json(checkpoint_path / "config.json") |
| for key in ("vocab_size", "graph_relation_types", "node_type_vocab_size"): |
| if int(ckpt_config.get(key, -1)) != int(config.get(key, -2)): |
| raise ValueError( |
| f"{entry.name}: checkpoint/config mismatch for {key}: " |
| f"{ckpt_config.get(key)} != {config.get(key)}" |
| ) |
| tokenizer_path = _resolve(root, entry.tokenizer_model) |
| if not tokenizer_path.exists(): |
| raise ValueError(f"{entry.name}: missing tokenizer model {tokenizer_path}") |
| elif not checkpoint_path.exists(): |
| warnings.append(f"{entry.name}: checkpoint path not present on this machine: {checkpoint_path}") |
|
|
| if canonical_count != 1: |
| raise ValueError(f"expected exactly one canonical entry, found {canonical_count}") |
| return warnings |
|
|
|
|
| __all__ = [ |
| "ALLOWED_STATUS", |
| "DEFAULT_REGISTRY_PATH", |
| "ModelRegistry", |
| "ModelRegistryEntry", |
| "load_model_registry", |
| "validate_model_registry", |
| ] |
|
|