diff --git a/dm_env/specs.py b/dm_env/specs.py index 0dc989a..2577e8a 100644 --- a/dm_env/specs.py +++ b/dm_env/specs.py @@ -32,6 +32,10 @@ _MAXIMUM_INCOMPATIBLE_WITH_SHAPE = '`maximum` is incompatible with `shape`' +def _random_generator(rng): + return np.random.default_rng() if rng is None else rng + + class Array: """Describes a numpy array or scalar shape and dtype. @@ -117,6 +121,35 @@ def generate_value(self): """Generate a test value which conforms to this spec.""" return np.zeros(shape=self.shape, dtype=self.dtype) + def sample(self, rng=None): + """Samples a random value of the specified shape and dtype. + + Args: + rng: Optional `numpy.random.Generator` for reproducible sampling. + By default, a new generator is used. + + Returns: + A NumPy array conforming to this spec. Unbounded floating and complex + arrays use standard normal draws; integers cover their dtype's range. + Unsupported non-numeric dtypes raise TypeError. + """ + rng = _random_generator(rng) + if np.issubdtype(self.dtype, np.bool_): + return rng.integers(0, 2, size=self.shape, dtype=np.uint8).astype( + self.dtype) + if np.issubdtype(self.dtype, np.integer): + info = np.iinfo(self.dtype) + return rng.integers( + info.min, info.max, size=self.shape, dtype=self.dtype, + endpoint=True) + if np.issubdtype(self.dtype, np.floating): + return np.asarray(rng.standard_normal(self.shape), dtype=self.dtype) + if np.issubdtype(self.dtype, np.complexfloating): + sample = (rng.standard_normal(self.shape) + + 1j * rng.standard_normal(self.shape)) + return np.asarray(sample, dtype=self.dtype) + raise TypeError('Cannot sample non-numeric dtype {}'.format(self.dtype)) + def _get_constructor_kwargs(self): """Returns constructor kwargs for instantiating a new copy of this spec.""" # Get the names and kinds of the constructor parameters. @@ -259,6 +292,43 @@ def generate_value(self): return (np.ones(shape=self.shape, dtype=self.dtype) * self.dtype.type(self.minimum)) + def sample(self, rng=None): + """Samples from the inclusive bounds, broadcasting per-element limits. + + Integer and boolean bounds are sampled uniformly, including the maximum. + Finite floating bounds use a uniform distribution. Infinite floating + bounds use a normal or an exponential tail to produce finite samples when + possible. A fixed bound always produces that exact value. + """ + rng = _random_generator(rng) + minimum = np.broadcast_to(self.minimum, self.shape) + maximum = np.broadcast_to(self.maximum, self.shape) + if np.issubdtype(self.dtype, np.bool_): + return rng.integers( + minimum.astype(np.int8), maximum.astype(np.int8), + size=self.shape, dtype=np.int8, endpoint=True).astype(self.dtype) + if np.issubdtype(self.dtype, np.integer): + return rng.integers( + minimum, maximum, size=self.shape, dtype=self.dtype, endpoint=True) + if not np.issubdtype(self.dtype, np.floating): + raise TypeError('Cannot sample bounded dtype {}'.format(self.dtype)) + + lower = minimum.astype(np.float64) + upper = maximum.astype(np.float64) + u = rng.random(self.shape) + normal = rng.standard_normal(self.shape) + exponential = rng.exponential(size=self.shape) + with np.errstate(invalid='ignore', over='ignore'): + result = np.where( + np.isfinite(lower) & np.isfinite(upper), + (1.0 - u) * lower + u * upper, + np.where( + np.isfinite(lower), lower + exponential, + np.where(np.isfinite(upper), upper - exponential, normal))) + result = np.where(lower == upper, lower, result) + result = np.clip(result, lower, upper) + return np.asarray(result, dtype=self.dtype) + def __reduce__(self): return BoundedArray, (self._shape, self._dtype, self._minimum, self._maximum, self._name) @@ -397,6 +467,19 @@ def generate_value(self): empty_string = self.string_type() # pylint: disable=not-callable return np.full(shape=self.shape, dtype=self.dtype, fill_value=empty_string) + def sample(self, rng=None): + """Samples short random ASCII strings, preserving str/bytes element types.""" + rng = _random_generator(rng) + alphabet = 'abcdefghijklmnopqrstuvwxyz' + count = int(np.prod(self.shape)) + values = [] + for _ in range(count): + length = int(rng.integers(1, 9)) + indices = rng.integers(0, len(alphabet), size=length) + value = ''.join(alphabet[index] for index in indices) + values.append(value if self.string_type is str else value.encode('ascii')) + return np.asarray(values, dtype=object).reshape(self.shape) + def __repr__(self): return self._REPR_TEMPLATE.format(self=self) # pytype: disable=duplicate-keyword-argument diff --git a/dm_env/specs_sampling_test.py b/dm_env/specs_sampling_test.py new file mode 100644 index 0000000..d3e10b0 --- /dev/null +++ b/dm_env/specs_sampling_test.py @@ -0,0 +1,174 @@ +# Copyright 2026 The dm_env Authors. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Regression tests for dtype- and bounds-aware random spec sampling.""" + +from absl.testing import absltest +from absl.testing import parameterized +from dm_env import specs +import numpy as np + + +class SamplingTest(parameterized.TestCase): + + @parameterized.named_parameters( + ('int8', np.int8), + ('uint8', np.uint8), + ('int64', np.int64), + ('uint64', np.uint64), + ('float16', np.float16), + ('float32', np.float32), + ('float64', np.float64), + ('complex64', np.complex64), + ('complex128', np.complex128), + ('bool', np.bool_), + ) + def test_array_samples_match_dtype_and_shape(self, dtype): + spec = specs.Array((4, 7), dtype) + result = spec.sample(np.random.default_rng(17)) + self.assertEqual(result.dtype, spec.dtype) + self.assertEqual(result.shape, (4, 7)) + np.testing.assert_array_equal(spec.validate(result), result) + + def test_array_scalar_and_empty_shapes(self): + for shape in ((), (0,), (2, 0, 3)): + with self.subTest(shape=shape): + spec = specs.Array(shape, np.float32) + sample = spec.sample() + self.assertEqual(sample.shape, shape) + spec.validate(sample) + + def test_array_reproducible_with_seeded_generator(self): + spec = specs.Array((128,), np.float32) + first = spec.sample(np.random.default_rng(123)) + second = spec.sample(np.random.default_rng(123)) + np.testing.assert_array_equal(first, second) + self.assertGreater(np.unique(first).size, 50) + + def test_unbounded_integer_sampling_can_use_full_uint64_dtype(self): + spec = specs.Array((20,), np.uint64) + values = spec.sample(np.random.default_rng(23)) + spec.validate(values) + self.assertGreater(values.max(), np.iinfo(np.int64).max) + + def test_unbounded_object_dtype_requires_specialized_spec(self): + spec = specs.Array((2,), object) + with self.assertRaisesRegex(TypeError, 'non-numeric dtype'): + spec.sample(np.random.default_rng(1)) + + def test_bounded_integer_samples_respect_broadcast_bounds(self): + spec = specs.BoundedArray( + shape=(400, 3), + dtype=np.int16, + minimum=[-10, 3, 7], + maximum=[12, 3, 8], + ) + sample = spec.sample(np.random.default_rng(7)) + spec.validate(sample) + self.assertTrue(np.all(sample[:, 1] == 3)) + self.assertEqual(set(np.unique(sample[:, 2])), {7, 8}) + self.assertTrue(np.any(sample[:, 0] == -10)) + self.assertTrue(np.any(sample[:, 0] == 12)) + + def test_bounded_unsigned_integer_including_uint64_maximum(self): + max_value = np.iinfo(np.uint64).max + spec = specs.BoundedArray( + (40,), np.uint64, minimum=max_value - 1, maximum=max_value + ) + sample = spec.sample(np.random.default_rng(44)) + spec.validate(sample) + self.assertEqual(set(sample.tolist()), {max_value - 1, max_value}) + + def test_bounded_float_samples_match_heterogeneous_bounds(self): + spec = specs.BoundedArray( + (500, 3), + np.float32, + minimum=[-0.1, 2.0, -100.0], + maximum=[0.2, 2.0, 0.0], + ) + sample = spec.sample(np.random.default_rng(14)) + spec.validate(sample) + self.assertTrue(np.all(sample[:, 1] == 2.0)) + self.assertTrue(np.all(sample[:, 0] >= -0.1)) + self.assertTrue(np.all(sample[:, 0] <= 0.2)) + self.assertGreater(np.unique(sample[:, 0]).size, 100) + + def test_floating_infinite_limits_produce_finite_values(self): + spec = specs.BoundedArray( + (300, 4), + np.float64, + minimum=[-np.inf, 2.0, -np.inf, -3.0], + maximum=[np.inf, np.inf, -2.0, 3.0], + ) + sample = spec.sample(np.random.default_rng(9)) + spec.validate(sample) + self.assertTrue(np.isfinite(sample).all()) + self.assertTrue(np.all(sample[:, 1] >= 2)) + self.assertTrue(np.all(sample[:, 2] <= -2)) + self.assertTrue(np.all(sample[:, 3] <= 3)) + + def test_fixed_infinite_float_bounds_are_exact(self): + for value in (np.inf, -np.inf): + with self.subTest(value=value): + spec = specs.BoundedArray((), np.float64, value, value) + sample = spec.sample(np.random.default_rng(2)) + spec.validate(sample) + self.assertEqual(sample.item(), value) + + def test_boolean_boundaries_and_discrete_actions(self): + for spec in ( + specs.Array((32,), np.bool_), + specs.BoundedArray((32,), np.bool_, minimum=False, maximum=True), + specs.DiscreteArray(4, dtype=np.int32), + ): + with self.subTest(spec=spec): + sample = spec.sample(np.random.default_rng(24)) + spec.validate(sample) + fixed = specs.BoundedArray((5,), np.bool_, minimum=True, maximum=True) + self.assertTrue(fixed.sample().all()) + + def test_discrete_array_uses_inclusive_action_range(self): + spec = specs.DiscreteArray(5, dtype=np.int8) + values = [ + int(spec.sample(np.random.default_rng(seed))) for seed in range(80) + ] + self.assertEqual(min(values), 0) + self.assertEqual(max(values), 4) + self.assertTrue(all(0 <= value < 5 for value in values)) + + @parameterized.named_parameters(('str', str), ('bytes', bytes)) + def test_string_array_preserves_element_type(self, string_type): + spec = specs.StringArray((3, 4), string_type=string_type) + sample = spec.sample(np.random.default_rng(21)) + spec.validate(sample) + self.assertTrue(all(isinstance(item, string_type) for item in sample.flat)) + self.assertTrue(all(bool(item) for item in sample.flat)) + self.assertGreater(len(set(sample.flat)), 1) + + def test_string_array_scalar_and_empty_shapes(self): + for shape in ((), (0,), (2, 0)): + with self.subTest(shape=shape): + spec = specs.StringArray(shape, bytes) + result = spec.sample(np.random.default_rng(6)) + spec.validate(result) + self.assertEqual(result.shape, shape) + + def test_sampling_does_not_change_existing_generate_value(self): + spec = specs.BoundedArray((2,), np.int32, 3, 7) + np.testing.assert_array_equal(spec.generate_value(), [3, 3]) + _ = spec.sample(np.random.default_rng(5)) + np.testing.assert_array_equal(spec.generate_value(), [3, 3]) + + +if __name__ == '__main__': + absltest.main() diff --git a/docs/index.md b/docs/index.md index 3a12abc..975d012 100644 --- a/docs/index.md +++ b/docs/index.md @@ -165,5 +165,32 @@ bounds and name of the corresponding action or observation array. Note: Actions should almost always specify bounds, e.g. they should use the [`BoundedArray` spec][specs] subclass. +### Sampling values from specs + +Use `spec.sample()` to generate a random value with the declared shape, +dtype and bounds. It is available on `Array`, `BoundedArray`, `DiscreteArray` +and `StringArray` specs. Pass a NumPy random generator for reproducible +sampling: + +```python +import numpy as np +from dm_env import specs + +action_spec = specs.BoundedArray( + shape=(3,), dtype=np.float32, minimum=-1.0, maximum=1.0 +) +rng = np.random.default_rng(seed=42) +action = action_spec.sample(rng) +action_spec.validate(action) +``` + +Integer bounds are inclusive, including for `DiscreteArray` action indices. +Finite floating bounds are sampled uniformly; infinite bounds use finite +normal or exponential samples where possible. Unbounded numeric `Array` +specs are sampled according to their dtype, and string specs produce short +random ASCII strings. Unsupported nonnumeric `Array` dtypes raise +`TypeError`. The existing `generate_value()` method remains deterministic +and is unchanged. + [numpy_array]: https://docs.scipy.org/doc/numpy/reference/generated/numpy.array.html [specs]: ../dm_env/specs.py