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
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
import shutil

from ..storage import DiskStore
from ._identity import canonical_json, preprocessor_config, validate_identity, tree_config
from ._identity import canonical_json, compatibility_config, preprocessor_config, validate_identity, tree_config
from ._persistent_preparation import statistics_contract, statistics_identity
from .preparation_result import DiskPreparationResult

Expand Down Expand Up @@ -78,7 +78,9 @@ def _header(store, preprocessor, manifest, source_version, preprocessing_version
"preprocessing_version": preprocessing_version, "split": split}
if count is not None:
expected["sample_count"] = count
differences = [name for name, value in expected.items() if row[name] != value]
differences = [name for name, value in expected.items()
if (compatibility_config(json.loads(row[name])) != compatibility_config(json.loads(value))
if name == "config" else row[name] != value)]
if differences:
raise ValueError(f"Manifest {manifest!r} is incompatible: {', '.join(differences)} changed. "
"Use the matching contract or a new manifest name; cache metadata must not be rewritten.")
Expand Down
32 changes: 20 additions & 12 deletions mtlearn/python/mtlearn/layers/cfp/preparation/_identity.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,8 @@
"""Versioned, primitive-only identities for local persistent preparation."""
from functools import lru_cache
import hashlib
import json
from pathlib import Path

from .... import morphology
from ...._native import load_bindings
from ..specs import FeatureSpec, TreeSpec

FORMAT_VERSION = 4
Expand All @@ -24,12 +21,20 @@ def digest_file(path):
return digest.hexdigest()


@lru_cache(maxsize=1)
def implementation_identity():
# Binary identity is deliberately conservative across backend builds. The
# semantic version must change when Python preparation/quantization changes.
return {"semantics": PREPARATION_SEMANTICS,
"native_sha256": digest_file(Path(load_bindings().__file__))}
return {"semantics": PREPARATION_SEMANTICS}


def compatible_implementation(implementation):
return (isinstance(implementation, dict)
and set(implementation) in ({"semantics"}, {"semantics", "native_sha256"})
and implementation["semantics"] == PREPARATION_SEMANTICS)


def compatibility_config(config):
if not isinstance(config, dict) or not compatible_implementation(config.get("implementation")):
raise ValueError("Incompatible persistent preparation implementation or format.")
return dict(config, implementation=implementation_identity())


