Buckets:
| #!/usr/bin/env python | |
| """Differentially private synthetic data for a conversation-rubric CSV. | |
| Stage 1 -- selection (AIM). Measures every 1-way marginal, then adaptively | |
| buys 2-way marginals: each round the exponential mechanism picks the pair the | |
| current model fits worst, and the Gaussian mechanism measures it. | |
| Stage 2 -- estimation. Fits a mixture of products to the noisy marginals and | |
| samples synthetic records from it. | |
| Every read of the raw data happens inside the PRIVACY-CRITICAL SECTION below. | |
| Everything after it sees only noisy measurements. | |
| Usage: | |
| ./dp_synth.py --in data.csv --out production/ | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import collections | |
| import csv | |
| import functools | |
| import itertools | |
| import os | |
| import sys | |
| import time | |
| import jax | |
| import jax.numpy as jnp | |
| import numpy as np | |
| import opendp.combinators as odp_comb | |
| import opendp.domains as odp_dom | |
| import opendp.measurements as odp_meas | |
| import opendp.measures as odp_mes | |
| import opendp.metrics as odp_met | |
| import opendp.transformations as odp_tr | |
| import optax | |
| import pandas as pd | |
| from opendp.mod import enable_features | |
| from schema import (APPLICABILITY, BINARY, BUCKET, COLUMNS, LAB, NUMERIC, | |
| OTHER_BINARY, SIZES, cell_label, decode, discretize) | |
| from smooth import smooth | |
| enable_features("contrib", "honest-but-curious") | |
| jax.config.update("jax_enable_compilation_cache", False) | |
| # =========================================================================== | |
| # Privacy-critical section. | |
| # | |
| # The raw data is read only inside SensitiveData. It leaves this class only as | |
| # noised marginals and as the identity of the pair that selection picked. | |
| # | |
| # This uses unbounded (add/remove-one-record) zCDP, converted to (ε,δ)-DP by | |
| # OpenDP at the end. | |
| # - Each marginal query is an OpenDP Transformation from the dataset to that | |
| # marginal, composed with make_gaussian. | |
| # - Selection is a Transformation to the vector of candidate scores composed | |
| # with make_noisy_max (which uses the exponential mechanism). | |
| # - The list of successive privacy budgets that will be spent by queries is | |
| # declared in advance, using adaptive composition. *Which* marginals get | |
| # chosen depends on earlier answers. | |
| # That list is SCHEDULE, drawn up outside this section. SensitiveData | |
| # prices it through OpenDP and refuses it if it exceeds RHO. | |
| # | |
| # Two stability maps are asserted rather than derived by OpenDP, which is what | |
| # enable_features("honest-but-curious") allows: | |
| # | |
| # 1. _marginal() one record moves exactly one cell of the marginal by 1, so | |
| # symmetric distance d maps to L2 distance d. | |
| # | |
| # 2. _scores() one record moves any pair's L1 gap by at most 1, so | |
| # symmetric distance d maps to L-infinity distance d. | |
| # | |
| # =========================================================================== | |
| # OpenDP needs a domain to type the transformations' input, and a metric. The | |
| # domain is "fake" in that we're not actually encapsulating the data in OpenDP | |
| # for performance reasons. Its predicate is never called, so it's arbitrary. | |
| # We use symmetric_distance to match the add-or-remove neighboring relation, | |
| # even though it's also arbitrary. | |
| _DATASET = odp_dom.user_domain("Dataset", lambda _: True) | |
| _RECORDS = odp_met.symmetric_distance() | |
| # Marginal measurements use OpenDP's native vector domain and L2 distance to | |
| # work well with Gaussian noise and zCDP. | |
| _MARGINAL = odp_dom.vector_domain(odp_dom.atom_domain(T=float, nan=False)) | |
| _L2 = odp_met.l2_distance(T=float) | |
| # Selection scores also use OPenDP's native vector domain and LInf distance. | |
| # Scores can move both up and down when a record is added, so we can't set | |
| # monotonic=True. | |
| _SCORES = odp_dom.vector_domain(odp_dom.atom_domain(T=float, nan=False)) | |
| _LINF = odp_met.linf_distance(T=float, monotonic=False) | |
| EPSILON = 10.0 | |
| DELTA = 1e-5 | |
| def rho(epsilon, delta): | |
| """Largest rho whose conversion to (epsilon, delta)-DP stays in budget.""" | |
| def epsilon_at(rho): | |
| gaussian = odp_meas.make_gaussian(_MARGINAL, _L2, scale=1 / np.sqrt(2 * rho)) | |
| converted = odp_comb.make_zCDP_to_approxDP(gaussian) | |
| return odp_comb.make_fix_delta(converted, delta).map(1.0)[0] | |
| low, high = 1e-9, 1e3 | |
| for _ in range(200): | |
| mid = (low + high) / 2 | |
| low, high = (mid, high) if epsilon_at(mid) <= epsilon else (low, mid) | |
| return low | |
| RHO = rho(EPSILON, DELTA) | |
| class SensitiveData: | |
| """Encapsulates all calls to the raw data so the code is easier to audit. | |
| It is handed a schedule of query costs. Every query must then match that | |
| schedule, and the whole schedule must fit in RHO -- so no sequence of | |
| calls can overspend, whatever the caller planned. | |
| """ | |
| def __init__(self, path, schedule): | |
| # Read the data and discretize all columns. | |
| df = pd.read_csv(path, usecols=COLUMNS) | |
| data = discretize(df) | |
| self._data = data | |
| # Every 2-way contingency table, computed once. Stacking the one-hot | |
| # encodings of all columns into M makes M.T @ M a block matrix whose | |
| # (i, j) block is the marginal of columns i and j, and whose (i, i) | |
| # block is diag(1-way marginal of i). Selection needs all 18,528 of | |
| # them every round, so it is memoised; the per-query measurements read | |
| # the dataset directly. | |
| shape = np.array([SIZES[c] for c in COLUMNS]) | |
| self.offsets = np.concatenate([[0], np.cumsum(shape)]) | |
| rows = np.arange(len(df))[:, None] | |
| onehot = np.zeros((len(df), int(self.offsets[-1])), dtype=np.float32) | |
| onehot[rows, self.offsets[:-1][None, :] + | |
| np.stack([data[c] for c in COLUMNS], 1)] = 1.0 | |
| onehot = jnp.asarray(onehot) | |
| self._gram = np.asarray(jax.block_until_ready(onehot.T @ onehot), np.float64) | |
| # The pairs selection chooses between, in one fixed order so that | |
| # scoring and reporting agree on it. Held as two index arrays, which is | |
| # what indexing the score matrix with it wants. | |
| num_columns = len(COLUMNS) | |
| pairs = [(i, j) for i in range(num_columns) | |
| for j in range(i + 1, num_columns)] | |
| self._candidate_pairs = tuple(np.array(index) for index in zip(*pairs)) | |
| composed = odp_comb.make_adaptive_composition( | |
| _DATASET, _RECORDS, odp_mes.zero_concentrated_divergence(), | |
| 1, schedule) | |
| # OpenDP's own accounting of the whole schedule, against the budget. | |
| # This is the check the guarantee rests on: the plan only proposes | |
| # prices, and one that asks for more than RHO stops here. | |
| self.spent = composed.map(1) | |
| if self.spent > RHO + 1e-12: | |
| raise RuntimeError( | |
| f"schedule needs {self.spent:.9f} of a {RHO:.9f} rho budget") | |
| self._queryable = composed(data) | |
| def _marginal(self, clique): | |
| """A Transformation from the dataset to one flattened marginal. | |
| Stability: adding or removing one record moves exactly one cell of the | |
| marginal by 1, whatever the clique's order or the sizes involved, so a | |
| symmetric distance of d maps to an L2 distance of d. This is the claim | |
| OpenDP takes on trust; everything downstream of it is derived. | |
| """ | |
| sizes = [SIZES[c] for c in clique] | |
| def function(data): | |
| flat = np.zeros(len(data[clique[0]]), np.int64) | |
| for col, size in zip(clique, sizes): | |
| flat = flat * size + data[col] | |
| return np.bincount(flat, minlength=int(np.prod(sizes))).astype(float) | |
| return odp_tr.make_user_transformation( | |
| _DATASET, _RECORDS, _MARGINAL, _L2, | |
| function=function, stability_map=lambda d_in: float(d_in)) | |
| def measure(self, clique, sigma): | |
| """Gaussian mechanism on one marginal.""" | |
| transformation = self._marginal(clique) | |
| measurement = transformation >> odp_meas.make_gaussian( | |
| transformation.output_domain, transformation.output_metric, | |
| scale=float(sigma)) | |
| return self._queryable(measurement) | |
| def _scores(self, model_gram, sigma): | |
| """A Transformation from the dataset to every candidate pair's score. | |
| The score is the L1 gap between the true marginal and the model's, | |
| minus the L1 error fresh Gaussian noise would itself introduce, so | |
| re-measuring a pair we already know is not rewarded. | |
| Stability: adding or removing one record moves any pair's L1 gap by at | |
| most 1, so a symmetric distance of d maps to an L-infinity distance of | |
| d. The bias term is a constant and ``model_gram`` derives only from | |
| noisy measurements, so neither affects sensitivity. | |
| """ | |
| starts = self.offsets[:-1] | |
| cells = np.outer(np.diff(self.offsets), np.diff(self.offsets)) | |
| bias = np.sqrt(2 / np.pi) * sigma * cells | |
| def function(data): | |
| gap = np.abs(self._gram - model_gram) | |
| l1 = np.add.reduceat(np.add.reduceat(gap, starts, 0), starts, 1) | |
| return list((l1 - bias)[self._candidate_pairs]) | |
| return odp_tr.make_user_transformation( | |
| _DATASET, _RECORDS, _SCORES, _LINF, | |
| function=function, stability_map=lambda d_in: float(d_in)) | |
| def select_pair(self, model_gram, scale, sigma): | |
| """Privately pick the pair the model fits worst. | |
| This uses the exponential mechanism behind the scenes under zCDP. | |
| """ | |
| transformation = self._scores(model_gram, sigma) | |
| measurement = transformation >> odp_meas.make_noisy_max( | |
| transformation.output_domain, transformation.output_metric, | |
| odp_mes.zero_concentrated_divergence(), | |
| scale=float(scale), negate=False) | |
| index = self._queryable(measurement) | |
| rows, cols = self._candidate_pairs | |
| return (COLUMNS[rows[index]], COLUMNS[cols[index]]) | |
| # =========================================================================== | |
| # END PRIVACY-CRITICAL SECTION | |
| # Everything below reads only noisy measurements. | |
| # =========================================================================== | |
| # --- The query plan -------------------------------------------------------- | |
| # | |
| # Which marginals the run buys, how the budget splits over them, and the noise | |
| # level each share works out to. | |
| # How the zCDP budget is divided; must sum to 1. The fixed-clique share is | |
| # divided further between the sets by CLIQUE_SHARE. | |
| BUDGET_ONE_WAY = 0.02 | |
| BUDGET_FIXED_CLIQUES = 0.68 | |
| BUDGET_ADAPTIVE_PAIRS = 0.25 | |
| BUDGET_SELECTION = 0.05 | |
| if abs(BUDGET_ONE_WAY + BUDGET_FIXED_CLIQUES + BUDGET_ADAPTIVE_PAIRS | |
| + BUDGET_SELECTION - 1) > 1e-12: | |
| raise SystemExit("the BUDGET_* shares do not sum to 1") | |
| # Adaptively selected 2-way marginals. | |
| ROUNDS = 3000 | |
| # The marginals bought outright, before AIM picks anything. | |
| CLIQUE_SETS = { | |
| "applic": list(itertools.combinations(APPLICABILITY, 2)), | |
| # One bucket per applicability criterion; a record can be in several. | |
| "bucket": list(itertools.combinations(BUCKET, 2)), | |
| "lab-bucket-applic": [(LAB, b, a) for b in BUCKET for a in APPLICABILITY], | |
| "applic-other": [(a, b) for a in APPLICABILITY for b in OTHER_BINARY], | |
| "lab-applic-other": [(LAB, a, b) for a in APPLICABILITY for b in OTHER_BINARY], | |
| "numeric-other": [(n, b) for n in NUMERIC for b in BINARY], | |
| "numeric-numeric": list(itertools.combinations(NUMERIC, 2)), | |
| "lab-numeric": [(LAB, n) for n in NUMERIC], | |
| } | |
| # How much budget fraction we allocate to each set. | |
| CLIQUE_SHARE = { | |
| "applic": 0.01, | |
| "bucket": 0.01, | |
| "lab-bucket-applic": 0.04, | |
| "applic-other": 0.12, | |
| "lab-applic-other": 0.45, | |
| "numeric-other": 0.03, | |
| "numeric-numeric": 0.01, | |
| "lab-numeric": 0.01, | |
| } | |
| if CLIQUE_SETS.keys() != CLIQUE_SHARE.keys(): | |
| raise SystemExit(f"CLIQUE_SETS names {sorted(CLIQUE_SETS)}, " | |
| f"CLIQUE_SHARE names {sorted(CLIQUE_SHARE)}") | |
| if abs(sum(CLIQUE_SHARE.values()) - BUDGET_FIXED_CLIQUES) > 1e-12: | |
| raise SystemExit(f"CLIQUE_SHARE sums to {sum(CLIQUE_SHARE.values())}, " | |
| f"not the {BUDGET_FIXED_CLIQUES} set aside for it") | |
| # Cliques must be in column order and unique: the schedule prices each one | |
| # once, and AIM returns its picks in column order. | |
| ALL_NAMED_CLIQUES = [clique for cliques in CLIQUE_SETS.values() | |
| for clique in cliques] | |
| _order = {c: i for i, c in enumerate(COLUMNS)} | |
| if any(list(c) != sorted(c, key=_order.__getitem__) for c in ALL_NAMED_CLIQUES): | |
| raise SystemExit("a clique in CLIQUE_SETS is not in column order") | |
| if len(ALL_NAMED_CLIQUES) != len(set(ALL_NAMED_CLIQUES)): | |
| raise SystemExit("CLIQUE_SETS buys the same clique twice") | |
| # A query costs rho = 1 / (2 s^2), so a share of the budget spread over count | |
| # queries buys each of them this sigma. | |
| def _sigma(count, share): | |
| return np.sqrt(count / (2 * RHO * share)) | |
| SIGMA_ONE_WAY = _sigma(len(COLUMNS), BUDGET_ONE_WAY) | |
| SIGMA_PAIR = _sigma(ROUNDS, BUDGET_ADAPTIVE_PAIRS) | |
| SELECTION_SCALE = _sigma(ROUNDS, BUDGET_SELECTION) | |
| SIGMA_NAMED = {name: _sigma(len(cliques), CLIQUE_SHARE[name]) | |
| for name, cliques in CLIQUE_SETS.items()} | |
| # The price of each query, in the order the queries must be asked: every | |
| # 1-way marginal, then the named cliques, then one selection and one | |
| # measurement per round. Priced by OpenDP itself so the numbers match. | |
| def _gaussian_price(sigma): | |
| return odp_meas.make_gaussian(_MARGINAL, _L2, scale=float(sigma)).map(1) | |
| def _selection_price(scale): | |
| return odp_meas.make_noisy_max( | |
| _SCORES, _LINF, odp_mes.zero_concentrated_divergence(), | |
| scale=float(scale), negate=False).map(1) | |
| SCHEDULE = ( | |
| [_gaussian_price(SIGMA_ONE_WAY)] * len(COLUMNS) | |
| + [price for name, cliques in CLIQUE_SETS.items() | |
| for price in [_gaussian_price(SIGMA_NAMED[name])] * len(cliques)] | |
| + [_selection_price(SELECTION_SCALE), _gaussian_price(SIGMA_PAIR)] * ROUNDS) | |
| def print_noise_levels(total): | |
| """Sigma on each measurement, against the average cell it measures.""" | |
| print(f" {'measurement':<22}{'cliques':>8}{'cells':>7}{'avg cell':>10}" | |
| f"{'sigma':>7}{'noise':>8}") | |
| rows = [("one-way", len(COLUMNS), 1, SIGMA_ONE_WAY)] | |
| for name, cliques in CLIQUE_SETS.items(): | |
| group_sigma = SIGMA_NAMED[name] | |
| # One row per clique shape within the set. | |
| shapes = collections.Counter( | |
| int(np.prod([SIZES[c] for c in clique])) for clique in cliques) | |
| for cells, count in sorted(shapes.items()): | |
| label = name if len(shapes) == 1 else f"{name} ({cells} cells)" | |
| rows.append((label, count, cells, group_sigma)) | |
| # Adaptive pairs are counted at 2x2, which nearly all of them are. | |
| rows.append(("adaptive pairs", ROUNDS, 4, SIGMA_PAIR)) | |
| for name, count, cells, group_sigma in rows: | |
| if name == "one-way": | |
| cells = int(np.mean(list(SIZES.values()))) | |
| average = total / cells | |
| print(f" {name:<22}{count:>8}{cells:>7}{average:>10.0f}" | |
| f"{group_sigma:>7.1f}{100 * group_sigma / average:>7.1f}%") | |
| def write_marginals(measured, destination): | |
| """Every noisy marginal the run measured, one row per cell.""" | |
| path = destination | |
| width = max(len(clique) for _, clique, _, _ in measured) | |
| header = ["marginal", "source", "sigma", "arity"] | |
| for i in range(width): | |
| header += [f"column{i + 1}", f"value{i + 1}"] | |
| header.append("noisy_count") | |
| with open(path, "w", newline="") as handle: | |
| writer = csv.writer(handle) | |
| writer.writerow(header) | |
| for index, (source, clique, sigma, values) in enumerate(measured): | |
| block = np.asarray(values).reshape([SIZES[c] for c in clique]) | |
| for cell in np.ndindex(block.shape): | |
| row = [index, source, round(float(sigma), 4), len(clique)] | |
| for i in range(width): | |
| if i < len(clique): | |
| row += [clique[i], cell_label(clique[i], cell[i])] | |
| else: | |
| row += ["", ""] | |
| row.append(round(float(block[cell]), 3)) | |
| writer.writerow(row) | |
| return path | |
| # --- Fitting the model ----------------------------------------------------- | |
| SEED = 0 # model init, the record draw and the repair step. The | |
| # measurement noise is drawn from OS entropy and never seeded. | |
| COMPONENTS = 2400 | |
| SCORING_WARMUP_ITERS = 1000 | |
| SCORING_ROUND_ITERS = 100 | |
| FINAL_ITERS = 20000 | |
| class ScoringModel: | |
| """Mixture of products: p(x) = (1/C) sum_k prod_j p_{k,j}(x_j). | |
| Scores every candidate pair for AIM each round, and is the final | |
| generative model. Measured marginals are evaluated by one batched einsum | |
| over a fixed number of slots, so XLA compiles the graph once. | |
| """ | |
| def __init__(self, total, capacity, seed=0, components=None, extra=()): | |
| self.total = float(total) | |
| self.components = COMPONENTS if components is None else components | |
| self.index = {a: i for i, a in enumerate(COLUMNS)} | |
| shape = np.array([SIZES[c] for c in COLUMNS]) | |
| self.max_size = int(shape.max()) | |
| num = len(shape) | |
| self.valid = jnp.asarray(np.arange(self.max_size)[None, :] < shape[:, None]) | |
| self.logits = 0.25 * jax.random.normal( | |
| jax.random.PRNGKey(seed), (num, self.components, self.max_size)) | |
| # Gather indices that flatten (attribute, value) into gram block order. | |
| self.flat = jnp.asarray(np.concatenate( | |
| [i * self.max_size + np.arange(n) for i, n in enumerate(shape)])) | |
| # One slot group per exact clique shape, each sized to the most | |
| # cliques of that shape the run can ask for. | |
| needed = collections.Counter( | |
| (int(shape[i]), int(shape[j])) | |
| for i in range(num) for j in range(i + 1, num)) | |
| for clique in extra: | |
| if len(clique) != 2: | |
| needed[tuple(SIZES[c] for c in clique)] += 1 | |
| self._slots = {form: min(count, capacity) | |
| for form, count in needed.items()} | |
| self.groups = {} | |
| self.slot = {} | |
| self.y1 = np.zeros((num, self.max_size), np.float32) | |
| self.w1 = np.zeros((num, self.max_size), np.float32) | |
| self.optimizer = optax.adam(0.05) | |
| self.opt_state = self.optimizer.init(self.logits) | |
| def _merge(y_old, w_old, y_new, w_new): | |
| """Combine two measurements of one marginal: inverse-variance | |
| weighted mean, precisions added.""" | |
| precision = w_old**2 + w_new**2 | |
| mean = np.divide(y_old * w_old**2 + y_new * w_new**2, precision, | |
| out=np.zeros_like(precision), where=precision > 0) | |
| return mean, np.sqrt(precision) | |
| def _group(self, shape): | |
| """Slots for cliques of this exact shape, created on first use.""" | |
| if shape not in self.groups: | |
| n = self._slots[shape] | |
| self.groups[shape] = dict( | |
| index=np.zeros((len(shape), n), np.int32), | |
| y=np.zeros((n,) + shape, np.float32), | |
| w=np.zeros((n,) + shape, np.float32), used=0) | |
| return self.groups[shape] | |
| def add(self, values, clique, stddev): | |
| y = np.asarray(values, np.float32) | |
| weight = 1.0 / stddev | |
| if len(clique) == 1: | |
| i = self.index[clique[0]] | |
| n = y.size | |
| self.y1[i, :n], self.w1[i, :n] = self._merge( | |
| self.y1[i, :n], self.w1[i, :n], y, weight) | |
| return | |
| shape = tuple(SIZES[c] for c in clique) | |
| group = self._group(shape) | |
| if clique not in self.slot: | |
| if group["used"] >= len(group["y"]): | |
| raise RuntimeError(f"no slot left for {clique}") | |
| self.slot[clique] = group["used"] | |
| for axis, col in enumerate(clique): | |
| group["index"][axis, group["used"]] = self.index[col] | |
| group["used"] += 1 | |
| k = self.slot[clique] | |
| group["y"][k], group["w"][k] = self._merge( | |
| group["y"][k], group["w"][k], y.reshape(shape), weight) | |
| def fit(self, iters): | |
| pairs = tuple((jnp.asarray(g["index"]), jnp.asarray(g["y"]), | |
| jnp.asarray(g["w"])) for g in self.groups.values()) | |
| one_way = tuple(jnp.asarray(x) for x in (self.y1, self.w1)) | |
| self.logits, self.opt_state = _train( | |
| self.logits, self.opt_state, self.valid, self.total, | |
| self.components, pairs, one_way, iters, self.optimizer) | |
| def all_pair_marginals(self): | |
| """Model estimate of every 2-way marginal, in gram block layout.""" | |
| return np.asarray(jax.block_until_ready( | |
| _gram_of(self.logits, self.valid, self.flat, self.total, | |
| self.components)), np.float64) | |
| def synthetic_data_rounded(self, rows, rng): | |
| """Draw records by randomised rounding. | |
| Every component gets the same number of rows, and within one, each | |
| column's expected counts are rounded to integers and shuffled. The | |
| rows are therefore not an independent sample. | |
| """ | |
| probs = np.asarray(_probs(self.logits, self.valid)) | |
| per = rows // self.components + 1 | |
| blocks = {attr: [] for attr in COLUMNS} | |
| for k in range(self.components): | |
| for i, attr in enumerate(COLUMNS): | |
| size = SIZES[attr] | |
| counts = probs[i, k, :size] | |
| counts = counts * per / counts.sum() | |
| frac, whole = np.modf(counts) | |
| whole = whole.astype(int) | |
| short = per - whole.sum() | |
| if short > 0: | |
| weights = frac / frac.sum() if frac.sum() > 0 else None | |
| picked = rng.choice(size, short, replace=size < short, | |
| p=weights) | |
| np.add.at(whole, picked, 1) | |
| elif short < 0: | |
| order = np.argsort(frac) | |
| for j in order[:(-short)]: | |
| if whole[j] > 0: | |
| whole[j] -= 1 | |
| values = np.repeat(np.arange(size), whole) | |
| rng.shuffle(values) | |
| blocks[attr].append(values) | |
| out = {a: np.concatenate(v)[:rows] for a, v in blocks.items()} | |
| order = rng.permutation(rows) | |
| return {a: v[order] for a, v in out.items()} | |
| def _probs(logits, valid): | |
| return jax.nn.softmax(jnp.where(valid[:, None, :], logits, -jnp.inf), axis=-1) | |
| def _gram_of(logits, valid, flat, total, components): | |
| probs = _probs(logits, valid) | |
| num, comp, size = probs.shape | |
| stacked = probs.transpose(1, 0, 2).reshape(comp, num * size)[:, flat] | |
| return stacked.T @ stacked * (total / components) | |
| def _loss(logits, valid, total, components, pairs, one_way): | |
| probs = _probs(logits, valid) | |
| loss = 0.0 | |
| for index, y, w in pairs: | |
| # "tka,tkb,...->tab...": for each slot t, contract the per-component | |
| # tables of the clique's columns. | |
| letters = "abcdefgh"[:index.shape[0]] | |
| formula = ",".join(f"tk{c}" for c in letters) + "->t" + letters | |
| tables = [probs[index[axis], :, :y.shape[axis + 1]] | |
| for axis in range(index.shape[0])] | |
| fitted = jnp.einsum(formula, *tables) * (total / components) | |
| loss = loss + jnp.sum((w * (fitted - y)) ** 2) | |
| y1, w1 = one_way | |
| fitted1 = probs.mean(1) * total | |
| return loss + jnp.sum((w1 * (fitted1 - y1)) ** 2) | |
| def _train(logits, opt_state, valid, total, components, pair, one_way, iters, | |
| optimizer): | |
| def body(carry, _): | |
| params, state = carry | |
| grad = jax.grad(_loss)(params, valid, total, components, pair, one_way) | |
| updates, state = optimizer.update(grad, state, params) | |
| return (optax.apply_updates(params, updates), state), None | |
| (logits, opt_state), _ = jax.lax.scan(body, (logits, opt_state), None, | |
| length=iters) | |
| return logits, opt_state | |
| def aim_select(sensitive): | |
| """Measure 1-way marginals and the named cliques, then adaptively buy | |
| ROUNDS 2-way marginals, in the order SCHEDULE declares them.""" | |
| t0 = time.perf_counter() | |
| measured = [] # (source, clique, sigma, values) per query | |
| one_way = [sensitive.measure((attr,), SIGMA_ONE_WAY) | |
| for attr in COLUMNS] | |
| named_values = [[sensitive.measure(clique, SIGMA_NAMED[name]) | |
| for clique in cliques] | |
| for name, cliques in CLIQUE_SETS.items()] | |
| # The record count, from the noisy 1-way marginals: each sums to the count | |
| # plus noise of variance sigma^2 * cells, so combine them by 1 / cells. | |
| totals = np.array([np.sum(values) for values in one_way]) | |
| cells = np.array([np.size(values) for values in one_way]) | |
| total = max(1.0, float(np.average(totals, weights=1 / cells))) | |
| print(f" estimated records {total:.0f}") | |
| print_noise_levels(total) | |
| model = ScoringModel(total, ROUNDS + len(ALL_NAMED_CLIQUES), SEED, | |
| COMPONENTS, extra=ALL_NAMED_CLIQUES) | |
| for attr, values in zip(COLUMNS, one_way): | |
| model.add(values, (attr,), SIGMA_ONE_WAY) | |
| measured.append(("one-way", (attr,), SIGMA_ONE_WAY, values)) | |
| for (name, cliques), values in zip(CLIQUE_SETS.items(), named_values): | |
| group_sigma = SIGMA_NAMED[name] | |
| for clique, y in zip(cliques, values): | |
| model.add(y, clique, group_sigma) | |
| measured.append(("named", clique, group_sigma, y)) | |
| model.fit(SCORING_WARMUP_ITERS) | |
| print(f" 1-way measured, scoring model warmed up " | |
| f"[{time.perf_counter() - t0:.1f}s]") | |
| t0 = time.perf_counter() | |
| for round_index in range(ROUNDS): | |
| pair = sensitive.select_pair( | |
| model.all_pair_marginals(), SELECTION_SCALE, SIGMA_PAIR) | |
| values = sensitive.measure(pair, SIGMA_PAIR) | |
| model.add(values, pair, SIGMA_PAIR) | |
| measured.append(("adaptive", pair, SIGMA_PAIR, values)) | |
| model.fit(SCORING_ROUND_ITERS) | |
| if (round_index + 1) % max(1, ROUNDS // 20) == 0: | |
| print(f" {round_index + 1}/{ROUNDS} pairs " | |
| f"[{time.perf_counter() - t0:.0f}s]", flush=True) | |
| return model, total, measured | |
| class _Tee: | |
| """Write to a stream and a file at once.""" | |
| def __init__(self, stream, path): | |
| self.stream = stream | |
| self.file = open(path, "w") | |
| def write(self, text): | |
| self.stream.write(text) | |
| self.file.write(text) | |
| self.file.flush() | |
| def flush(self): | |
| self.stream.flush() | |
| self.file.flush() | |
| def main(): | |
| parser = argparse.ArgumentParser( | |
| description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) | |
| parser.add_argument("--in", dest="source", required=True, metavar="FILE", | |
| help="the input CSV") | |
| parser.add_argument("--out", required=True, metavar="DIR", | |
| help="directory to write into: synthetic.csv," | |
| " synthetic_unsmoothed.csv, marginals.csv," | |
| " model.npz and run.log") | |
| args = parser.parse_args() | |
| os.makedirs(args.out, exist_ok=True) | |
| out = lambda name: os.path.join(args.out, name) | |
| sys.stdout = _Tee(sys.stdout, out("run.log")) | |
| print(f"target ({EPSILON}, {DELTA:g})-DP -> rho = {RHO:.4f} zCDP") | |
| start = time.perf_counter() | |
| path = os.path.expanduser(args.source) | |
| print(f"modelling the {len(COLUMNS)} columns schema.py declares") | |
| print(f" schedule: {len(COLUMNS)} one-way at sigma={SIGMA_ONE_WAY:.1f}, " | |
| + "".join(f"{len(c)} {n} at sigma={SIGMA_NAMED[n]:.1f}, " | |
| for n, c in CLIQUE_SETS.items()) | |
| + f"{ROUNDS} pairs at sigma={SIGMA_PAIR:.1f}, " | |
| f"{ROUNDS} selections at scale={SELECTION_SCALE:.1f}") | |
| sensitive = SensitiveData(path, SCHEDULE) | |
| print(f" OpenDP prices it at rho {sensitive.spent:.6f} of {RHO:.6f} " | |
| f"and refuses anything beyond it") | |
| print("stage 1: AIM selection") | |
| model, total, measured = aim_select(sensitive) | |
| print(f"stage 2: mixture-of-products estimation ({COMPONENTS} components)") | |
| t0 = time.perf_counter() | |
| model.fit(FINAL_ITERS) | |
| rows = int(round(total)) | |
| rng = np.random.default_rng(SEED) | |
| codes = model.synthetic_data_rounded(rows, rng) | |
| print(f" fit and generated {rows} records in {time.perf_counter() - t0:.0f}s") | |
| np.savez_compressed( | |
| out("model.npz"), | |
| logits=np.asarray(model.logits), total=np.float64(total), | |
| attributes=np.array(COLUMNS, dtype=object), | |
| sizes=np.array([SIZES[c] for c in COLUMNS])) | |
| print(f" wrote the fitted model to {out('model.npz')}") | |
| write_marginals(measured, out("marginals.csv")) | |
| print(f" wrote {len(measured):,} marginals to {out('marginals.csv')}") | |
| frame = pd.DataFrame(decode(codes, rng)) | |
| # The model's own output, before any repair, so the two can be compared. | |
| frame.to_csv(out("synthetic_unsmoothed.csv"), index=False) | |
| print(f" wrote {out('synthetic_unsmoothed.csv')}") | |
| frame, report = smooth(frame, SEED) | |
| repaired = sum(count for _, _, count in report) | |
| print(f" repaired {repaired} impossible flag combinations") | |
| frame.to_csv(out("synthetic.csv"), index=False) | |
| print(f"wrote {out('synthetic.csv')} " | |
| f"total wall time {time.perf_counter() - start:.0f}s") | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 29.8 kB
- Xet hash:
- 842960ed13d0b2d172d4203cd9177c87855a95c58688ed9e57771396b159cf63
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.