diff --git a/sdv/data_processing/data_processor.py b/sdv/data_processing/data_processor.py index cdf6c5630..4887b36cd 100644 --- a/sdv/data_processing/data_processor.py +++ b/sdv/data_processing/data_processor.py @@ -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 diff --git a/sdv/metadata/_single_table.py b/sdv/metadata/_single_table.py index 6a7e2df9d..3d30beeb1 100644 --- a/sdv/metadata/_single_table.py +++ b/sdv/metadata/_single_table.py @@ -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']), @@ -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 @@ -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': @@ -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 @@ -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) def _detect_columns( self, data, table_name=None, infer_sdtypes=True, infer_keys='primary_only', verbose=False diff --git a/sdv/metadata/metadata.py b/sdv/metadata/metadata.py index da7b3de01..acae4cdad 100644 --- a/sdv/metadata/metadata.py +++ b/sdv/metadata/metadata.py @@ -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: @@ -739,7 +738,6 @@ 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 ) @@ -747,9 +745,8 @@ def _detect_foreign_keys_by_column_name(self, data, verbose=False): 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: diff --git a/sdv/metadata/utils.py b/sdv/metadata/utils.py index be2849305..f93d1933e 100644 --- a/sdv/metadata/utils.py +++ b/sdv/metadata/utils.py @@ -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") diff --git a/tests/integration/metadata/test_metadata.py b/tests/integration/metadata/test_metadata.py index fd02cde9f..6251b4291 100644 --- a/tests/integration/metadata/test_metadata.py +++ b/tests/integration/metadata/test_metadata.py @@ -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 @@ -1747,6 +1747,7 @@ def test_detect_from_dataframes_small_dataset(): 'sdtype': 'categorical', 'range_is_nullable': False, 'range_values': [True, False], + 'high_cardinality': False, }, }, }, @@ -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', @@ -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 diff --git a/tests/unit/metadata/test__single_table.py b/tests/unit/metadata/test__single_table.py index 94af6cc78..01c8f69b0 100644 --- a/tests/unit/metadata/test__single_table.py +++ b/tests/unit/metadata/test__single_table.py @@ -291,24 +291,66 @@ def test__validate_datetime_with_datetime_ranges(self, invalid_kwargs, expected_ instance._validate_datetime('start_date', **invalid_kwargs) @pytest.mark.parametrize('sdtype', ['ordinal', 'categorical']) - def test__validate_categorical(self, sdtype): - """Test the ``_validate_categorical`` method.""" + @pytest.mark.parametrize( + ('kwargs', 'expected_error'), + [ + ( + {'high_cardinality': 'True'}, + "Invalid `high_cardinality` value provided for {sdtype} column 'name'. " + 'The `high_cardinality` must be a boolean value.', + ), + ( + {'high_cardinality': True, 'range_values': ['a', 'b', 'c']}, + 'Invalid combination of `high_cardinality` and `range_values` for {sdtype} ' + "column 'name'. If high_cardinality is set to True, then range_values is not " + 'allowed to be set for the column.', + ), + ( + {'range_values': 'a'}, + "Invalid `range_values` value provided for {sdtype} column 'name'. " + 'The `range_values` must be a list with 1 or more elements.', + ), + ( + {'range_values': []}, + "Invalid `range_values` value provided for {sdtype} column 'name'. " + 'The `range_values` must be a list with 1 or more elements.', + ), + ( + {'range_values': ['a', None]}, + "Invalid `range_values` value provided for {sdtype} column 'name'. " + 'The `range_values` list must not contain null values, use the ' + '`range_is_nullable` parameter instead.', + ), + ], + ) + def test__validate_categorical_and_ordinal(self, sdtype, kwargs, expected_error): + """Test the ``_validate_categorical_and_ordinal`` method.""" # Setup instance = _SingleTableMetadata() + expected_error = re.escape(expected_error.format(sdtype=sdtype)) - # Run / Assert - instance._validate_categorical('name', sdtype=sdtype) - instance._validate_categorical('name', sdtype=sdtype, range_values=['a', 'b', 'c']) + # Run and Assert + with pytest.raises(InvalidMetadataError, match=expected_error): + instance._validate_categorical_and_ordinal('name', sdtype=sdtype, **kwargs) - error_msg_range_values = re.escape( - f"Invalid `range_values` value provided for {sdtype} column 'name'. " - 'The `range_values` must be a list with 1 or more elements.' - ) - with pytest.raises(InvalidMetadataError, match=error_msg_range_values): - instance._validate_categorical('name', sdtype=sdtype, range_values='a') + @pytest.mark.parametrize('sdtype', ['ordinal', 'categorical']) + @pytest.mark.parametrize( + 'kwargs', + [ + {}, + {'range_values': ['a', 'b', 'c']}, + {'high_cardinality': False}, + {'high_cardinality': True}, + {'high_cardinality': False, 'range_values': ['a', 'b', 'c']}, + ], + ) + def test__validate_categorical_and_ordinal_valid(self, sdtype, kwargs): + """Test the ``_validate_categorical_and_ordinal`` method with valid arguments.""" + # Setup + instance = _SingleTableMetadata() - with pytest.raises(InvalidMetadataError, match=error_msg_range_values): - instance._validate_categorical('name', sdtype=sdtype, range_values=[]) + # Run and Assert + instance._validate_categorical_and_ordinal('name', sdtype=sdtype, **kwargs) def test__validate_id(self): """Test the ``_validate_id`` method. @@ -483,7 +525,7 @@ def test__validate_column_numerical(self, mock__validate_numerical, mock__valida mock__validate_numerical.assert_called_once_with('age', range_min=0.0) @patch('sdv.metadata._single_table._SingleTableMetadata._validate_unexpected_kwargs') - @patch('sdv.metadata._single_table._SingleTableMetadata._validate_categorical') + @patch('sdv.metadata._single_table._SingleTableMetadata._validate_categorical_and_ordinal') def test__validate_column_categorical(self, mock__validate_categorical, mock__validate_kwargs): """Test ``_validate_column`` method. @@ -499,10 +541,10 @@ def test__validate_column_categorical(self, mock__validate_categorical, mock__va Mock: - ``_validate_unexpected_kwargs`` - - ``_validate_categorical`` function from ``_SingleTableMetadata``. + - ``_validate_categorical_and_ordinal`` function from ``_SingleTableMetadata``. Side effects: - - ``_validate_categorical`` has been called once. + - ``_validate_categorical_and_ordinal`` has been called once. """ # Setup instance = _SingleTableMetadata() @@ -1315,11 +1357,13 @@ def test__detect_ranges(self, mock_learn_rounding_digits): 'sdtype': 'categorical', 'range_values': ['a', 'b'], 'range_is_nullable': True, + 'high_cardinality': False, } assert instance.columns['ordinal'] == { 'sdtype': 'ordinal', 'range_values': [1, 2], 'range_is_nullable': True, + 'high_cardinality': False, } assert instance.columns['boolean'] == { 'sdtype': 'boolean', @@ -1352,6 +1396,7 @@ def test__detect_ranges_does_not_add_range_values_with_500_unique_values(self): assert instance.columns['categorical'] == { 'sdtype': 'categorical', 'range_is_nullable': False, + 'high_cardinality': True, } def test__detect_columns(self, data): @@ -1668,6 +1713,7 @@ def test_detect_from_dataframe(self, mock_log): 'sdtype': 'categorical', 'range_is_nullable': True, 'range_values': ['cat', 'dog'], + 'high_cardinality': False, }, 'date': { 'sdtype': 'datetime', @@ -1693,6 +1739,7 @@ def test_detect_from_dataframe(self, mock_log): 'sdtype': 'categorical', 'range_is_nullable': True, 'range_values': [True, False], + 'high_cardinality': False, }, } @@ -1930,6 +1977,7 @@ def test_detect_from_csv(self, mock_log, tmp_path): 'sdtype': 'categorical', 'range_is_nullable': True, 'range_values': ['cat', 'dog', 'tiger'], + 'high_cardinality': False, }, 'date': { 'datetime_format': '%Y-%m-%d', @@ -1956,6 +2004,7 @@ def test_detect_from_csv(self, mock_log, tmp_path): 'sdtype': 'categorical', 'range_is_nullable': True, 'range_values': [True, False], + 'high_cardinality': False, }, } @@ -2006,6 +2055,7 @@ def test_detect_from_csv_with_kwargs(self, mock_log, tmp_path): 'sdtype': 'categorical', 'range_is_nullable': True, 'range_values': ['cat', 'dog', 'tiger'], + 'high_cardinality': False, }, 'date': { 'sdtype': 'datetime', @@ -2031,6 +2081,7 @@ def test_detect_from_csv_with_kwargs(self, mock_log, tmp_path): 'sdtype': 'categorical', 'range_is_nullable': True, 'range_values': [True, False], + 'high_cardinality': False, }, } @@ -4397,11 +4448,11 @@ def test__detect_columns_verbose(self, data, capsys): "- Column 'alternate_id': sdtype='id', range_is_nullable=False\n" "- Column 'alternate_id_string': sdtype='id', range_is_nullable=False\n" "- Column 'categorical': sdtype='categorical', range_is_nullable=False, " - "range_values=['a', 'b']\n" + "range_values=['a', 'b'], high_cardinality=False\n" "- Column 'bool': sdtype='categorical', range_is_nullable=False, " - 'range_values=[True, False]\n' + 'range_values=[True, False], high_cardinality=False\n' "- Column 'unknown': sdtype='categorical', range_is_nullable=True, range_values=[" - "'a', 'b', 'c', 1, 2.2, 'd', 'e', 'f']\n" + "'a', 'b', 'c', 1, 2.2, 'd', 'e', 'f'], high_cardinality=False\n" "- Column 'first_name': sdtype='first_name', pii=True, range_is_nullable=False\n" '\nDetecting primary key:\n' "- primary_key='id'\n" @@ -4418,10 +4469,7 @@ def test__detect_columns_verbose_infer_sdtypes_false(self, data, capsys): """Test the ``_detect_columns`` method with verbose (only print PK).""" # Setup instance = _SingleTableMetadata() - expected_output = ( - "\nDetecting primary key:\n- primary_key='id' " - "(updating sdtype to 'id', removing 'pii' field)\n" - ) + expected_output = "\nDetecting primary key:\n- primary_key='id'\n" # Run instance._detect_columns(data, infer_sdtypes=False, verbose=True) @@ -4444,11 +4492,11 @@ def test__detect_columns_verbose_infer_keys_none(self, data, capsys): "- Column 'alternate_id': sdtype='id', range_is_nullable=False\n" "- Column 'alternate_id_string': sdtype='id', range_is_nullable=False\n" "- Column 'categorical': sdtype='categorical', range_is_nullable=False, " - "range_values=['a', 'b']\n" + "range_values=['a', 'b'], high_cardinality=False\n" "- Column 'bool': sdtype='categorical', range_is_nullable=False, " - 'range_values=[True, False]\n' + 'range_values=[True, False], high_cardinality=False\n' "- Column 'unknown': sdtype='categorical', range_is_nullable=True, " - "range_values=['a', 'b', 'c', 1, 2.2, 'd', 'e', 'f']\n" + "range_values=['a', 'b', 'c', 1, 2.2, 'd', 'e', 'f'], high_cardinality=False\n" "- Column 'first_name': sdtype='first_name', pii=True, range_is_nullable=False\n" ) diff --git a/tests/unit/metadata/test_metadata.py b/tests/unit/metadata/test_metadata.py index 171fe9cc8..912300341 100644 --- a/tests/unit/metadata/test_metadata.py +++ b/tests/unit/metadata/test_metadata.py @@ -818,6 +818,7 @@ def test_add_relationship_child_key_is_primary_key(self): 'sdtype': 'categorical', 'range_is_nullable': False, 'range_values': ['a', 'b', 'c'], + 'high_cardinality': False, }, }, }, @@ -836,6 +837,7 @@ def test_add_relationship_child_key_is_primary_key(self): 'sdtype': 'categorical', 'range_is_nullable': False, 'range_values': ['a', 'b', 'c'], + 'high_cardinality': False, }, }, }, @@ -1459,6 +1461,7 @@ def test_validate_child_key_is_primary_key(self): 'sdtype': 'categorical', 'range_is_nullable': False, 'range_values': ['a', 'b', 'c'], + 'high_cardinality': False, }, }, 'primary_key': 'pk', @@ -1477,6 +1480,7 @@ def test_validate_child_key_is_primary_key(self): 'sdtype': 'categorical', 'range_is_nullable': False, 'range_values': ['a', 'b', 'c'], + 'high_cardinality': False, }, }, 'primary_key': 'pk', @@ -4102,7 +4106,7 @@ def test__get_table_info(self, mock_columns_node, mock_summarized_columns_node): @pytest.mark.parametrize( 'initial_fk_sdtype,expected_suffix', [ - ('categorical', " (updating sdtype to 'id')"), + ('categorical', ''), ('id', ''), ], ) diff --git a/tests/unit/metadata/test_utils.py b/tests/unit/metadata/test_utils.py index c351222a5..50ae14f88 100644 --- a/tests/unit/metadata/test_utils.py +++ b/tests/unit/metadata/test_utils.py @@ -83,23 +83,13 @@ def test__format_column_metadata_mixed_value_types(): assert result == "sdtype='datetime', datetime_format=None, pii=True" -@pytest.mark.parametrize( - 'sdtype_updated,pii_removed,expected', - [ - (False, False, "- primary_key='user_id'\n"), - (True, False, "- primary_key='user_id' (updating sdtype to 'id')\n"), - (False, True, "- primary_key='user_id' (removing 'pii' field)\n"), - ( - True, - True, - "- primary_key='user_id' (updating sdtype to 'id', removing 'pii' field)\n", - ), - ], -) -def test__print_primary_key_detection(capsys, sdtype_updated, pii_removed, expected): - """Test ``_print_primary_key_detection`` prints the PK with and without notes.""" +def test__print_primary_key_detection(capsys): + """Test ``_print_primary_key_detection`` prints the PK.""" + # Setup + expected = "- primary_key='user_id'\n" + # Run - _print_primary_key_detection('user_id', sdtype_updated, pii_removed) + _print_primary_key_detection('user_id') # Assert assert capsys.readouterr().out == expected @@ -108,7 +98,7 @@ def test__print_primary_key_detection(capsys, sdtype_updated, pii_removed, expec def test__print_primary_key_detection_no_pk(capsys): """Test ``_print_primary_key_detection`` prints a fallback message when no PK .""" # Run - _print_primary_key_detection(None, False, False) + _print_primary_key_detection(None) # Assert assert capsys.readouterr().out == '- No primary key found\n' diff --git a/tests/utils.py b/tests/utils.py index 8739aebaf..6e56bc2cd 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -19,6 +19,7 @@ 'range_max', 'range_values', 'decimal_places', + 'high_cardinality', }