nima1's picture
Publish verified checkpoint and losslessly compressed study evidence
4be6a52 verified
Raw History Blame Contribute Delete
5.05 kB
import copy
import hashlib
import json
import pytest
from stackcraft.data import (
DatasetConfig,
audit_dataset,
canonical_json,
generate_dataset,
observation_hash,
write_dataset,
)
COMMIT = "412be98cf90bc7e3e60cd39d8bf9d479bcf65d27"
@pytest.fixture
def bundle():
return generate_dataset(
DatasetConfig(
train_seeds=(10000, 10001, 10002), validation_seeds=(20000, 20001, 20002), max_pieces=2
),
COMMIT,
)
def refresh_hash(records, manifest, split):
data = "".join(canonical_json(row) + "\n" for row in records[split])
manifest["splits"][split]["sha256"] = hashlib.sha256(data.encode()).hexdigest()
manifest["splits"][split]["records"] = len(records[split])
def test_dataset_is_reproducible_and_never_collects_test_episodes(bundle) -> None:
assert bundle == generate_dataset(DatasetConfig(**bundle.manifest["config"]), COMMIT)
assert audit_dataset(bundle.records, bundle.manifest) == {
split: len(rows) for split, rows in bundle.records.items()
}
assert bundle.manifest["test_trajectories_generated"] is False
assert bundle.manifest["seed_pools"]["reserved_test"] == list(range(30000, 30200))
assert all(
row["seed"] not in range(30000, 30200) for rows in bundle.records.values() for row in rows
)
assert set(bundle.manifest["source_hashes"]) == {
"engine.py",
"pieces.py",
"schema.py",
"players/__init__.py",
"expert.py",
"data.py",
}
def test_observation_hash_ignores_colors_not_geometry(bundle) -> None:
observation = copy.deepcopy(bundle.records["train"][1]["observation"])
colored = copy.deepcopy(observation)
colored["board"] = [[7 if cell else 0 for cell in row] for row in colored["board"]]
assert observation_hash(observation) == observation_hash(colored)
colored["board"][0][0] = 1
assert observation_hash(observation) != observation_hash(colored)
def test_collector_deduplicates_identical_starting_observations() -> None:
# At least two of 60 starts share one of the 49 possible current/preview pairs.
result = generate_dataset(
DatasetConfig(
train_seeds=tuple(range(10000, 10060)), validation_seeds=(20000,), max_pieces=1
),
COMMIT,
)
assert result.manifest["deduplication"]["exclusions"]["train"]["within_split"] > 0
hashes = [row["observation_hash"] for rows in result.records.values() for row in rows]
assert len(hashes) == len(set(hashes))
@pytest.mark.parametrize(
"field,value", [("action_id", "not-legal"), ("seed", 30000), ("source_commit", "bad")]
)
def test_audit_rejects_invalid_provenance_or_labels_even_after_rehash(bundle, field, value) -> None:
rows, manifest = copy.deepcopy(bundle.records), copy.deepcopy(bundle.manifest)
rows["train"][0][field] = value
refresh_hash(rows, manifest, "train")
with pytest.raises(ValueError):
audit_dataset(rows, manifest)
def test_audit_rejects_hidden_fields(bundle) -> None:
rows, manifest = copy.deepcopy(bundle.records), copy.deepcopy(bundle.manifest)
rows["train"][0]["observation"]["seed"] = 10000
refresh_hash(rows, manifest, "train")
with pytest.raises(ValueError, match="hidden fields"):
audit_dataset(rows, manifest)
def test_audit_detects_cross_split_identical_observation(bundle) -> None:
rows, manifest = copy.deepcopy(bundle.records), copy.deepcopy(bundle.manifest)
source = copy.deepcopy(rows["train"][0])
target = rows["validation"][0]
for field in ("observation", "observation_hash", "action_values", "action_id"):
target[field] = source[field]
refresh_hash(rows, manifest, "validation")
with pytest.raises(ValueError, match="duplicate"):
audit_dataset(rows, manifest)
def test_written_hashes_match_and_artifacts_are_not_overwritten(bundle, tmp_path) -> None:
write_dataset(bundle, tmp_path)
for split in bundle.records:
assert (
hashlib.sha256((tmp_path / f"{split}.jsonl").read_bytes()).hexdigest()
== bundle.manifest["splits"][split]["sha256"]
)
assert json.loads((tmp_path / "manifest.json").read_text()) == bundle.manifest
with pytest.raises(FileExistsError):
write_dataset(bundle, tmp_path)
@pytest.mark.parametrize(
"kwargs",
[
{"train_seeds": (0,)},
{"train_seeds": (1000,)},
{"validation_seeds": (10000,)},
{"train_seeds": (True,)},
{"max_pieces": 0},
{"behavior_cycle": ("unknown",)},
],
)
def test_dataset_pool_and_config_validation(kwargs) -> None:
with pytest.raises(ValueError):
DatasetConfig(**kwargs)
def test_working_tree_provenance_is_explicit() -> None:
result = generate_dataset(
DatasetConfig(train_seeds=(10000,), validation_seeds=(20000,), max_pieces=1),
COMMIT + "+working-tree",
)
assert result.manifest["source_commit"] == COMMIT + "+working-tree"
assert result.manifest["source_hashes"]["data.py"]