diff --git a/rdt/transformers/id.py b/rdt/transformers/id.py index 1ae3a32a..31722abf 100644 --- a/rdt/transformers/id.py +++ b/rdt/transformers/id.py @@ -2,6 +2,7 @@ import logging import warnings +from itertools import count import numpy as np import pandas as pd @@ -132,6 +133,32 @@ def __getstate__(self): state.pop('generator') return state + def _create_numerical_fallback_generator(self): + """Create fallback generator.""" + regex_generator, _ = strings_from_regex(self.regex_format) + last_template = '' + for _ in range(self.generated): + last_template = next(regex_generator) + + try: + # count is like range, but generates infinite values + for value in count(int(last_template) + 1): + yield str(value) + + except ValueError: + # Generate values like A(0), B(0), A(1), B(1), etc + for counter in count(): + regex_generator, _ = strings_from_regex(self.regex_format) + for value in regex_generator: + yield f'{value}({counter})' + + def _create_fallback_generator(self): + """Create fallback generator. + + NOTE: this is necessary to be overwritten in Enterprise. + """ + return self._create_numerical_fallback_generator() + def __setstate__(self, state): """Set the generator when pickling.""" generator_size = state.get('generator_size') @@ -141,13 +168,18 @@ def __setstate__(self, state): state['generator_size'] = size if generated is None: state['generated'] = 0 - - if generated: + if '_num_fallback_samples_generated' not in state: + state['_num_fallback_samples_generated'] = 0 + if generated and not state['_num_fallback_samples_generated']: for _ in range(generated): next(generator) state['generator'] = generator self.__dict__ = state + if self._num_fallback_samples_generated: + self.generator = self._create_fallback_generator() + for _ in range(self._num_fallback_samples_generated): + next(self.generator) def __init__( self, @@ -180,12 +212,14 @@ def __init__( # Used otherwise self.generator_size = None self.generated = None + self._num_fallback_samples_generated = 0 def reset_randomization(self): """Create a new generator and reset the generated values counter.""" super().reset_randomization() self.generator, self.generator_size = strings_from_regex(self.regex_format) self.generated = 0 + self._num_fallback_samples_generated = 0 if hasattr(self, 'cardinality_rule') and self.cardinality_rule == 'scale': self._remaining_samples['repetitions'] = 0 @@ -247,6 +281,9 @@ def _warn_not_enough_unique_values(self, sample_size, unique_condition, match_ca match_cardinality (bool): Whether or not to match the cardinality of the data. """ + if unique_condition and self._num_fallback_samples_generated: + return + warned = False warn_msg = ( f"The regex for '{self.get_input_column()}' can only generate " @@ -389,16 +426,28 @@ def _sample(self, num_samples, unique_condition): if self.cardinality_rule == 'scale': return self._sample_scale(num_samples) - # If there aren't enough values left in the generator, reset it - if num_samples > self.generator_size - self.generated: + if unique_condition and self._num_fallback_samples_generated: + samples = [next(self.generator) for _ in range(num_samples)] + self._num_fallback_samples_generated += len(samples) + return samples + + # If there aren't enough values left in the generator, reset it if cardinality_rule!=unique + remaining = max(self.generator_size - self.generated, 0) + if num_samples > remaining and not unique_condition: self.reset_randomization() samples = self._sample_from_generator(num_samples) + + # Need more samples than the generator can produce if num_samples > len(samples): if unique_condition: - new_samples = self._sample_fallback(num_samples - len(samples), samples) + fallback_size = num_samples - len(samples) + self.generator = self._create_fallback_generator() + new_samples = [next(self.generator) for _ in range(fallback_size)] + self._num_fallback_samples_generated += len(new_samples) else: new_samples = self._sample_from_template(num_samples - len(samples), samples) + samples.extend(new_samples) return samples diff --git a/tests/integration/transformers/test_id.py b/tests/integration/transformers/test_id.py index b4d34ead..a34780ac 100644 --- a/tests/integration/transformers/test_id.py +++ b/tests/integration/transformers/test_id.py @@ -243,6 +243,27 @@ def test_called_multiple_times_cardinality_rule_unique(self): pd.testing.assert_frame_equal(first_reverse_transform, expected_first_reverse_transform) pd.testing.assert_frame_equal(second_reverse_transform, expected_second_reverse_transform) + def test_called_multiple_times_cardinality_rule_unique_numerical_fallback(self): + """Test calling multiple times when ``cardinality_rule=unique`` with numerical fallback.""" + # Setup + data = pd.DataFrame({'my_column': np.arange(10)}) + generator = RegexGenerator(regex_format=r'\d', cardinality_rule='unique') + + # Run + transformed = generator.fit_transform(data, 'my_column') + first_reverse_transform = generator.reverse_transform(transformed) + second_reverse_transform = generator.reverse_transform(transformed) + + # Assert + expected_first_reverse_transform = pd.DataFrame({ + 'my_column': [str(value) for value in range(10)] + }) + expected_second_reverse_transform = pd.DataFrame({ + 'my_column': [str(value) for value in range(10, 20)] + }) + pd.testing.assert_frame_equal(first_reverse_transform, expected_first_reverse_transform) + pd.testing.assert_frame_equal(second_reverse_transform, expected_second_reverse_transform) + def test_pickled(self, tmpdir): """Test that ensures that ``RegexGenerator`` can be pickled.""" # Setup diff --git a/tests/unit/transformers/test_id.py b/tests/unit/transformers/test_id.py index 0f7ff0e9..156db006 100644 --- a/tests/unit/transformers/test_id.py +++ b/tests/unit/transformers/test_id.py @@ -1,5 +1,6 @@ """Test for ID transformers.""" +import pickle import re import warnings from string import ascii_uppercase @@ -183,6 +184,7 @@ def test___getstate__(self): '_data_cardinality': None, '_data_cardinality_scale': None, '_remaining_samples': {'value': None, 'repetitions': 0}, + '_num_fallback_samples_generated': 0, } @patch('rdt.transformers.id.strings_from_regex') @@ -239,6 +241,21 @@ def test___setstate__(self, mock_strings_from_regex): assert instance.generator_size == 26 mock_strings_from_regex.assert_called_once_with('[A-Za-z]{5}') + def test__sample_fallback_after_pickling(self): + """Test continuing fallback sampling after pickling.""" + # Setup + instance = RegexGenerator('A', cardinality_rule='unique') + instance.reset_randomization() + first_samples = instance._sample(2, unique_condition=True) + + # Run + restored = pickle.loads(pickle.dumps(instance)) + second_samples = restored._sample(2, unique_condition=True) + + # Assert + assert first_samples == ['A', 'A(0)'] + assert second_samples == ['A(1)', 'A(2)'] + def test___init__default(self): """Test the default instantiation of the transformer. @@ -730,6 +747,7 @@ def test__reverse_transform_not_enough_unique_values_cardniality_rule(self, mock # Run out = instance._reverse_transform(columns_data) + instance._reverse_transform(columns_data) # Assert mock_warnings.warn.assert_called_once_with( @@ -762,7 +780,7 @@ def test__reverse_transform_not_enough_unique_values_numerical(self, mock_warnin # Setup instance = RegexGenerator('[1-3]', cardinality_rule='unique') instance.data_length = 6 - generator = AsciiGenerator(5) + generator = iter(['1', '2', '3']) instance.generator = generator instance.generator_size = 3 instance.generated = 0