diff --git a/mtlearn/python/mtlearn/layers/cfp/preparation/_disk_preparation.py b/mtlearn/python/mtlearn/layers/cfp/preparation/_disk_preparation.py index bd0dfae..73bbf88 100644 --- a/mtlearn/python/mtlearn/layers/cfp/preparation/_disk_preparation.py +++ b/mtlearn/python/mtlearn/layers/cfp/preparation/_disk_preparation.py @@ -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 @@ -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.") diff --git a/mtlearn/python/mtlearn/layers/cfp/preparation/_identity.py b/mtlearn/python/mtlearn/layers/cfp/preparation/_identity.py index b0bed23..85a1c3c 100644 --- a/mtlearn/python/mtlearn/layers/cfp/preparation/_identity.py +++ b/mtlearn/python/mtlearn/layers/cfp/preparation/_identity.py @@ -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 @@ -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): @@ -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]} @@ -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): @@ -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): diff --git a/mtlearn/python/mtlearn/layers/cfp/preparation/_persistent_preparation.py b/mtlearn/python/mtlearn/layers/cfp/preparation/_persistent_preparation.py index 159cc83..3377a0e 100644 --- a/mtlearn/python/mtlearn/layers/cfp/preparation/_persistent_preparation.py +++ b/mtlearn/python/mtlearn/layers/cfp/preparation/_persistent_preparation.py @@ -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 @@ -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): @@ -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() diff --git a/mtlearn/python/mtlearn/layers/cfp/storage/_manifest.py b/mtlearn/python/mtlearn/layers/cfp/storage/_manifest.py index cd50fff..2c99a23 100644 --- a/mtlearn/python/mtlearn/layers/cfp/storage/_manifest.py +++ b/mtlearn/python/mtlearn/layers/cfp/storage/_manifest.py @@ -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); @@ -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() diff --git a/mtlearn/tests/python/test_cfp_cache_compatibility.py b/mtlearn/tests/python/test_cfp_cache_compatibility.py new file mode 100644 index 0000000..593faa5 --- /dev/null +++ b/mtlearn/tests/python/test_cfp_cache_compatibility.py @@ -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"]