GNN4Colliders / validation /artifacts.py
ho22joshua's picture
rewriting codebase (#7)
916755e
Raw History Blame Contribute Delete
5.15 kB
"""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()