Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 15 additions & 14 deletions dpsynth/adapters/beam.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,6 @@
from dpsynth import domain
from dpsynth.discrete_mechanisms import common as dm_common
from dpsynth.local_mode import initialization
from dpsynth.local_mode import primitives
import mbi
import numpy as np

Expand Down Expand Up @@ -193,6 +192,8 @@ def run_from_summary(
sparse_stats: dict[str, list[tuple[Any, int]]],
initializers: dict[str, CalibratedInitializer],
rng: np.random.Generator,
*,
num_rows: int | None = None,
) -> dict[str, initialization.ColumnMeasurement]:
"""Converts materialized sparse stats to ColumnMeasurements on the driver.

Expand All @@ -204,6 +205,8 @@ def run_from_summary(
produced by ``ComputeSufficientStats``.
initializers: Calibrated initializers keyed by column name.
rng: NumPy random generator for DP noise.
num_rows: Optional total row count, used to compute out-of-domain counts for
numerical columns with ``clip_to_range=False``.

Returns:
Per-column ``ColumnMeasurement`` results.
Expand All @@ -213,7 +216,12 @@ def run_from_summary(
sparse = sparse_stats[column]
if isinstance(init, initialization.NumericalInitializer):
counts = _sparse_to_dense_numerical(sparse, init.grid_spec[2])
results[column] = init.from_summary(rng, counts)
ood_count = (
float(num_rows - counts.sum())
if (not init.attribute.clip_to_range and num_rows is not None)
else 0.0
)
results[column] = init.from_summary(rng, counts, ood_count=ood_count)
elif isinstance(init, initialization.CategoricalInitializer):
counts = _sparse_to_dense_categorical(sparse, init.attribute.size)
results[column] = init.from_summary(rng, counts)
Expand Down Expand Up @@ -382,7 +390,6 @@ def generate_from_marginals(
rng: np.random.Generator,
column_measurements: dict[str, initialization.ColumnMeasurement],
marginals: mbi.CliqueVector,
total_measurement: mbi.LinearMeasurement,
) -> data_generation_v3.DataGenerationResult:
"""Runs the discrete mechanism and decoding from pre-computed marginals.

