diff --git a/dpsynth/__init__.py b/dpsynth/__init__.py index ddf4392..5dbc2a8 100644 --- a/dpsynth/__init__.py +++ b/dpsynth/__init__.py @@ -16,6 +16,7 @@ # pylint: disable=g-importing-member __version__ = '0.4.0' +from absl import logging from dpsynth import api from dpsynth import constraints from dpsynth import discrete_mechanisms @@ -35,6 +36,16 @@ from dpsynth.domain import Schema from dpsynth.serialize import from_yaml from dpsynth.serialize import to_yaml +import mbi + +# Route MBI callback logs (which use print-style formatting) to +# absl.logging.info. +if hasattr(mbi, 'callbacks') and hasattr(mbi.callbacks, 'set_log_fn'): + mbi.callbacks.set_log_fn( + lambda *args, sep=' ', **kwargs: logging.info( + sep.join(str(a) for a in args) + ) + ) ForeignKeyRelation = relational.ForeignKeyRelation MultiDataGenerationResult = relational.MultiDataGenerationResult diff --git a/dpsynth/api.py b/dpsynth/api.py index 4027960..8dd2e41 100644 --- a/dpsynth/api.py +++ b/dpsynth/api.py @@ -35,6 +35,7 @@ import abc from collections.abc import Callable +import dataclasses import functools from typing import Any @@ -118,6 +119,24 @@ class MechanismConfig(abc.ABC): _registry: dict[str, type[MechanismConfig]] = {} + @property + def working_dir(self) -> str | None: + """Base directory path for checkpointing intermediate mechanism state.""" + return None + + def with_working_dir(self, working_dir: str | None) -> MechanismConfig: + """Returns a copy of the config with working_dir set if supported and unset.""" + if self.working_dir is not None or working_dir is None: + return self + if dataclasses.is_dataclass(self): + try: + return dataclasses.replace( # pyrefly: ignore[bad-specialization] + self, working_dir=working_dir + ) + except (TypeError, ValueError): + return self + return self + 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 3d98c8f..c797a29 100644 --- a/dpsynth/data_generation_v3.py +++ b/dpsynth/data_generation_v3.py @@ -313,6 +313,7 @@ def __call__( m.compress(mappings, discrete.domain) # pyrefly: ignore[bad-argument-type] for m in initial_measurements ] + logging.info('[DPSynth]: Compressed discrete domain:\n%s', discrete.domain) cfg = self.config.discrete_mechanism if hasattr(cfg, 'supporting_cliques'): @@ -386,6 +387,8 @@ class TabularConfig(api.MechanismConfig): mbi.extensions.precompute_marginals) to compute marginals from Dataset. use_jax_for_generation: Whether to use JAX-accelerated generation (via mbi.extensions.synthetic_data) to generate synthetic data from the model. + 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 @@ -396,6 +399,7 @@ class TabularConfig(api.MechanismConfig): compress_columns: bool = False use_jax_for_bincount: bool = False use_jax_for_generation: bool = False + working_dir: str | None = None def _compute_per_col_deltas(self, domains, delta): # Split delta across open-set columns, analogous to splitting zcdp_rho. @@ -504,7 +508,11 @@ def configure( for col, init in inits.items() } - calibrated_discrete = self.discrete_mechanism.configure( + discrete_mechanism = self.discrete_mechanism.with_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/aim.py b/dpsynth/discrete_mechanisms/aim.py index adbafd7..d5eb714 100644 --- a/dpsynth/discrete_mechanisms/aim.py +++ b/dpsynth/discrete_mechanisms/aim.py @@ -61,7 +61,7 @@ def _filter_candidates( def _worst_approximated( rng: np.random.Generator, candidates: Mapping[mbi.Clique, float], - answers: mbi.CliqueVector, + data: mbi.Dataset | mbi.CliqueVector, estimates: mbi.CliqueVector, eps: float, sigma: float, @@ -72,7 +72,7 @@ def _worst_approximated( errors = {} for cl in candidates: wgt = candidates[cl] - diff = answers[cl].datavector() - estimates[cl].datavector() + diff = data.project(cl).datavector() - estimates[cl].datavector() bias = jnp.sqrt(2 / jnp.pi) * max_records_per_user * sigma * domain.size(cl) errors[cl] = wgt * (jnp.linalg.norm(diff, ord=1) - bias) @@ -171,13 +171,11 @@ def __call__( rho_per_round = self.zcdp_rho / max_rounds ######################################################################### - # Compile workload into candidate measurements, and precompute answers. # + # Compile workload into candidate measurements. # ######################################################################### candidates = common.compiled_workload( data.domain, self.config.workload, self.config.max_marginal_size ) - answers = mbi.CliqueVector.from_projectable(data, list(candidates)) # pyrefly: ignore[bad-argument-type] - logging.info('[AIM]: Calculated workload-query answers.') estimator = mbi.estimation.MirrorDescent(self.config.marginal_oracle) model = estimator.estimate( @@ -215,7 +213,7 @@ def __call__( marginal_query = _worst_approximated( rng, small_candidates, - answers, + data, estimates, epsilon, sigma, diff --git a/dpsynth/discrete_mechanisms/aim_gdp.py b/dpsynth/discrete_mechanisms/aim_gdp.py index 9a379df..44bc6f7 100644 --- a/dpsynth/discrete_mechanisms/aim_gdp.py +++ b/dpsynth/discrete_mechanisms/aim_gdp.py @@ -64,24 +64,22 @@ def expected_size(cl): def _compute_dp_errors( rng: np.random.Generator, - answers: mbi.CliqueVector, + data: mbi.Dataset | mbi.CliqueVector, estimates: mbi.CliqueVector, gdp_budget: float, - subset: Iterable[mbi.Clique] | None = None, + subset: Iterable[mbi.Clique], max_records_per_user: int = 1, ) -> dict[mbi.Clique, float]: """Compute L1 error between the model answers and the true answers with DP.""" - if subset is None: - subset = answers.cliques - + clique_list = list(subset) # The L1 error of a marginal changes by at most ``max_records_per_user`` when # a single user (contributing up to that many records) is added or removed. per_candidate_sigma = max_records_per_user * accounting.gdp_gaussian_sigma( - gdp_budget / len(subset) # pyrefly: ignore[bad-argument-type] + gdp_budget / len(clique_list) # pyrefly: ignore[bad-argument-type] ) result = {} - for cl in subset: - actual = answers[cl].datavector(flatten=True) + for cl in clique_list: + actual = data.project(cl).datavector(flatten=True) estimate = estimates[cl].datavector(flatten=True) error = jnp.linalg.norm(actual - estimate, ord=1) noise = rng.normal(loc=0, scale=per_candidate_sigma) @@ -93,7 +91,7 @@ def _worst_approximated( rng: np.random.Generator, candidates: Mapping[mbi.Clique, float], errors: dict[mbi.Clique, float], # will be updated in-place. - answers: mbi.CliqueVector, # derived from sensitive data. + data: mbi.Dataset | mbi.CliqueVector, # sensitive data. model: mbi.MarkovRandomField, select_budget: float, # satisfies select_budget-GDP. measure_sigma: float, @@ -118,10 +116,10 @@ def _worst_approximated( estimates = mbi.marginal_oracles.bulk_variable_elimination( model.potentials, subset, model.total # pyrefly: ignore[bad-argument-type] ) - # Only step that uses "answers", satisfies DP. + # Only step that uses "data", satisfies DP. current_errors = _compute_dp_errors( rng, - answers, + data, estimates, select_budget, subset, @@ -240,13 +238,11 @@ def __call__( budget_per_round = budget_remaining / max_rounds ######################################################################### - # Compile workload into candidate measurements, and precompute answers. # + # Compile workload into candidate measurements. # ######################################################################### candidates = common.compiled_workload( data.domain, self.config.workload, self.config.max_marginal_size ) - answers = mbi.CliqueVector.from_projectable(data, candidates) # pyrefly: ignore[bad-argument-type] - logging.info('[AIM] Calculated workload-query answers.') domain = data.domain estimator = mbi.estimation.MirrorDescent(self.config.marginal_oracle) @@ -297,7 +293,7 @@ def __call__( rng, candidates=small_candidates, errors=errors, - answers=answers, + data=data, model=model, select_budget=select_budget, measure_sigma=measure_sigma, diff --git a/dpsynth/discrete_mechanisms/common.py b/dpsynth/discrete_mechanisms/common.py index f2698a6..b51d558 100644 --- a/dpsynth/discrete_mechanisms/common.py +++ b/dpsynth/discrete_mechanisms/common.py @@ -151,7 +151,9 @@ def precompute_marginals( *, use_jax: bool = False, ) -> mbi.CliqueVector: - """Computes marginals over cliques from a dataset, optionally using JAX.""" + """Computes marginals over cliques from a Dataset, optionally using JAX.""" + if not cliques: + return mbi.CliqueVector(data.domain, [], {}) if use_jax: return mbi.extensions.precompute_marginals( data, cliques # pyrefly: ignore[bad-argument-type] @@ -231,8 +233,8 @@ def exponential_mechanism( def measure_marginals_with_noise( rng: np.random.Generator, - data: mbi.Projectable, - marginal_queries: list[tuple[str, ...]], + data: mbi.Dataset | mbi.CliqueVector, + marginal_queries: Sequence[mbi.Clique], gdp_sigma: float, weights: np.ndarray | None = None, max_records_per_user: int = 1, @@ -418,7 +420,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 = [tuple(cl) for cl in workload.keys()] else: @@ -456,7 +461,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 = {tuple(cl): 1.0 for cl in workload} @@ -477,7 +485,7 @@ def score(cl): def compute_independence_errors( - data: mbi.Projectable, + data: mbi.Dataset | mbi.CliqueVector, model: mbi.MarkovRandomField, cliques: Sequence[mbi.Clique], ) -> dict[mbi.Clique, float]: diff --git a/dpsynth/discrete_mechanisms/discrete.py b/dpsynth/discrete_mechanisms/discrete.py index 355f0ed..b0c0a12 100644 --- a/dpsynth/discrete_mechanisms/discrete.py +++ b/dpsynth/discrete_mechanisms/discrete.py @@ -26,8 +26,10 @@ 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 @@ -49,6 +51,8 @@ class DiscreteConfig(api.MechanismConfig): mbi.extensions.precompute_marginals) to compute marginals from Dataset. use_jax_for_generation: Whether to use JAX-accelerated generation (via mbi.extensions.synthetic_data) to generate synthetic data from the model. + working_dir: Base directory path for intermediate checkpoints. If None, + checkpointing is disabled. """ mechanism: api.MechanismConfig = mst.MSTConfig() @@ -57,14 +61,17 @@ class DiscreteConfig(api.MechanismConfig): constraints: Sequence[mbi.Constraint] = () use_jax_for_bincount: bool = False use_jax_for_generation: bool = False + working_dir: str | 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.with_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, @@ -133,20 +140,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( + '[DPSynth]: 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 = [] @@ -159,6 +176,7 @@ def __call__( if mappings and isinstance(data, mbi.Dataset): data = data.compress(mappings) # pyrefly: ignore[bad-argument-type] measurements = [m.compress(mappings, data.domain) for m in measurements] # pyrefly: ignore[bad-argument-type] + logging.info('[DPSynth]: Compressed discrete domain:\n%s', data.domain) cfg = self.config.mechanism if isinstance(data, mbi.Dataset) and hasattr(cfg, 'supporting_cliques'): diff --git a/dpsynth/discrete_mechanisms/swift.py b/dpsynth/discrete_mechanisms/swift.py index 126f1ad..4c7483f 100644 --- a/dpsynth/discrete_mechanisms/swift.py +++ b/dpsynth/discrete_mechanisms/swift.py @@ -26,15 +26,18 @@ from __future__ import annotations from collections.abc import Iterable, Mapping, Sequence +import concurrent.futures import dataclasses import functools import itertools import math import time +from typing import Any 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 @@ -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: str | None = None def supporting_cliques(self, domain: mbi.Domain) -> list[mbi.Clique]: """Returns the workload cliques filtered by max_marginal_size.""" @@ -98,23 +104,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[Any] | None, + concurrent.futures.Future[Any] | 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,9 +130,6 @@ 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] - with common.timed(phase_times, 'initial_mirror_descent'): estimator = mbi.estimation.MirrorDescent(self.config.marginal_oracle) model = estimator.estimate( @@ -135,15 +139,11 @@ 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, - answers, # pyrefly: ignore[bad-argument-type] + data, model, # pyrefly: ignore[bad-argument-type] list(candidates), select_gdp_budget, @@ -156,18 +156,16 @@ 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 + oracle = self.config.marginal_oracle or mbi.marginal_oracles.default_oracle( + all_cliques, data.domain, has_constraints=bool(constraints) ) + closed_oracle = functools.partial(oracle, jtree=jtree) estimator = mbi.estimation.MirrorDescent(marginal_oracle=closed_oracle) rows = mbi.estimation.minimum_variance_unbiased_total(initial_measurements) # pyrefly: ignore[bad-argument-type] rows = int(max(rows, 1)) @@ -180,46 +178,153 @@ 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( rng, - answers, + data, selected, measure_gdp_budget, 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 # - ######################################################## - with common.timed(phase_times, 'estimation'): - t0 = time.time() - pgm_future.result() - logging.info('[SWIFT] PGM precompile wait: %.2fs', time.time() - t0) + 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[Any] | None = None, + ) -> mbi.Model: + """Estimates the MRF model from measurements using MirrorDescent.""" + with common.timed(phase_times, 'estimation'): + if pgm_future is not None: + t0 = time.time() + pgm_future.result() + logging.info('[SWIFT] PGM precompile wait: %.2fs', time.time() - t0) + + all_cliques = list(jtree.nodes) + oracle = ( + self.config.marginal_oracle + or mbi.marginal_oracles.default_oracle( + all_cliques, domain, has_constraints=bool(constraints) + ) + ) + closed_oracle = functools.partial(oracle, 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[Any] | 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( - measurements=measurements, + synthetic_data=syn, + 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 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, ) @@ -332,7 +437,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 @@ -348,6 +453,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/dpsynth/postprocessing.py b/dpsynth/postprocessing.py index 80d07a3..570d6b9 100644 --- a/dpsynth/postprocessing.py +++ b/dpsynth/postprocessing.py @@ -140,7 +140,7 @@ def generate_synthetic_data_from_marginals( log: bool = False, nrows: int | None = None, estimator: mbi.Estimator | None = None, - marginal_oracle: mbi.marginal_oracles.MarginalOracle = mbi.marginal_oracles.message_passing_stable, # pyrefly: ignore[bad-function-definition] + marginal_oracle: mbi.marginal_oracles.MarginalOracle | None = None, exact_marginals: Sequence[pd.DataFrame] | None = None, cross_attribute_constraints: Sequence[constraints.Constraint] = (), extra_domain_elements: dict[str, Sequence[Any]] | None = None, diff --git a/tests/checkpoint_test.py b/tests/checkpoint_test.py new file mode 100644 index 0000000..89133ef --- /dev/null +++ b/tests/checkpoint_test.py @@ -0,0 +1,160 @@ +# 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 dataclasses +import pathlib +from absl.testing import absltest +from dpsynth import api +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')) + + def test_with_working_dir_propagates_if_supported(self): + @dataclasses.dataclass(frozen=True) + class MockConfig(api.MechanismConfig): + working_dir: epath.PathLike | None = None + + def configure(self, *args, **kwargs): + pass + + cfg = MockConfig() + self.assertIsNone(cfg.working_dir) + + updated = cfg.with_working_dir('/tmp/test_dir') + self.assertEqual(updated.working_dir, '/tmp/test_dir') + + # Existing working_dir is preserved (not overwritten). + preserved = updated.with_working_dir('/tmp/other_dir') + self.assertEqual(preserved.working_dir, '/tmp/test_dir') + + def test_with_working_dir_noop_on_unsupported_config(self): + @dataclasses.dataclass(frozen=True) + class MockConfigWithoutWorkingDir(api.MechanismConfig): + param: int = 42 + + def configure(self, *args, **kwargs): + pass + + cfg = MockConfigWithoutWorkingDir() + updated = cfg.with_working_dir('/tmp/test_dir') + self.assertIs(updated, cfg) + + def test_with_working_dir_none_returns_self(self): + @dataclasses.dataclass(frozen=True) + class MockConfig(api.MechanismConfig): + working_dir: epath.PathLike | None = None + + def configure(self, *args, **kwargs): + pass + + cfg = MockConfig() + self.assertIs(cfg.with_working_dir(None), cfg) + + +if __name__ == '__main__': + absltest.main() diff --git a/tests/data_generation_v3_test.py b/tests/data_generation_v3_test.py index 77e7755..167b9a1 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 @@ -578,6 +580,48 @@ def test_use_jax_for_generation_and_bincount(self): self.assertIsInstance(result.synthetic_data, pd.DataFrame) self.assertListEqual(result.synthetic_data.columns.tolist(), ['A', 'B']) + def test_mbi_callbacks_logging_configured(self): + if not hasattr(mbi, 'callbacks') or not hasattr( + mbi.callbacks, 'set_log_fn' + ): + self.skipTest( + 'mbi.callbacks.set_log_fn not supported in this mbi version' + ) + import dpsynth # pylint: disable=g-import-not-at-top,unused-import + + with self.assertLogs(level='INFO') as logs: + mbi.callbacks.log('test', 'message', sep=' | ') + self.assertTrue(any('test | message' in output for output in logs.output)) + + 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_mechanisms_test.py b/tests/discrete_mechanisms/discrete_mechanisms_test.py index 2e1e86a..3eeffd4 100644 --- a/tests/discrete_mechanisms/discrete_mechanisms_test.py +++ b/tests/discrete_mechanisms/discrete_mechanisms_test.py @@ -70,7 +70,7 @@ def test_mechanism_runs_on_precomputed_marginals(self, mechanism): calibrated = mechanism.configure(zcdp_rho=_ZCDP_RHO) cliques = mechanism.supporting_cliques(domain) - precomputed = mbi.CliqueVector.from_projectable(data, cliques) + precomputed = common.precompute_marginals(data, cliques) result = calibrated(rng, precomputed) self.assertIsInstance(result, common.DiscreteMechanismResult) diff --git a/tests/discrete_mechanisms/discrete_test.py b/tests/discrete_mechanisms/discrete_test.py index 0468ec9..905fe3a 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 @@ -158,6 +159,28 @@ def test_use_jax_for_bincount_and_generation(self): self.assertEqual(result.synthetic_data.domain, domain) self.assertEqual(result.synthetic_data.records, 200) + 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..251881e 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,72 @@ 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 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_measurements = epath.Path(temp_dir) / 'measurements.npz' + ckpt_model = epath.Path(temp_dir) / 'model.npz' + + 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_measurements(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 to produce measurements and model + 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 model, keep measurements + (epath.Path(temp_dir) / 'model.npz').unlink() + + # Second run: should reuse measurements 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()