From 88bff240dd86c8cee0bfa87e223a943d8d54e81a Mon Sep 17 00:00:00 2001 From: Ryan McKenna Date: Sat, 10 Oct 2026 18:54:09 -0700 Subject: [PATCH] Pass workload at configuration time and introduce canonical Workload class. - Add dpsynth.Workload (mapping cliques to positive float weights, normalizing unweighted sequences to 1.0) with schema validation (error on unknown attributes, warning on uncovered attributes) and YAML serialization support. - Move workload from DirectConfig.prespecified_marginal_queries, AIMConfig.workload, and SWIFTConfig.workload to configure(..., workload=...) and dpsynth.calibrate(..., workload=...). - Move supporting_cliques onto the calibrated discrete mechanisms. PiperOrigin-RevId: 997260163 --- docs/api_reference.rst | 2 + dpsynth/__init__.py | 2 + dpsynth/_calibration.py | 19 ++-- dpsynth/adapters/beam.py | 20 ++-- dpsynth/api.py | 9 +- dpsynth/data_generation_v3.py | 25 +++-- dpsynth/discrete_mechanisms/README.md | 24 ++--- dpsynth/discrete_mechanisms/aim.py | 32 +++--- dpsynth/discrete_mechanisms/direct.py | 38 ++++--- dpsynth/discrete_mechanisms/discrete.py | 36 +++++-- dpsynth/discrete_mechanisms/independent.py | 7 +- dpsynth/discrete_mechanisms/mst.py | 7 +- dpsynth/discrete_mechanisms/swift.py | 30 +++--- dpsynth/domain.py | 98 ++++++++++++++++++- dpsynth/serialize.py | 33 +++++++ tests/adapters/beam_test.py | 13 ++- tests/data_generation_v3_test.py | 28 ++++-- tests/discrete_mechanisms/aim_test.py | 4 +- tests/discrete_mechanisms/direct_test.py | 16 ++- .../discrete_mechanisms_test.py | 22 +++-- tests/domain_test.py | 49 ++++++++++ tests/serialize_test.py | 37 ++++++- 22 files changed, 424 insertions(+), 127 deletions(-) diff --git a/docs/api_reference.rst b/docs/api_reference.rst index 751f23dc..91fefd13 100644 --- a/docs/api_reference.rst +++ b/docs/api_reference.rst @@ -48,6 +48,8 @@ argument to :class:`~dpsynth.TabularConfig`. NumericalAttribute OpenSetCategoricalAttribute FreeFormTextAttribute + Schema + Workload ---- diff --git a/dpsynth/__init__.py b/dpsynth/__init__.py index 88777805..8dc4e786 100644 --- a/dpsynth/__init__.py +++ b/dpsynth/__init__.py @@ -38,6 +38,7 @@ from dpsynth.domain import NumericalAttribute from dpsynth.domain import OpenSetCategoricalAttribute from dpsynth.domain import Schema +from dpsynth.domain import Workload from dpsynth.reporting import PrivacyReport from dpsynth.serialize import from_yaml from dpsynth.serialize import to_yaml @@ -70,6 +71,7 @@ 'TabularConfig', 'TabularMechanism', 'TabularSynthesizer', + 'Workload', 'api', 'calibrate', 'checkpoint', diff --git a/dpsynth/_calibration.py b/dpsynth/_calibration.py index 8ccb9639..bff89953 100644 --- a/dpsynth/_calibration.py +++ b/dpsynth/_calibration.py @@ -72,6 +72,7 @@ def calibrate( *, epsilon: float, delta: float, + workload: Any = None, delta_split: float = 0.5, poisson_sampling_prob: float = 1.0, group_size: int = 1, @@ -89,6 +90,8 @@ def calibrate( domain: Optional domain specification, forwarded to ``config.configure()``. epsilon: Target epsilon for (epsilon, delta)-DP. delta: Target delta for (epsilon, delta)-DP. + workload: Optional workload specification, forwarded to + ``config.configure()``. delta_split: Fraction of ``delta`` passed to ``config.configure()`` for sub-mechanisms that consume approximate DP budget directly (e.g. open-set partition selection). Defaults to 0.5. @@ -123,12 +126,12 @@ def calibrate( if group_size < 1: raise ValueError(f'group_size < 1: {group_size}.') + configure_kwargs: dict[str, Any] = {'delta': delta * delta_split} + if workload is not None: + configure_kwargs['workload'] = workload + def make_event_fn(rho: float) -> dp_accounting.DpEvent: - base = config.configure( - domain, - budget=rho, - delta=delta * delta_split, - ).dp_event + base = config.configure(domain, budget=rho, **configure_kwargs).dp_event base = with_group_size(base, group_size) sampled = dp_accounting.PoissonSampledDpEvent(poisson_sampling_prob, base) return base if poisson_sampling_prob == 1.0 else sampled @@ -169,8 +172,4 @@ def make_event_fn(rho: float) -> dp_accounting.DpEvent: ) optimal_rho = max(rhos.values()) - return config.configure( - domain, - budget=optimal_rho, - delta=delta * delta_split, - ) + return config.configure(domain, budget=optimal_rho, **configure_kwargs) diff --git a/dpsynth/adapters/beam.py b/dpsynth/adapters/beam.py index add3bc3c..bcd7c988 100644 --- a/dpsynth/adapters/beam.py +++ b/dpsynth/adapters/beam.py @@ -481,8 +481,8 @@ def _run_two_pass( column_measurements, synth.schema ).mbi_domain - assert hasattr(synth.config.discrete_mechanism, 'supporting_cliques') - workload = synth.config.discrete_mechanism.supporting_cliques(mbi_domain) + assert hasattr(synth.base_mechanism, 'supporting_cliques') + workload = synth.base_mechanism.supporting_cliques(mbi_domain) # Pass 2: compute the marginal workload. with beam.Pipeline(**pipeline_kwargs) as p: @@ -558,20 +558,20 @@ class BeamTabularConfig(api.MechanismConfig): temp_location: str | None = None pipeline_options: beam.options.pipeline_options.PipelineOptions | None = None - def __post_init__(self): - if not hasattr(self.synthesizer.discrete_mechanism, 'supporting_cliques'): - raise ValueError( - 'self.synthesizer.discrete_mechanism must have a supporting_cliques' - ' method.' - ) - - def configure(self, schema=None, *, budget, delta=0) -> BeamTabularMechanism: + def configure( + self, schema=None, *, budget, delta=0, workload=None + ) -> BeamTabularMechanism: """Returns a copy whose synthesizer is configured with the given budget.""" synthesizer = self.synthesizer.configure( schema, budget=budget, delta=delta, + workload=workload, ) + if not hasattr(synthesizer.base_mechanism, 'supporting_cliques'): + raise ValueError( + 'synthesizer.base_mechanism must have a supporting_cliques method.' + ) return BeamTabularMechanism( synthesizer=synthesizer, temp_location=self.temp_location, diff --git a/dpsynth/api.py b/dpsynth/api.py index e9f0eaa2..39698c2c 100644 --- a/dpsynth/api.py +++ b/dpsynth/api.py @@ -140,7 +140,9 @@ def get_subclass(cls, name: str) -> type[MechanismConfig] | None: return cls._registry.get(name) @abc.abstractmethod - def configure(self, domain=None, *, budget, delta=0) -> CalibratedMechanism: + def configure( + self, domain=None, *, budget, delta=0, workload=None + ) -> CalibratedMechanism: """Returns a calibrated mechanism for the given dummy budget. Converts the budget into the mechanism's natural privacy parameter @@ -166,6 +168,9 @@ def configure(self, domain=None, *, budget, delta=0) -> CalibratedMechanism: delta: Approximate DP delta consumed by the mechanism itself (e.g., for thresholding). Defaults to 0 (pure zCDP). Mechanisms that need delta should raise if it is 0. + workload: Optional workload specification (e.g., a ``dpsynth.Workload``, + sequence of attribute tuples, or mapping from attribute tuples to + weights). Mechanisms that do not use a workload ignore this argument. Returns: A calibrated, runnable mechanism. @@ -178,6 +183,7 @@ def calibrate( *, epsilon: float, delta: float, + workload: Any = None, delta_split: float = 0.5, poisson_sampling_prob: float = 1.0, group_size: int = 1, @@ -197,6 +203,7 @@ def calibrate( domain, epsilon=epsilon, delta=delta, + workload=workload, delta_split=delta_split, poisson_sampling_prob=poisson_sampling_prob, group_size=group_size, diff --git a/dpsynth/data_generation_v3.py b/dpsynth/data_generation_v3.py index bf3fa2c1..a8b2dc23 100644 --- a/dpsynth/data_generation_v3.py +++ b/dpsynth/data_generation_v3.py @@ -323,9 +323,8 @@ def from_summary( ] logging.info('[DPSynth]: Compressed discrete domain:\n%s', data.domain) - cfg = self.config.discrete_mechanism - if hasattr(cfg, 'supporting_cliques'): - cliques = cfg.supporting_cliques(data.domain) + if hasattr(self.base_mechanism, 'supporting_cliques'): + cliques = self.base_mechanism.supporting_cliques(data.domain) data = _checkpoint.get_or_compute( 'precomputed_marginals', dm_common.precompute_marginals, @@ -443,6 +442,7 @@ def configure( *, budget: float, delta: float = 0.0, + workload: domain.WorkloadInput | None = None, ) -> TabularMechanism: """Returns a calibrated mechanism configured with the given privacy budget. @@ -460,12 +460,17 @@ def configure( delta: Approximate DP delta allocated to partition selection for open-set columns (split evenly across open-set columns). Must be positive when open-set categorical attributes are present. + workload: Optional workload specification (e.g. a ``dpsynth.Workload``, + sequence of attribute tuples, or mapping from attribute tuples to + weights) validated against ``schema`` and forwarded to the discrete + mechanism. Returns: A calibrated TabularMechanism ready to be run on tabular data. Raises: - ValueError: If open-set attributes exist but delta is 0. + ValueError: If open-set attributes exist but delta is 0, or if the + workload contains attributes not present in ``schema``. """ if schema is not None and isinstance(schema, domain.Schema): pass @@ -483,6 +488,7 @@ def configure( ' construction time.' ) + resolved_workload = domain.Workload.from_any(workload, schema) per_col_deltas = self._compute_per_col_deltas(schema, delta) inits = create_initializers( @@ -503,9 +509,14 @@ def configure( for col, init in inits.items() } - calibrated_discrete = self.discrete_mechanism.configure( - budget=discrete_rho, - ) + if resolved_workload is not None: + calibrated_discrete = self.discrete_mechanism.configure( + schema, budget=discrete_rho, workload=resolved_workload + ) + else: + calibrated_discrete = self.discrete_mechanism.configure( + schema, budget=discrete_rho + ) return TabularMechanism( config=self, diff --git a/dpsynth/discrete_mechanisms/README.md b/dpsynth/discrete_mechanisms/README.md index c9380bc2..f27af4fe 100644 --- a/dpsynth/discrete_mechanisms/README.md +++ b/dpsynth/discrete_mechanisms/README.md @@ -51,17 +51,17 @@ does not model relationships between columns. **Public API:** `IndependentConfig` -## `direct.py` — Caller-Defined Workload +## `direct.py` - Caller-Defined Workload -Implements `DirectConfig`, which measures the cliques supplied through -`prespecified_marginal_queries`. It performs no data-dependent selection and -does not create its own one-way measurements, so the full budget is available -for the specified workload. Initial measurements supplied by another layer are -included when fitting the final model. +Implements `DirectConfig`, which measures the cliques supplied through the +`workload` parameter of `configure()` / `dpsynth.calibrate()`. It performs no +data-dependent selection and does not create its own one-way measurements, so +the full budget is available for the specified workload. Initial measurements +supplied by another layer are included when fitting the final model. -**Public API:** `DirectConfig(prespecified_marginal_queries=...)` +**Public API:** `DirectConfig` -## `mst.py` — Private Pairwise Spanning-Tree Selection +## `mst.py` - Private Pairwise Spanning-Tree Selection Implements `MSTConfig`, the default general-purpose mechanism for preserving pairwise relationships. It begins with one-way marginals, privately selects @@ -75,27 +75,27 @@ selection and measurement; `_select()` calls the spanning-tree selection logic. `dp_maximum_spanning_tree()` and `_select_two_way_marginal_queries()` implement the private pairwise-selection step. -## `aim.py` — Adaptive Iterative Selection +## `aim.py` - Adaptive Iterative Selection Implements `AIMConfig`, an adaptive workload-based mechanism. Instead of selecting cliques once, it repeatedly finds a marginal that the current model approximates poorly, measures it, and updates the model. -**Public API:** `AIMConfig(workload=...)` +**Public API:** `AIMConfig` **Internal behavior:** `_one_way_cliques()` limits initial measurements to the workload; `_allocate_budget()` reserves rho for the adaptive loop; `_run()` replaces the standard base execution path. Helper functions filter valid candidates and privately choose the worst-approximated marginal. -## `swift.py` — Workload and Clique-Tree Mechanism +## `swift.py` - Workload and Clique-Tree Mechanism Implements `SWIFT`, a workload-informed mechanism that selects marginals while controlling clique-tree complexity. It uses a custom junction-tree-aware estimation and sampling path rather than the standard one-pass implementation in `base.py`. -**Public API:** `SWIFTConfig(workload=...)` +**Public API:** `SWIFTConfig` **Internal behavior:** `_allocate_budget()` splits rho between selection and measurement; `_run()` compiles the workload, selects supported cliques, builds a diff --git a/dpsynth/discrete_mechanisms/aim.py b/dpsynth/discrete_mechanisms/aim.py index c6931b35..7d5007df 100644 --- a/dpsynth/discrete_mechanisms/aim.py +++ b/dpsynth/discrete_mechanisms/aim.py @@ -14,12 +14,13 @@ """Implementation of the Adaptive+Iterative Mechanism (AIM).""" -from collections.abc import Iterable, Mapping +from collections.abc import Mapping from collections.abc import Sequence import dataclasses from absl import logging import dp_accounting from dpsynth import api +from dpsynth import domain as domain_lib from dpsynth.discrete_mechanisms import accounting from dpsynth.discrete_mechanisms import common from dpsynth.local_mode import primitives @@ -104,8 +105,6 @@ class AIMConfig(api.MechanismConfig): max_model_size >= 80. Attributes: - workload: A collection of marginal queries (and weights) the synthetic data - should be tailored to. max_rounds: The maximum number of rounds to run the mechanism. max_model_size: The maximum size of the graphical model in megabytes. Controls the utility/runtime trade-off. @@ -115,7 +114,6 @@ class AIMConfig(api.MechanismConfig): selecting two-way marginal queries. """ - workload: Mapping[mbi.Clique, float] | Iterable[mbi.Clique] | None = None max_rounds: int | None = None max_model_size: int = 80 max_marginal_size: float = 1e6 @@ -124,16 +122,19 @@ class AIMConfig(api.MechanismConfig): pgm_iters: int = 1000 marginal_oracle: mbi.MarginalOracle | None = None - def supporting_cliques(self, domain: mbi.Domain) -> list[mbi.Clique]: - """Returns the workload cliques filtered by max_marginal_size.""" - return common.supporting_cliques( - domain, self.workload, self.max_marginal_size - ) - - def configure(self, _=None, *, budget, delta=0): + def configure( + self, + domain=None, + *, + budget: float, + delta: float = 0.0, + workload: domain_lib.WorkloadInput | None = None, + ) -> 'AIM': + del delta return AIM( config=self, zcdp_rho=budget, + workload=domain_lib.Workload.from_any(workload, domain), ) @@ -143,6 +144,13 @@ class AIM(api.CalibratedMechanism): config: AIMConfig zcdp_rho: float + workload: domain_lib.Workload | None = None + + def supporting_cliques(self, domain: mbi.Domain) -> list[mbi.Clique]: + """Returns the workload cliques filtered by max_marginal_size.""" + return common.supporting_cliques( + domain, self.workload, self.config.max_marginal_size + ) @property def dp_event(self) -> dp_accounting.DpEvent: @@ -171,7 +179,7 @@ def __call__( # Compile workload into candidate measurements. # ######################################################################### candidates = common.compiled_workload( - data.domain, self.config.workload, self.config.max_marginal_size + data.domain, self.workload, self.config.max_marginal_size ) estimator = mbi.estimation.MirrorDescent(self.config.marginal_oracle) diff --git a/dpsynth/discrete_mechanisms/direct.py b/dpsynth/discrete_mechanisms/direct.py index 407a650f..72463cb2 100644 --- a/dpsynth/discrete_mechanisms/direct.py +++ b/dpsynth/discrete_mechanisms/direct.py @@ -19,6 +19,7 @@ from absl import logging import dp_accounting from dpsynth import api +from dpsynth import domain as domain_lib from dpsynth.discrete_mechanisms import accounting from dpsynth.discrete_mechanisms import common import mbi @@ -29,23 +30,24 @@ class DirectConfig(api.MechanismConfig): """Config for the direct mechanism that measures prespecified marginals.""" - def configure(self, _=None, *, budget, delta=0): - return Direct( - config=self, - gdp_budget=accounting.zcdp_to_gdp(budget), - ) - estimator: mbi.Estimator = mbi.estimation.MirrorDescent() marginal_oracle: mbi.MarginalOracle | None = None pgm_iters: int = 5000 - prespecified_marginal_queries: list[tuple[str, ...]] = dataclasses.field( - default_factory=list - ) - def supporting_cliques(self, domain: mbi.Domain) -> list[mbi.Clique]: - """Returns the prespecified marginal queries.""" - del domain # Unused. - return list(self.prespecified_marginal_queries) + def configure( + self, + domain=None, + *, + budget: float, + delta: float = 0.0, + workload: domain_lib.WorkloadInput | None = None, + ) -> 'Direct': + del delta + return Direct( + config=self, + gdp_budget=accounting.zcdp_to_gdp(budget), + workload=domain_lib.Workload.from_any(workload, domain), + ) @dataclasses.dataclass(frozen=True) @@ -54,6 +56,12 @@ class Direct(api.CalibratedMechanism): config: DirectConfig gdp_budget: float + workload: domain_lib.Workload | None = None + + def supporting_cliques(self, domain: mbi.Domain) -> list[mbi.Clique]: + """Returns the prespecified marginal queries.""" + del domain # Unused. + return list(self.workload.keys()) if self.workload else [] @property def dp_event(self) -> dp_accounting.DpEvent: @@ -73,7 +81,8 @@ def __call__( """Selects, measures, estimates, and generates in the compressed domain.""" common.validate_initial_measurements(initial_measurements) phase_times = {} - selected = list(self.config.prespecified_marginal_queries) + selected = list(self.workload.keys()) if self.workload else [] + weights = np.array(list(self.workload.values())) if self.workload else None all_cliques = [m.clique for m in initial_measurements] + list(selected) summary = mbi.summarize(data.domain, all_cliques) @@ -94,6 +103,7 @@ def __call__( data=data, # pyrefly: ignore[bad-argument-type] marginal_queries=selected, gdp_sigma=accounting.gdp_gaussian_sigma(self.gdp_budget), + weights=weights, ) measurements = list(initial_measurements) + new_measurements diff --git a/dpsynth/discrete_mechanisms/discrete.py b/dpsynth/discrete_mechanisms/discrete.py index e51ad838..d830b5ac 100644 --- a/dpsynth/discrete_mechanisms/discrete.py +++ b/dpsynth/discrete_mechanisms/discrete.py @@ -29,6 +29,7 @@ from absl import logging import dp_accounting from dpsynth import api +from dpsynth import domain as domain_lib from dpsynth.discrete_mechanisms import accounting from dpsynth.discrete_mechanisms import common from dpsynth.discrete_mechanisms import mst @@ -59,14 +60,26 @@ class DiscreteConfig(api.MechanismConfig): use_jax_for_bincount: bool = False use_jax_for_generation: bool = False - def configure(self, _=None, *, budget, delta=0): + def configure( + self, + domain=None, + *, + budget: float, + delta: float = 0.0, + workload: domain_lib.WorkloadInput | None = None, + ) -> DiscreteMechanism: """Configures the synthesizer with a zCDP budget.""" + resolved_workload = domain_lib.Workload.from_any(workload, domain) one_way_rho = budget * self.one_way_budget_fraction remaining_rho = budget * (1 - self.one_way_budget_fraction) - inner = self.mechanism.configure( - budget=remaining_rho, - delta=delta, - ) + if resolved_workload is not None: + inner = self.mechanism.configure( + domain, budget=remaining_rho, delta=delta, workload=resolved_workload + ) + else: + inner = self.mechanism.configure( + domain, budget=remaining_rho, delta=delta + ) return DiscreteMechanism( config=self, base_mechanism=inner, @@ -82,6 +95,12 @@ class DiscreteMechanism(api.CalibratedMechanism): base_mechanism: api.CalibratedMechanism one_way_gdp_budget: float + def supporting_cliques(self, domain: mbi.Domain) -> list[mbi.Clique]: + """Returns the cliques needed by the inner mechanism.""" + if hasattr(self.base_mechanism, 'supporting_cliques'): + return self.base_mechanism.supporting_cliques(domain) + return [(a,) for a in domain.attributes] + @property def dp_event(self) -> dp_accounting.DpEvent: """Composes one-way measurement event with the inner mechanism's event.""" @@ -155,9 +174,10 @@ def __call__( 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'): - cliques = cfg.supporting_cliques(data.domain) + if isinstance(data, mbi.Dataset) and hasattr( + self.base_mechanism, 'supporting_cliques' + ): + cliques = self.base_mechanism.supporting_cliques(data.domain) data = common.precompute_marginals( data, cliques, # pyrefly: ignore[bad-argument-type] diff --git a/dpsynth/discrete_mechanisms/independent.py b/dpsynth/discrete_mechanisms/independent.py index 77da78ca..80fdb431 100644 --- a/dpsynth/discrete_mechanisms/independent.py +++ b/dpsynth/discrete_mechanisms/independent.py @@ -29,7 +29,8 @@ class IndependentConfig(api.MechanismConfig): pgm_iters: int = 5000 - def configure(self, _=None, *, budget, delta=0.0): + def configure(self, _=None, *, budget, delta=0.0, workload=None): + del budget, delta, workload return Independent(config=self) def supporting_cliques(self, domain: mbi.Domain) -> list[mbi.Clique]: @@ -43,6 +44,10 @@ class Independent(api.CalibratedMechanism): config: IndependentConfig + def supporting_cliques(self, domain: mbi.Domain) -> list[mbi.Clique]: + """Returns the one-way marginals this mechanism expects to process.""" + return self.config.supporting_cliques(domain) + @property def dp_event(self) -> dp_accounting.DpEvent: """Returns a zero-cost DP event (no new measurements).""" diff --git a/dpsynth/discrete_mechanisms/mst.py b/dpsynth/discrete_mechanisms/mst.py index 36edb755..20271a28 100644 --- a/dpsynth/discrete_mechanisms/mst.py +++ b/dpsynth/discrete_mechanisms/mst.py @@ -188,7 +188,8 @@ def supporting_cliques(self, domain: mbi.Domain) -> list[mbi.Clique]: self.maximum_marginal_size, ) - def configure(self, _=None, *, budget, delta=0): + def configure(self, _=None, *, budget, delta=0, workload=None): + del delta, workload return MST( config=self, zcdp_rho=budget, @@ -202,6 +203,10 @@ class MST(api.CalibratedMechanism): config: MSTConfig zcdp_rho: float + def supporting_cliques(self, domain: mbi.Domain) -> list[mbi.Clique]: + """Returns all pairwise marginals within the size limit.""" + return self.config.supporting_cliques(domain) + @property def dp_event(self) -> dp_accounting.DpEvent: """Returns the DP event for the MST mechanism.""" diff --git a/dpsynth/discrete_mechanisms/swift.py b/dpsynth/discrete_mechanisms/swift.py index 9c1d000e..b8541cb2 100644 --- a/dpsynth/discrete_mechanisms/swift.py +++ b/dpsynth/discrete_mechanisms/swift.py @@ -37,6 +37,7 @@ import dp_accounting from dpsynth import _checkpoint from dpsynth import api +from dpsynth import domain as domain_lib from dpsynth.discrete_mechanisms import accounting from dpsynth.discrete_mechanisms import clique_tree from dpsynth.discrete_mechanisms import common @@ -56,8 +57,6 @@ class SWIFTConfig(api.MechanismConfig): prevent long compilation times. Attributes: - workload: The set of marginals to consider for the mechanism. Can be a - mapping from cliques to their weights or just an iterable of cliques. max_clique_size: The maximum size (domain product) allowed for any clique in the junction tree. This is the main knob to tune to improve utility for a given compute cost. @@ -70,7 +69,6 @@ class SWIFTConfig(api.MechanismConfig): marginals to measure. """ - workload: Mapping[mbi.Clique, float] | Iterable[mbi.Clique] | None = None max_clique_size: float = 1e7 max_marginal_size: float = 1e6 pgm_iters: int = 10_000 @@ -80,16 +78,19 @@ class SWIFTConfig(api.MechanismConfig): use_jax_for_bincount: bool = True use_jax_for_generation: bool = True - def supporting_cliques(self, domain: mbi.Domain) -> list[mbi.Clique]: - """Returns the workload cliques filtered by max_marginal_size.""" - return common.supporting_cliques( - domain, self.workload, self.max_marginal_size - ) - - def configure(self, _=None, *, budget, delta=0): + def configure( + self, + domain=None, + *, + budget: float, + delta: float = 0.0, + workload: domain_lib.WorkloadInput | None = None, + ) -> SWIFT: + del delta return SWIFT( config=self, gdp_budget=accounting.zcdp_to_gdp(budget), + workload=domain_lib.Workload.from_any(workload, domain), ) @@ -99,6 +100,13 @@ class SWIFT(api.CalibratedMechanism): config: SWIFTConfig gdp_budget: float + workload: domain_lib.Workload | None = None + + def supporting_cliques(self, domain: mbi.Domain) -> list[mbi.Clique]: + """Returns the workload cliques filtered by max_marginal_size.""" + return common.supporting_cliques( + domain, self.workload, self.config.max_marginal_size + ) @property def dp_event(self) -> dp_accounting.DpEvent: @@ -126,7 +134,7 @@ def select_and_measure_queries( with common.timed(phase_times, 'compiled_workload'): candidates = common.compiled_workload( data.domain, - self.config.workload, + self.workload, self.config.max_marginal_size, ) logging.info('[SWIFT] %d candidates.', len(candidates)) diff --git a/dpsynth/domain.py b/dpsynth/domain.py index 6445e13d..135bd9ed 100644 --- a/dpsynth/domain.py +++ b/dpsynth/domain.py @@ -43,7 +43,7 @@ values when none should exist. """ -from collections.abc import Iterator, Mapping, Sequence +from collections.abc import Iterable, Iterator, Mapping, Sequence import dataclasses import functools import math @@ -357,6 +357,102 @@ def items(self): return self.attributes.items() +@dataclasses.dataclass(frozen=True, init=False) +class Workload(Mapping[tuple[str, ...], float]): + """Canonical representation of a marginal query workload. + + Accepts either a mapping from attribute tuples to positive float weights, or + an iterable of attribute sequences (normalized to unit weight 1.0). + + Attributes: + cliques: Mapping from attribute tuple to positive float weight. + """ + + cliques: Mapping[tuple[str, ...], float] + + def __init__( + self, + cliques: Mapping[Sequence[str], float] | Iterable[Sequence[str]] = (), + ): + if isinstance(cliques, Mapping): + raw_items = cliques.items() + else: + raw_items = ((cl, 1.0) for cl in cliques) + + normalized: dict[tuple[str, ...], float] = {} + for cl, weight in raw_items: + if isinstance(cl, str) or not isinstance(cl, Sequence): + raise ValueError(f'Clique must be a sequence of strings, got {cl!r}.') + clique = tuple(cl) + if not clique: + raise ValueError('Workload cliques must be non-empty.') + if not all(isinstance(col, str) for col in clique): + raise ValueError(f'Clique elements must be strings, got {clique!r}.') + if len(set(clique)) != len(clique): + raise ValueError(f'Clique has duplicate attributes: {clique!r}.') + w = float(weight) + if not math.isfinite(w) or w <= 0: + raise ValueError(f'Weight for {clique!r} must be positive: {weight}.') + normalized[clique] = w + + object.__setattr__(self, 'cliques', normalized) + + @classmethod + def from_any( + cls, + workload: 'WorkloadInput | None', + attributes: Iterable[str] | None = None, + ) -> 'Workload | None': + """Coerces a raw workload into a Workload and optionally validates it.""" + if workload is None: + return None + resolved = workload if isinstance(workload, cls) else cls(workload) + if attributes is not None: + resolved.validate(attributes) + return resolved + + def __getitem__(self, key: tuple[str, ...]) -> float: + return self.cliques[key] + + def __iter__(self) -> Iterator[tuple[str, ...]]: + return iter(self.cliques) + + def __len__(self) -> int: + return len(self.cliques) + + def keys(self): + return self.cliques.keys() + + def values(self): + return self.cliques.values() + + def items(self): + return self.cliques.items() + + def validate(self, attributes: Iterable[str]) -> None: + """Validates that the workload is consistent with schema attributes. + + Args: + attributes: Collection of attribute names in the schema or domain. + + Raises: + ValueError: If any clique references an attribute not in `attributes`. + """ + attr_set = set(attributes) + workload_attrs = set().union(*self.keys()) if self.cliques else set() + unknown = workload_attrs - attr_set + if unknown: + raise ValueError(f'Unknown workload attributes: {sorted(unknown)}.') + uncovered = attr_set - workload_attrs + if uncovered: + logging.warning('Attributes not in workload: %s', sorted(uncovered)) + + +WorkloadInput: TypeAlias = ( + Workload | Mapping[Sequence[str], float] | Iterable[Sequence[str]] +) + + def to_yaml_file(domain: Mapping[str, AttributeType], filepath: str | PathType): """Writes a dictionary of Attribute objects to a YAML file.""" warnings.warn( diff --git a/dpsynth/serialize.py b/dpsynth/serialize.py index 8e3f59ee..7ad0d554 100644 --- a/dpsynth/serialize.py +++ b/dpsynth/serialize.py @@ -57,6 +57,33 @@ def _resolve_type(type_name: str) -> type[Any] | None: return None +def _unstructure_workload(obj: domain.Workload) -> dict[str, Any]: + cliques = [list(cl) for cl in obj.keys()] + weights = list(obj.values()) + result: dict[str, Any] = {'type': 'Workload', 'cliques': cliques} + if any(w != 1.0 for w in weights): + result['weights'] = weights + return result + + +def _structure_workload(data: Any, _: Any) -> domain.Workload: + """Structures a YAML representation into a domain.Workload.""" + if isinstance(data, domain.Workload): + return data + if isinstance(data, Sequence) and not isinstance(data, str): + return domain.Workload(data) + if isinstance(data, Mapping) and 'cliques' in data: + cliques = data['cliques'] + weights = data.get('weights') + if weights is not None: + if len(cliques) != len(weights): + raise ValueError('Workload cliques and weights must have same length.') + weighted = {tuple(cl): float(w) for cl, w in zip(cliques, weights)} + return domain.Workload(weighted) + return domain.Workload(cliques) + raise ValueError(f'Cannot structure {data!r} as Workload.') + + def _unstructure_dataclass(cl: type[Any], conv: cattrs.Converter) -> Any: base_fn = cattrs.gen.make_dict_unstructure_fn( cl, conv, _cattrs_omit_if_default=True @@ -67,6 +94,8 @@ def _unstructure_dataclass(cl: type[Any], conv: cattrs.Converter) -> Any: def _structure_polymorphic(data: Any, _: Any, conv: cattrs.Converter) -> Any: if isinstance(data, Mapping) and 'type' in data: cls = _resolve_type(data['type']) + if cls is domain.Workload: + return _structure_workload(data, cls) if cls is not None: return cattrs.gen.make_dict_structure_fn(cls, conv)(data, cls) raise ValueError(f"Unknown type: '{data['type']}'") @@ -83,6 +112,7 @@ def _make_converter() -> cattrs.Converter: conv.register_unstructure_hook_factory( dataclasses.is_dataclass, lambda cl: _unstructure_dataclass(cl, conv) ) + conv.register_unstructure_hook(domain.Workload, _unstructure_workload) conv.register_unstructure_hook( api.MechanismConfig, lambda obj: _unstructure_dataclass(obj.__class__, conv)(obj), @@ -105,6 +135,7 @@ def _make_converter() -> cattrs.Converter: ) # 2. Polymorphic structuring for abstract base classes and unions + conv.register_structure_hook(domain.Workload, _structure_workload) conv.register_structure_hook( api.MechanismConfig, lambda data, _: _structure_polymorphic(data, _, conv) ) @@ -201,6 +232,8 @@ def from_yaml( if isinstance(data, Mapping) and 'type' in data: type_name = data['type'] target_cls = _resolve_type(type_name) + if target_cls is domain.Workload: + return _structure_workload(data, target_cls) if target_cls is not None: return cattrs.gen.make_dict_structure_fn(target_cls, converter)( data, target_cls diff --git a/tests/adapters/beam_test.py b/tests/adapters/beam_test.py index b3066738..45bdb178 100644 --- a/tests/adapters/beam_test.py +++ b/tests/adapters/beam_test.py @@ -381,20 +381,19 @@ def test_end_to_end_mixed_types(self): self.assertCountEqual(result.synthetic_data.columns, ['age', 'grade']) @parameterized.named_parameters( - ('mst', discrete_mechanisms.MSTConfig(pgm_iters=250)), + ('mst', discrete_mechanisms.MSTConfig(pgm_iters=250), None), ( 'independent', discrete_mechanisms.IndependentConfig(), + None, ), ( 'direct', - discrete_mechanisms.DirectConfig( - prespecified_marginal_queries=[('a',), ('b',), ('a', 'b')], - pgm_iters=250, - ), + discrete_mechanisms.DirectConfig(pgm_iters=250), + [('a',), ('b',), ('a', 'b')], ), ) - def test_runs_across_mechanisms(self, mechanism): + def test_runs_across_mechanisms(self, mechanism, workload): """The pipeline generalizes to any mechanism via supporting_cliques.""" domains = { 'a': domain.CategoricalAttribute(possible_values=['x', 'y']), @@ -402,7 +401,7 @@ def test_runs_across_mechanisms(self, mechanism): } synth = data_generation_v3.TabularConfig(discrete_mechanism=mechanism) beam_synth = beam_adapter.BeamTabularConfig(synth).configure( - domains, budget=100.0 + domains, budget=100.0, workload=workload ) rows = [ {'a': 'x', 'b': 'p'}, diff --git a/tests/data_generation_v3_test.py b/tests/data_generation_v3_test.py index 9ff26deb..9a3523d4 100644 --- a/tests/data_generation_v3_test.py +++ b/tests/data_generation_v3_test.py @@ -73,7 +73,9 @@ def _discrete_workload_mechanism_baseline_errors( rng = np.random.default_rng(0) data = _make_discrete_data(rng) - mechanism_result = config.configure(budget=budget)(rng, data) + mechanism_result = config.configure(budget=budget, workload=workload)( + rng, data + ) baseline_result = baseline_config.configure(budget=budget)(rng, data) mechanism_error = np.mean([ @@ -101,9 +103,9 @@ def _mixed_workload_mechanism_baseline_errors( numerical_bins=numerical_bins, ) - mechanism_result = mechanism_synth.configure(domains, budget=budget)( - rng, data - ) + mechanism_result = mechanism_synth.configure( + domains, budget=budget, workload=workload + )(rng, data) baseline_result = baseline_synth.configure(domains, budget=budget)(rng, data) mechanism_error = np.mean([ @@ -314,7 +316,7 @@ def test_heterogeneous_input_dataframe_with_none(self): def test_discrete_workload_regression_with_aim(self): workload = [('a',), ('b',), ('c',), ('a', 'b'), ('a', 'c'), ('b', 'c')] - config = aim.AIMConfig(workload=workload, max_rounds=4, pgm_iters=500) + config = aim.AIMConfig(max_rounds=4, pgm_iters=500) baseline_config = IndependentConfig(pgm_iters=500) mechanism_error, baseline_error = ( _discrete_workload_mechanism_baseline_errors( @@ -325,13 +327,27 @@ def test_discrete_workload_regression_with_aim(self): def test_mixed_workload_regression_with_aim(self): workload = [('a',), ('b',), ('c',), ('a', 'b'), ('a', 'c'), ('b', 'c')] - config = aim.AIMConfig(workload=workload, max_rounds=4, pgm_iters=500) + config = aim.AIMConfig(max_rounds=4, pgm_iters=500) baseline_config = IndependentConfig(pgm_iters=500) mechanism_error, baseline_error = _mixed_workload_mechanism_baseline_errors( config, baseline_config, workload ) self.assertLess(mechanism_error, 0.05 * baseline_error) + def test_configure_validates_workload_against_schema(self): + domains = { + 'a': domain.CategoricalAttribute(possible_values=['x', 'y']), + 'b': domain.CategoricalAttribute(possible_values=['p', 'q']), + } + config = TabularConfig(discrete_mechanism=aim.AIMConfig(pgm_iters=10)) + with self.assertRaisesRegex(ValueError, 'Unknown workload attributes'): + config.configure(domains, budget=1.0, workload=[('a', 'missing')]) + with self.assertLogs(level='WARNING') as logs: + config.configure(domains, budget=1.0, workload=[('a',)]) + self.assertTrue( + any('Attributes not in workload' in msg for msg in logs.output) + ) + def test_empty_dataset(self): """Tests that DPSynth works without crashing on empty datasets, and outputs noisy rows.""" domains = { diff --git a/tests/discrete_mechanisms/aim_test.py b/tests/discrete_mechanisms/aim_test.py index 9dd8b9ed..eb6f8fab 100644 --- a/tests/discrete_mechanisms/aim_test.py +++ b/tests/discrete_mechanisms/aim_test.py @@ -59,9 +59,9 @@ class AIMTest(absltest.TestCase): def test_fits_one_way_marginals_with_aim(self): data = mbi.Dataset.synthetic(mbi.Domain(["a", "b", "c"], [3, 4, 5]), N=1000) workload = [("a",), ("b",), ("c",)] - config = aim.AIMConfig(workload=workload, max_rounds=4, pgm_iters=500) + config = aim.AIMConfig(max_rounds=4, pgm_iters=500) - calibrated = config.configure(budget=10000) + calibrated = config.configure(budget=10000, workload=workload) result = calibrated(np.random.default_rng(0), data) self.assertIsInstance(result, common.DiscreteMechanismResult) diff --git a/tests/discrete_mechanisms/direct_test.py b/tests/discrete_mechanisms/direct_test.py index 9eace302..967ac75e 100644 --- a/tests/discrete_mechanisms/direct_test.py +++ b/tests/discrete_mechanisms/direct_test.py @@ -25,11 +25,9 @@ def test_fits_one_way_marginals(self): data = mbi.Dataset.synthetic(mbi.Domain(['a', 'b', 'c'], [3, 4, 5]), N=1000) prespecified_queries = [('a', 'b'), ('a', 'c'), ('b', 'c')] - config = direct.DirectConfig( - prespecified_marginal_queries=prespecified_queries, - pgm_iters=500, - ) - result = config.configure(budget=10000)(np.random.default_rng(0), data) + config = direct.DirectConfig(pgm_iters=500) + calibrated = config.configure(budget=10000, workload=prespecified_queries) + result = calibrated(np.random.default_rng(0), data) self.assertIsInstance(result, common.DiscreteMechanismResult) self.assertLen(result.measurements, len(prespecified_queries)) @@ -46,11 +44,11 @@ def test_custom_estimator(self): data = mbi.Dataset.synthetic(mbi.Domain(['a', 'b', 'c'], [3, 4, 5]), N=1000) prespecified_queries = [('a', 'b'), ('a', 'c'), ('b', 'c')] config = direct.DirectConfig( - prespecified_marginal_queries=prespecified_queries, estimator=mbi.estimation.InteriorGradient(), pgm_iters=500, ) - result = config.configure(budget=10000)(np.random.default_rng(0), data) + calibrated = config.configure(budget=10000, workload=prespecified_queries) + result = calibrated(np.random.default_rng(0), data) for col in data.domain: expected = data.project([col]).datavector() @@ -68,12 +66,12 @@ def tracking_oracle(potentials, total=1, constraints=()): ) config = direct.DirectConfig( - prespecified_marginal_queries=[('a', 'b')], estimator=mbi.estimation.InteriorGradient(), marginal_oracle=tracking_oracle, pgm_iters=10, ) - config.configure(budget=100)(np.random.default_rng(0), data) + calibrated = config.configure(budget=100, workload=[('a', 'b')]) + calibrated(np.random.default_rng(0), data) self.assertNotEmpty(calls) diff --git a/tests/discrete_mechanisms/discrete_mechanisms_test.py b/tests/discrete_mechanisms/discrete_mechanisms_test.py index c687b6a6..2ba432e1 100644 --- a/tests/discrete_mechanisms/discrete_mechanisms_test.py +++ b/tests/discrete_mechanisms/discrete_mechanisms_test.py @@ -36,13 +36,11 @@ _WORKLOAD = [('a', 'b'), ('b', 'c'), ('a',), ('b',), ('c',)] _MECHANISMS = { - 'AIM': aim.AIMConfig(workload=_WORKLOAD, max_rounds=4, pgm_iters=500), + 'AIM': aim.AIMConfig(max_rounds=4, pgm_iters=500), 'MST': mst.MSTConfig(pgm_iters=500), - 'SWIFT': swift.SWIFTConfig(workload=_WORKLOAD, pgm_iters=500), + 'SWIFT': swift.SWIFTConfig(pgm_iters=500), 'Independent': independent.IndependentConfig(), - 'Direct': direct.DirectConfig( - prespecified_marginal_queries=_WORKLOAD, pgm_iters=500 - ), + 'Direct': direct.DirectConfig(pgm_iters=500), } @@ -63,8 +61,8 @@ def test_mechanism_runs_on_precomputed_marginals(self, mechanism): data = mbi.Dataset.synthetic(domain, N=500) rng = np.random.default_rng(42) - calibrated = mechanism.configure(budget=_ZCDP_RHO) - cliques = mechanism.supporting_cliques(domain) + calibrated = mechanism.configure(budget=_ZCDP_RHO, workload=_WORKLOAD) + cliques = calibrated.supporting_cliques(domain) precomputed = common.precompute_marginals(data, cliques) @@ -86,7 +84,9 @@ def test_compression_restores_domain(self, config): data = _make_skewed_dataset(rng) original_domain = data.domain - result = synth_config.configure(budget=_ZCDP_RHO)(rng, data) + result = synth_config.configure(budget=_ZCDP_RHO, workload=_WORKLOAD)( + rng, data + ) self.assertEqual(result.synthetic_data.domain, original_domain) @@ -103,7 +103,7 @@ def test_compression_with_initial_measurements(self, config): rng, data, [('a',), ('b',)], gdp_sigma=1.0 ) - mechanism = synth_config.configure(budget=_ZCDP_RHO) + mechanism = synth_config.configure(budget=_ZCDP_RHO, workload=_WORKLOAD) result = mechanism(rng, data, initial_measurements=initial_measurements) self.assertEqual(result.synthetic_data.domain, original_domain) @@ -119,7 +119,9 @@ def test_low_epsilon_calibration(self, mechanism): return rng = np.random.default_rng(0) data = _make_skewed_dataset(rng) - result = mechanism.calibrate(epsilon=1e-3, delta=1e-5)(rng, data) + result = mechanism.calibrate(epsilon=1e-3, delta=1e-5, workload=_WORKLOAD)( + rng, data + ) self.assertIsInstance(result, common.DiscreteMechanismResult) diff --git a/tests/domain_test.py b/tests/domain_test.py index 89cc6274..5f4b4c39 100644 --- a/tests/domain_test.py +++ b/tests/domain_test.py @@ -292,5 +292,54 @@ def test_with_constraints(self): self.assertSequenceEqual(s.constraints, (mock_constraint,)) +class WorkloadTest(parameterized.TestCase): + + def test_from_unweighted_list_of_tuples_and_lists(self): + w = domain.Workload([('a', 'b'), ['b', 'c']]) + self.assertLen(w, 2) + self.assertEqual(dict(w), {('a', 'b'): 1.0, ('b', 'c'): 1.0}) + self.assertEqual(w[('a', 'b')], 1.0) + self.assertIn(('b', 'c'), w) + self.assertListEqual(list(w.keys()), [('a', 'b'), ('b', 'c')]) + self.assertListEqual(list(w.values()), [1.0, 1.0]) + + def test_from_weighted_mapping(self): + w = domain.Workload({('a', 'b'): 2.5, ('c',): 0.5}) + self.assertEqual(dict(w), {('a', 'b'): 2.5, ('c',): 0.5}) + + def test_from_any(self): + self.assertIsNone(domain.Workload.from_any(None)) + w = domain.Workload([('a', 'b')]) + self.assertIs(domain.Workload.from_any(w), w) + self.assertEqual(domain.Workload.from_any([('a', 'b')]), w) + + @parameterized.named_parameters( + dict(testcase_name='bare_string', cliques=['ab']), + dict(testcase_name='empty_clique', cliques=[()]), + dict(testcase_name='non_string_element', cliques=[(1, 2)]), + dict(testcase_name='duplicate_in_clique', cliques=[('a', 'a')]), + dict(testcase_name='zero_weight', cliques={('a',): 0.0}), + dict(testcase_name='negative_weight', cliques={('a',): -1.0}), + dict(testcase_name='nan_weight', cliques={('a',): math.nan}), + dict(testcase_name='inf_weight', cliques={('a',): math.inf}), + ) + def test_invalid_workload_raises(self, cliques): + with self.assertRaises(ValueError): + domain.Workload(cliques) + + def test_validate_unknown_attribute_raises(self): + w = domain.Workload([('a', 'b'), ('b', 'unknown')]) + with self.assertRaisesRegex(ValueError, 'Unknown workload attributes'): + w.validate(['a', 'b', 'c']) + + def test_validate_uncovered_attribute_logs_warning(self): + w = domain.Workload([('a', 'b')]) + with self.assertLogs(level='WARNING') as logs: + w.validate(['a', 'b', 'c']) + self.assertTrue( + any('Attributes not in workload' in msg for msg in logs.output) + ) + + if __name__ == '__main__': absltest.main() diff --git a/tests/serialize_test.py b/tests/serialize_test.py index 53246a09..07d5cb19 100644 --- a/tests/serialize_test.py +++ b/tests/serialize_test.py @@ -84,10 +84,7 @@ def test_independent_config_roundtrip(self): self.assertEqual(loaded, config) def test_direct_config_roundtrip(self): - config = direct.DirectConfig( - pgm_iters=4000, - prespecified_marginal_queries=[('a', 'b'), ('c',)], - ) + config = direct.DirectConfig(pgm_iters=4000) yaml_str = serialize.to_yaml(config) raw_dict = yaml.safe_load(yaml_str) self.assertNotIn('estimator', raw_dict) @@ -97,7 +94,6 @@ def test_direct_config_roundtrip(self): def test_direct_config_with_estimator_roundtrip(self): config = direct.DirectConfig( pgm_iters=1500, - prespecified_marginal_queries=[('a', 'b')], estimator=mbi.estimation.UniversalAcceleratedMethod(linesearch=True), ) yaml_str = serialize.to_yaml(config) @@ -109,6 +105,37 @@ def test_direct_config_with_estimator_roundtrip(self): loaded = serialize.from_yaml(yaml_str) self.assertEqual(loaded, config) + def test_workload_unweighted_roundtrip(self): + workload = domain.Workload([('a', 'b'), ('b', 'c')]) + yaml_str = dpsynth.to_yaml(workload) + raw_dict = yaml.safe_load(yaml_str) + self.assertEqual( + raw_dict, + {'type': 'Workload', 'cliques': [['a', 'b'], ['b', 'c']]}, + ) + loaded = dpsynth.from_yaml(yaml_str) + self.assertEqual(loaded, workload) + + def test_workload_weighted_roundtrip(self): + workload = domain.Workload({('a', 'b'): 2.5, ('c',): 0.5}) + yaml_str = dpsynth.to_yaml(workload) + raw_dict = yaml.safe_load(yaml_str) + self.assertEqual( + raw_dict, + { + 'type': 'Workload', + 'cliques': [['a', 'b'], ['c']], + 'weights': [2.5, 0.5], + }, + ) + loaded = dpsynth.from_yaml(yaml_str) + self.assertEqual(loaded, workload) + + def test_workload_from_raw_list_yaml(self): + raw_yaml = '- [a, b]\n- [b, c]\n' + loaded = dpsynth.from_yaml(raw_yaml, domain.Workload) + self.assertEqual(loaded, domain.Workload([('a', 'b'), ('b', 'c')])) + def test_discrete_config_roundtrip(self): config = discrete.DiscreteConfig( mechanism=aim.AIMConfig(pgm_iters=400),