1717)
1818from sklearn .preprocessing import MinMaxScaler
1919
20+ from .anonymizer import generate_root_buffers , hash_strings
2021from .bucket import Buckets
21- from .common import ColumnId , ColumnType , Value , check_column_names_or_ids
22+ from .common import (
23+ AnonymizationContext ,
24+ AnonymizationParams ,
25+ ColumnId ,
26+ ColumnType ,
27+ Value ,
28+ check_column_names_or_ids ,
29+ )
2230from .interval import Interval , Intervals
2331from .tree import Branch , Leaf , Node
2432
3341
3442
3543class DataConvertor (ABC ):
36- def __init__ (self ) -> None :
44+ def __init__ (self , column : str , anonymization_params : AnonymizationParams ) -> None :
3745 self .scaler : Optional [MinMaxScaler ] = None
3846 self .value_safe_flag : bool = False
3947
48+ base_seed = hash_strings (iter ([str (column )]))
49+ self .lower_buffer , self .upper_buffer = generate_root_buffers (
50+ AnonymizationContext (base_seed , anonymization_params )
51+ )
52+
4053 @abstractmethod
4154 def column_type (self ) -> ColumnType :
4255 pass
@@ -64,8 +77,8 @@ def denormalize_safe_values(self) -> None:
6477
6578
6679class BooleanConvertor (DataConvertor ):
67- def __init__ (self ) -> None :
68- super ().__init__ ()
80+ def __init__ (self , column : str , anonymization_params : AnonymizationParams ) -> None :
81+ super ().__init__ (column , anonymization_params )
6982
7083 def column_type (self ) -> ColumnType :
7184 return ColumnType .BOOLEAN
@@ -84,13 +97,12 @@ def create_value_safe_set(self, values: pd.Series) -> None:
8497
8598
8699class RealConvertor (DataConvertor ):
87- def __init__ (self , values : Iterable [Value ]) -> None :
88- super ().__init__ ()
89- # Fit up to 0.9999 so that the max bucket range is [0-1)
90- self .scaler = MinMaxScaler (feature_range = (0.0 , 0.9999 )) # type: ignore
100+ def __init__ (self , column : str , anonymization_params : AnonymizationParams , values : Iterable [Value ]) -> None :
101+ super ().__init__ (column , anonymization_params )
102+ self .scaler = MinMaxScaler (feature_range = (self .lower_buffer , self .upper_buffer )) # type: ignore
91103 # This value-neutral fitting is only for passing unit tests, gets overridden
92104 # later by fit_transform().
93- self .scaler .fit (np .array ([[0.0 ], [0.9999 ]]))
105+ self .scaler .fit (np .array ([[self . lower_buffer ], [self . upper_buffer ]]))
94106 self .final_round_precision = _get_round_precision (cast (Iterable [float ], values ))
95107
96108 def column_type (self ) -> ColumnType :
@@ -116,13 +128,12 @@ def create_value_safe_set(self, values: pd.Series) -> None:
116128
117129
118130class IntegerConvertor (DataConvertor ):
119- def __init__ (self ) -> None :
120- super ().__init__ ()
121- # Fit up to 0.9999 so that the max bucket range is [0-1)
122- self .scaler = MinMaxScaler (feature_range = (0.0 , 0.9999 )) # type: ignore
131+ def __init__ (self , column : str , anonymization_params : AnonymizationParams ) -> None :
132+ super ().__init__ (column , anonymization_params )
133+ self .scaler = MinMaxScaler (feature_range = (self .lower_buffer , self .upper_buffer )) # type: ignore
123134 # This value-neutral fitting is only for passing unit tests, gets overridden
124135 # later by fit_transform().
125- self .scaler .fit (np .array ([[0.0 ], [0.9999 ]]))
136+ self .scaler .fit (np .array ([[self . lower_buffer ], [self . upper_buffer ]]))
126137
127138 def column_type (self ) -> ColumnType :
128139 return ColumnType .INTEGER
@@ -146,13 +157,12 @@ def create_value_safe_set(self, values: pd.Series) -> None:
146157
147158
148159class TimestampConvertor (DataConvertor ):
149- def __init__ (self ) -> None :
150- super ().__init__ ()
151- # Fit up to 0.9999 so that the max bucket range is [0-1)
152- self .scaler = MinMaxScaler (feature_range = (0.0 , 0.9999 )) # type: ignore
160+ def __init__ (self , column : str , anonymization_params : AnonymizationParams ) -> None :
161+ super ().__init__ (column , anonymization_params )
162+ self .scaler = MinMaxScaler (feature_range = (self .lower_buffer , self .upper_buffer )) # type: ignore
153163 # This value-neutral fitting is only for passing unit tests, gets overridden
154164 # later by fit_transform().
155- self .scaler .fit (np .array ([[0.0 ], [0.9999 ]]))
165+ self .scaler .fit (np .array ([[self . lower_buffer ], [self . upper_buffer ]]))
156166
157167 def column_type (self ) -> ColumnType :
158168 return ColumnType .TIMESTAMP
@@ -177,8 +187,8 @@ def create_value_safe_set(self, values: pd.Series) -> None:
177187
178188
179189class StringConvertor (DataConvertor ):
180- def __init__ (self , values : Iterable [Value ]) -> None :
181- super ().__init__ ()
190+ def __init__ (self , column : str , anonymization_params : AnonymizationParams , values : Iterable [Value ]) -> None :
191+ super ().__init__ (column , anonymization_params )
182192 unique_values = set ()
183193 for v in values :
184194 if not pd .isna (v ):
@@ -196,7 +206,6 @@ def __init__(self, values: Iterable[Value]) -> None:
196206
197207 # Note that self.safe_values is only used if self.value_safe_flag is False
198208 self .safe_values : Set [float ] = set ()
199- # Fit up to 0.9999 so that the max bucket range is [0-1)
200209 self .scaler = MinMaxScaler (feature_range = (0.0 , 0.9999 )) # type: ignore
201210 # This value-neutral fitting is only for passing unit tests, gets overridden
202211 # later by fit_transform().
@@ -231,7 +240,7 @@ def _map_interval(self, interval: Interval, rng: Random) -> MicrodataValue:
231240 min_value = int (interval .min )
232241 # max_value is inclusive
233242 max_value = min (int (interval .max ) - 1 , len (self .value_map ) - 1 )
234- # The latter term in the above line can 0 (not sure why TODO: check)
243+ # The latter term in the above line can be 0 (not sure why TODO: check)
235244 max_value = max (min_value , max_value )
236245 value = rng .randint (min_value , max_value )
237246 if self .value_safe_flag is True or value in self .safe_values :
@@ -343,19 +352,19 @@ def _microdata_row_generator(
343352 yield [_generate (i , c , nm , rng ) for i , c , nm in zip (intervals , convertors , null_mappings )]
344353
345354
346- def get_convertor (df : pd .DataFrame , column : str ) -> DataConvertor :
355+ def get_convertor (df : pd .DataFrame , column : str , anonymization_params : AnonymizationParams ) -> DataConvertor :
347356 dtype = df .dtypes [column ]
348357 if is_integer_dtype (dtype ):
349- return IntegerConvertor ()
358+ return IntegerConvertor (column , anonymization_params )
350359 elif is_float_dtype (dtype ):
351- return RealConvertor (df [column ])
360+ return RealConvertor (column , anonymization_params , df [column ])
352361 elif is_bool_dtype (dtype ):
353- return BooleanConvertor ()
362+ return BooleanConvertor (column , anonymization_params )
354363 elif is_datetime64_dtype (dtype ):
355- return TimestampConvertor ()
364+ return TimestampConvertor (column , anonymization_params )
356365 elif is_string_dtype (dtype ):
357366 # Note above is `True` for `object` dtype, but `StringConvertor` will assert values are `str`.
358- return StringConvertor (df [column ])
367+ return StringConvertor (column , anonymization_params , df [column ])
359368 else :
360369 raise TypeError (f"Dtype { dtype } is not supported." )
361370
0 commit comments