From 48156769959b36313b8036a01b4d9d1ae7e7ba35 Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Thu, 17 Sep 2026 15:01:16 +0100 Subject: [PATCH 01/12] def 2990 --- sdv/data_processing/data_processor.py | 9 +++++++-- sdv/metadata/_single_table.py | 29 ++++++++++++++++++++++----- 2 files changed, 31 insertions(+), 7 deletions(-) diff --git a/sdv/data_processing/data_processor.py b/sdv/data_processing/data_processor.py index cdf6c5630..cbf4798d0 100644 --- a/sdv/data_processing/data_processor.py +++ b/sdv/data_processing/data_processor.py @@ -316,8 +316,9 @@ def _get_categorical_transformer(self, parameters, transformer): rdt.transformers.BaseTransformer: A categorical transformer. """ - if 'range_values' in parameters: - parameters.pop('range_values') + for parameter in ['high_cardinality', 'range_values']: + if parameter in parameters: + parameters.pop(parameter) return self._get_transformer_with_parameters(parameters, transformer) @@ -338,6 +339,9 @@ def _get_ordinal_transformer(self, parameters, transformer): order = parameters.pop('range_values') parameters['order'] = order + if 'high_cardinality' in parameters: + parameters.pop('high_cardinality') + return self._get_transformer_with_parameters(parameters, transformer) def _get_numerical_transformer(self, parameters, transformer): @@ -404,6 +408,7 @@ def _get_transformer_instance(self, sdtype, column_metadata): '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..6f91335ef 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,23 @@ def _validate_datetime(column_name, **kwargs): ) @staticmethod - def _validate_categorical(column_name, sdtype, **kwargs): + def _validate_categorical_and_ordinal(column_name, sdtype, **kwargs): + 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 +342,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 +864,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 From 50b554787cf47389fd4f3b842c50fe1d3b99a10a Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Thu, 17 Sep 2026 15:01:33 +0100 Subject: [PATCH 02/12] unit tests --- tests/unit/metadata/test__single_table.py | 95 +++++++++++++++++------ tests/unit/metadata/test_metadata.py | 4 + tests/utils.py | 1 + 3 files changed, 78 insertions(+), 22 deletions(-) diff --git a/tests/unit/metadata/test__single_table.py b/tests/unit/metadata/test__single_table.py index 94af6cc78..c2df6dd27 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 / 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" @@ -4444,11 +4495,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..19f2c8ae4 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', 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', } From 884fbdff7b6edef1d3cee4934f9dc30dba2cd20b Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Thu, 17 Sep 2026 15:11:07 +0100 Subject: [PATCH 03/12] integration tests --- tests/integration/metadata/test_metadata.py | 41 ++++++++++++++++++++- 1 file changed, 40 insertions(+), 1 deletion(-) diff --git a/tests/integration/metadata/test_metadata.py b/tests/integration/metadata/test_metadata.py index fd02cde9f..d19700c31 100644 --- a/tests/integration/metadata/test_metadata.py +++ b/tests/integration/metadata/test_metadata.py @@ -1681,7 +1681,7 @@ def test_detect_from_dataframes_verbose_updates_fk_sdtype(capsys): "- 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' @@ -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 From 460cb2d32422a8d5d1e9a05aefe18098bd8d9f53 Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Thu, 17 Sep 2026 17:08:37 +0100 Subject: [PATCH 04/12] update metadata detection prints --- sdv/metadata/_single_table.py | 2 +- sdv/metadata/metadata.py | 5 +---- sdv/metadata/utils.py | 11 ++-------- tests/integration/metadata/test_metadata.py | 4 ++-- tests/unit/metadata/test__single_table.py | 5 +---- tests/unit/metadata/test_metadata.py | 2 +- tests/unit/metadata/test_utils.py | 24 ++++++--------------- 7 files changed, 15 insertions(+), 38 deletions(-) diff --git a/sdv/metadata/_single_table.py b/sdv/metadata/_single_table.py index 6f91335ef..779cff007 100644 --- a/sdv/metadata/_single_table.py +++ b/sdv/metadata/_single_table.py @@ -881,7 +881,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 d19700c31..6251b4291 100644 --- a/tests/integration/metadata/test_metadata.py +++ b/tests/integration/metadata/test_metadata.py @@ -1676,7 +1676,7 @@ 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, " @@ -1686,7 +1686,7 @@ def test_detect_from_dataframes_verbose_updates_fk_sdtype(capsys): "- 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 diff --git a/tests/unit/metadata/test__single_table.py b/tests/unit/metadata/test__single_table.py index c2df6dd27..d9ad58e57 100644 --- a/tests/unit/metadata/test__single_table.py +++ b/tests/unit/metadata/test__single_table.py @@ -4469,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) diff --git a/tests/unit/metadata/test_metadata.py b/tests/unit/metadata/test_metadata.py index 19f2c8ae4..912300341 100644 --- a/tests/unit/metadata/test_metadata.py +++ b/tests/unit/metadata/test_metadata.py @@ -4106,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' From e3b1b0410ad7e97d5a39a2775c2c585b945e7b3a Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Thu, 17 Sep 2026 17:24:06 +0100 Subject: [PATCH 05/12] dosctring --- sdv/metadata/_single_table.py | 1 + tests/unit/metadata/test__single_table.py | 2 +- 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/sdv/metadata/_single_table.py b/sdv/metadata/_single_table.py index 779cff007..3d30beeb1 100644 --- a/sdv/metadata/_single_table.py +++ b/sdv/metadata/_single_table.py @@ -227,6 +227,7 @@ def _validate_datetime(column_name, **kwargs): @staticmethod 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: diff --git a/tests/unit/metadata/test__single_table.py b/tests/unit/metadata/test__single_table.py index d9ad58e57..01c8f69b0 100644 --- a/tests/unit/metadata/test__single_table.py +++ b/tests/unit/metadata/test__single_table.py @@ -349,7 +349,7 @@ def test__validate_categorical_and_ordinal_valid(self, sdtype, kwargs): # Setup instance = _SingleTableMetadata() - # Run / Assert + # Run and Assert instance._validate_categorical_and_ordinal('name', sdtype=sdtype, **kwargs) def test__validate_id(self): From 5771ce620614d800b3a9f35a26e82f586bb50240 Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Fri, 18 Sep 2026 13:43:01 +0100 Subject: [PATCH 06/12] use lazy import for sdv version in _utils.py --- sdv/_utils.py | 7 ++++++- tests/unit/test__utils.py | 22 +++++++++++----------- 2 files changed, 17 insertions(+), 12 deletions(-) diff --git a/sdv/_utils.py b/sdv/_utils.py index 963e8ed31..14eb183e4 100644 --- a/sdv/_utils.py +++ b/sdv/_utils.py @@ -15,7 +15,6 @@ from pandas.core.tools.datetimes import _guess_datetime_format_for_array from rdt.transformers.utils import _GENERATORS, strings_from_regex -from sdv import version from sdv.errors import InvalidDataTypeError, SDVVersionWarning, SynthesizerInputError, VersionError try: @@ -351,6 +350,8 @@ def check_sdv_versions_and_warn(synthesizer): If the current SDV or SDV Enterprise version does not match the version used to fit the synthesizer. """ + from sdv import version + current_community_version = getattr(version, 'community', None) current_enterprise_version = getattr(version, 'enterprise', None) if getattr(synthesizer, '_fitted', False): @@ -447,6 +448,8 @@ def check_synthesizer_version(synthesizer, is_fit_method=False, compare_operator VersionError: If the current version of the software is lower than the synthesizer's version. """ + from sdv import version + current_community_version = getattr(version, 'community', None) current_enterprise_version = getattr(version, 'enterprise', None) static_message = 'Downgrading your SDV version is not supported.' @@ -519,6 +522,8 @@ def generate_synthesizer_id(synthesizer): ID: A unique identifier for this synthesizer. """ + from sdv import version + class_name = synthesizer.__class__.__name__ synth_version = version.community unique_id = ''.join(str(uuid.uuid4()).split('-')) diff --git a/tests/unit/test__utils.py b/tests/unit/test__utils.py index 706456f0d..4ee0ecd56 100644 --- a/tests/unit/test__utils.py +++ b/tests/unit/test__utils.py @@ -448,7 +448,7 @@ def test_check_sdv_versions_and_warn_community_mismatch(): check_sdv_versions_and_warn(synthesizer) -@patch('sdv._utils.version') +@patch('sdv.version') def test_check_sdv_versions_and_warn_enterprise_mismatch(mock_version): """Test that warnings is raised when enterprise version is mismatched.""" # Setup @@ -470,7 +470,7 @@ def test_check_sdv_versions_and_warn_enterprise_mismatch(mock_version): check_sdv_versions_and_warn(synthesizer) -@patch('sdv._utils.version') +@patch('sdv.version') def test_check_sdv_versions_and_warn_community_and_enterprise_mismatch(mock_version): """Test that warnings is raised when both community and enterprise version mismatch.""" # Setup @@ -531,7 +531,7 @@ def test__compare_versions_lower(): assert result is False -@patch('sdv._utils.version') +@patch('sdv.version') def test_check_synthesizer_version_community_and_enterprise_are_lower(mock_version): """Test that VersionError is raised when both community and enterprise version are higher.""" # Setup @@ -550,7 +550,7 @@ def test_check_synthesizer_version_community_and_enterprise_are_lower(mock_versi check_synthesizer_version(synthesizer) -@patch('sdv._utils.version') +@patch('sdv.version') def test_check_synthesizer_version_community_is_lower(mock_version): """Test that VersionError is raised when only community version is lower.""" # Setup @@ -568,7 +568,7 @@ def test_check_synthesizer_version_community_is_lower(mock_version): check_synthesizer_version(synthesizer) -@patch('sdv._utils.version') +@patch('sdv.version') def test_check_synthesizer_version_enterprise_is_lower(mock_version): """Test that VersionError is raised when only enterprise version is lower.""" # Setup @@ -586,7 +586,7 @@ def test_check_synthesizer_version_enterprise_is_lower(mock_version): check_synthesizer_version(synthesizer) -@patch('sdv._utils.version') +@patch('sdv.version') def test_check_synthesizer_version_enterprise_is_none(mock_version): """Test that no VersionError is raised enterprise is None on the synthesizer.""" # Setup @@ -615,7 +615,7 @@ def test__get_root_tables(): assert result == {'parent'} -@patch('sdv._utils.version') +@patch('sdv.version') def test_check_synthesizer_version_check_synthesizer_is_greater(mock_version): """Test that ``VersionError`` is raised when checking if synthesizer is greater. @@ -638,7 +638,7 @@ def test_check_synthesizer_version_check_synthesizer_is_greater(mock_version): check_synthesizer_version(synthesizer, is_fit_method=True, compare_operator=operator.lt) -@patch('sdv._utils.version') +@patch('sdv.version') def test_check_synthesizer_version_check_synthesizer_is_greater_equal(mock_version): """Test that no ``VersionError`` is raised when versions match.""" # Setup @@ -651,7 +651,7 @@ def test_check_synthesizer_version_check_synthesizer_is_greater_equal(mock_versi check_synthesizer_version(synthesizer, is_fit_method=True, compare_operator=operator.lt) -@patch('sdv._utils.version') +@patch('sdv.version') def test_check_synthesizer_version_check_synthesizer_is_greater_community_mismatch(mock_version): """Test that ``VersionError`` is raised when checking if synthesizer is greater. @@ -674,7 +674,7 @@ def test_check_synthesizer_version_check_synthesizer_is_greater_community_mismat check_synthesizer_version(synthesizer, is_fit_method=True, compare_operator=operator.lt) -@patch('sdv._utils.version') +@patch('sdv.version') def test_check_synthesizer_version_check_synthesizer_is_greater_both_mismatch(mock_version): """Test that ``VersionError`` is raised when community and enterprise are greater. @@ -698,7 +698,7 @@ def test_check_synthesizer_version_check_synthesizer_is_greater_both_mismatch(mo @patch('sdv._utils.uuid') -@patch('sdv._utils.version') +@patch('sdv.version') def test_generate_synthesizer_id(mock_version, mock_uuid): """Test that ``generate_synthesizer_id`` returns the expected id.""" # Setup From c57201dee261ebf4d55592fdb27c93cfab58e595 Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Fri, 18 Sep 2026 14:17:58 +0100 Subject: [PATCH 07/12] fix import for sdv-enterprise --- sdv/__init__.py | 1 + sdv/_utils.py | 7 +------ tests/unit/test__utils.py | 22 +++++++++++----------- 3 files changed, 13 insertions(+), 17 deletions(-) diff --git a/sdv/__init__.py b/sdv/__init__.py index 6ed72040e..593fe5b05 100644 --- a/sdv/__init__.py +++ b/sdv/__init__.py @@ -14,6 +14,7 @@ from importlib.metadata import entry_points from operator import attrgetter from types import ModuleType +from sdv import _utils from sdv import ( data_processing, diff --git a/sdv/_utils.py b/sdv/_utils.py index 14eb183e4..963e8ed31 100644 --- a/sdv/_utils.py +++ b/sdv/_utils.py @@ -15,6 +15,7 @@ from pandas.core.tools.datetimes import _guess_datetime_format_for_array from rdt.transformers.utils import _GENERATORS, strings_from_regex +from sdv import version from sdv.errors import InvalidDataTypeError, SDVVersionWarning, SynthesizerInputError, VersionError try: @@ -350,8 +351,6 @@ def check_sdv_versions_and_warn(synthesizer): If the current SDV or SDV Enterprise version does not match the version used to fit the synthesizer. """ - from sdv import version - current_community_version = getattr(version, 'community', None) current_enterprise_version = getattr(version, 'enterprise', None) if getattr(synthesizer, '_fitted', False): @@ -448,8 +447,6 @@ def check_synthesizer_version(synthesizer, is_fit_method=False, compare_operator VersionError: If the current version of the software is lower than the synthesizer's version. """ - from sdv import version - current_community_version = getattr(version, 'community', None) current_enterprise_version = getattr(version, 'enterprise', None) static_message = 'Downgrading your SDV version is not supported.' @@ -522,8 +519,6 @@ def generate_synthesizer_id(synthesizer): ID: A unique identifier for this synthesizer. """ - from sdv import version - class_name = synthesizer.__class__.__name__ synth_version = version.community unique_id = ''.join(str(uuid.uuid4()).split('-')) diff --git a/tests/unit/test__utils.py b/tests/unit/test__utils.py index 4ee0ecd56..706456f0d 100644 --- a/tests/unit/test__utils.py +++ b/tests/unit/test__utils.py @@ -448,7 +448,7 @@ def test_check_sdv_versions_and_warn_community_mismatch(): check_sdv_versions_and_warn(synthesizer) -@patch('sdv.version') +@patch('sdv._utils.version') def test_check_sdv_versions_and_warn_enterprise_mismatch(mock_version): """Test that warnings is raised when enterprise version is mismatched.""" # Setup @@ -470,7 +470,7 @@ def test_check_sdv_versions_and_warn_enterprise_mismatch(mock_version): check_sdv_versions_and_warn(synthesizer) -@patch('sdv.version') +@patch('sdv._utils.version') def test_check_sdv_versions_and_warn_community_and_enterprise_mismatch(mock_version): """Test that warnings is raised when both community and enterprise version mismatch.""" # Setup @@ -531,7 +531,7 @@ def test__compare_versions_lower(): assert result is False -@patch('sdv.version') +@patch('sdv._utils.version') def test_check_synthesizer_version_community_and_enterprise_are_lower(mock_version): """Test that VersionError is raised when both community and enterprise version are higher.""" # Setup @@ -550,7 +550,7 @@ def test_check_synthesizer_version_community_and_enterprise_are_lower(mock_versi check_synthesizer_version(synthesizer) -@patch('sdv.version') +@patch('sdv._utils.version') def test_check_synthesizer_version_community_is_lower(mock_version): """Test that VersionError is raised when only community version is lower.""" # Setup @@ -568,7 +568,7 @@ def test_check_synthesizer_version_community_is_lower(mock_version): check_synthesizer_version(synthesizer) -@patch('sdv.version') +@patch('sdv._utils.version') def test_check_synthesizer_version_enterprise_is_lower(mock_version): """Test that VersionError is raised when only enterprise version is lower.""" # Setup @@ -586,7 +586,7 @@ def test_check_synthesizer_version_enterprise_is_lower(mock_version): check_synthesizer_version(synthesizer) -@patch('sdv.version') +@patch('sdv._utils.version') def test_check_synthesizer_version_enterprise_is_none(mock_version): """Test that no VersionError is raised enterprise is None on the synthesizer.""" # Setup @@ -615,7 +615,7 @@ def test__get_root_tables(): assert result == {'parent'} -@patch('sdv.version') +@patch('sdv._utils.version') def test_check_synthesizer_version_check_synthesizer_is_greater(mock_version): """Test that ``VersionError`` is raised when checking if synthesizer is greater. @@ -638,7 +638,7 @@ def test_check_synthesizer_version_check_synthesizer_is_greater(mock_version): check_synthesizer_version(synthesizer, is_fit_method=True, compare_operator=operator.lt) -@patch('sdv.version') +@patch('sdv._utils.version') def test_check_synthesizer_version_check_synthesizer_is_greater_equal(mock_version): """Test that no ``VersionError`` is raised when versions match.""" # Setup @@ -651,7 +651,7 @@ def test_check_synthesizer_version_check_synthesizer_is_greater_equal(mock_versi check_synthesizer_version(synthesizer, is_fit_method=True, compare_operator=operator.lt) -@patch('sdv.version') +@patch('sdv._utils.version') def test_check_synthesizer_version_check_synthesizer_is_greater_community_mismatch(mock_version): """Test that ``VersionError`` is raised when checking if synthesizer is greater. @@ -674,7 +674,7 @@ def test_check_synthesizer_version_check_synthesizer_is_greater_community_mismat check_synthesizer_version(synthesizer, is_fit_method=True, compare_operator=operator.lt) -@patch('sdv.version') +@patch('sdv._utils.version') def test_check_synthesizer_version_check_synthesizer_is_greater_both_mismatch(mock_version): """Test that ``VersionError`` is raised when community and enterprise are greater. @@ -698,7 +698,7 @@ def test_check_synthesizer_version_check_synthesizer_is_greater_both_mismatch(mo @patch('sdv._utils.uuid') -@patch('sdv.version') +@patch('sdv._utils.version') def test_generate_synthesizer_id(mock_version, mock_uuid): """Test that ``generate_synthesizer_id`` returns the expected id.""" # Setup From 8657d287e569f432cc1788f8e9a43afa3c7c8584 Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Fri, 18 Sep 2026 15:05:11 +0100 Subject: [PATCH 08/12] Metadata lazy import for sdv.cag._utils --- sdv/cag/_utils.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/sdv/cag/_utils.py b/sdv/cag/_utils.py index c6e27774a..3a08e0762 100644 --- a/sdv/cag/_utils.py +++ b/sdv/cag/_utils.py @@ -12,7 +12,6 @@ from sdv._utils import _cast_to_datetime64, _cast_to_iterable from sdv.cag._errors import ConstraintNotMetError from sdv.errors import RefitWarning, SynthesizerInputError, TableNameError -from sdv.metadata import Metadata PRECISION_LEVELS = { '%Y': 1, # Year @@ -484,6 +483,8 @@ def _remove_columns_from_metadata(metadata, table_name, columns_to_drop): Returns: (sdv.metadata.Metadata): The new Metadata, with the columns removed. """ + from sdv.metadata import Metadata + if isinstance(metadata, Metadata): metadata = metadata.to_dict() column_set = set(columns_to_drop) From 012a93e0c6e84e60f8dafe2d1d10cdea628b2b7d Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Fri, 18 Sep 2026 15:45:48 +0100 Subject: [PATCH 09/12] import single_table module before multi_table --- sdv/__init__.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/sdv/__init__.py b/sdv/__init__.py index 593fe5b05..844c87984 100644 --- a/sdv/__init__.py +++ b/sdv/__init__.py @@ -24,10 +24,10 @@ logging, metadata, metrics, + single_table, multi_table, sampling, sequential, - single_table, version, utils, ) @@ -40,10 +40,10 @@ 'logging', 'metadata', 'metrics', + 'single_table', 'multi_table', 'sampling', 'sequential', - 'single_table', 'version', 'utils', ] From 3b1510568e122bf195cb7c86b3b5c358579fcf6c Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Mon, 21 Sep 2026 17:59:59 +0100 Subject: [PATCH 10/12] address comment --- pyproject.toml | 2 +- sdv/data_processing/data_processor.py | 9 ++------- 2 files changed, 3 insertions(+), 8 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 9002713a5..e85cef8f6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -44,7 +44,7 @@ dependencies = [ "ctgan>=0.12.0;python_version>='3.14'", "deepecho>=0.7.0;python_version<'3.14'", "deepecho>=0.8.0;python_version>='3.14'", - "rdt @ git+https://github.com/sdv-dev/RDT.git@rdt_2.0", + "rdt @ git+https://github.com/sdv-dev/RDT.git@fix-orderuniformencoder-nan-order", "sdmetrics>=0.30.0", 'platformdirs>=4.0', 'pyyaml>=6.0.1', diff --git a/sdv/data_processing/data_processor.py b/sdv/data_processing/data_processor.py index cbf4798d0..4887b36cd 100644 --- a/sdv/data_processing/data_processor.py +++ b/sdv/data_processing/data_processor.py @@ -316,9 +316,8 @@ def _get_categorical_transformer(self, parameters, transformer): rdt.transformers.BaseTransformer: A categorical transformer. """ - for parameter in ['high_cardinality', 'range_values']: - if parameter in parameters: - parameters.pop(parameter) + if 'range_values' in parameters: + parameters.pop('range_values') return self._get_transformer_with_parameters(parameters, transformer) @@ -339,9 +338,6 @@ def _get_ordinal_transformer(self, parameters, transformer): order = parameters.pop('range_values') parameters['order'] = order - if 'high_cardinality' in parameters: - parameters.pop('high_cardinality') - return self._get_transformer_with_parameters(parameters, transformer) def _get_numerical_transformer(self, parameters, transformer): @@ -407,7 +403,6 @@ def _get_transformer_instance(self, sdtype, column_metadata): 'range_is_nullable', 'pii', 'sdtype', - 'range_values', 'high_cardinality', ] parameters = { From 3fc83f69f806506970540d747a9bdd386fa69e55 Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Tue, 22 Sep 2026 10:51:16 +0100 Subject: [PATCH 11/12] undo changes for enterprise imports --- sdv/__init__.py | 5 ++--- sdv/cag/_utils.py | 3 +-- 2 files changed, 3 insertions(+), 5 deletions(-) diff --git a/sdv/__init__.py b/sdv/__init__.py index 844c87984..6ed72040e 100644 --- a/sdv/__init__.py +++ b/sdv/__init__.py @@ -14,7 +14,6 @@ from importlib.metadata import entry_points from operator import attrgetter from types import ModuleType -from sdv import _utils from sdv import ( data_processing, @@ -24,10 +23,10 @@ logging, metadata, metrics, - single_table, multi_table, sampling, sequential, + single_table, version, utils, ) @@ -40,10 +39,10 @@ 'logging', 'metadata', 'metrics', - 'single_table', 'multi_table', 'sampling', 'sequential', + 'single_table', 'version', 'utils', ] diff --git a/sdv/cag/_utils.py b/sdv/cag/_utils.py index 3a08e0762..c6e27774a 100644 --- a/sdv/cag/_utils.py +++ b/sdv/cag/_utils.py @@ -12,6 +12,7 @@ from sdv._utils import _cast_to_datetime64, _cast_to_iterable from sdv.cag._errors import ConstraintNotMetError from sdv.errors import RefitWarning, SynthesizerInputError, TableNameError +from sdv.metadata import Metadata PRECISION_LEVELS = { '%Y': 1, # Year @@ -483,8 +484,6 @@ def _remove_columns_from_metadata(metadata, table_name, columns_to_drop): Returns: (sdv.metadata.Metadata): The new Metadata, with the columns removed. """ - from sdv.metadata import Metadata - if isinstance(metadata, Metadata): metadata = metadata.to_dict() column_set = set(columns_to_drop) From 256f3a35df7b809d65176088af9be2c179e8e978 Mon Sep 17 00:00:00 2001 From: R-Palazzo Date: Wed, 23 Sep 2026 10:00:29 +0100 Subject: [PATCH 12/12] point to rdt_2.0 branch --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index e85cef8f6..9002713a5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -44,7 +44,7 @@ dependencies = [ "ctgan>=0.12.0;python_version>='3.14'", "deepecho>=0.7.0;python_version<'3.14'", "deepecho>=0.8.0;python_version>='3.14'", - "rdt @ git+https://github.com/sdv-dev/RDT.git@fix-orderuniformencoder-nan-order", + "rdt @ git+https://github.com/sdv-dev/RDT.git@rdt_2.0", "sdmetrics>=0.30.0", 'platformdirs>=4.0', 'pyyaml>=6.0.1',