"""Portable event-level artifacts used by the parity campaign.""" from __future__ import annotations import hashlib import json from dataclasses import dataclass, field from pathlib import Path from typing import Any import numpy as np SCHEMA_VERSION = 1 @dataclass class ValidationArtifact: """Normalized representation of one extraction run. Variable-sized node and edge arrays are stored flattened with offsets. ``sample_id`` is the join key and is never inferred from row position. """ sample_id: np.ndarray labels: np.ndarray folds: np.ndarray weights: np.ndarray globals: np.ndarray node_features_flat: np.ndarray node_offsets: np.ndarray edge_src_flat: np.ndarray edge_dst_flat: np.ndarray edge_features_flat: np.ndarray edge_offsets: np.ndarray logits: np.ndarray | None = None scores: np.ndarray | None = None predictions: np.ndarray | None = None manifest: dict[str, Any] = field(default_factory=dict) @property def event_count(self) -> int: return int(self.sample_id.shape[0]) def validate(self) -> None: n = self.event_count for name in ("labels", "folds", "weights", "globals"): if getattr(self, name).shape[0] != n: raise ValueError(f"{name} does not contain one row per sample") for name in ("node_offsets", "edge_offsets"): offsets = getattr(self, name) if offsets.shape != (n + 1,) or offsets[0] != 0: raise ValueError(f"{name} must have shape ({n + 1},) and start at zero") if np.any(offsets[1:] < offsets[:-1]): raise ValueError(f"{name} must be monotonic") if self.node_offsets[-1] != len(self.node_features_flat): raise ValueError("node offsets do not describe node_features_flat") if self.edge_offsets[-1] != len(self.edge_src_flat): raise ValueError("edge offsets do not describe edge_src_flat") if len(self.edge_src_flat) != len(self.edge_dst_flat): raise ValueError("edge source and destination arrays differ in length") if self.edge_offsets[-1] != len(self.edge_features_flat): raise ValueError("edge offsets do not describe edge_features_flat") if len(np.unique(self.sample_id)) != n: raise ValueError("sample_id contains duplicates") def event_nodes(self, index: int) -> np.ndarray: start, stop = self.node_offsets[index : index + 2] return self.node_features_flat[start:stop] def event_edges(self, index: int) -> tuple[np.ndarray, np.ndarray, np.ndarray]: start, stop = self.edge_offsets[index : index + 2] return ( self.edge_src_flat[start:stop], self.edge_dst_flat[start:stop], self.edge_features_flat[start:stop], ) def _optional_arrays(artifact: ValidationArtifact) -> dict[str, np.ndarray]: return { name: value for name in ("logits", "scores", "predictions") if (value := getattr(artifact, name)) is not None } def save_artifact(artifact: ValidationArtifact, directory: str | Path) -> Path: artifact.validate() directory = Path(directory) directory.mkdir(parents=True, exist_ok=True) arrays = { "sample_id": artifact.sample_id, "labels": artifact.labels, "folds": artifact.folds, "weights": artifact.weights, "globals": artifact.globals, "node_features_flat": artifact.node_features_flat, "node_offsets": artifact.node_offsets, "edge_src_flat": artifact.edge_src_flat, "edge_dst_flat": artifact.edge_dst_flat, "edge_features_flat": artifact.edge_features_flat, "edge_offsets": artifact.edge_offsets, **_optional_arrays(artifact), } np.savez_compressed(directory / "artifact.npz", **arrays) manifest = { "schema_version": SCHEMA_VERSION, "event_count": artifact.event_count, "arrays": sorted(arrays), **artifact.manifest, } (directory / "manifest.json").write_text( json.dumps(manifest, indent=2, sort_keys=True) + "\n" ) return directory def load_artifact(directory: str | Path) -> ValidationArtifact: directory = Path(directory) manifest = json.loads((directory / "manifest.json").read_text()) if manifest.get("schema_version") != SCHEMA_VERSION: raise ValueError( f"unsupported validation artifact schema: {manifest.get('schema_version')}" ) with np.load(directory / "artifact.npz", allow_pickle=False) as data: values = {key: data[key] for key in data.files} optional = { name: values.pop(name, None) for name in ("logits", "scores", "predictions") } artifact = ValidationArtifact(**values, **optional, manifest=manifest) artifact.validate() return artifact def file_sha256(path: str | Path) -> str: digest = hashlib.sha256() with Path(path).open("rb") as stream: for chunk in iter(lambda: stream.read(1024 * 1024), b""): digest.update(chunk) return digest.hexdigest()