Skip to content
Merged
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
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -110,7 +110,7 @@ Config:
"transformers": {
"last_login": "UnixTimestampEncoder()",
"email_optin": "BinaryEncoder()",
"credit_card": "FrequencyEncoder()",
"credit_card": "UniformEncoder()",
"age": "FloatFormatter()",
"dollars_spent": "FloatFormatter()"
}
Expand Down
96 changes: 27 additions & 69 deletions rdt/hyper_transformer.py
Original file line number Diff line number Diff line change
Expand Up @@ -288,68 +288,33 @@ def set_config(self, config):
warnings.warn(self._REFIT_MESSAGE)

def _validate_update_transformers_by_sdtype(
self, sdtype, transformer, transformer_name, transformer_parameters
self, sdtype, transformer_name, transformer_parameters
):
if not self.field_sdtypes:
raise ConfigNotSetError(
'Nothing to update. Use the `detect_initial_config` method to '
'pre-populate all the sdtypes and transformers from your dataset.'
)

if transformer_name is None:
if transformer is None:
raise InvalidConfigError("Missing required parameter 'transformer_name'.")

if not isinstance(transformer, BaseTransformer):
raise InvalidConfigError(
'Invalid transformer. Please input an rdt transformer object.'
)

if sdtype not in transformer.get_supported_sdtypes():
raise InvalidConfigError(
"The transformer you've assigned is incompatible with the sdtype."
)

else:
if (
transformer_name not in get_class_by_transformer_name()
or sdtype
not in get_class_by_transformer_name()[transformer_name].get_supported_sdtypes()
):
raise InvalidConfigError(
f"Invalid transformer name '{transformer_name}' for the '{sdtype}' sdtype."
)

if transformer_parameters is not None:
transformer = get_class_by_transformer_name()[transformer_name]
valid = inspect.signature(transformer).parameters
invalid_parameters = {arg for arg in transformer_parameters if arg not in valid}
if invalid_parameters:
raise TransformerInputError(
f'Invalid parameters {tuple(sorted(invalid_parameters))} '
f"for the '{transformer_name}'."
)

def _warn_update_transformers_by_sdtype(self, transformer, transformer_name):
if self._fitted:
warnings.warn(self._REFIT_MESSAGE)
if (
transformer_name not in get_class_by_transformer_name()
or sdtype
not in get_class_by_transformer_name()[transformer_name].get_supported_sdtypes()
):
raise InvalidConfigError(
f"Invalid transformer name '{transformer_name}' for the '{sdtype}' sdtype."
)

