Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 32 additions & 5 deletions rdt/transformers/id.py
Original file line number Diff line number Diff line change
Expand Up @@ -136,6 +136,7 @@ def __setstate__(self, state):
"""Set the generator when pickling."""
generator_size = state.get('generator_size')
generated = state.get('generated')
last_generated_value = state.get('_last_generated_value')
generator, size = strings_from_regex(state.get('regex_format'))
if generator_size is None:
state['generator_size'] = size
Expand All @@ -144,7 +145,13 @@ def __setstate__(self, state):

if generated:
for _ in range(generated):
next(generator)
regex_value = next(generator)

if last_generated_value is None:
state['_last_generated_value'] = regex_value

elif '_last_generated_value' not in state:
state['_last_generated_value'] = None

state['generator'] = generator
self.__dict__ = state
Expand Down Expand Up @@ -180,12 +187,14 @@ def __init__(
# Used otherwise
self.generator_size = None
self.generated = None
self._last_generated_value = None

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._last_generated_value = None

if hasattr(self, 'cardinality_rule') and self.cardinality_rule == 'scale':
self._remaining_samples['repetitions'] = 0
Expand Down Expand Up @@ -297,6 +306,9 @@ def _sample_from_generator(self, num_samples):
except (RuntimeError, StopIteration):
pass

if samples:
self._last_generated_value = samples[-1]

return samples

def _sample_from_template(self, num_samples, template_samples):
Expand Down Expand Up @@ -389,18 +401,33 @@ 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 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)
regex_sample_size = min(num_samples, remaining) if unique_condition else num_samples
samples = self._sample_from_generator(regex_sample_size)

# 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)
last_generated_value = getattr(self, '_last_generated_value')
template_samples = samples or [
'' if last_generated_value is None else last_generated_value
]
new_samples = self._sample_fallback(
num_samples - len(samples),
template_samples,
)
else:
new_samples = self._sample_from_template(num_samples - len(samples), samples)

samples.extend(new_samples)

if samples:
self._last_generated_value = samples[-1]

return samples

def _generate_unique_regexes(self):
Expand Down
24 changes: 23 additions & 1 deletion tests/integration/transformers/test_id.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -422,8 +443,9 @@ def test_cardinality_rule_match_empty_regex(self):
reverse_transform_none = instance_none.reverse_transform(transformed_none)

# Assert
expected_unique = pd.DataFrame({'id': ['', '(0)', '(1)', '(2)', '(3)']})
expected = pd.DataFrame({'id': ['', '', '', '', '']})
pd.testing.assert_frame_equal(reverse_transform_unique, expected)
pd.testing.assert_frame_equal(reverse_transform_unique, expected_unique)
pd.testing.assert_frame_equal(reverse_transform_match, expected)
pd.testing.assert_frame_equal(reverse_transform_none, expected)

Expand Down
11 changes: 9 additions & 2 deletions tests/unit/transformers/test_id.py
Original file line number Diff line number Diff line change
Expand Up @@ -183,6 +183,7 @@ def test___getstate__(self):
'_data_cardinality': None,
'_data_cardinality_scale': None,
'_remaining_samples': {'value': None, 'repetitions': 0},
'_last_generated_value': None,
}

@patch('rdt.transformers.id.strings_from_regex')
Expand Down Expand Up @@ -777,7 +778,10 @@ def test__reverse_transform_not_enough_unique_values_numerical(self, mock_warnin
"The regex for 'a' can only generate 3 "
'unique values. Additional values may not exactly follow the provided regex.'
)
np.testing.assert_array_equal(out, np.array(['1', '2', '3', '4', '5', '6']))
np.testing.assert_array_equal(
out,
np.array(['A', 'B', 'C', 'A(0)', 'B(0)', 'C(0)']),
)

@patch('rdt.transformers.id.warnings')
def test__reverse_transform_unique_not_enough_remaining(self, mock_warnings):
Expand All @@ -800,7 +804,10 @@ def test__reverse_transform_unique_not_enough_remaining(self, mock_warnings):
'The regex generator is not able to generate 6 new unique '
'values (only 1 unique values left).'
)
np.testing.assert_array_equal(out, np.array(['A', 'B', 'C', 'D', 'E', 'F']))
np.testing.assert_array_equal(
out,
np.array(['A', 'A(0)', 'A(1)', 'A(2)', 'A(3)', 'A(4)']),
)

@patch('rdt.transformers.id.LOGGER')
def test__reverse_transform_info_message(self, mock_logger):
Expand Down