Skip to content
Open
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
28 changes: 20 additions & 8 deletions dpsynth/text/dp_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,7 @@
import dataclasses
import json
import math
import typing

from absl import logging
import dp_accounting
Expand Down Expand Up @@ -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
Expand All @@ -111,19 +112,25 @@ 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)

@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

Expand All @@ -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(
Expand Down
18 changes: 18 additions & 0 deletions tests/text/dp_trainer_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()
Loading