Expand All @@ -391,7 +398,6 @@ def generate_from_marginals(
rng: NumPy random generator for the discrete mechanism's DP noise.
column_measurements: Per-column results from pass 1 initialization.
marginals: The exact joint marginals computed by pass 2.
total_measurement: The DP-noised total-count measurement (clique ``()``).

Returns:
A DataGenerationResult containing the synthetic DataFrame.
Expand All @@ -402,7 +408,7 @@ def generate_from_marginals(
column_measurements, synth.schema
)

initial_measurements = [total_measurement, *codec.one_way_measurements()]
initial_measurements = codec.one_way_measurements()
logging.info('[DPSynth/Beam]: Running discrete mechanism.')
# pyrefly: ignore[missing-attribute,not-callable]
mechanism_result = synth.base_mechanism(
Expand Down Expand Up @@ -439,7 +445,6 @@ def _run_two_pass(
) -> data_generation_v3.DataGenerationResult:
"""Two-pass Beam pipeline that delegates to a local TabularConfig."""

sigma = synth.total_count_sigma
inits = cast(dict[str, CalibratedInitializer], synth.initializers)
if pipeline_kwargs is None:
pipeline_kwargs = {}
Expand All @@ -465,15 +470,11 @@ def _run_two_pass(
_ = count | 'WriteRowCount' >> beam.Map(_write, path=count_path)
# We run this on the driver so we don't have to track worker-side RNGs.
sparse_stats = _read(summary_path)
column_measurements = run_from_summary(sparse_stats, inits, rng)
num_rows = int(_read(count_path))
logging.info('[DPSynth/Beam]: Pass 1 complete.')
# pyrefly: ignore[missing-attribute]
total = primitives.add_gaussian_noise(
rng, float(num_rows), sigma, cast(int, synth.max_records_per_user)
column_measurements = run_from_summary(
sparse_stats, inits, rng, num_rows=num_rows
)
total = float(max(1.0, total))
total_measurement = mbi.LinearMeasurement(np.array([total]), (), sigma)
logging.info('[DPSynth/Beam]: Pass 1 complete.')

# Ask the configured discrete mechanism which marginals it needs.
mbi_domain = data_generation_v3.TabularCodec.from_measurements(
Expand All @@ -499,7 +500,7 @@ def _run_two_pass(

# Run the discrete mechanism and decode on the driver.
return generate_from_marginals(
synth, rng, column_measurements, clique_vector, total_measurement
synth, rng, column_measurements, clique_vector
)
finally:
# Only remove a temp dir we created; never a user-supplied temp_location.
Expand Down
50 changes: 9 additions & 41 deletions dpsynth/data_generation_v3.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,6 @@
from dpsynth import domain
from dpsynth.discrete_mechanisms import common as dm_common
from dpsynth.local_mode import initialization
from dpsynth.local_mode import primitives
from dpsynth.local_mode import vectorized_transformations as vtx
import mbi
import numpy as np
Expand Down Expand Up @@ -157,8 +156,6 @@ def one_way_measurements(self) -> list[mbi.LinearMeasurement]:
cm = c.column_measurement
if cm.noisy_counts is None:
continue
elif isinstance(cm, initialization.NumericalMeasurement):
query = mbi.DatavectorQuery(use_for_total_estimation=False)
elif isinstance(cm, initialization.OpenSetMeasurement):
query = mbi.SlicedQuery(start=1)
else:
Expand Down Expand Up @@ -208,7 +205,6 @@ class TabularMechanism(api.CalibratedMechanism):
schema: The dataset schema and constraints.
base_mechanism: The calibrated discrete mechanism.
initializers: Per-column calibrated initializers.
total_count_sigma: Sigma for the total-count mechanism.
max_records_per_user: Assumed upper bound on the number of records a single
user contributes.
"""
Expand All @@ -217,16 +213,12 @@ class TabularMechanism(api.CalibratedMechanism):
schema: domain.Schema
base_mechanism: discrete_mechanisms.CalibratedMechanism
initializers: dict[str, api.CalibratedMechanism]
total_count_sigma: float = dataclasses.field(repr=False)
max_records_per_user: int = 1

@property
def dp_event(self) -> dp_accounting.DpEvent:
"""Returns the composed DpEvent for all sub-mechanisms."""
events = [init.dp_event for init in self.initializers.values()]
events.append(
dp_accounting.GaussianDpEvent(noise_multiplier=self.total_count_sigma)
)
events.append(self.base_mechanism.dp_event)
events = [e for e in events if not isinstance(e, dp_accounting.NoOpDpEvent)]

Expand Down Expand Up @@ -267,32 +259,13 @@ def __call__(
)

# Phase 1: Per-column initialization.
# Measure total count first, then run per-column initializers.
def _run_initializers():
noisy_total = primitives.add_gaussian_noise(
rng,
len(data),
self.total_count_sigma,
self.max_records_per_user,
)
total = max(1.0, noisy_total)
total_measurement = mbi.LinearMeasurement(
noisy_measurement=np.array([total]),
clique=(),
stddev=self.max_records_per_user * self.total_count_sigma,
)
return {
col: init(rng, data[col].values)
for col, init in self.initializers.items()
}

results: dict[str, initialization.ColumnMeasurement] = {}
for col, init in self.initializers.items():
if isinstance(init, initialization.NumericalInitializer):
results[col] = init(
rng, data[col].values, estimated_total=float(total)
)
else:
results[col] = init(rng, data[col].values)
return total_measurement, results

total_measurement, results = _checkpoint.get_or_compute(
results = _checkpoint.get_or_compute(
'column_measurements', _run_initializers
)

Expand All @@ -306,7 +279,6 @@ def _run_initializers():
rng,
results,
discrete,
total_measurement,
cross_attribute_constraints=cross_attribute_constraints,
num_rows=num_rows,
)
Expand All @@ -316,7 +288,6 @@ def from_summary(
rng: np.random.Generator,
column_measurements: Mapping[str, initialization.ColumnMeasurement],
data: mbi.Dataset | mbi.CliqueVector,
total_measurement: mbi.LinearMeasurement,
*,
cross_attribute_constraints: Sequence[constraints.Constraint] = (),
num_rows: int | None = None,
Expand All @@ -328,11 +299,11 @@ def from_summary(
)
mbi_constraints = tuple(c.to_mbi() for c in cross_attribute_constraints)

# Feed the noisy total (clique ()) and one-way column measurements as
# initial measurements so the mechanism does not re-measure them.
# Feed one-way column measurements as initial measurements so the mechanism
# does not re-measure them.
codec = TabularCodec.from_measurements(column_measurements, self.schema)
one_way_measurements = codec.one_way_measurements()
initial_measurements = [total_measurement, *one_way_measurements]
initial_measurements = list(one_way_measurements)

mappings = {}
if isinstance(data, mbi.Dataset):
Expand Down Expand Up @@ -539,10 +510,8 @@ def configure(
schema, self.numerical_bins, self.numerical_epsilon_ratio
)
init_rho = self.init_budget_fraction * zcdp_rho
# +1 for the DPGaussianCount that always measures the total.
per_col_rho = init_rho / (len(inits) + 1)
per_col_rho = init_rho / len(inits)
discrete_rho = (1 - self.init_budget_fraction) * zcdp_rho
total_count_sigma = (0.5 / per_col_rho) ** 0.5

calibrated_inits: dict[str, api.CalibratedMechanism]

Expand All @@ -566,7 +535,6 @@ def configure(
schema=schema,
base_mechanism=calibrated_discrete,
initializers=calibrated_inits,
total_count_sigma=total_count_sigma,
max_records_per_user=max_records_per_user,
)

Expand Down
Loading
Loading