Download code/tests/test_data.py from nima1/stackcraft-clef-flash-lora: direct link, hf CLI and curl.
- Browser
- Download file 5.05 kB
-
https://huggingface.co/nima1/stackcraft-clef-flash-lora/resolve/main/code/tests/test_data.py
- Command line
-
hf download hf://nima1/stackcraft-clef-flash-lora/code/tests/test_data.py
-
curl -L -o test_data.py https://huggingface.co/nima1/stackcraft-clef-flash-lora/resolve/main/code/tests/test_data.py
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" | |
| 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)) | |
| 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) | |
| 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"] | |