Skip to content
2 changes: 1 addition & 1 deletion sdv/data_processing/data_processor.py
Original file line number Diff line number Diff line change
Expand Up @@ -403,7 +403,7 @@ def _get_transformer_instance(self, sdtype, column_metadata):
'range_is_nullable',
'pii',
'sdtype',
'range_values',
'high_cardinality',
]
parameters = {
key: value for key, value in column_metadata.items() if key not in non_param_keys
Expand Down
32 changes: 26 additions & 6 deletions sdv/metadata/_single_table.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,8 +60,8 @@ class _SingleTableMetadata:
_SDTYPE_KWARGS = {
'numerical': frozenset(['range_min', 'range_max', 'range_is_nullable', 'decimal_places']),
'datetime': frozenset(['datetime_format', 'range_min', 'range_max', 'range_is_nullable']),
'categorical': frozenset(['range_values', 'range_is_nullable']),
'ordinal': frozenset(['range_values', 'range_is_nullable']),
'categorical': frozenset(['range_values', 'range_is_nullable', 'high_cardinality']),
'ordinal': frozenset(['range_values', 'range_is_nullable', 'high_cardinality']),
'boolean': frozenset(['range_is_nullable']),
'id': frozenset(['regex_format', 'range_is_nullable']),
'unknown': frozenset(['pii', 'range_is_nullable']),
Expand Down Expand Up @@ -226,7 +226,24 @@ def _validate_datetime(column_name, **kwargs):
)

@staticmethod
def _validate_categorical(column_name, sdtype, **kwargs):
def _validate_categorical_and_ordinal(column_name, sdtype, **kwargs):
"""Validate the metadata keys for the categorical and ordinal sdtypes."""
high_cardinality = kwargs.get('high_cardinality')
range_values = kwargs.get('range_values')
if high_cardinality is not None:
if not isinstance(high_cardinality, bool):
raise InvalidMetadataError(
f'Invalid `high_cardinality` value provided for {sdtype} column '
f"'{column_name}'. The `high_cardinality` must be a boolean value."
)

if high_cardinality and range_values is not None:
raise InvalidMetadataError(
f'Invalid combination of `high_cardinality` and `range_values` for {sdtype} '
f"column '{column_name}'. If high_cardinality is set to True, then range_values"
' is not allowed to be set for the column.'
)

range_values = kwargs.get('range_values')
if range_values is not None and (
not isinstance(range_values, list) or len(range_values) == 0
Expand Down Expand Up @@ -326,9 +343,9 @@ def _validate_column_args(self, column_name, sdtype, **kwargs):
self._validate_unexpected_kwargs(column_name, sdtype, **kwargs)
self._validate_null_range(column_name, **kwargs)
if sdtype == 'categorical':
self._validate_categorical(column_name, sdtype='categorical', **kwargs)
self._validate_categorical_and_ordinal(column_name, sdtype='categorical', **kwargs)
if sdtype == 'ordinal':
self._validate_categorical(column_name, sdtype='ordinal', **kwargs)
self._validate_categorical_and_ordinal(column_name, sdtype='ordinal', **kwargs)
elif sdtype == 'numerical':
self._validate_numerical(column_name, **kwargs)
elif sdtype == 'datetime':
Expand Down Expand Up @@ -848,6 +865,9 @@ def _detect_ranges(self, data):
range_values = self._detect_range_values(column_data)
if range_values is not None:
column_metadata['range_values'] = range_values
column_metadata['high_cardinality'] = False
else:
column_metadata['high_cardinality'] = True

def _print_detection(
self, table_name, data, infer_sdtypes, infer_keys, chosen_pk, sdtype_updated, pii_removed
Expand All @@ -862,7 +882,7 @@ def _print_detection(
if infer_keys == 'primary_only':
table_str = f" for table '{table_name}'" if table_name else ''
sys.stdout.write(f'\nDetecting primary key{table_str}:\n')
_print_primary_key_detection(chosen_pk, sdtype_updated, pii_removed)
_print_primary_key_detection(chosen_pk)

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Now that all the detection results are printed together at the end, I don't think we need to mention when the sdtype was updated anymore.

Before, it made sense because the output was printed during detection. Now it can be confusing to see sdtype='id' (updated to 'id') when we're already showing the final state.

Let me know if it makes sense this way.

def _print_detection(


def _detect_columns(
self, data, table_name=None, infer_sdtypes=True, infer_keys='primary_only', verbose=False
Expand Down
5 changes: 1 addition & 4 deletions sdv/metadata/metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -726,7 +726,6 @@ def _detect_foreign_keys_by_column_name(self, data, verbose=False):
continue

try:
sdtype_updated = False
if pk_sdtype == 'id' and original_fk_sdtype != 'id':
update_kwargs = {'sdtype': 'id'}
if 'range_is_nullable' in original_fk_meta:
Expand All @@ -739,17 +738,15 @@ def _detect_foreign_keys_by_column_name(self, data, verbose=False):
column_name=primary_key,
**update_kwargs,
)
sdtype_updated = True
self.add_relationship(
parent_candidate, child_candidate, primary_key, primary_key
)
is_foreign_keys_found = True
if verbose:
child_col = f"'{child_candidate}.{primary_key}'"
parent_col = f"'{parent_candidate}.{primary_key}'"
suffix = " (updating sdtype to 'id')" if sdtype_updated else ''
sys.stdout.write(
f'- Column {child_col} refers to column {parent_col}{suffix}\n'
f'- Column {child_col} refers to column {parent_col}\n'
)

except InvalidMetadataError:
Expand Down
11 changes: 2 additions & 9 deletions sdv/metadata/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,16 +68,9 @@ def _format_column_metadata(sdtype_info):
return ', '.join(parts)


def _print_primary_key_detection(chosen_pk, sdtype_updated, pii_removed):
def _print_primary_key_detection(chosen_pk):
if not chosen_pk:
sys.stdout.write('- No primary key found\n')
return

notes = []
if sdtype_updated:
notes.append("updating sdtype to 'id'")
if pii_removed:
notes.append("removing 'pii' field")

suffix = f' ({", ".join(notes)})' if notes else ''
sys.stdout.write(f"- primary_key='{chosen_pk}'{suffix}\n")
sys.stdout.write(f"- primary_key='{chosen_pk}'\n")
45 changes: 42 additions & 3 deletions tests/integration/metadata/test_metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -1676,17 +1676,17 @@ def test_detect_from_dataframes_verbose_updates_fk_sdtype(capsys):
"\nDetecting table 'users':\n"
"- Column 'account': sdtype='id'\n\n"
"Detecting primary key for table 'users':\n"
"- primary_key='account' (updating sdtype to 'id')\n\n"
"- primary_key='account'\n\n"
"Detecting table 'transactions':\n"
"- Column 'transaction_id': sdtype='id'\n"
"- Column 'account': sdtype='categorical', range_is_nullable=False, "
"range_values=['acct_0', 'acct_1', 'acct_2', 'acct_3', 'acct_4', 'acct_5', "
"'acct_6', 'acct_7', 'acct_8', 'acct_9']\n\n"
"'acct_6', 'acct_7', 'acct_8', 'acct_9'], high_cardinality=False\n\n"
"Detecting primary key for table 'transactions':\n"
"- primary_key='transaction_id'\n\n"
'Detecting foreign keys:\n'
"- Column 'transactions.account' refers to column "
"'users.account' (updating sdtype to 'id')\n"
"'users.account'\n"
)

# Run
Expand Down Expand Up @@ -1747,6 +1747,7 @@ def test_detect_from_dataframes_small_dataset():
'sdtype': 'categorical',
'range_is_nullable': False,
'range_values': [True, False],
'high_cardinality': False,
},
},
},
Expand All @@ -1764,11 +1765,13 @@ def test_detect_from_dataframes_small_dataset():
'sdtype': 'categorical',
'range_is_nullable': False,
'range_values': ['food', 'travel'],
'high_cardinality': False,
},
'rating': {
'sdtype': 'ordinal',
'range_is_nullable': False,
'range_values': [1, 2, 3, 4, 5],
'high_cardinality': False,
},
'amount': {
'sdtype': 'numerical',
Expand Down Expand Up @@ -1801,6 +1804,42 @@ def test_detect_from_dataframes_small_dataset():
assert metadata.to_dict() == expected_metadata


def test_detect_from_dataframes_high_cardinality():
"""Test the detection of high cardinality columns."""
# Setup
data = {
'users': pd.DataFrame({
'user_id': range(3000),
'categorical': [f'user_{i}' for i in range(600)] * 5,
})
}
expected_metadata = {
'tables': {
'users': {
'primary_key': 'user_id',
'columns': {
'user_id': {
'sdtype': 'id',
},
'categorical': {
'sdtype': 'categorical',
'range_is_nullable': False,
'high_cardinality': True,
},
},
},
},
'relationships': [],
'METADATA_SPEC_VERSION': 'V2',
}

# Run
metadata = Metadata.detect_from_dataframes(data)

# Assert
assert metadata.to_dict() == expected_metadata


def test_detect_from_dataframes_verbose_no_pk_found(capsys):
"""Test 'detect_from_dataframes' verbose output when no PK found."""
# Setup
Expand Down
Loading
Loading