diff --git a/dpsynth/local_mode/vectorized_transformations.py b/dpsynth/local_mode/vectorized_transformations.py index 5144c081..830d7bc4 100644 --- a/dpsynth/local_mode/vectorized_transformations.py +++ b/dpsynth/local_mode/vectorized_transformations.py @@ -218,8 +218,10 @@ def undiscretize( else: raise ValueError(f'Unsupported interval_handling: {handling}') - if attribute_domain.dtype == 'int' and attribute_domain.clip_to_range: - result = np.ceil(result).astype(int) + if attribute_domain.dtype == 'int': + result = np.ceil(result) + if attribute_domain.clip_to_range or not np.isnan(sentinel): + result = result.astype(int) return result diff --git a/tests/local_mode/vectorized_transformations_test.py b/tests/local_mode/vectorized_transformations_test.py index 6f7fb17b..32d1d55f 100644 --- a/tests/local_mode/vectorized_transformations_test.py +++ b/tests/local_mode/vectorized_transformations_test.py @@ -298,6 +298,27 @@ def test_custom_sentinel_sample(self): ) self.assertEqual(result[0], -1) + @parameterized.parameters(('midpoint', None), ('sample', -1)) + def test_integer_dtype_no_clip(self, interval_handling, sentinel): + rng = np.random.default_rng(0) + attr = domain.NumericalAttribute( + min_value=0, + max_value=3, + dtype='int', + clip_to_range=False, + sentinel=sentinel, + interval_handling=interval_handling, + ) + result = vectorized_transformations.undiscretize( + np.array([0, 1, 2, 3, 4]), np.array([0.0, 1.0, 2.0]), attr, rng=rng + ) + np.testing.assert_array_equal(result[1:], [0, 1, 2, 3]) + if sentinel is None: + self.assertTrue(np.isnan(result[0])) + else: + self.assertTrue(np.issubdtype(result.dtype, np.integer)) + self.assertEqual(result[0], sentinel) + class MergeRareValuesTest(absltest.TestCase):