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
14 changes: 7 additions & 7 deletions dpsynth/adapters/beam.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@

from __future__ import annotations

from collections.abc import Callable
from collections.abc import Callable, Mapping, Sequence
import dataclasses
import io
import math
Expand Down Expand Up @@ -227,9 +227,9 @@ class _EncodeAndProject(beam.DoFn):

def __init__(
self,
column_measurements: dict[str, initialization.ColumnMeasurement],
domains: dict[str, Any],
workload: list[mbi.Clique],
column_measurements: Mapping[str, initialization.ColumnMeasurement],
domains: Mapping[str, Any],
workload: Sequence[mbi.Clique],
):
super().__init__()
# Reuse the shared per-column codec so Beam encoding matches the in-memory
Expand Down Expand Up @@ -298,9 +298,9 @@ class ComputeMarginals(beam.PTransform):

def __init__(
self,
column_measurements: dict[str, initialization.ColumnMeasurement],
domains: dict[str, Any],
workload: list[mbi.Clique],
column_measurements: Mapping[str, initialization.ColumnMeasurement],
domains: Mapping[str, Any],
workload: Sequence[mbi.Clique],
):
super().__init__()
self._column_measurements = column_measurements
Expand Down
5 changes: 3 additions & 2 deletions dpsynth/adapters/pydantic_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@

"""Pydantic <--> DataFrame conversion utilities for TabularSynthesizer."""

