diff --git a/sdmetrics/single_table/new_row_synthesis.py b/sdmetrics/single_table/new_row_synthesis.py index f6692152..153bdacc 100644 --- a/sdmetrics/single_table/new_row_synthesis.py +++ b/sdmetrics/single_table/new_row_synthesis.py @@ -136,10 +136,11 @@ def compute_breakdown( f'{abs(numerical_match_tolerance * row[field])}' ) elif field in categorical_fields: - if real_data[field].dtype == 'O': - field_filter = f'`{field}` == {repr(row[field])}' + value = row[field] + if isinstance(value, str): + field_filter = f'`{field}` == {repr(value)}' else: - field_filter = f'`{field}` == {row[field]}' + field_filter = f'`{field}` == {value}' row_filter.append(field_filter) diff --git a/tests/unit/single_table/test_new_row_synthesis.py b/tests/unit/single_table/test_new_row_synthesis.py index 6a2b79b6..0c294eb5 100644 --- a/tests/unit/single_table/test_new_row_synthesis.py +++ b/tests/unit/single_table/test_new_row_synthesis.py @@ -80,6 +80,72 @@ def test_compute_breakdown_multi_line(self): # Assert assert score == 0.5 + def test_compute_breakdown_with_category_dtype(self): + """Test ``compute_breakdown`` when a categorical column uses pandas category dtype. + + Expect that string categories are matched instead of raising UndefinedVariableError. + """ + # Setup + real_data = pd.DataFrame({ + 'gender': pd.Series(['F', 'M', 'F'], dtype='category'), + }) + synthetic_data = pd.DataFrame({ + 'gender': ['F', 'X'], + }) + metadata = { + 'tables': { + 'table': { + 'columns': { + 'gender': {'sdtype': 'categorical'}, + }, + } + } + } + metric = NewRowSynthesis() + + # Run + result = metric.compute_breakdown(real_data, synthetic_data, metadata) + + # Assert + assert result == { + 'score': 0.5, + 'num_new_rows': 1, + 'num_matched_rows': 1, + } + + def test_compute_breakdown_with_string_dtype(self): + """Test ``compute_breakdown`` when a categorical column uses pandas string dtype. + + Expect that string values are matched instead of raising UndefinedVariableError. + """ + # Setup + real_data = pd.DataFrame({ + 'gender': pd.Series(['F', 'M', 'F'], dtype='string'), + }) + synthetic_data = pd.DataFrame({ + 'gender': ['F', 'X'], + }) + metadata = { + 'tables': { + 'table': { + 'columns': { + 'gender': {'sdtype': 'categorical'}, + }, + } + } + } + metric = NewRowSynthesis() + + # Run + result = metric.compute_breakdown(real_data, synthetic_data, metadata) + + # Assert + assert result == { + 'score': 0.5, + 'num_new_rows': 1, + 'num_matched_rows': 1, + } + def test_compute_with_sample_size(self): """Test the ``compute`` method with a sample size.