Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
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
83 changes: 83 additions & 0 deletions dm_env/specs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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

Expand Down
174 changes: 174 additions & 0 deletions dm_env/specs_sampling_test.py
Original file line number Diff line number Diff line change
@@ -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()
27 changes: 27 additions & 0 deletions docs/index.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading