Buckets:
| #!/usr/bin/env python | |
| """Generate records from a saved model, without refitting. | |
| dp_synth.py writes model.npz alongside its output: the fitted mixture of | |
| products, which is what the privacy budget actually bought. Drawing records | |
| from it is post-processing, so any number of records, drawn any way, costs | |
| nothing further and cannot weaken the guarantee. | |
| Two samplers, because they are not interchangeable: | |
| rounded every component contributes the same number of rows, and within | |
| one, each column's expected counts are rounded to integers. The | |
| realised marginals then match the model's closely. This is what | |
| dp_synth.py writes, and it is the better choice for reading rates | |
| off the data. | |
| iid each row is drawn independently: pick a component, then each column | |
| from it. Noisier -- at this many rows the sampling noise is about | |
| the size of the errors the evaluation measures -- but the rows are | |
| a genuine independent sample, so a standard error computed on them | |
| means what it usually means. Under `rounded` the marginals are | |
| pinned to the model and anything assuming independence will read as | |
| more certain than it should. | |
| Usage: | |
| ./generate.py --model production/model.npz --out more.csv --rows 50000 | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import numpy as np | |
| import pandas as pd | |
| from schema import decode | |
| from smooth import smooth | |
| def load(path): | |
| """The fitted per-column, per-component probabilities.""" | |
| blob = np.load(path, allow_pickle=True) | |
| logits, sizes = blob["logits"], blob["sizes"] | |
| attributes = [str(a) for a in blob["attributes"]] | |
| probs = np.empty_like(logits) | |
| for i, size in enumerate(sizes): | |
| row = logits[i, :, :size] | |
| row = np.exp(row - row.max(axis=1, keepdims=True)) | |
| probs[i, :, :size] = row / row.sum(axis=1, keepdims=True) | |
| return probs, attributes, dict(zip(attributes, sizes)), float(blob["total"]) | |
| def draw_iid(probs, attributes, sizes, rows, rng): | |
| which = rng.integers(0, probs.shape[1], rows) | |
| uniform = rng.random((rows, 1)) | |
| out = {} | |
| for i, attr in enumerate(attributes): | |
| size = sizes[attr] | |
| table = probs[i, which, :size] | |
| out[attr] = (table.cumsum(1) < uniform).sum(1).clip(0, size - 1) | |
| return out | |
| def draw_rounded(probs, attributes, sizes, rows, rng): | |
| components = probs.shape[1] | |
| per = rows // components + 1 | |
| blocks = {attr: [] for attr in attributes} | |
| for k in range(components): | |
| for i, attr in enumerate(attributes): | |
| size = sizes[attr] | |
| counts = probs[i, k, :size] * per / probs[i, k, :size].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 | |
| np.add.at(whole, rng.choice(size, short, replace=size < short, | |
| p=weights), 1) | |
| elif short < 0: | |
| for j in np.argsort(frac)[:(-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 main(): | |
| parser = argparse.ArgumentParser( | |
| description=__doc__, | |
| formatter_class=argparse.RawDescriptionHelpFormatter) | |
| parser.add_argument("--model", required=True, metavar="FILE", | |
| help="the model.npz a run wrote") | |
| parser.add_argument("--out", required=True, metavar="FILE") | |
| parser.add_argument("--rows", type=int, default=None, | |
| help="how many records (default: as many as the run" | |
| " estimated the input had)") | |
| parser.add_argument("--sampler", choices=("rounded", "iid"), | |
| default="rounded") | |
| parser.add_argument("--seed", type=int, default=0) | |
| parser.add_argument("--no-repair", action="store_true", | |
| help="skip the repair of impossible flag combinations") | |
| args = parser.parse_args() | |
| probs, attributes, sizes, total = load(args.model) | |
| rows = args.rows if args.rows is not None else int(round(total)) | |
| rng = np.random.default_rng(args.seed) | |
| draw = draw_rounded if args.sampler == "rounded" else draw_iid | |
| frame = pd.DataFrame(decode(draw(probs, attributes, sizes, rows, rng), rng)) | |
| print(f"{rows:,} records, {args.sampler} sampler, " | |
| f"{probs.shape[1]} components") | |
| if not args.no_repair: | |
| frame, report = smooth(frame, args.seed) | |
| print(f" repaired {sum(c for _, _, c in report)} impossible " | |
| f"flag combinations") | |
| frame.to_csv(args.out, index=False) | |
| print(f"wrote {args.out}") | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 5.07 kB
- Xet hash:
- 4c0f0adf5b7ece48bc4243c3c984ec0b454b4c9ebb5d0c93718da2eef313445a
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.