Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 7 additions & 4 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -85,10 +85,13 @@ before decoding, and scores and parquet output are rebuilt through the normal
recorder path. A new run ID starts a new evaluation.

The checkpoint freezes sample selection (including the robustness pass), source
content hashes, model/evaluator identity, seed (42 when omitted), and scoring
settings. Changed inputs fail closed. Preparing a new run reads the selected files
to fingerprint them; resumed runs validate only pending samples. Model setup and
uncommitted work may repeat after interruption.
paths and file-generation metadata (size, nanosecond modification/change times),
model/evaluator content hashes, seed (42 when omitted), and scoring settings.
Startup does not read media or augmentation payloads. Pending inputs are checked
against their frozen metadata before processing; changed inputs fail closed.
Input storage must be trusted and preserve these timestamps across mounts, as
Modal Volumes do. Metadata checks do not protect against a storage owner who can
forge file timestamps. Model setup and uncommitted work may repeat after interruption.

Checkpoint storage must survive the process or container being replaced. The
Python API accepts `checkpoint_persist(directory)` for filesystems requiring an
Expand Down
3 changes: 2 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"

[project]
name = "gasbench"
version = "0.10.0"
version = "0.10.1"
description = "GASBench - ML model benchmark evaluation package"
readme = "README.md"
requires-python = ">=3.10"
Expand Down Expand Up @@ -63,6 +63,7 @@ where = ["src"]
"" = "src"

[tool.setuptools.package-data]
"gasbench" = ["checkpoint_compatibility.json"]
"gasbench.dataset" = ["configs/*.yaml"]

[tool.black]
Expand Down
95 changes: 73 additions & 22 deletions src/gasbench/benchmarks/_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,8 +14,13 @@

import hashlib
import json
import logging
import os
import tempfile
import time
from collections import deque
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from copy import deepcopy
from pathlib import Path
from typing import Callable, Iterable, Mapping, Optional
Expand All @@ -42,6 +47,44 @@ def _read(path: Path):
raise CheckpointError(f"Cannot read checkpoint {path.name}") from exc


# Limit both active I/O and queued decoded batches on remote filesystems.
_READ_WORKERS = 16
_READ_AHEAD = 32


@contextmanager
def _prefetch_batches(paths):
"""Read ahead in parallel, but expose batches in commit order.

Keep only a bounded window of futures alive. The context owns the executor
so validation failure also cancels queued reads and joins active readers.
"""
if not paths:
yield iter(())
return
executor = ThreadPoolExecutor(max_workers=min(_READ_WORKERS, len(paths)))
pending = deque()
remaining = iter(paths)

def submit_next():
path = next(remaining, None)
if path is not None:
pending.append((path, executor.submit(_read, path)))

def ordered():
while pending:
path, future = pending.popleft()
yield path, future.result()
submit_next()

try:
for _ in range(min(_READ_AHEAD, len(paths))):
submit_next()
yield ordered()
finally:
executor.shutdown(wait=True, cancel_futures=True)


def prediction_key(row: Mapping) -> tuple:
"""Identity of one planned prediction, using the existing recorder schema."""
fields = ("run_id", "dataset_name", "sample_id")
Expand Down Expand Up @@ -149,35 +192,43 @@ def __init__(
raise CheckpointError("Invalid or incomplete checkpoint commit head")
# Files newer than the head were never acknowledged. They can be safely
# replaced when the interrupted batch is replayed.
for path in batches[: head["last_batch"] + 1]:
if path.name != self._batch_name(self._next_batch):
committed = batches[: head["last_batch"] + 1]
for index, path in enumerate(committed):
if path.name != self._batch_name(index):
raise CheckpointError("Checkpoint batch sequence is incomplete")
envelope = _read(path)
if not isinstance(envelope, dict):
raise CheckpointError(f"Invalid checkpoint batch {path.name}")
payload = envelope.get("payload")
try:
valid = (
isinstance(payload, dict)
and envelope.get("sha256") == _digest(payload)
and payload.get("manifest_sha256") == self._manifest_digest
and payload.get("batch_index") == self._next_batch
)
except (TypeError, ValueError):
valid = False
if not valid:
raise CheckpointError(f"Invalid checkpoint batch {path.name}")
records = self._validate_records(payload.get("records"))
if any(prediction_key(row) in self._records for row in records):
raise CheckpointError("Duplicate work in committed checkpoint batches")
self._records.update((prediction_key(row), row) for row in records)
self._next_batch += 1
started = time.monotonic()
with _prefetch_batches(committed) as loaded:
for path, envelope in loaded:
if not isinstance(envelope, dict):
raise CheckpointError(f"Invalid checkpoint batch {path.name}")
payload = envelope.get("payload")
try:
valid = (
isinstance(payload, dict)
and envelope.get("sha256") == _digest(payload)
and payload.get("manifest_sha256") == self._manifest_digest
and payload.get("batch_index") == self._next_batch
)
except (TypeError, ValueError):
valid = False
if not valid:
raise CheckpointError(f"Invalid checkpoint batch {path.name}")
records = self._validate_records(payload.get("records"))
if any(prediction_key(row) in self._records for row in records):
raise CheckpointError("Duplicate work in committed checkpoint batches")
self._records.update((prediction_key(row), row) for row in records)
self._next_batch += 1
if (
batches
and head["last_batch"] >= 0
and envelope["sha256"] != head.get("sha256")
):
raise CheckpointError("Checkpoint tail differs from commit head")
if committed:
logging.getLogger(__name__).info(
"Validated %s checkpoint batches (%s rows) in %.2fs",
len(committed), len(self._records), time.monotonic() - started,
)

@staticmethod
def _batch_name(index: int) -> str:
Expand Down
77 changes: 72 additions & 5 deletions src/gasbench/benchmarks/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import hashlib
import uuid
import platform
import stat
from importlib.metadata import version, PackageNotFoundError
from pathlib import Path
from typing import Callable
Expand Down Expand Up @@ -274,6 +275,20 @@ def fingerprint_files(paths, root):
return digest.hexdigest()


def compatible_evaluator_identity(evaluator_dir, current, previous):
"""Allow only explicitly audited, exact-build checkpoint reader upgrades."""
if current == previous:
return True
registry = Path(evaluator_dir) / "checkpoint_compatibility.json"
if not registry.is_file():
return False
try:
pairs = json.loads(registry.read_text())
except (OSError, ValueError):
return False
return isinstance(pairs, dict) and isinstance(pairs.get(current), list) and previous in pairs[current]


def runtime_versions():
versions = {
"python": platform.python_version(),
Expand Down Expand Up @@ -315,9 +330,36 @@ def sample_digest(sample):
return fingerprint_files(paths, paths[0].parent)


def file_metadata_digest(paths):
"""Identify cached file generations without reading media payloads.

Inputs must be on a trusted filesystem that preserves modification/change
timestamps across mounts. This detects replacement and in-place edits,
including same-size edits with restored mtime (ctime still changes).
It is not a content hash or protection against a malicious storage owner.
Sandbox inputs are read-only; checkpoint batch checksums remain unchanged.
"""
paths = sorted(paths)
if not paths:
raise CheckpointError("A selected sample is missing from the cache")
entries = []
for path in paths:
try:
info = path.stat()
except OSError as exc:
raise CheckpointError("Cannot stat a selected input") from exc
if not stat.S_ISREG(info.st_mode):
raise CheckpointError("A selected input is not a regular file")
entries.append((str(path), info.st_size, info.st_mtime_ns, info.st_ctime_ns))
return hashlib.sha256(json.dumps(entries, separators=(",", ":")).encode()).hexdigest()


def verify_sample(sample):
"""Reject changed pending input before preprocessing; completed rows need no I/O."""
if "content_sha256" in sample:
if "file_metadata_sha256" in sample:
if file_metadata_digest(sample_files(sample)) != sample["file_metadata_sha256"]:
raise CheckpointError("A selected sample changed since this run started")
elif "content_sha256" in sample:
try:
if sample_digest(sample) != sample["content_sha256"]:
raise CheckpointError(
Expand All @@ -330,6 +372,16 @@ def verify_sample(sample):
def use_augmentation_cache(sample, path):
"""Use only the derived artifact selected for this run, or regenerate it."""
path = Path(path)
if "augmentation_cache_metadata_sha256" in sample:
expected = sample["augmentation_cache_metadata_sha256"]
if expected is None:
return False
try:
if file_metadata_digest([path]) != expected:
raise CheckpointError("Selected augmentation cache changed or disappeared")
except CheckpointError as exc:
raise CheckpointError("Selected augmentation cache changed or disappeared") from exc
return True
if "augmentation_cache_sha256" not in sample:
return path.is_file()
expected = sample["augmentation_cache_sha256"]
Expand Down Expand Up @@ -399,6 +451,7 @@ def create_tracker(
)
}
context = {
"input_identity": "file-metadata-v1",
"settings": settings,
"seed": seed,
"runtime": runtime_versions(),
Expand All @@ -422,6 +475,12 @@ def create_tracker(
context = json.loads(json.dumps(context))
if saved is not None:
previous = saved.get("benchmark", {})
if compatible_evaluator_identity(
evaluator_dir, context["evaluator_sha256"], previous.get("evaluator_sha256")
):
# Preserve the original manifest identity for both restored and new batches.
# Every other context field is still compared below without exceptions.
context["evaluator_sha256"] = previous["evaluator_sha256"]
if {k: v for k, v in previous.items() if k != "samples"} != context:
raise CheckpointError(
"Checkpoint configuration or model differs from this run"
Expand All @@ -430,7 +489,10 @@ def create_tracker(
raise CheckpointError("Checkpoint has no frozen sample selection")
plan.samples = previous["samples"]
else:
for dataset in plan.available_datasets:
logger = get_logger(__name__)
started = time.monotonic()
logger.info("Freezing sample selection using file metadata (no media payload scan)")
for index, dataset in enumerate(plan.available_datasets, 1):
selections = {}
for pass_name, cap in (
("base", plan.sampling_plan[dataset.name]),
Expand All @@ -455,7 +517,7 @@ def create_tracker(
)
selections[pass_name] = list(iterator)
for sample in selections[pass_name]:
sample["content_sha256"] = sample_digest(sample)
sample["file_metadata_sha256"] = file_metadata_digest(sample_files(sample))
if pass_name == "aug" and config.aug_cache_dir:
from .aug_cache import img_aug_cache_path, vid_aug_cache_path

Expand All @@ -473,12 +535,17 @@ def create_tracker(
)
# Newly generated cache files are outputs of this attempt;
# only artifacts present in the frozen plan may be inputs.
sample["augmentation_cache_sha256"] = (
fingerprint_files([path], path.parent)
sample["augmentation_cache_metadata_sha256"] = (
file_metadata_digest([path])
if path.is_file()
else None
)
plan.samples[dataset.name] = selections
if index == 1 or index % 10 == 0 or index == len(plan.available_datasets):
logger.info(
"Frozen dataset %s/%s (%s) in %.1fs",
index, len(plan.available_datasets), dataset.name, time.monotonic() - started,
)
context["samples"] = plan.samples
tracker = BenchmarkRunRecorder(
run_id=config.run_id,
Expand Down
5 changes: 5 additions & 0 deletions src/gasbench/checkpoint_compatibility.json
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
{
"cefcbe4740db068e9e34d3f4184cbde9d65b441482bc2ecdbcfeb98584634288": [
"ec5643b53236c144c9ed296f653a1dd39b852ee37ce4feaefa2179650b1ad2cb"
]
}
4 changes: 2 additions & 2 deletions src/gasbench/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,15 +4,15 @@
#
# Image is 3-class (no rendered): 0=real, 1=synthetic, 2=semisynthetic.
# Video is 3-class: 0=real, 1=synthetic, 2=semisynthetic.
# Classical rendered video is real; rendering provenance lives in dataset metadata.
# Non-AI rendered images and video are real; rendering provenance lives in metadata.
# Audio stays binary: 0=real, 1=synthetic (semisynthetic collapsed).
#
# Semisynthetic media retains materially captured visual content alongside
# spatially localized generated or replaced visual content. Fully synthesized
# output remains synthetic, even when captured media conditions its generation;
# modifying exclusively synthetic or rendered media does not make it semisynthetic.
#
# Image has no rendered class: CGI/game-engine stills are excluded from image configs.
# Image has no rendered class: non-AI CGI, charts, and rendered text map to real.
IMAGE_MEDIA_TYPE_TO_LABEL = {
"real": 0,
"synthetic": 1,
Expand Down
Loading
Loading