from collections.abc import Mapping
import enum
import inspect
import math
Expand Down Expand Up @@ -141,7 +142,7 @@ def infer_domain_from_model(

def models_to_dataframe(
records: list[RecordT],
domains_dict: dict[str, domain.AttributeType],
domains_dict: Mapping[str, domain.AttributeType],
) -> pd.DataFrame:
"""Converts a list of pydantic models to a TabularSynthesizer-compatible DataFrame.

Expand All @@ -166,7 +167,7 @@ def models_to_dataframe(
def dataframe_to_models(
df: pd.DataFrame | data_generation_v3.DataGenerationResult,
model_cls: type[RecordT],
domains_dict: dict[str, domain.AttributeType],
domains_dict: Mapping[str, domain.AttributeType],
) -> list[RecordT]:
"""Converts a synthetic DataFrame back to pydantic model instances.

Expand Down
6 changes: 3 additions & 3 deletions dpsynth/discrete_mechanisms/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -197,8 +197,8 @@ def exponential_mechanism(

def measure_marginals_with_noise(
rng: np.random.Generator,
data: mbi.Projectable,
marginal_queries: list[tuple[str, ...]],
data: mbi.Projectable | mbi.Dataset,
marginal_queries: Sequence[tuple[str, ...]],
gdp_sigma: float,
weights: np.ndarray | None = None,
max_records_per_user: int = 1,
Expand Down Expand Up @@ -443,7 +443,7 @@ def score(cl):


def compute_independence_errors(
data: mbi.Projectable,
data: mbi.Projectable | mbi.Dataset,
model: mbi.MarkovRandomField,
cliques: Sequence[mbi.Clique],
) -> dict[mbi.Clique, float]:
Expand Down
2 changes: 1 addition & 1 deletion dpsynth/discrete_mechanisms/direct.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,7 @@ def configure(self, _=None, *, zcdp_rho, delta=0, max_records_per_user=1):

marginal_oracle: mbi.MarginalOracle | None = None
pgm_iters: int = 5000
prespecified_marginal_queries: list[tuple[str, ...]] = dataclasses.field(
prespecified_marginal_queries: Sequence[tuple[str, ...]] = dataclasses.field(
default_factory=list
)

Expand Down
8 changes: 5 additions & 3 deletions dpsynth/domain.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,7 @@

PathType = pathlib.Path

CategoricalValue: TypeAlias = bool | int | str
CategoricalValue: TypeAlias = bool | int | float | str

IntervalHandling = Literal['midpoint', 'sample', 'interval']

Expand Down Expand Up @@ -193,7 +193,7 @@ class NumericalAttribute:
min_value: float
max_value: float
clip_to_range: bool = True
sentinel: float | int | str | None = None
sentinel: float | int | str | np.integer | np.floating | None = None
dtype: str = 'float'
interval_handling: str = 'midpoint'
description: str | None = None
Expand Down Expand Up @@ -232,7 +232,9 @@ def __post_init__(self):
)

@property
def resolved_sentinel(self) -> float | int | str:
def resolved_sentinel(
self,
) -> float | int | str | np.integer | np.floating:
"""Returns the effective sentinel, with mode-appropriate defaults."""
if self.sentinel is not None:
return self.sentinel
Expand Down
31 changes: 31 additions & 0 deletions dpsynth/local_mode/primitives.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@

from __future__ import annotations

from typing import overload

import numpy as np
import scipy.stats
Expand Down Expand Up @@ -321,6 +322,36 @@ def _select_partitions_sips(
# ---------------------------------------------------------------------------


@overload
def add_gaussian_noise(
rng: np.random.Generator,
counts: float | int,
sigma: float,
max_records_per_user: int = 1,
) -> float:
...


@overload
def add_gaussian_noise(
rng: np.random.Generator,
counts: np.ndarray,
sigma: float,
max_records_per_user: int = 1,
) -> np.ndarray:
...


@overload
def add_gaussian_noise(
rng: np.random.Generator,
counts: np.ndarray | float | int,
sigma: float,
max_records_per_user: int = 1,
) -> float | np.ndarray:
...


def add_gaussian_noise(
rng: np.random.Generator,
counts: np.ndarray | float | int,
Expand Down
8 changes: 6 additions & 2 deletions dpsynth/relational/domain.py
Original file line number Diff line number Diff line change
Expand Up @@ -232,7 +232,9 @@ def from_yaml_file(


def to_dict(
table_domains: Mapping[str, domain.Schema],
table_domains: Mapping[
str, domain.Schema | Mapping[str, domain.AttributeType]
],
foreign_keys: Sequence[ForeignKeyRelation] = (),
) -> dict[str, Any]:
"""Converts multi-table schemas and foreign keys to a dictionary.
Expand Down Expand Up @@ -261,7 +263,9 @@ def to_dict(


def to_yaml_file(
table_domains: Mapping[str, domain.Schema],
table_domains: Mapping[
str, domain.Schema | Mapping[str, domain.AttributeType]
],
foreign_keys: Sequence[ForeignKeyRelation],
filepath: str | PathType,
) -> None:
Expand Down
16 changes: 10 additions & 6 deletions dpsynth/relational/synthesizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@
from collections.abc import Collection, Hashable, Mapping, Sequence
import dataclasses
import math
from typing import Any, Literal
from typing import Any, Literal, TypeAlias

from absl import logging
import dp_accounting
Expand All @@ -42,9 +42,13 @@
_LOGGING_UNUSED = logging
# pylint: enable=unused-import

TableDomains: TypeAlias = Mapping[
str, domain.Schema | Mapping[str, domain.AttributeType]
]


def _validate_input_table_columns(
domains: Mapping[str, domain.Schema],
domains: TableDomains,
foreign_keys: Sequence[rel_domain.ForeignKeyRelation],
table_columns: Mapping[str, Collection[str]],
) -> None:
Expand Down Expand Up @@ -264,7 +268,7 @@ def _run_table_initializers(


def _encode_and_compress_tables(
domains: Mapping[str, domain.Schema],
domains: TableDomains,
table_measurements: Mapping[
str, Mapping[str, initialization.ColumnMeasurement]
],
Expand Down Expand Up @@ -419,7 +423,7 @@ def _run_table_preprocessing(


def _create_table_initializers(
domains: Mapping[str, domain.Schema],
domains: TableDomains,
numerical_bins: int,
) -> dict[str, dict[str, api.MechanismConfig]]:
"""Creates per-table and per-column initializers from relational schemas."""
Expand All @@ -430,7 +434,7 @@ def _create_table_initializers(


def _compute_table_col_deltas(
domains: Mapping[str, domain.Schema],
domains: TableDomains,
delta: float,
init_budget_fraction: float,
) -> dict[str, dict[str, float]]:
Expand Down Expand Up @@ -847,7 +851,7 @@ def _decompress_synthetic_datasets(
def _decode_synthetic_tables(
decompressed_datasets: Mapping[str, mbi.Dataset],
column_codecs: Mapping[str, data_generation_v3.TabularCodec],
domains: Mapping[str, domain.Schema],
domains: TableDomains,
rng: np.random.Generator,
) -> dict[str, pd.DataFrame]:
"""Decodes discrete datasets into continuous/categorical DataFrames.
Expand Down
4 changes: 2 additions & 2 deletions dpsynth/text/bulk_inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@
import random
import re
import time
from typing import Protocol, TypeVar
from typing import Any, Protocol, TypeVar

from absl import logging
from dpsynth import domain
Expand Down Expand Up @@ -376,7 +376,7 @@ def _categorical_json_type(

def domain_to_json_schema(
domain_spec: Mapping[str, domain.AttributeType],
) -> dict[str, object]:
) -> dict[str, Any]:
"""Converts a dpsynth Domain to a JSON schema dict for GenAI."""
properties = {}
for name, attr in domain_spec.items():
Expand Down
40 changes: 21 additions & 19 deletions dpsynth/transformations.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,20 +16,22 @@

import bisect
from collections.abc import Callable, Mapping, Sequence
import dataclasses
import math
from typing import Any, Generic, TypeAlias, TypeVar

import attr
from dpsynth import domain
import numpy as np
import pandas as pd

CategoricalValue: TypeAlias = bool | int | float | str
R, T, S = TypeVar('R'), TypeVar('T'), TypeVar('S')
R = TypeVar('R')
T = TypeVar('T')
S = TypeVar('S')


@attr.define(frozen=True)
class DataTransformation(Generic[R, T]): # pyrefly: ignore[not-a-type]
@dataclasses.dataclass(frozen=True)
class DataTransformation(Generic[R, T]):
"""Dataclass for transforming data from one domain to another.

DataTransformations are both reversible (via inverse) and composable (via @).
Expand All @@ -48,22 +50,22 @@ class DataTransformation(Generic[R, T]): # pyrefly: ignore[not-a-type]
1
"""

transform: Callable[[R], T] | Mapping[R, T] = attr.field() # pyrefly: ignore[not-a-type]
inverse_transform: Callable[[T], R] | Mapping[T, R] = attr.field() # pyrefly: ignore[not-a-type]
transform: Callable[[R], T] | Mapping[R, T]
inverse_transform: Callable[[T], R] | Mapping[T, R]

def __call__(self, value: R) -> T: # pyrefly: ignore[not-a-type]
def __call__(self, value: R) -> T:
if isinstance(self.transform, Mapping):
return self.transform[value]
return self.transform(value)

@property
def inverse(self) -> 'DataTransformation[T, R]': # pyrefly: ignore[not-a-type]
def inverse(self) -> 'DataTransformation[T, R]':
"""The reverse transformation of this instance."""
return DataTransformation(self.inverse_transform, self.transform) # pyrefly: ignore[bad-argument-count]
return DataTransformation(self.inverse_transform, self.transform)

def __matmul__(
self, other: 'DataTransformation[T, S]' # pyrefly: ignore[not-a-type]
) -> 'DataTransformation[R, S]': # pyrefly: ignore[not-a-type]
self, other: 'DataTransformation[S, R]'
) -> 'DataTransformation[S, T]':
"""Returns a DataTransformation that composes this instance with other.

Example Usage:
Expand All @@ -82,8 +84,8 @@ def __matmul__(
A DataTransformation that composes this instance with other.
"""
return DataTransformation(
lambda x: self(other(x)), # pyrefly: ignore[bad-argument-count]
lambda x: other.inverse(self.inverse(x)),
lambda x: self(other(x)), # pyrefly: ignore[bad-argument-type]
lambda x: other.inverse(self.inverse(x)), # pyrefly: ignore[bad-argument-type]
)


Expand Down Expand Up @@ -129,10 +131,10 @@ def discrete_encoder(
ood = attribute_domain.out_of_domain_index
transform = lambda v: index_map.get(value_type(v), ood)
reverse = dict(enumerate(attribute_domain.possible_values))
return DataTransformation(transform, reverse) # pyrefly: ignore[bad-argument-count]
return DataTransformation(transform, reverse)


@attr.define(frozen=True)
@dataclasses.dataclass(frozen=True)
class _Interval:
"""A numeric interval with a string representation."""

Expand Down Expand Up @@ -199,7 +201,7 @@ def create_discretize_transformation(
attribute_domain.max_value,
]
intervals = [
_Interval(left, right, closed_left=(i == 0)) # pyrefly: ignore[bad-argument-count, unexpected-keyword]
_Interval(left, right, closed_left=(i == 0))
for i, (left, right) in enumerate(zip(bin_edges[:-1], bin_edges[1:]))
]
interval_strs = [str(iv) for iv in intervals]
Expand All @@ -216,7 +218,7 @@ def transform(value: Any) -> str:
idx = bisect.bisect_left(inner_edges, value)
return interval_strs[idx]

def reverse(value: str) -> float | str:
def reverse(value: str) -> float | int | str | np.integer | np.floating:
if value == ood_sentinel:
return sentinel
idx = interval_strs.index(value)
Expand All @@ -231,8 +233,8 @@ def reverse(value: str) -> float | str:
return math.ceil(result)
return result

new_domain = domain.CategoricalAttribute(possible_values) # pyrefly: ignore[bad-argument-count]
transformation = DiscretizeTransformation(transform, reverse) # pyrefly: ignore[bad-argument-count]
new_domain = domain.CategoricalAttribute(possible_values)
transformation = DiscretizeTransformation(transform, reverse)
return new_domain, transformation


Expand Down
Loading
Loading