def tree_config(spec):
Expand All @@ -55,7 +60,8 @@ def preprocessor_config(preprocessor):
from .cfp_preprocessor import CFPPreprocessor
if type(preprocessor) is not CFPPreprocessor or preprocessor.morphology is not morphology:
raise ValueError("Persistent preparation currently requires the standard CFPPreprocessor implementation.")
return {"format_version": FORMAT_VERSION, "implementation": dict(implementation_identity()),
implementation = getattr(preprocessor, "_implementation", implementation_identity())
return {"format_version": FORMAT_VERSION, "implementation": dict(implementation),
"attribute_dtype": preprocessor.attribute_dtype.name,
"trees": [{"tree": tree_config(spec),
"attributes": [a.name for a in preprocessor.features[key].attributes]}
Expand All @@ -64,14 +70,16 @@ def preprocessor_config(preprocessor):

def preprocessor_from_config(config):
from .cfp_preprocessor import CFPPreprocessor
if config["format_version"] != FORMAT_VERSION or config["implementation"] != implementation_identity():
if config["format_version"] != FORMAT_VERSION or not compatible_implementation(config["implementation"]):
raise ValueError("Incompatible persistent preparation implementation or format.")
trees, features = {}, {}
for entry in config["trees"]:
tree = tree_from_config(entry["tree"])
trees[tree.cache_key()] = tree
features[tree.cache_key()] = FeatureSpec(tuple(getattr(morphology.AttributeType, a) for a in entry["attributes"]))
return CFPPreprocessor(tree_specs=trees, features=features, attribute_dtype=config["attribute_dtype"])
preprocessor = CFPPreprocessor(tree_specs=trees, features=features, attribute_dtype=config["attribute_dtype"])
preprocessor._implementation = dict(config["implementation"])
return preprocessor


def image_identity(preprocessor, image, tree_key, *, pixel_digest=None):
Expand All @@ -91,7 +99,7 @@ def validate_identity(identity):
expected = {"format_version", "implementation", "pixels_sha256", "shape", "tree", "attributes", "attribute_dtype"}
if set(identity) != expected or identity["format_version"] != FORMAT_VERSION:
raise ValueError("Unsupported preparation identity format.")
if identity["implementation"] != implementation_identity():
if not compatible_implementation(identity["implementation"]):
raise ValueError("Preparation backend/semantics do not match this runtime.")
digest = identity["pixels_sha256"]
if not isinstance(digest, str) or len(digest) != 64 or any(c not in "0123456789abcdef" for c in digest):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from ..normalization import AttributeNormalizer
from ..normalization.statistics_snapshot import StatisticsSnapshot
from ..runtime.cache_input_contract import validate_cfp_cache_batch_x
from ._identity import canonical_json, image_identity, identity_key, preprocessor_config, preprocessor_from_config
from ._identity import canonical_json, compatibility_config, image_identity, identity_key, preprocessor_config, preprocessor_from_config
from .preparation_result import PreparationResult
from ._preparation_progress import emit

Expand Down Expand Up @@ -63,13 +63,19 @@ def _begin_manifest(store, name, preprocessor, count, source_version, preprocess
if any(not isinstance(v, str) or not v for v in (name, source_version, preprocessing_version)):
raise ValueError("Persistent preparation requires manifest, source_version and preprocessing_version strings.")
config = canonical_json(preprocessor_config(preprocessor))
values = (config, source_version, preprocessing_version, split, count)
row = store._db.execute("SELECT * FROM manifests WHERE name=?", (name,)).fetchone()
if row is not None:
previous = json.loads(row["config"])
if compatibility_config(previous) == compatibility_config(json.loads(config)):
preprocessor = preprocessor_from_config(previous)
config = row["config"]
values = (config, source_version, preprocessing_version, split, count)
if row is not None and tuple(row[k] for k in ("config", "source_version", "preprocessing_version", "split", "sample_count")) != values:
raise ValueError("Manifest contract/source version changed; use a new manifest name to reuse compatible content safely.")
with store._db:
store._db.execute("""INSERT INTO manifests VALUES (?,?,?,?,?,?,'preparing',NULL)
ON CONFLICT(name) DO UPDATE SET state='preparing',error=NULL""", (name, *values))
return preprocessor


def _bind_sample(store, name, position, sample_id, bindings, shape):
Expand Down Expand Up @@ -146,7 +152,7 @@ def prepare_source(preprocessor, source, *, store=None, manifest=None, source_ve
persistent = getattr(store, "persistent", False)
if persistent:
store._check(write=True)
_begin_manifest(store, manifest, preprocessor, count, source_version, preprocessing_version, split)
preprocessor = _begin_manifest(store, manifest, preprocessor, count, source_version, preprocessing_version, split)
# A resumed pass must revalidate durable files, even if this process has
# already borrowed valid CPU buffers from a now damaged/missing file.
store.clear()
Expand Down
7 changes: 5 additions & 2 deletions mtlearn/python/mtlearn/layers/cfp/storage/_manifest.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
import json
import sqlite3

from ..preparation._identity import FORMAT_VERSION, canonical_json, implementation_identity
from ..preparation._identity import FORMAT_VERSION, canonical_json, implementation_identity, compatible_implementation

SCHEMA = """
CREATE TABLE metadata (key TEXT PRIMARY KEY, value TEXT NOT NULL);
Expand Down Expand Up @@ -55,7 +55,10 @@ def connect(path, *, readonly):
db.rollback()
raise
row = db.execute("SELECT value FROM metadata WHERE key='contract'").fetchone()
if row is None or json.loads(row[0]) != expected:
contract = json.loads(row[0]) if row is not None else None
if (not isinstance(contract, dict) or set(contract) != set(expected)
or contract["format_version"] != FORMAT_VERSION
or not compatible_implementation(contract["implementation"])):
raise ValueError("DiskStore format/backend is incompatible; use a separate store directory.")
except BaseException:
db.close()
Expand Down
101 changes: 101 additions & 0 deletions mtlearn/tests/python/test_cfp_cache_compatibility.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,101 @@
import json

import numpy as np
import pytest
import torch

import mtlearn
from mtlearn import _native
from mtlearn.layers.cfp import CFPPreprocessor, DiskStore, PreparedDataset
from mtlearn.layers.cfp.preparation import _identity
from test_cfp_disk_cache import QUOTA, forbid, model, prepare


@pytest.mark.parametrize("legacy", [False, True])
@pytest.mark.parametrize("workers", [0, 1])
def test_cache_reuse_across_binary_and_package_versions(tmp_path, monkeypatch, legacy, workers):
preprocessor = CFPPreprocessor.from_layer(model())
source = [torch.tensor([[[0, 1], [2, 3]]], dtype=torch.uint8)]
current = _identity.implementation_identity
with monkeypatch.context() as previous:
if legacy:
previous.setattr(_identity, "implementation_identity", lambda: dict(current(), native_sha256="a" * 64))
previous.setattr(mtlearn, "__version__", "1.2.0")
with DiskStore(tmp_path, max_disk_bytes=QUOTA) as store:
prepare(preprocessor, source, store)
config = json.loads(store._db.execute("SELECT config FROM manifests").fetchone()[0])
entries = [tuple(row) for row in store._db.execute("SELECT * FROM entries")]
expected = PreparedDataset(store, "train", source)[0]
expected_attributes = next(iter(expected.samples[0][0].values())).raw_attributes
expected_attributes = {key: value.clone() for key, value in expected_attributes.items()}
del expected
files = {path: path.read_bytes() for path in tmp_path.rglob("*.pt")}
assert files
monkeypatch.setattr(mtlearn, "__version__", "1.2.1")
monkeypatch.setattr(_native, "load_bindings", forbid)
monkeypatch.setattr(CFPPreprocessor, "prepare_u8", forbid)
assert _identity.implementation_identity() == {"semantics": _identity.PREPARATION_SEMANTICS}
assert CFPPreprocessor.from_config(config).get_config() == config
for readonly in (True, False):
with DiskStore(tmp_path, readonly=readonly, max_disk_bytes=None if readonly else QUOTA) as store:
batch = PreparedDataset(store, "train", source)[0]
actual = next(iter(batch.samples[0][0].values())).raw_attributes
for key in expected_attributes:
torch.testing.assert_close(actual[key], expected_attributes[key])
del batch, actual
if not readonly:
assert prepare(preprocessor, source, store, num_workers=workers).status == "complete"
assert [tuple(row) for row in store._db.execute("SELECT * FROM entries")] == entries
result = preprocessor.prepare_or_reuse(path=tmp_path, manifest="train", source_version="source-v1",
preprocessing_version="uint8-v1", mode="reuse")
assert result.status == "complete"
assert result.prepared_entries == 0
assert {path: path.read_bytes() for path in files} == files


@pytest.mark.parametrize("field,value", [("semantics", "incompatible-preparation"), ("unknown", "value")])
def test_incompatible_implementation_rejected_at_all_boundaries(tmp_path, field, value):
preprocessor = CFPPreprocessor.from_layer(model())
config = preprocessor.get_config()
config["implementation"][field] = value
with pytest.raises(ValueError, match="Incompatible"):
CFPPreprocessor.from_config(config)
identity = _identity.image_identity(preprocessor, np.array([[0, 1]], dtype=np.uint8),
next(iter(preprocessor.tree_specs)))
identity["implementation"][field] = value
with pytest.raises(ValueError, match="backend/semantics"):
_identity.validate_identity(identity)
with DiskStore(tmp_path, max_disk_bytes=QUOTA) as store:
store._db.execute("UPDATE metadata SET value=? WHERE key='contract'", (_identity.canonical_json({
"format_version": _identity.FORMAT_VERSION, "implementation": config["implementation"]}),))
store._db.commit()
with pytest.raises(ValueError, match="incompatible"):
DiskStore(tmp_path, readonly=True)


@pytest.mark.parametrize("workers", [0, 1])
def test_resume_legacy_cache_preserves_original_keys(tmp_path, monkeypatch, workers):
preprocessor = CFPPreprocessor.from_layer(model())
source = [torch.tensor([[[0, 1], [2, 3]]], dtype=torch.uint8),
torch.tensor([[[3, 2], [0, 1]]], dtype=torch.uint8)]
current = _identity.implementation_identity
with monkeypatch.context() as previous:
previous.setattr(_identity, "implementation_identity", lambda: dict(current(), native_sha256="b" * 64))
with DiskStore(tmp_path, max_disk_bytes=QUOTA) as store:
result = prepare(preprocessor, source, store,
cancel=lambda: store._db.execute("SELECT count(*) FROM samples").fetchone()[0] == 1)
assert result.status == "cancelled"
config = json.loads(store._db.execute("SELECT config FROM manifests").fetchone()[0])
key = store._db.execute("SELECT key FROM entries").fetchone()[0]
result = preprocessor.prepare_or_reuse(source, path=tmp_path, manifest="train", source_version="source-v1",
preprocessing_version="uint8-v1", mode="prepare_missing", max_disk_bytes=QUOTA, num_workers=workers)
assert result.status == "complete"
with DiskStore(tmp_path, readonly=True) as store:
assert json.loads(store._db.execute("SELECT config FROM manifests").fetchone()[0]) == config
assert store.get(key) is not None
dataset = PreparedDataset(store, "train", source)
assert len(dataset) == 2
for index in range(len(dataset)):
assert dataset[index].shape == (1, 1, 2, 2)
for row in store._db.execute("SELECT identity FROM entries"):
assert json.loads(row[0])["implementation"] == config["implementation"]
Loading