if transformer_name is not None:
if transformer is not None:
warnings.warn(
"The 'transformer' parameter will no longer be supported in future versions "
"of the RDT. Using the 'transformer_name' parameter instead.",
FutureWarning,
if transformer_parameters is not None:
transformer = get_class_by_transformer_name()[transformer_name]
valid = inspect.signature(transformer).parameters
invalid_parameters = {arg for arg in transformer_parameters if arg not in valid}
if invalid_parameters:
raise TransformerInputError(
f'Invalid parameters {tuple(sorted(invalid_parameters))} '
f"for the '{transformer_name}'."
)

else:
warnings.warn(
"The 'transformer' parameter will no longer be supported in future versions "
"of the RDT. Please use the 'transformer_name' and 'transformer_parameters' "
'parameters instead.',
FutureWarning,
)

def _remove_column_in_multi_column_fields(self, column):
"""Remove a column that is part of a multi-column field.

Expand Down Expand Up @@ -401,42 +366,35 @@ def _update_multi_column_transformer(self):
def update_transformers_by_sdtype(
self,
sdtype,
transformer=None,
transformer_name=None,
transformer_parameters=None,
):
"""Update the transformers for the specified ``sdtype``.

Given an ``sdtype`` and a ``transformer``, change all the fields of the ``sdtype``
Given an ``sdtype`` and a ``transformer_name``, change all the fields of the ``sdtype``
to use the given transformer.

Args:
sdtype (str):
Semantic data type for the transformer.
transformer (rdt.transformers.BaseTransformer):
Transformer class or instance to be used for the given ``sdtype``.
Note: this parameter is deprecated, use ``transformer_name`` and
``transformer_parameters`` instead.
transformer_name (str):
A string with the class name of the transformer.
transformer_parameters (dict):
A dict of the kwargs of the transformer.
"""
self._validate_update_transformers_by_sdtype(
sdtype, transformer, transformer_name, transformer_parameters
sdtype, transformer_name, transformer_parameters
)
self._warn_update_transformers_by_sdtype(transformer, transformer_name)

transformer_instance = transformer
if self._fitted:
warnings.warn(self._REFIT_MESSAGE)

if transformer_name is not None:
if transformer_parameters is not None:
transformer_instance = get_class_by_transformer_name()[transformer_name](
**transformer_parameters
)
if transformer_parameters is not None:
transformer_instance = get_class_by_transformer_name()[transformer_name](
**transformer_parameters
)

else:
transformer_instance = get_class_by_transformer_name()[transformer_name]()
else:
transformer_instance = get_class_by_transformer_name()[transformer_name]()

for field, field_sdtype in self.field_sdtypes.items():
if field_sdtype == sdtype:
Expand Down
2 changes: 0 additions & 2 deletions rdt/performance/datasets/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,6 @@
id,
numerical,
pii,
text,
)
from rdt.performance.datasets.base import BaseDatasetGenerator

Expand All @@ -20,7 +19,6 @@
'id',
'numerical',
'pii',
'text',
'BaseDatasetGenerator',
]

Expand Down
57 changes: 0 additions & 57 deletions rdt/performance/datasets/text.py

This file was deleted.

23 changes: 9 additions & 14 deletions rdt/transformers/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,6 @@
from rdt.transformers.base import BaseMultiColumnTransformer, BaseTransformer
from rdt.transformers.boolean import BinaryEncoder
from rdt.transformers.categorical import (
FrequencyEncoder,
LabelEncoder,
OneHotEncoder,
UniformEncoder,
Expand All @@ -19,7 +18,7 @@
OptimizedTimestampEncoder,
UnixTimestampEncoder,
)
from rdt.transformers.id import IDGenerator, IndexGenerator, RegexGenerator
from rdt.transformers.id import IndexGenerator, RegexGenerator
from rdt.transformers.null import NullTransformer
from rdt.transformers.numerical import (
ClusterBasedNormalizer,
Expand All @@ -36,7 +35,6 @@
AnonymizedFaker,
PseudoAnonymizedFaker,
)
from rdt.transformers.utils import WarnDict

__all__ = [
'BaseTransformer',
Expand All @@ -45,7 +43,6 @@
'ClusterBasedNormalizer',
'OrderedLabelEncoder',
'FloatFormatter',
'FrequencyEncoder',
'GaussianNormalizer',
'LabelEncoder',
'LogScaler',
Expand All @@ -56,7 +53,6 @@
'RegexGenerator',
'AnonymizedFaker',
'PseudoAnonymizedFaker',
'IDGenerator',
'IndexGenerator',
'get_transformer_name',
'get_transformer_class',
Expand Down Expand Up @@ -94,15 +90,14 @@ def get_transformer_name(transformer):
for transformer in BaseTransformer.get_subclasses()
}

DEFAULT_TRANSFORMERS = WarnDict(
boolean=UniformEncoder(),
categorical=UniformEncoder(),
datetime=UnixTimestampEncoder(),
id=RegexGenerator(),
numerical=FloatFormatter(),
pii=AnonymizedFaker(),
text=RegexGenerator(),
)
DEFAULT_TRANSFORMERS = {
'boolean': UniformEncoder(),
'categorical': UniformEncoder(),
'datetime': UnixTimestampEncoder(),
'id': RegexGenerator(),
'numerical': FloatFormatter(),
'pii': AnonymizedFaker(),
}


@lru_cache()
Expand Down
42 changes: 1 addition & 41 deletions rdt/transformers/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,16 +111,6 @@ def reset_randomization(self):
'reverse_transform': np.random.RandomState(self.random_seed + 1),
}

@property
def model_missing_values(self):
"""Whether or not a new column is being used to model missing values."""
warnings.warn(
"Future versions of RDT will not support the 'model_missing_values' parameter. "
"Please switch to using the 'missing_value_generation' parameter instead.",
FutureWarning,
)
return self.missing_value_generation == 'from_column'

def _set_missing_value_generation(self, missing_value_generation):
if missing_value_generation not in (None, 'from_column', 'random'):
raise TransformerInputError(
Expand All @@ -130,18 +120,6 @@ def _set_missing_value_generation(self, missing_value_generation):

self.missing_value_generation = missing_value_generation

def _set_model_missing_values(self, model_missing_values):
warnings.warn(
"Future versions of RDT will not support the 'model_missing_values' parameter. "
"Please switch to using the 'missing_value_generation' parameter to select your "
'strategy.',
FutureWarning,
)
if model_missing_values is True:
self._set_missing_value_generation('from_column')
elif model_missing_values is False:
self._set_missing_value_generation('random')

@classmethod
def get_name(cls):
"""Return transformer name.
Expand Down Expand Up @@ -169,20 +147,6 @@ def get_subclasses(cls):

return subclasses

@classmethod
def get_input_sdtype(cls):
"""Return the input sdtype supported by the transformer.

Returns:
string:
Accepted input sdtype of the transformer.
"""
warnings.warn(
'`get_input_sdtype` is deprecated. Please use `get_supported_sdtypes` instead.',
FutureWarning,
)
return cls.get_supported_sdtypes()[0]

@classmethod
def get_supported_sdtypes(cls):
"""Return the supported sdtypes by the transformer.
Expand Down Expand Up @@ -327,11 +291,7 @@ def __repr__(self):
custom_args = []
args = inspect.getfullargspec(self.__init__)
keys = args.args[1:]
instanced = {
key: getattr(self, key)
for key in keys
if key != 'model_missing_values' and hasattr(self, key) # Remove after deprecation
}
instanced = {key: getattr(self, key) for key in keys if hasattr(self, key)}

default_values_list = args.defaults or []
default_arg_to_default_value = {}
Expand Down
8 changes: 0 additions & 8 deletions rdt/transformers/boolean.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,11 +20,6 @@ class BinaryEncoder(BaseTransformer):
Indicate what to replace the null values with. If the string ``'mode'`` is given,
replace them with the most common value.
Defaults to ``mode``.
model_missing_values (bool):
**DEPRECATED** Whether to create a new column to indicate which values were null or
not. The column will be created only if there are null values. If ``True``, create
the new column if there are null values. If ``False``, do not create the new column
even if there are null values. Defaults to ``False``.
missing_value_generation (str or None):
The way missing values are being handled. There are three strategies:

Expand All @@ -42,14 +37,11 @@ class BinaryEncoder(BaseTransformer):
def __init__(
self,
missing_value_replacement='mode',
model_missing_values=None,
missing_value_generation='random',
):
super().__init__()
self.missing_value_replacement = missing_value_replacement
self._set_missing_value_generation(missing_value_generation)
if model_missing_values is not None:
self._set_model_missing_values(model_missing_values)

def _fit(self, data):
"""Fit the transformer to the data.
Expand Down
Loading