From 7beb2adaa077894b25f0fe244e8284aeadf1a0d1 Mon Sep 17 00:00:00 2001 From: DPSynth Team Date: Tue, 30 Jun 2026 03:27:15 -0700 Subject: [PATCH] examples: add end-to-end training example for TabularTransformer on Adult dataset PiperOrigin-RevId: 940357136 --- dpsynth/text/dp_trainer.py | 28 ++++++++++++++++++++-------- tests/text/dp_trainer_test.py | 18 ++++++++++++++++++ 2 files changed, 38 insertions(+), 8 deletions(-) diff --git a/dpsynth/text/dp_trainer.py b/dpsynth/text/dp_trainer.py index baa91df9..b71baa65 100644 --- a/dpsynth/text/dp_trainer.py +++ b/dpsynth/text/dp_trainer.py @@ -51,6 +51,7 @@ import dataclasses import json import math +import typing from absl import logging import dp_accounting @@ -92,7 +93,7 @@ class DPTrainer(primitives.DPMechanism): init_params: training.Params loss_fn: training.LossFn - mechanism_config: execution_plan.BandMFConfig + mechanism_config: execution_plan.ExecutionPlanConfig optimizer: optax.GradientTransformation performance_flags: execution_plan.PerformanceFlags | None = None callback: training.CallbackFn | None = None @@ -111,11 +112,14 @@ def configure(self, *, zcdp_rho: float, delta: float = 0.0) -> DPTrainer: # pyr Returns: A new ``DPTrainer`` with calibrated ``config.noise_multiplier``. """ - num_bands = len(self.mechanism_config.strategy) # pyrefly: ignore[bad-argument-type] - rounds = math.ceil(self.mechanism_config.iterations / num_bands) + if isinstance(self.mechanism_config, execution_plan.NonPrivateConfig): + return self + cfg = typing.cast(execution_plan.BandMFConfig, self.mechanism_config) + num_bands = len(cfg.strategy) # pyrefly: ignore[bad-argument-type] + rounds = math.ceil(cfg.iterations / num_bands) noise_multiplier = math.sqrt(rounds / (2.0 * zcdp_rho)) calibrated_config = dataclasses.replace( - self.mechanism_config, + cfg, noise_multiplier=noise_multiplier, ) return dataclasses.replace(self, mechanism_config=calibrated_config) @@ -123,7 +127,10 @@ def configure(self, *, zcdp_rho: float, delta: float = 0.0) -> DPTrainer: # pyr @property def dp_event(self) -> dp_accounting.DpEvent: """The DpEvent characterizing the privacy cost of DP-SGD training.""" - if self.mechanism_config.noise_multiplier is None: + if ( + hasattr(self.mechanism_config, 'noise_multiplier') + and self.mechanism_config.noise_multiplier is None + ): raise ValueError('noise_multiplier is not set. Call calibrate() first.') return self._make_plan().dp_event @@ -141,11 +148,16 @@ def __call__(self, rng: int, data: training.Batch) -> training.TrainingState: Returns: Final ``TrainingState`` containing the trained parameters. """ - if self.mechanism_config.noise_multiplier is None: + if ( + hasattr(self.mechanism_config, 'noise_multiplier') + and self.mechanism_config.noise_multiplier is None + ): raise ValueError('noise_multiplier is not set. Call calibrate() first.') - d = dataclasses.asdict(self.mechanism_config) - d['strategy'] = self.mechanism_config.strategy.tolist() # JSON/numpy hack. # pyrefly: ignore[missing-attribute] + cfg = typing.cast(typing.Any, self.mechanism_config) + d = dataclasses.asdict(cfg) + if hasattr(self.mechanism_config, 'strategy') and cfg.strategy is not None: + d['strategy'] = cfg.strategy.tolist() # JSON/numpy hack. logging.info('DPTrainer config:\n%s', json.dumps(d, indent=2)) dp_trainer = training.DPTrainer( diff --git a/tests/text/dp_trainer_test.py b/tests/text/dp_trainer_test.py index 43ca586f..0be08e2e 100644 --- a/tests/text/dp_trainer_test.py +++ b/tests/text/dp_trainer_test.py @@ -14,6 +14,7 @@ from absl.testing import absltest +import dp_accounting from dpsynth.text import dp_trainer from flax import nnx import jax.numpy as jnp @@ -109,6 +110,23 @@ def loss_fn(params, batch, prng): train_state = trainer(rng=42, data={'x': jnp.ones((10, 4, 4))}) self.assertIsNotNone(train_state) + def test_non_private_config(self): + params, loss_fn = _dummy_params_and_loss() + non_private_config = jax_privacy.execution_plan.NonPrivateConfig( + iterations=5, + batch_size=2, + ) + trainer = dp_trainer.DPTrainer( + init_params=params, + loss_fn=loss_fn, + mechanism_config=non_private_config, + optimizer=optax.adamw(1e-4), + ).configure(zcdp_rho=float('inf')) + + self.assertIsInstance(trainer.dp_event, dp_accounting.NonPrivateDpEvent) + train_state = trainer(rng=42, data={'x': jnp.ones((10, 4, 4))}) + self.assertIsNotNone(train_state) + if __name__ == '__main__': absltest.main()