diff --git a/dpsynth/api.py b/dpsynth/api.py index 4027960..7c89923 100644 --- a/dpsynth/api.py +++ b/dpsynth/api.py @@ -39,6 +39,7 @@ from typing import Any import dp_accounting +from etils import epath class CalibratedMechanism(abc.ABC): @@ -118,6 +119,11 @@ class MechanismConfig(abc.ABC): _registry: dict[str, type[MechanismConfig]] = {} + @property + def working_dir(self) -> epath.PathLike | None: + """Base directory path for checkpointing intermediate mechanism state.""" + return None + def __init_subclass__(cls, **kwargs: Any): super().__init_subclass__(**kwargs) MechanismConfig._registry[cls.__name__] = cls diff --git a/dpsynth/checkpoint.py b/dpsynth/checkpoint.py new file mode 100644 index 0000000..e82a0b5 --- /dev/null +++ b/dpsynth/checkpoint.py @@ -0,0 +1,93 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Checkpointing utilities for long-running mechanism synthesis. + +Provides :class:`Checkpointer`, which serializes and deserializes intermediate +mechanism state (e.g. exact marginals, noisy measurements, graphical models) +using :mod:`mbi` pytree serialization on top of :mod:`etils.epath`. +""" + +from __future__ import annotations + +import dataclasses +import io +from typing import Any + +from etils import epath +import mbi + + +@dataclasses.dataclass(frozen=True) +class Checkpointer: + """Saves and restores intermediate mechanism state as .npz checkpoints. + + When ``working_dir`` is None (the default), all save/load operations are + no-ops, allowing callers to disable checkpointing without branching. + When ``working_dir`` is provided, intermediate mechanism state is persisted + directly under that directory as .npz files using ``mbi.save`` and + ``mbi.load``. + + Attributes: + working_dir: Base directory path for checkpoint files (supports local, + Cloud, and remote paths via epath.Path). If None, checkpointing is + disabled. + """ + + working_dir: epath.PathLike | None = None + + @property + def path(self) -> epath.Path | None: + """The resolved working directory path, or None if disabled.""" + return ( + epath.Path(self.working_dir) if self.working_dir is not None else None + ) + + def save(self, name: str, obj: Any) -> None: + """Saves an object to the working directory (no-op if disabled). + + Args: + name: Filename to write the object to (e.g. 'model.npz'). + obj: A JAX pytree to serialize (e.g. a CliqueVector, model, or list of + measurements). + """ + if self.path is None: + return + self.path.mkdir(parents=True, exist_ok=True) + buf = io.BytesIO() + mbi.save(obj, buf) + (self.path / name).write_bytes(buf.getvalue()) + + def load(self, name: str) -> Any | None: + """Loads an object from the working directory, or None if absent/disabled. + + Args: + name: Filename of the checkpointed object. + + Returns: + The deserialized object, or None if checkpointing is disabled or the + file does not exist. + """ + if self.path is None: + return None + target = self.path / name + if not target.exists(): + return None + return mbi.load(io.BytesIO(target.read_bytes())) + + def exists(self, name: str) -> bool: + """Returns True if the named checkpoint file exists.""" + if self.path is None: + return False + return (self.path / name).exists() diff --git a/dpsynth/data_generation_v3.py b/dpsynth/data_generation_v3.py index 576ee1a..956da72 100644 --- a/dpsynth/data_generation_v3.py +++ b/dpsynth/data_generation_v3.py @@ -30,6 +30,7 @@ from dpsynth.local_mode import initialization from dpsynth.local_mode import primitives from dpsynth.local_mode import vectorized_transformations as vtx +from etils import epath import mbi import numpy as np import pandas as pd @@ -368,6 +369,8 @@ class TabularConfig(api.MechanismConfig): cross_attribute_constraints: Constraints to enforce on generated data. compress_columns: Whether to compress rare categories (< 3*sigma) for CategoricalAttribute columns not present in constraints. + working_dir: Base directory path for intermediate checkpoints (passed down + to the underlying discrete mechanism). If None, checkpointing is disabled. """ domains: Mapping[str, domain.AttributeType] | None = None @@ -376,6 +379,7 @@ class TabularConfig(api.MechanismConfig): init_budget_fraction: float = 0.1 cross_attribute_constraints: Sequence[constraints.Constraint] = () compress_columns: bool = False + working_dir: epath.PathLike | None = None def _compute_per_col_deltas(self, domains, delta): # Split delta across open-set columns, analogous to splitting zcdp_rho. @@ -484,7 +488,21 @@ def configure( for col, init in inits.items() } - calibrated_discrete = self.discrete_mechanism.configure( + discrete_mechanism = self.discrete_mechanism + if ( + self.working_dir is not None + and dataclasses.is_dataclass(discrete_mechanism) + and any( + f.name == 'working_dir' + for f in dataclasses.fields(discrete_mechanism) + ) + and discrete_mechanism.working_dir is None + ): + discrete_mechanism = dataclasses.replace( # pyrefly: ignore[bad-specialization] + discrete_mechanism, working_dir=self.working_dir + ) + + calibrated_discrete = discrete_mechanism.configure( max_records_per_user=max_records_per_user, zcdp_rho=discrete_rho, ) diff --git a/dpsynth/discrete_mechanisms/common.py b/dpsynth/discrete_mechanisms/common.py index 49eddcb..0f834a1 100644 --- a/dpsynth/discrete_mechanisms/common.py +++ b/dpsynth/discrete_mechanisms/common.py @@ -384,7 +384,10 @@ def supporting_cliques( A list of cliques from the workload whose domain size is within the limit. """ if workload is None: - cliques = list(itertools.combinations(domain.attributes, 3)) + k = min(len(domain.attributes), 3) + cliques = ( + list(itertools.combinations(domain.attributes, k)) if k > 0 else [] + ) elif isinstance(workload, Mapping): cliques = list(workload.keys()) else: @@ -422,7 +425,10 @@ def compiled_workload( """ if workload is None: - workload = list(itertools.combinations(domain.attributes, 3)) + k = min(len(domain.attributes), 3) + workload = ( + list(itertools.combinations(domain.attributes, k)) if k > 0 else [] + ) if not isinstance(workload, Mapping): workload = {cl: 1.0 for cl in workload} diff --git a/dpsynth/discrete_mechanisms/discrete.py b/dpsynth/discrete_mechanisms/discrete.py index c8c1229..152f513 100644 --- a/dpsynth/discrete_mechanisms/discrete.py +++ b/dpsynth/discrete_mechanisms/discrete.py @@ -26,11 +26,14 @@ from collections.abc import Sequence import dataclasses +from absl import logging import dp_accounting from dpsynth import api +from dpsynth import checkpoint as checkpoint_lib from dpsynth.discrete_mechanisms import accounting from dpsynth.discrete_mechanisms import common from dpsynth.discrete_mechanisms import mst +from etils import epath import mbi import numpy as np @@ -45,20 +48,36 @@ class DiscreteConfig(api.MechanismConfig): one_way_budget_fraction: Fraction of zCDP budget for one-way marginals. constraints: Default MBI constraints to enforce. Can be overridden at call time via the ``constraints`` kwarg on ``DiscreteMechanism.__call__``. + working_dir: Base directory path for intermediate checkpoints. If None, + checkpointing is disabled. """ mechanism: api.MechanismConfig = mst.MSTConfig() compress_columns: bool | Sequence[str] = False one_way_budget_fraction: float = 0.1 constraints: Sequence[mbi.Constraint] = () + working_dir: epath.PathLike | None = None def configure(self, _=None, *, zcdp_rho, delta=0, max_records_per_user=1): """Configures the synthesizer with a zCDP budget.""" api.validate_max_records_per_user(max_records_per_user) + inner_mechanism = self.mechanism + if ( + self.working_dir is not None + and dataclasses.is_dataclass(inner_mechanism) + and any( + f.name == 'working_dir' for f in dataclasses.fields(inner_mechanism) + ) + and inner_mechanism.working_dir is None + ): + inner_mechanism = dataclasses.replace( # pyrefly: ignore[bad-specialization] + inner_mechanism, working_dir=self.working_dir + ) + one_way_rho = zcdp_rho * self.one_way_budget_fraction - remaining_rho = zcdp_rho * (1 - self.one_way_budget_fraction) - inner = self.mechanism.configure( + remaining_rho = zcdp_rho - one_way_rho + inner = inner_mechanism.configure( zcdp_rho=remaining_rho, delta=delta, max_records_per_user=max_records_per_user, @@ -127,20 +146,30 @@ def __call__( if constraints is None: constraints = self.config.constraints + checkpointer = checkpoint_lib.Checkpointer(self.config.working_dir) + if initial_measurements is not None: measurements = list(initial_measurements) elif self.one_way_gdp_budget > 0: - one_way_cliques = [(a,) for a in data.domain] - if hasattr(data, 'cliques'): - supported = common.downward_closure(data.cliques) - one_way_cliques = [cl for cl in one_way_cliques if cl in supported] - measurements = common.measure_marginals_with_noise( - rng=rng, - data=data, # pyrefly: ignore[bad-argument-type] - marginal_queries=one_way_cliques, # pyrefly: ignore[bad-argument-type] - gdp_sigma=accounting.gdp_gaussian_sigma(self.one_way_gdp_budget), - max_records_per_user=self.max_records_per_user, - ) + if checkpointer.exists('one_way_measurements.npz'): + logging.info( + '[DiscreteMechanism] Resuming one-way measurements from checkpoint.' + ) + measurements = checkpointer.load('one_way_measurements.npz') + assert measurements is not None + else: + one_way_cliques = [(a,) for a in data.domain] + if hasattr(data, 'cliques'): + supported = common.downward_closure(data.cliques) + one_way_cliques = [cl for cl in one_way_cliques if cl in supported] + measurements = common.measure_marginals_with_noise( + rng=rng, + data=data, # pyrefly: ignore[bad-argument-type] + marginal_queries=one_way_cliques, # pyrefly: ignore[bad-argument-type] + gdp_sigma=accounting.gdp_gaussian_sigma(self.one_way_gdp_budget), + max_records_per_user=self.max_records_per_user, + ) + checkpointer.save('one_way_measurements.npz', measurements) else: measurements = [] diff --git a/dpsynth/discrete_mechanisms/swift.py b/dpsynth/discrete_mechanisms/swift.py index 7bcdf6a..c382f98 100644 --- a/dpsynth/discrete_mechanisms/swift.py +++ b/dpsynth/discrete_mechanisms/swift.py @@ -26,6 +26,7 @@ from __future__ import annotations from collections.abc import Iterable, Mapping, Sequence +import concurrent.futures import dataclasses import functools import itertools @@ -35,10 +36,12 @@ from absl import logging import dp_accounting from dpsynth import api +from dpsynth import checkpoint as checkpoint_lib from dpsynth.discrete_mechanisms import accounting from dpsynth.discrete_mechanisms import clique_tree from dpsynth.discrete_mechanisms import common from dpsynth.discrete_mechanisms import swift_utils +from etils import epath import mbi import networkx as nx import numpy as np @@ -59,6 +62,8 @@ class SWIFTConfig(api.MechanismConfig): pgm_iters: Number of mirror descent iterations for PGM estimation. select_budget_frac: Fraction of the total budget used for selecting which marginals to measure. + working_dir: Base directory path for intermediate checkpoints (e.g. exact + marginals, noisy measurements, model). If None, checkpointing is disabled. """ workload: Mapping[mbi.Clique, float] | Iterable[mbi.Clique] | None = None @@ -67,6 +72,7 @@ class SWIFTConfig(api.MechanismConfig): pgm_iters: int = 10_000 marginal_oracle: mbi.MarginalOracle | None = None select_budget_frac: float = 0.1 + working_dir: epath.PathLike | None = None def supporting_cliques(self, domain: mbi.Domain) -> list[mbi.Clique]: """Returns the workload cliques filtered by max_marginal_size.""" @@ -83,6 +89,17 @@ def configure(self, _=None, *, zcdp_rho, delta=0, max_records_per_user=1): ) +def _compute_marginals( + data: mbi.Dataset | mbi.CliqueVector, + candidates: Mapping[mbi.Clique, float] | Sequence[mbi.Clique], +) -> mbi.CliqueVector: + """Computes exact marginals using accelerated precompute or projectable.""" + clique_list = list(candidates) + if hasattr(data, 'data'): + return mbi.extensions.precompute_marginals(data, clique_list) # pyrefly: ignore[bad-argument-type] + return mbi.CliqueVector.from_projectable(data, clique_list) # pyrefly: ignore[bad-argument-type] + + @dataclasses.dataclass(frozen=True, kw_only=True) class SWIFT(api.CalibratedMechanism): """Calibrated SWIFT instance.""" @@ -98,23 +115,24 @@ def dp_event(self) -> dp_accounting.DpEvent: accounting.gdp_gaussian_sigma(self.gdp_budget) ) - def __call__( + def _select_and_measure( self, rng: np.random.Generator, data: mbi.Dataset | mbi.CliqueVector, - *, - initial_measurements: Sequence[mbi.LinearMeasurement] = (), - constraints: Sequence[mbi.Constraint] = (), - ) -> common.DiscreteMechanismResult: - common.validate_initial_measurements(initial_measurements) - phase_times = {} - + checkpointer: checkpoint_lib.Checkpointer, + phase_times: dict[str, float], + initial_measurements: Sequence[mbi.LinearMeasurement], + constraints: Sequence[mbi.Constraint], + ) -> tuple[ + list[mbi.LinearMeasurement], + nx.Graph, + concurrent.futures.Future | None, + concurrent.futures.Future | None, + ]: + """Selects and measures candidate marginals, returning measurements and jtree.""" select_gdp_budget = self.gdp_budget * self.config.select_budget_frac measure_gdp_budget = self.gdp_budget - select_gdp_budget - ######################################################################### - # Compile workload into candidate measurements, and precompute answers. # - ######################################################################### with common.timed(phase_times, 'compiled_workload'): candidates = common.compiled_workload( data.domain, @@ -123,8 +141,14 @@ def __call__( ) logging.info('[SWIFT] %d candidates.', len(candidates)) - with common.timed(phase_times, 'from_projectable'): - answers = mbi.CliqueVector.from_projectable(data, candidates) # pyrefly: ignore[bad-argument-type] + if checkpointer.exists('marginals.npz'): + logging.info('[SWIFT] Resuming exact marginals from checkpoint.') + answers = checkpointer.load('marginals.npz') + assert answers is not None + else: + with common.timed(phase_times, 'from_projectable'): + answers = _compute_marginals(data, candidates) + checkpointer.save('marginals.npz', answers) with common.timed(phase_times, 'initial_mirror_descent'): estimator = mbi.estimation.MirrorDescent(self.config.marginal_oracle) @@ -135,11 +159,7 @@ def __call__( constraints=constraints, ) - ########################################### - # Select subset of candidates to measure. # - ########################################### with common.timed(phase_times, 'selection'): - with common.timed(phase_times, 'compute_initial_errors'): noisy_errors = _compute_initial_errors( rng, @@ -156,15 +176,12 @@ def __call__( candidates, data.domain, self.config.max_clique_size, - measure_gdp_budget, # budget is not consumed (no data dependence) + measure_gdp_budget, ) all_cliques = [m.clique for m in initial_measurements] + list(selected) logging.info(mbi.summarize(data.domain, all_cliques, jtree)) - ######################################################## - # Precompile MirrorDescent + synth while measuring. # - ######################################################## closed_oracle = functools.partial( mbi.marginal_oracles.message_passing_stable, jtree=jtree ) @@ -180,9 +197,6 @@ def __call__( ) logging.info('[SWIFT] Started precompilation of MirrorDescent + synth.') - ########################################## - # Measure the selected marginal queries. # - ########################################## with common.timed(phase_times, 'measurement'): logging.info('[SWIFT] Starting measurements.') new_measurements = _measure_selected_marginals( @@ -193,36 +207,138 @@ def __call__( max_records_per_user=self.max_records_per_user, ) measurements = list(initial_measurements) + new_measurements + checkpointer.save('measurements.npz', measurements) logging.info('[SWIFT] Finished measurements.') - ######################################################## - # Estimate the model using all measurements # - ######################################################## + return measurements, jtree, pgm_future, synth_future + + def _estimate_model( + self, + domain: mbi.Domain, + measurements: Sequence[mbi.LinearMeasurement], + jtree: nx.Graph, + checkpointer: checkpoint_lib.Checkpointer, + phase_times: dict[str, float], + constraints: Sequence[mbi.Constraint], + pgm_future: concurrent.futures.Future | None = None, + ) -> mbi.Model: + """Estimates the MRF model from measurements using MirrorDescent.""" with common.timed(phase_times, 'estimation'): - t0 = time.time() - pgm_future.result() - logging.info('[SWIFT] PGM precompile wait: %.2fs', time.time() - t0) + if pgm_future is not None: + t0 = time.time() + pgm_future.result() + logging.info('[SWIFT] PGM precompile wait: %.2fs', time.time() - t0) + closed_oracle = functools.partial( + mbi.marginal_oracles.message_passing_stable, jtree=jtree + ) + estimator = mbi.estimation.MirrorDescent(marginal_oracle=closed_oracle) final_model = estimator.estimate( - data.domain, - measurements, + domain, + list(measurements), iters=self.config.pgm_iters, - callback_fn=mbi.callbacks.default(measurements, data.domain), + callback_fn=mbi.callbacks.default(list(measurements), domain), constraints=constraints, ) + checkpointer.save('model.npz', final_model) logging.info('[SWIFT] Estimated final model.') + return final_model - t0 = time.time() - synth_future.result() - logging.info('[SWIFT] Synth precompile wait: %.2fs', time.time() - t0) + def _synthesize_result( + self, + final_model: mbi.Model, + measurements: Sequence[mbi.LinearMeasurement], + initial_measurements: Sequence[mbi.LinearMeasurement], + phase_times: dict[str, float], + synth_future: concurrent.futures.Future | None = None, + ) -> common.DiscreteMechanismResult: + """Synthesizes dataset records from model and builds mechanism result.""" + if synth_future is not None: + t0 = time.time() + synth_future.result() + logging.info('[SWIFT] Synth precompile wait: %.2fs', time.time() - t0) + total_src = initial_measurements if initial_measurements else measurements + rows = mbi.estimation.minimum_variance_unbiased_total(total_src) # pyrefly: ignore[bad-argument-type] + rows = int(round(max(rows, 1))) syn = mbi.extensions.synthetic_data(final_model, rows) # pyrefly: ignore[bad-argument-type] logging.info('[SWIFT] Generated %d synthetic records.', rows) + + diagnostics = common.clique_stats(final_model) + diagnostics.phase_times = phase_times return common.DiscreteMechanismResult( synthetic_data=syn, - measurements=measurements, + measurements=list(measurements), model=final_model, - diagnostics=common.clique_stats(final_model), + diagnostics=diagnostics, + ) + + def __call__( + self, + rng: np.random.Generator, + data: mbi.Dataset | mbi.CliqueVector, + *, + initial_measurements: Sequence[mbi.LinearMeasurement] = (), + constraints: Sequence[mbi.Constraint] = (), + ) -> common.DiscreteMechanismResult: + common.validate_initial_measurements(initial_measurements) + phase_times = {} + checkpointer = checkpoint_lib.Checkpointer(self.config.working_dir) + + # 1. Full resume: if model and measurements already exist, skip to synthesis. + if checkpointer.exists('model.npz') and checkpointer.exists( + 'measurements.npz' + ): + logging.info('[SWIFT] Resuming from checkpointed model and measurements.') + final_model = checkpointer.load('model.npz') + measurements = checkpointer.load('measurements.npz') + assert final_model is not None and measurements is not None + return self._synthesize_result( + final_model, measurements, initial_measurements, phase_times + ) + + # 2. Stage 1: Measurements + if checkpointer.exists('measurements.npz'): + logging.info('[SWIFT] Resuming from checkpointed measurements.') + measurements = checkpointer.load('measurements.npz') + assert measurements is not None + jtree, _ = mbi.junction_tree.make_junction_tree( + data.domain, [m.clique for m in measurements] + ) + pgm_future, synth_future = None, None + else: + measurements, jtree, pgm_future, synth_future = self._select_and_measure( + rng, + data, + checkpointer, + phase_times, + initial_measurements, + constraints, + ) + + # 3. Stage 2: Model Estimation + if checkpointer.exists('model.npz'): + logging.info('[SWIFT] Resuming from checkpointed model.') + final_model = checkpointer.load('model.npz') + assert final_model is not None + else: + final_model = self._estimate_model( + data.domain, + measurements, + jtree, + checkpointer, + phase_times, + constraints, + pgm_future, + ) + + # 4. Stage 3: Synthesis & Diagnostics + return self._synthesize_result( + final_model, + measurements, + initial_measurements, + phase_times, + synth_future=synth_future, ) @@ -335,7 +451,7 @@ def build_best_clique_tree( if len(cl) == 2 and tuple(sorted(cl)) in supported ) - if score > best_score: + if score > best_score or best_tree is None: best_score = score best_tree = tree assert best_tree is not None @@ -351,6 +467,8 @@ def _compute_initial_errors( max_records_per_user: int = 1, ) -> dict[mbi.Clique, float]: """Computes DP initial errors for the SWIFT mechanism.""" + if not cliques: + return {} budget_per_clique = gdp_budget / len(cliques) sigma_per_clique = max_records_per_user * accounting.gdp_gaussian_sigma( budget_per_clique diff --git a/tests/checkpoint_test.py b/tests/checkpoint_test.py new file mode 100644 index 0000000..990cd5b --- /dev/null +++ b/tests/checkpoint_test.py @@ -0,0 +1,117 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Unit tests for dpsynth.checkpoint.""" + +import pathlib +from absl.testing import absltest +from dpsynth import checkpoint as checkpoint_lib +from etils import epath +import jax.numpy as jnp +import mbi +import numpy as np + + +class CheckpointerTest(absltest.TestCase): + + def test_noop_when_disabled(self): + ckpt = checkpoint_lib.Checkpointer(working_dir=None) + self.assertIsNone(ckpt.path) + self.assertFalse(ckpt.exists('model.npz')) + self.assertIsNone(ckpt.load('model.npz')) + + # Saving should be a no-op and not raise. + domain = mbi.Domain.fromdict({'a': 2, 'b': 3}) + cliques = [('a',), ('b',)] + potentials = mbi.CliqueVector.zeros(domain, cliques) + ckpt.save('model.npz', potentials) + self.assertFalse(ckpt.exists('model.npz')) + self.assertIsNone(ckpt.load('model.npz')) + + def test_save_and_load_roundtrip(self): + working_dir = self.create_tempdir().full_path + ckpt = checkpoint_lib.Checkpointer(working_dir=working_dir) + + domain = mbi.Domain.fromdict({'a': 2, 'b': 3}) + cliques = [('a',), ('a', 'b')] + potentials = mbi.CliqueVector.zeros(domain, cliques) + potentials[('a',)] = jnp.array([1.0, 2.0]) + marginals = mbi.CliqueVector.zeros(domain, cliques) + mrf = mbi.MarkovRandomField( + potentials=potentials, marginals=marginals, total=10.0 + ) + + self.assertFalse(ckpt.exists('model.npz')) + ckpt.save('model.npz', mrf) + self.assertTrue(ckpt.exists('model.npz')) + + loaded = ckpt.load('model.npz') + self.assertIsInstance(loaded, mbi.MarkovRandomField) + np.testing.assert_allclose(loaded.potentials[('a',)], potentials[('a',)]) + self.assertEqual(loaded.total, 10.0) + + def test_save_and_load_linear_measurements(self): + working_dir = self.create_tempdir().full_path + ckpt = checkpoint_lib.Checkpointer(working_dir=working_dir) + + measurements = [ + mbi.LinearMeasurement(np.array([5.0, 10.0]), ('a',), stddev=1.0), + mbi.LinearMeasurement( + np.array([1.0, 2.0, 3.0, 4.0, 5.0, 6.0]), ('a', 'b'), stddev=0.5 + ), + ] + ckpt.save('measurements.npz', measurements) + self.assertTrue(ckpt.exists('measurements.npz')) + + loaded = ckpt.load('measurements.npz') + self.assertLen(loaded, 2) + self.assertEqual(loaded[0].clique, ('a',)) + self.assertEqual(loaded[0].stddev, 1.0) + np.testing.assert_allclose( + loaded[0].noisy_measurement, measurements[0].noisy_measurement + ) + self.assertEqual(loaded[1].clique, ('a', 'b')) + self.assertEqual(loaded[1].stddev, 0.5) + np.testing.assert_allclose( + loaded[1].noisy_measurement, measurements[1].noisy_measurement + ) + + def test_load_nonexistent_returns_none(self): + working_dir = self.create_tempdir().full_path + ckpt = checkpoint_lib.Checkpointer(working_dir=working_dir) + self.assertIsNone(ckpt.load('nonexistent.npz')) + + def test_accepts_different_path_types(self): + temp_dir = self.create_tempdir().full_path + + # str + ckpt_str = checkpoint_lib.Checkpointer(working_dir=temp_dir) + ckpt_str.save('test_str.npz', {'v': jnp.array([1, 2, 3])}) + self.assertTrue(ckpt_str.exists('test_str.npz')) + + # pathlib.Path + ckpt_pathlib = checkpoint_lib.Checkpointer( + working_dir=pathlib.Path(temp_dir) + ) + self.assertTrue(ckpt_pathlib.exists('test_str.npz')) + loaded = ckpt_pathlib.load('test_str.npz') + np.testing.assert_array_equal(loaded['v'], [1, 2, 3]) + + # epath.Path + ckpt_epath = checkpoint_lib.Checkpointer(working_dir=epath.Path(temp_dir)) + self.assertTrue(ckpt_epath.exists('test_str.npz')) + + +if __name__ == '__main__': + absltest.main() diff --git a/tests/data_generation_v3_test.py b/tests/data_generation_v3_test.py index f20c166..2568966 100644 --- a/tests/data_generation_v3_test.py +++ b/tests/data_generation_v3_test.py @@ -19,12 +19,14 @@ from absl.testing import absltest from absl.testing import parameterized import dp_accounting +from dpsynth import checkpoint as checkpoint_lib from dpsynth import constraints from dpsynth import data_generation_v3 from dpsynth import discrete_mechanisms from dpsynth import domain from dpsynth.discrete_mechanisms import aim from dpsynth.discrete_mechanisms import aim_gdp +from dpsynth.discrete_mechanisms import swift from dpsynth.discrete_mechanisms.independent import IndependentConfig import mbi import numpy as np @@ -560,6 +562,35 @@ def test_compress_columns_respects_constraints(self): result = calibrated(rng, df) self.assertNotIn('A', result.discrete_mechanism_result.mappings) + def test_working_dir_propagates_and_checkpoints(self): + working_dir = self.create_tempdir().full_path + domains = { + 'A': domain.CategoricalAttribute( + possible_values=['a', 'b', 'c'], out_of_domain_index=0 + ), + 'B': domain.CategoricalAttribute( + possible_values=['x', 'y', 'z'], out_of_domain_index=0 + ), + } + df = pd.DataFrame({'A': ['a', 'b', 'c'] * 20, 'B': ['x', 'y', 'z'] * 20}) + rng = np.random.default_rng(0) + + config = data_generation_v3.TabularConfig( + discrete_mechanism=swift.SWIFTConfig(pgm_iters=100), + working_dir=working_dir, + ) + calibrated = config.configure(domains, zcdp_rho=100.0) + result1 = calibrated(rng, df) + self.assertIsInstance(result1.synthetic_data, pd.DataFrame) + + checkpointer = checkpoint_lib.Checkpointer(working_dir) + self.assertTrue(checkpointer.exists('model.npz')) + self.assertTrue(checkpointer.exists('measurements.npz')) + + # Second run should resume from checkpointed model/measurements. + result2 = calibrated(rng, df) + self.assertIsInstance(result2.synthetic_data, pd.DataFrame) + if __name__ == '__main__': absltest.main() diff --git a/tests/discrete_mechanisms/discrete_test.py b/tests/discrete_mechanisms/discrete_test.py index ab245b5..43b2c43 100644 --- a/tests/discrete_mechanisms/discrete_test.py +++ b/tests/discrete_mechanisms/discrete_test.py @@ -16,6 +16,7 @@ from unittest import mock from absl.testing import absltest +from dpsynth import checkpoint as checkpoint_lib from dpsynth.discrete_mechanisms import accounting from dpsynth.discrete_mechanisms import common from dpsynth.discrete_mechanisms import discrete @@ -131,6 +132,28 @@ def test_converts_dataset_to_clique_vector(self): synth(np.random.default_rng(0), data) self.assertIsInstance(mock_call.call_args.args[1], mbi.CliqueVector) + def test_checkpoint_saves_and_resumes_one_way_measurements(self): + working_dir = self.create_tempdir().full_path + domain = mbi.Domain(['a', 'b', 'c'], [3, 4, 5]) + data = mbi.Dataset.synthetic(domain, N=200) + rng = np.random.default_rng(0) + + config = DiscreteConfig( + mechanism=MSTConfig(pgm_iters=500), + working_dir=working_dir, + ) + synth = config.configure(zcdp_rho=100.0) + result1 = synth(rng, data) + self.assertIsInstance(result1, common.DiscreteMechanismResult) + + # Verify one_way_measurements.npz exists. + checkpointer = checkpoint_lib.Checkpointer(working_dir) + self.assertTrue(checkpointer.exists('one_way_measurements.npz')) + + # Second run should resume from checkpointed one-way measurements. + result2 = synth(rng, data) + self.assertIsInstance(result2, common.DiscreteMechanismResult) + if __name__ == '__main__': absltest.main() diff --git a/tests/discrete_mechanisms/swift_test.py b/tests/discrete_mechanisms/swift_test.py index b11afb4..27cf8e2 100644 --- a/tests/discrete_mechanisms/swift_test.py +++ b/tests/discrete_mechanisms/swift_test.py @@ -19,6 +19,7 @@ from dpsynth.discrete_mechanisms import common from dpsynth.discrete_mechanisms import swift from dpsynth.discrete_mechanisms import swift_utils +from etils import epath import mbi import networkx as nx import numpy as np @@ -146,6 +147,75 @@ def test_fits_one_way_marginals(self): actual = result.model.project([col]).datavector() np.testing.assert_allclose(actual, expected, atol=1) + def test_checkpointing_saves_and_resumes(self): + temp_dir = self.create_tempdir().full_path + data = mbi.Dataset.synthetic(mbi.Domain(['a', 'b', 'c'], [2, 3, 4]), N=500) + initial = [ + mbi.LinearMeasurement( + data.project((c,)).datavector(), (c,), stddev=0.01 + ) + for c in data.domain + ] + + # 1. Cold run: should save marginals, measurements, and model + config = swift.SWIFTConfig(pgm_iters=100, working_dir=temp_dir).configure( + zcdp_rho=1000 + ) + result1 = config( + np.random.default_rng(0), data, initial_measurements=initial + ) + + ckpt_marginals = epath.Path(temp_dir) / 'marginals.npz' + ckpt_measurements = epath.Path(temp_dir) / 'measurements.npz' + ckpt_model = epath.Path(temp_dir) / 'model.npz' + + self.assertTrue(ckpt_marginals.exists()) + self.assertTrue(ckpt_measurements.exists()) + self.assertTrue(ckpt_model.exists()) + + # 2. Resume run: should load model and measurements and skip to synthesis + result2 = config( + np.random.default_rng(1), data, initial_measurements=initial + ) + self.assertEqual(result2.synthetic_data.records, 500) + for cl in result1.model.potentials.cliques: + np.testing.assert_allclose( + result2.model.potentials[cl].values, + result1.model.potentials[cl].values, + ) + + def test_checkpointing_resumes_from_marginals(self): + temp_dir = self.create_tempdir().full_path + data = mbi.Dataset.synthetic(mbi.Domain(['a', 'b', 'c'], [2, 3, 4]), N=500) + initial = [ + mbi.LinearMeasurement( + data.project((c,)).datavector(), (c,), stddev=0.01 + ) + for c in data.domain + ] + + # First run up to marginals + config1 = swift.SWIFTConfig(pgm_iters=100, working_dir=temp_dir).configure( + zcdp_rho=1000 + ) + config1(np.random.default_rng(0), data, initial_measurements=initial) + + # Delete measurements and model, keep marginals + (epath.Path(temp_dir) / 'measurements.npz').unlink() + (epath.Path(temp_dir) / 'model.npz').unlink() + + # Second run: should reuse marginals and re-estimate + config2 = swift.SWIFTConfig(pgm_iters=100, working_dir=temp_dir).configure( + zcdp_rho=1000 + ) + result = config2( + np.random.default_rng(0), data, initial_measurements=initial + ) + + self.assertTrue((epath.Path(temp_dir) / 'measurements.npz').exists()) + self.assertTrue((epath.Path(temp_dir) / 'model.npz').exists()) + self.assertEqual(result.synthetic_data.records, 500) + if __name__ == '__main__': absltest.main()