Add protocols for authoring simulation environments. This is for preview only, we discourage users from using this in production code.
PiperOrigin-RevId: 897857927 Change-Id: I82ce0a23465429e7025334e83ba00dfbe5085649
This commit is contained in:
+876
@@ -0,0 +1,876 @@
|
||||
---
|
||||
name: mujoco-python
|
||||
description: >
|
||||
Build, manipulate, and simulate MuJoCo physics models using the Python
|
||||
bindings (MjSpec, MjModel, MjData). Covers choosing between compute
|
||||
backends (C++ for full features and noslip, MJWarp for GPU batch RL). Use when constructing scenes
|
||||
programmatically via the spec API, compiling and stepping simulations,
|
||||
reading sensor/body/geom data, attaching sub-models, composing specs with
|
||||
prefixed names, using contact sensors for fixed-size observation spaces,
|
||||
configuring collision filtering, offscreen rendering (context management,
|
||||
cameras, depth/segmentation), or performing spatial math (quaternion, pose,
|
||||
rotation conversions via mju_). Covers gotchas around compilation
|
||||
lifecycle, named indexing vs bind, geom size semantics, camera conventions,
|
||||
and orientation representations.
|
||||
---
|
||||
|
||||
# MuJoCo Python Bindings
|
||||
|
||||
## Compilation Lifecycle
|
||||
|
||||
```
|
||||
MjSpec ──spec.compile()──▶ MjModel ──MjData(model)──▶ MjData
|
||||
│ │ │
|
||||
│ (mutable blueprint) │ (compiled, mostly frozen) │ (simulation state)
|
||||
│ │ │
|
||||
└── spec.recompile(m, d) ─────┴────────────────────────────┘
|
||||
```
|
||||
|
||||
1. **MjSpec** — mutable data structure you edit to define the simulation.
|
||||
2. **`spec.compile()`** — produces `MjModel` + you create `MjData(model)`.
|
||||
After this, changing the spec has **no effect** until you recompile.
|
||||
3. **Most `MjModel` fields are unsafe to mutate.** Changing them requires
|
||||
`spec.recompile(model, data)`, which returns **new** model and data objects
|
||||
(preserving physics state for existing elements).
|
||||
|
||||
```python
|
||||
import mujoco
|
||||
|
||||
spec = mujoco.MjSpec()
|
||||
body = spec.worldbody.add_body(pos=[0, 0, 1])
|
||||
geom = body.add_geom(type=mujoco.mjtGeom.mjGEOM_SPHERE, size=[0.1])
|
||||
body.add_freejoint()
|
||||
|
||||
model = spec.compile()
|
||||
data = mujoco.MjData(model)
|
||||
mujoco.mj_forward(model, data)
|
||||
|
||||
# Later: add another body, recompile keeping state
|
||||
body2 = spec.worldbody.add_body(pos=[1, 0, 1])
|
||||
body2.add_geom(size=[0.1])
|
||||
body2.add_freejoint()
|
||||
model, data = spec.recompile(model, data) # state preserved
|
||||
```
|
||||
|
||||
> [!CAUTION]
|
||||
> `recompile` returns **new** objects. Always reassign: `model, data = spec.recompile(model, data)`.
|
||||
|
||||
### Loading and Serializing
|
||||
|
||||
```python
|
||||
spec = mujoco.MjSpec() # empty
|
||||
spec = mujoco.MjSpec.from_string(xml_string) # from XML string
|
||||
spec = mujoco.MjSpec.from_file('/path/to.xml') # from file
|
||||
model = mujoco.MjModel.from_xml_string(xml) # direct to model (no spec)
|
||||
|
||||
xml_out = spec.to_xml() # serialize back
|
||||
```
|
||||
|
||||
### Compile Error Debugging
|
||||
|
||||
Use the `.info` field on spec elements for traceability:
|
||||
|
||||
```python
|
||||
geom = spec.worldbody.add_geom()
|
||||
geom.info = 'created at my_file.py:42'
|
||||
spec.compile() # Error: "size 0 must be positive in geom\nElement name '', id 0, created at my_file.py:42"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Compute Backends
|
||||
|
||||
MuJoCo has two compute backends. **Choose early** — the backend determines
|
||||
which features, solvers, and APIs are available.
|
||||
|
||||
| | C++ (default) | MJWarp (NVIDIA GPU) |
|
||||
|---|---|---|
|
||||
| **Import** | `import mujoco` | `import mujoco_warp as mjw` |
|
||||
| **Optimized for** | Latency (single scene) | Throughput (big batches) |
|
||||
| **Hardware** | CPU | NVIDIA GPU |
|
||||
| **Solvers** | All (Newton, CG, PGS, **noslip**) | All except PGS, **noslip**, islands |
|
||||
| **Plugins** | ✅ All | SDF only |
|
||||
| **Precision** | float64 | float32 |
|
||||
| **Named access / bind** | ✅ | Via wrapper libraries or MJX `bind()` |
|
||||
| **Contact sensors** | ✅ | ✅ |
|
||||
| **Sparse Jacobians** | ✅ | ❌ (dense only) |
|
||||
| **Batch rendering** | ❌ | ✅ (BVH ray tracing) |
|
||||
|
||||
### When to use which
|
||||
|
||||
- **C++ (default)**: Real-time control, model predictive control, interactive
|
||||
visualization, any workflow needing full feature support (noslip solver,
|
||||
PGS, islands, plugins, ellipsoidal fluid model, sparse Jacobians). Also the
|
||||
only backend with native `bind()`. Use this unless you need massive
|
||||
parallelism. (For MJWarp, named access is available via wrapper
|
||||
libraries or MJX `bind()`.)
|
||||
|
||||
- **MJWarp**: Reinforcement learning with large batch sizes on NVIDIA GPUs.
|
||||
Scales better for contact-rich scenes and large meshes than the legacy
|
||||
MJX-JAX backend. Not differentiable. May degrade for scenes beyond ~60 DoFs.
|
||||
|
||||
> [!IMPORTANT]
|
||||
> The `noslip` solver (post-constraint velocity correction for exact zero
|
||||
> slip at contacts) is **only available in the C++ backend**. If your task
|
||||
> requires accurate friction modeling without any tangential sliding at
|
||||
> contacts, you must use C++.
|
||||
|
||||
> [!WARNING]
|
||||
> MJWarp uses **float32**, which can cause numerical differences vs C++
|
||||
> (float64). Solver convergence, small friction values, and long rollouts
|
||||
> may be sensitive to this. If you see NaNs or instability on GPU, try
|
||||
> increasing solver iterations or simplifying the model.
|
||||
|
||||
---
|
||||
|
||||
## Building Models with MjSpec
|
||||
|
||||
### Adding elements
|
||||
|
||||
Most `add_*` methods accept keyword arguments matching MJCF XML attributes:
|
||||
|
||||
```python
|
||||
spec = mujoco.MjSpec()
|
||||
|
||||
body = spec.worldbody.add_body(name='arm', pos=[0, 0, 1], quat=[1, 0, 0, 0])
|
||||
geom = body.add_geom(
|
||||
name='arm_geom',
|
||||
type=mujoco.mjtGeom.mjGEOM_CAPSULE,
|
||||
size=[0.05, 0.3],
|
||||
rgba=[1, 0, 0, 1],
|
||||
)
|
||||
joint = body.add_joint(
|
||||
name='hinge1',
|
||||
type=mujoco.mjtJoint.mjJNT_HINGE,
|
||||
axis=[0, 1, 0],
|
||||
range=[-1.57, 1.57],
|
||||
)
|
||||
site = body.add_site(name='sensor_site', pos=[0, 0, 0.3])
|
||||
cam = body.add_camera(name='arm_cam', pos=[0, -2, 0], xyaxes=[1,0,0, 0,0,1])
|
||||
```
|
||||
|
||||
### Orientation alternatives
|
||||
|
||||
In addition to `quat`, you can specify orientation with `euler`, `axisangle`,
|
||||
`xyaxes`, or `zaxis`. Only one can be set at a time:
|
||||
|
||||
```python
|
||||
body.add_geom(euler=[0, 90, 0]) # Euler angles (degrees by default)
|
||||
body.add_geom(axisangle=[0, 1, 0, 1.57]) # axis + angle
|
||||
body.add_geom(zaxis=[0, 1, 0]) # minimal rotation to align Z
|
||||
body.add_geom(xyaxes=[1,0,0, 0,0,1]) # explicit X and Y axes
|
||||
```
|
||||
|
||||
### Top-level elements
|
||||
|
||||
Sensors, actuators, tendons, materials, textures, and meshes are added directly
|
||||
to the spec (not to bodies):
|
||||
|
||||
```python
|
||||
spec.add_material(name='red', rgba=[1, 0, 0, 1])
|
||||
spec.add_actuator(name='motor', joint=joint.name, gear=[1, 0, 0, 0, 0, 0])
|
||||
spec.add_sensor(
|
||||
name='joint_pos',
|
||||
type=mujoco.mjtSensor.mjSENS_JOINTPOS,
|
||||
objtype=mujoco.mjtObj.mjOBJ_JOINT,
|
||||
objname=joint.name,
|
||||
)
|
||||
```
|
||||
|
||||
> [!IMPORTANT]
|
||||
> Always use `element.name` (e.g., `joint.name`, `geom.name`, `site.name`)
|
||||
> instead of hardcoded strings when referencing spec elements. This keeps
|
||||
> references correct if the element is renamed or attached with a prefix.
|
||||
|
||||
### Geom size semantics
|
||||
|
||||
| Type | Size params |
|
||||
| --------- | ----------------------------------------------- |
|
||||
| sphere | `[radius]` |
|
||||
| capsule | `[radius, half_length]` or `[radius]` + fromto |
|
||||
| cylinder | `[radius, half_length]` or `[radius]` + fromto |
|
||||
| box | `[half_x, half_y, half_z]` |
|
||||
| ellipsoid | `[radius_x, radius_y, radius_z]` |
|
||||
| plane | `[half_x, half_y, grid_spacing]` |
|
||||
|
||||
> [!WARNING]
|
||||
> Capsule/cylinder `size` changes meaning with `fromto`. Without `fromto`,
|
||||
> `size=[radius, half_length]`. With `fromto`, `size=[radius]` only — the
|
||||
> length is computed from the two endpoints.
|
||||
|
||||
---
|
||||
|
||||
## Accessing Compiled Data: Named Access vs Bind
|
||||
|
||||
There are **two** recommended ways to read/write compiled model and data fields.
|
||||
**Prefer `bind`** when working with spec elements; use **named access** otherwise.
|
||||
|
||||
### 1. Named Access (on MjModel / MjData)
|
||||
|
||||
```python
|
||||
model.geom('my_geom').size # → numpy view of geom_size for 'my_geom'
|
||||
data.body('torso').xpos # → numpy view of body_xpos
|
||||
data.joint('knee').qpos # → shape depends on joint type
|
||||
data.actuator('motor').ctrl = 1.0 # writable view
|
||||
```
|
||||
|
||||
Aliases: `joint` / `jnt`, `camera` / `cam`, `tendon` / `ten`, `material` / `mat`,
|
||||
`texture` / `tex`, `equality` / `eq`, `keyframe` / `key`.
|
||||
|
||||
> [!WARNING]
|
||||
> Named access returns **views, not copies.** After `mj_step`, old references
|
||||
> reflect new values. Use `.copy()` when logging:
|
||||
> `positions.append(data.body('torso').xpos.copy())`
|
||||
|
||||
### 2. Bind (bridges MjSpec elements → MjModel / MjData)
|
||||
|
||||
`bind()` connects spec elements (or lists of them) to their compiled
|
||||
counterparts. **Use `.set()` to write through bind:**
|
||||
|
||||
```python
|
||||
geom = spec.worldbody.add_geom(name='ball', size=[0.1], type=mujoco.mjtGeom.mjGEOM_SPHERE)
|
||||
joint = body.add_joint(name='j1', type=mujoco.mjtJoint.mjJNT_HINGE)
|
||||
model = spec.compile()
|
||||
data = mujoco.MjData(model)
|
||||
mujoco.mj_forward(model, data)
|
||||
|
||||
# Reading via bind
|
||||
model.bind(geom).size # → array([0.1, 0., 0.])
|
||||
data.bind(geom).xpos # → array([0., 0., 0.])
|
||||
|
||||
# Writing via bind — always use .set()
|
||||
data.bind(joint).set('qpos', 1.5) # sets the joint's qpos
|
||||
|
||||
# Bind a list of spec elements
|
||||
joints = [spec.joint('j1'), spec.joint('j2')]
|
||||
data.bind(joints).qpos # → concatenated array
|
||||
data.bind(joints).set('qpos', np.array([0.5, 1.0])) # write to both
|
||||
```
|
||||
|
||||
> [!CAUTION]
|
||||
> The spec must match the compiled model. If you modify the spec after
|
||||
> `compile()`, you must recompile before calling `bind()`, or you get:
|
||||
> `ValueError: 'The mjSpec does not match mjModel. Please recompile the mjSpec.'`
|
||||
|
||||
---
|
||||
|
||||
## Attachments: Composing Specs
|
||||
|
||||
Attach child specs/bodies to parent specs via frames or sites:
|
||||
|
||||
```python
|
||||
parent = mujoco.MjSpec()
|
||||
child = mujoco.MjSpec()
|
||||
child_body = child.worldbody.add_body(name='arm')
|
||||
child_body.add_geom(name='arm_geom', size=[0.05, 0.3], type=mujoco.mjtGeom.mjGEOM_CAPSULE)
|
||||
child_body.add_joint(name='arm_joint', type=mujoco.mjtJoint.mjJNT_HINGE)
|
||||
|
||||
frame = parent.worldbody.add_frame(pos=[0, 0, 1])
|
||||
frame.attach_body(child_body, prefix='left_')
|
||||
# 'arm' → 'left_arm', 'arm_geom' → 'left_arm_geom', 'arm_joint' → 'left_arm_joint'
|
||||
|
||||
# Or attach entire spec to a site
|
||||
site = parent.worldbody.add_site(name='attach_point', pos=[0, 0, 2])
|
||||
parent.attach(child, site=site, prefix='right_', suffix='_v2')
|
||||
```
|
||||
|
||||
> [!IMPORTANT]
|
||||
> **Cross-spec references require a shared parent.**
|
||||
> If you need to create an element (e.g., an equality constraint) that
|
||||
> references elements from *two different child specs*, you must first
|
||||
> attach both children to the same parent, then add the cross-referencing
|
||||
> element to the **parent** spec using the final prefixed/suffixed names:
|
||||
|
||||
```python
|
||||
# Two robot arms, each defined as a separate spec
|
||||
arm_spec = mujoco.MjSpec()
|
||||
arm_body = arm_spec.worldbody.add_body(name='hand')
|
||||
arm_body.add_geom(name='hand_geom', size=[0.05])
|
||||
wrist_joint = arm_body.add_joint(name='wrist', type=mujoco.mjtJoint.mjJNT_HINGE)
|
||||
|
||||
# Attach both to the parent with different prefixes
|
||||
parent = mujoco.MjSpec()
|
||||
left_prefix, right_prefix = 'left_', 'right_'
|
||||
|
||||
frame_l = parent.worldbody.add_frame(pos=[-0.5, 0, 1])
|
||||
frame_l.attach_body(arm_body, prefix=left_prefix) # left_wrist, left_hand, ...
|
||||
|
||||
frame_r = parent.worldbody.add_frame(pos=[0.5, 0, 1])
|
||||
frame_r.attach_body(arm_body, prefix=right_prefix) # right_wrist, right_hand, ...
|
||||
|
||||
# NOW add a constraint linking both arms — look up the prefixed joints
|
||||
# from the parent spec, don't hardcode the names
|
||||
left_wrist = parent.joint(f'{left_prefix}{wrist_joint.name}')
|
||||
right_wrist = parent.joint(f'{right_prefix}{wrist_joint.name}')
|
||||
parent.add_equality(type=mujoco.mjtEq.mjEQ_JOINT,
|
||||
name1=left_wrist.name, name2=right_wrist.name)
|
||||
model = parent.compile()
|
||||
```
|
||||
|
||||
### Attachment Transforms
|
||||
|
||||
When attaching to a site or frame, the child body's position is transformed
|
||||
relative to the parent's attachment point. Attachment also handles unit
|
||||
conversion (degrees vs radians) between parent and child specs automatically.
|
||||
|
||||
### Assets Get Renamed Too
|
||||
|
||||
Prefix/suffix changes apply to asset filenames:
|
||||
```python
|
||||
child.assets = {'mesh.obj': data}
|
||||
parent.attach(child, prefix='robot_')
|
||||
# Asset key becomes 'robot_mesh.obj' in parent
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Cameras
|
||||
|
||||
### Orientation
|
||||
|
||||
MuJoCo cameras look down the **negative Z axis**. The camera frame is:
|
||||
- **-Z** → forward (viewing direction)
|
||||
- **+X** → right
|
||||
- **+Y** → up
|
||||
|
||||
To point a camera downward (looking at the ground), set its Z axis to `[0, 0, 1]`:
|
||||
|
||||
```python
|
||||
body.add_camera(
|
||||
name='overhead',
|
||||
xyaxes=[1, 0, 0, 0, 1, 0], # x=[1,0,0], y=[0,1,0] → z=[0,0,1] → looks DOWN (-z)
|
||||
pos=[0, 0, 5],
|
||||
)
|
||||
```
|
||||
|
||||
### Geom Group Visibility
|
||||
|
||||
Each camera/viewer has 6 geom groups (0–5). Default visibility:
|
||||
|
||||
| Group | Default Visible | Typical Use |
|
||||
|-------|----------------|-------------|
|
||||
| 0 | ✅ Yes | Standard geoms (default group for new geoms) |
|
||||
| 1 | ✅ Yes | Secondary visual geoms |
|
||||
| 2 | ✅ Yes | Tertiary visual geoms |
|
||||
| 3 | ❌ No | Collision-only or debug geoms |
|
||||
| 4 | ❌ No | Hidden geoms |
|
||||
| 5 | ❌ No | Hidden geoms |
|
||||
|
||||
A newly created geom is in **group 0** by default. Toggle visibility at runtime
|
||||
via `mjvOption.geomgroup[i]`. The same 3-on/3-off default applies to sites,
|
||||
joints, tendons, actuators, flexes, and skins.
|
||||
|
||||
---
|
||||
|
||||
## Contacts: Use Sensors, Not the Contact Array
|
||||
|
||||
### The problem with `data.contact`
|
||||
|
||||
`data.contact` is a **variable-length** array that changes size every timestep
|
||||
depending on what's colliding. Iterating over it directly is fragile and
|
||||
**incompatible with learning-based agents** and fixed-size observation spaces.
|
||||
|
||||
```python
|
||||
# ❌ WRONG — don't iterate data.contact for reward/observation logic
|
||||
for c in data.contact:
|
||||
if c.geom1 == target_geom_id:
|
||||
force = ... # fragile, variable-length, non-deterministic order
|
||||
```
|
||||
|
||||
> [!CAUTION]
|
||||
> Never iterate `data.contact` to build observations or compute rewards.
|
||||
> The array's length and ordering can change between timesteps and even
|
||||
> between MuJoCo versions. Use **contact sensors** instead.
|
||||
|
||||
### Contact sensors: fixed-size, declarative contact queries
|
||||
|
||||
A `<contact>` sensor selects contacts via declarative matching criteria, reduces
|
||||
them to a fixed number of slots, and extracts requested data fields into
|
||||
`data.sensordata` — always the same size, every timestep.
|
||||
|
||||
The pipeline has three stages:
|
||||
1. **Matching** — filter contacts by geom, body, subtree, or site volume
|
||||
2. **Reduction** — keep the top `num` contacts (by order, min distance, max force, or net force)
|
||||
3. **Extraction** — copy requested fields (`found`, `force`, `torque`, `dist`, `pos`, `normal`, `tangent`)
|
||||
|
||||
### Example: detect contact force between a gripper and an object
|
||||
|
||||
```python
|
||||
import mujoco
|
||||
import numpy as np
|
||||
|
||||
spec = mujoco.MjSpec()
|
||||
|
||||
# Build a simple scene: floor + falling object
|
||||
floor = spec.worldbody.add_geom(
|
||||
name='floor', type=mujoco.mjtGeom.mjGEOM_PLANE, size=[1, 1, 0.01]
|
||||
)
|
||||
obj_body = spec.worldbody.add_body(name='obj', pos=[0, 0, 0.5])
|
||||
obj_body.add_freejoint()
|
||||
obj_geom = obj_body.add_geom(
|
||||
name='obj_geom', type=mujoco.mjtGeom.mjGEOM_SPHERE,
|
||||
size=[0.05], mass=0.1,
|
||||
)
|
||||
|
||||
# Add a contact sensor: report force for contacts involving obj_geom
|
||||
contact_sensor = spec.add_sensor(
|
||||
name='obj_contact',
|
||||
type=mujoco.mjtSensor.mjSENS_CONTACT,
|
||||
# Match any contact involving this geom — use .name, not a literal string:
|
||||
objname=obj_geom.name, objtype=mujoco.mjtObj.mjOBJ_GEOM,
|
||||
)
|
||||
|
||||
model = spec.compile()
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
# Step the simulation until the object lands
|
||||
mujoco.mj_step(model, data, nstep=500)
|
||||
mujoco.mj_forward(model, data)
|
||||
|
||||
# Read the contact sensor via bind — always fixed-size in data.sensordata
|
||||
contact_data = data.bind(contact_sensor).sensordata
|
||||
print(f'Contact sensor output: {contact_data}')
|
||||
```
|
||||
|
||||
### XML-based contact sensor (common pattern)
|
||||
|
||||
When loading from XML, contact sensors are even cleaner:
|
||||
|
||||
```xml
|
||||
<sensor>
|
||||
<!-- Is the gripper touching the object? Report force and normal for up to 3 contacts -->
|
||||
<contact name="grip_contact"
|
||||
body1="gripper" body2="object"
|
||||
num="3" data="found force normal"
|
||||
reduce="maxforce"/>
|
||||
|
||||
<!-- Total wrench from all contacts on a body -->
|
||||
<contact name="object_net"
|
||||
body1="object"
|
||||
data="force torque"
|
||||
reduce="netforce"/>
|
||||
</sensor>
|
||||
```
|
||||
|
||||
The output size is deterministic: `num × size(data fields)`. For `"found force
|
||||
normal"` with `num=3`, you get 3 × (1+3+3) = 21 numbers every timestep, padded
|
||||
with zeros if fewer contacts match.
|
||||
|
||||
### Touch sensor: simpler alternative for scalar normal force
|
||||
|
||||
If you only need a scalar "how hard is something pressing on this site", use a
|
||||
`touch` sensor instead:
|
||||
|
||||
```python
|
||||
site = body.add_site(name='fingertip', pos=[0, 0, 0.05], size=[0.02])
|
||||
spec.add_sensor(
|
||||
name='fingertip_touch',
|
||||
type=mujoco.mjtSensor.mjSENS_TOUCH,
|
||||
objname=site.name, objtype=mujoco.mjtObj.mjOBJ_SITE,
|
||||
)
|
||||
```
|
||||
|
||||
The touch sensor sums normal contact forces within the site volume — one scalar
|
||||
output, always present in `sensordata`.
|
||||
|
||||
---
|
||||
|
||||
## Spatial Math Utilities (mju_)
|
||||
|
||||
MuJoCo ships a library of spatial computation functions under the `mju_`
|
||||
namespace — quaternion algebra, rotation conversions, pose composition, and
|
||||
coordinate transforms. **Always check for an existing `mju_` function before
|
||||
implementing spatial math from scratch.** For basic vector arithmetic (add,
|
||||
subtract, dot product, norm), just use NumPy/JAX/Torch directly.
|
||||
|
||||
### Quaternion Operations
|
||||
|
||||
```python
|
||||
res = np.zeros(3)
|
||||
mujoco.mju_rotVecQuat(res, vec, quat) # rotate vector by quaternion
|
||||
|
||||
quat = np.zeros(4)
|
||||
mujoco.mju_mat2Quat(quat, mat3x3) # 3x3 rotation matrix → quaternion
|
||||
mujoco.mju_quat2Mat(mat, quat) # quaternion → 3x3 matrix
|
||||
mujoco.mju_axisAngle2Quat(quat, axis, angle) # axis-angle → quaternion
|
||||
mujoco.mju_euler2Quat(quat, euler, 'xyz') # Euler angles → quaternion
|
||||
mujoco.mju_mulQuat(res, q1, q2) # multiply quaternions
|
||||
mujoco.mju_negQuat(res, quat) # conjugate
|
||||
mujoco.mju_quatZ2Vec(quat, vec) # quat that rotates z-axis to vec
|
||||
mujoco.mju_quatIntegrate(quat, vel, scale) # integrate quat with angular velocity
|
||||
```
|
||||
|
||||
> [!TIP]
|
||||
> `mju_quatZ2Vec` is particularly useful: given a target direction vector, it
|
||||
> returns the quaternion that rotates the Z-axis to point in that direction.
|
||||
|
||||
### Pose Operations
|
||||
|
||||
```python
|
||||
mujoco.mju_mulPose(pos_res, quat_res, pos1, quat1, pos2, quat2) # compose poses
|
||||
mujoco.mju_negPose(pos_res, quat_res, pos, quat) # invert pose
|
||||
mujoco.mju_trnVecPose(res, pos, quat, vec) # transform vector by pose
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Common Gotchas
|
||||
|
||||
### 1. Computed fields are read-only
|
||||
|
||||
`data.xpos`, `data.xmat`, `data.xquat`, `data.geom_xpos` are **output** fields
|
||||
computed by `mj_forward()`. You cannot assign to them directly. Instead, modify
|
||||
input fields (`data.qpos`, `data.qvel`, `data.ctrl`) and call `mj_forward()` or
|
||||
`mj_step()`.
|
||||
|
||||
### 2. Duplicate names are forbidden
|
||||
|
||||
```python
|
||||
spec.add_material(name='yellow')
|
||||
spec.add_material(name='yellow') # ValueError: "repeated name 'yellow' in material"
|
||||
```
|
||||
|
||||
Names must be unique within each element type.
|
||||
|
||||
### 3. Orientation keywords are mutually exclusive
|
||||
|
||||
```python
|
||||
body.add_geom(axisangle=[1, 0, 0, 1.57], euler=[0, 0, 0])
|
||||
# ValueError: 'Only one of: axisangle, xyaxes, zaxis, or euler can be set.'
|
||||
```
|
||||
|
||||
Pick one orientation representation. Quaternion (`quat`) is the native format.
|
||||
|
||||
### 4. `size` must be positive for geoms
|
||||
|
||||
A geom with `size[0] == 0` will fail compilation. Always set at least
|
||||
`size=[radius]` for spheres/capsules, or `size=[hx, hy, hz]` for boxes.
|
||||
|
||||
### 5. `mj_step` with `nstep` repeats the same control
|
||||
|
||||
```python
|
||||
mujoco.mj_step(model, data, nstep=100) # 100 steps, same ctrl each step
|
||||
```
|
||||
|
||||
This is much faster than a Python loop and is fine for passive simulation or
|
||||
constant-control scenarios. But if you need to update `data.ctrl` between steps,
|
||||
you must step one at a time.
|
||||
|
||||
### 6. Euler sequence matters
|
||||
|
||||
`mju_euler2Quat` takes a 3-character sequence string. Lowercase = intrinsic
|
||||
rotations, uppercase = extrinsic:
|
||||
|
||||
```python
|
||||
mujoco.mju_euler2Quat(quat, [roll, pitch, yaw], 'xyz') # intrinsic x-y-z
|
||||
mujoco.mju_euler2Quat(quat, [roll, pitch, yaw], 'XYZ') # extrinsic X-Y-Z
|
||||
```
|
||||
|
||||
The sequence must be exactly 3 characters from `xyzXYZ`.
|
||||
|
||||
### 7. `copy()` vs view semantics
|
||||
|
||||
NumPy arrays from MjModel/MjData are **views** into C memory. `mj_step` changes
|
||||
them in-place. Always `.copy()` when storing values for later comparison.
|
||||
|
||||
### 8. Default class handling
|
||||
|
||||
```python
|
||||
main = spec.default # global default class (always named 'main')
|
||||
child_class = spec.add_default('high_friction', main)
|
||||
child_class.geom.friction = [1.5, 0.005, 0.0001]
|
||||
|
||||
geom = body.add_geom(child_class) # use specific default class
|
||||
geom = body.add_geom() # uses 'main' class implicitly
|
||||
```
|
||||
|
||||
### 9. Gravity is -Z by default
|
||||
|
||||
MuJoCo convention: **+Z is up**, gravity is `[0, 0, -9.81]`. The viewer and
|
||||
all built-in models assume this. Don't fight it — orient your scene accordingly.
|
||||
|
||||
### 10. Capsule/cylinder size with and without fromto
|
||||
|
||||
```python
|
||||
# With explicit pos/quat: size = [radius, half_length]
|
||||
body.add_geom(type=mujoco.mjtGeom.mjGEOM_CAPSULE, size=[0.05, 0.3])
|
||||
|
||||
# With fromto: size = [radius] only — length is inferred from endpoints
|
||||
body.add_geom(
|
||||
type=mujoco.mjtGeom.mjGEOM_CAPSULE,
|
||||
size=[0.05],
|
||||
fromto=[0, 0, 0, 0, 0, 0.6],
|
||||
)
|
||||
```
|
||||
|
||||
### 11. Collision filtering with contype/conaffinity
|
||||
|
||||
Two geoms collide only if `(g1.contype & g2.conaffinity) || (g2.contype & g1.conaffinity)`.
|
||||
By default both are `1`, so everything collides with everything.
|
||||
|
||||
```python
|
||||
# Visual-only geom: set contype=0, conaffinity=0 to disable collisions
|
||||
body.add_geom(size=[0.1], contype=0, conaffinity=0, group=1)
|
||||
|
||||
# Separate collision groups using bitmasks:
|
||||
robot_geom = body.add_geom(size=[0.05], contype=1, conaffinity=2)
|
||||
tool_geom = body.add_geom(size=[0.03], contype=2, conaffinity=1)
|
||||
# Robot and tool collide (1&1=0, but 2&2=0… wait):
|
||||
# contype=1 & conaffinity=1 → collide; contype=2 & conaffinity=2 → collide
|
||||
```
|
||||
|
||||
> [!TIP]
|
||||
> **`condim` and `friction` interact.** Each geom has `friction=[tangential, torsional, rolling]`
|
||||
> (default `[1, 0.005, 0.0001]`). The `condim` value controls which friction coefficients are
|
||||
> *active* in a contact:
|
||||
>
|
||||
> | condim | Active friction | Geom `friction` indices used |
|
||||
> |--------|----------------|------------------------------|
|
||||
> | 1 | None (frictionless, normal force only) | — |
|
||||
> | 3 | Tangential (opposes sliding) | `friction[0]` |
|
||||
> | 4 | Tangential + torsional (opposes sliding and twisting around contact normal) | `friction[0:2]` |
|
||||
> | 6 | Tangential + torsional + rolling (also opposes rolling around tangent axes) | `friction[0:3]` |
|
||||
>
|
||||
> Torsional friction models a surface contact patch resisting twist — useful for soft fingers.
|
||||
> Rolling friction dissipates energy from local deformations — useful for stopping balls from rolling
|
||||
> forever. Both torsional and rolling coefficients have **units of length** (roughly the contact
|
||||
> patch diameter or deformation depth).
|
||||
>
|
||||
> ```python
|
||||
> # A soft finger pad: enable torsional friction for stable grasping
|
||||
> finger_geom = body.add_geom(
|
||||
> type=mujoco.mjtGeom.mjGEOM_CAPSULE,
|
||||
> size=[0.01, 0.02],
|
||||
> condim=4,
|
||||
> friction=[1.0, 0.01, 0.0001], # tangential=1.0, torsional=0.01
|
||||
> )
|
||||
>
|
||||
> # A ball that should stop rolling on a surface
|
||||
> ball_geom = body.add_geom(
|
||||
> type=mujoco.mjtGeom.mjGEOM_SPHERE,
|
||||
> size=[0.05],
|
||||
> condim=6,
|
||||
> friction=[0.8, 0.005, 0.002], # tangential=0.8, torsional=0.005, rolling=0.002
|
||||
> )
|
||||
> ```
|
||||
|
||||
---
|
||||
|
||||
## Offscreen Rendering
|
||||
|
||||
Offscreen rendering produces images (RGB, depth, segmentation) without a
|
||||
display. It requires an OpenGL context — MuJoCo auto-detects the best
|
||||
available backend (EGL on headless Linux, GLFW on desktop, OSMesa as
|
||||
fallback).
|
||||
|
||||
### The `Renderer` class
|
||||
|
||||
`mujoco.Renderer` wraps GL context creation, scene management, and buffer
|
||||
readback. **Always use it as a context manager** to ensure GPU resources are
|
||||
freed:
|
||||
|
||||
```python
|
||||
import mujoco
|
||||
import numpy as np
|
||||
|
||||
# Define a camera in the spec and keep a reference
|
||||
overhead_cam = spec.worldbody.add_camera(
|
||||
name='overhead',
|
||||
pos=[0, 0, 3],
|
||||
quat=[0.707, 0.707, 0, 0], # looking down
|
||||
fovy=60,
|
||||
)
|
||||
|
||||
model = spec.compile()
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
# Create renderer — width/height must not exceed offscreen buffer (see below)
|
||||
with mujoco.Renderer(model, height=480, width=640) as renderer:
|
||||
mujoco.mj_forward(model, data)
|
||||
|
||||
# Use the spec element's .name — never a literal string
|
||||
renderer.update_scene(data, camera=overhead_cam.name)
|
||||
rgb = renderer.render() # → np.ndarray (H, W, 3), dtype=uint8
|
||||
|
||||
# Depth rendering
|
||||
renderer.enable_depth_rendering()
|
||||
renderer.update_scene(data, camera=overhead_cam.name)
|
||||
depth = renderer.render() # → np.ndarray (H, W), dtype=float32 (meters)
|
||||
renderer.disable_depth_rendering()
|
||||
|
||||
# Segmentation rendering
|
||||
renderer.enable_segmentation_rendering()
|
||||
renderer.update_scene(data, camera=overhead_cam.name)
|
||||
seg = renderer.render() # → np.ndarray (H, W, 2), dtype=int32
|
||||
# seg[:,:,0] = object ID, seg[:,:,1] = object type; background = (-1, -1)
|
||||
renderer.disable_segmentation_rendering()
|
||||
```
|
||||
|
||||
> [!WARNING]
|
||||
> Forgetting to close the renderer (or not using `with`) leaks GPU memory and
|
||||
> GL contexts. In loops, create the renderer **once** outside the loop.
|
||||
|
||||
### What the `Renderer` holds internally
|
||||
|
||||
When you create `mujoco.Renderer(model, height, width)`, it allocates three
|
||||
internal objects that must be freed together:
|
||||
|
||||
1. **`GLContext`** — an offscreen OpenGL context (EGL, GLFW, or OSMesa,
|
||||
auto-detected). Created with the requested `width × height`.
|
||||
2. **`MjrContext`** — MuJoCo's GPU rendering resources (shaders, textures,
|
||||
framebuffers), bound to the GLContext. Set to the offscreen framebuffer.
|
||||
3. **`MjvScene`** — geometry buffer holding the scene snapshot passed to the
|
||||
GPU each frame.
|
||||
|
||||
The context manager (`with Renderer(...) as r:`) calls `r.close()` on exit,
|
||||
which frees the MjrContext first and then the GLContext — **order matters**.
|
||||
If you use the renderer without `with`, call `renderer.close()` manually.
|
||||
|
||||
> [!CAUTION]
|
||||
> Internally, `MjrContext.free()` must be called **before** `GLContext.free()`.
|
||||
> Reversing the order leaks GPU resources or segfaults. The `Renderer` class
|
||||
> handles this automatically — prefer it over manual context management.
|
||||
|
||||
### Offscreen framebuffer size
|
||||
|
||||
The renderer cannot exceed the offscreen buffer dimensions. The defaults are
|
||||
640×480. Set larger buffers **before compilation** via `spec.visual`:
|
||||
|
||||
```python
|
||||
spec.visual.global_.offwidth = 1920
|
||||
spec.visual.global_.offheight = 1080
|
||||
model = spec.compile()
|
||||
|
||||
# Now you can render up to 1920×1080
|
||||
with mujoco.Renderer(model, height=1080, width=1920) as renderer:
|
||||
...
|
||||
```
|
||||
|
||||
> [!IMPORTANT]
|
||||
> Increasing offscreen buffer size consumes GPU memory. For batch rendering
|
||||
> of many cameras, keep the per-frame resolution modest.
|
||||
|
||||
### Cameras
|
||||
|
||||
MuJoCo has two camera systems: **fixed cameras** defined in the model, and
|
||||
the **free camera** for interactive viewing.
|
||||
|
||||
#### Defining cameras in MjSpec
|
||||
|
||||
Always store the return value of `add_camera` and use its `.name` or `.id`
|
||||
to reference the camera later — never hardcode literal strings:
|
||||
|
||||
```python
|
||||
# Fixed camera on worldbody — good for evaluation/recording
|
||||
overhead_cam = spec.worldbody.add_camera(
|
||||
name='overhead',
|
||||
pos=[0, 0, 3],
|
||||
quat=[0.707, 0.707, 0, 0],
|
||||
fovy=60,
|
||||
)
|
||||
|
||||
# Camera attached to a body — moves with the body
|
||||
wrist_cam = wrist_body.add_camera(
|
||||
name='wrist_cam',
|
||||
pos=[0.05, 0, 0],
|
||||
xyaxes=[0, -1, 0, 0, 0, -1],
|
||||
fovy=90,
|
||||
)
|
||||
```
|
||||
|
||||
#### Selecting a camera for rendering
|
||||
|
||||
`update_scene` accepts a camera **name** (str), **id** (int), or an
|
||||
`MjvCamera` object. Always derive from the spec element:
|
||||
|
||||
```python
|
||||
# By name via spec element (recommended — survives recompilation)
|
||||
renderer.update_scene(data, camera=overhead_cam.name)
|
||||
|
||||
# By id via spec element (after compile; matches model.cam_* arrays)
|
||||
renderer.update_scene(data, camera=overhead_cam.id)
|
||||
|
||||
# Free camera (default) — no camera argument needed
|
||||
renderer.update_scene(data)
|
||||
|
||||
# Custom free camera with explicit lookat/distance/angles
|
||||
cam = mujoco.MjvCamera()
|
||||
cam.type = mujoco.mjtCamera.mjCAMERA_FREE
|
||||
cam.lookat[:] = [0, 0, 0.5]
|
||||
cam.distance = 3.0
|
||||
cam.azimuth = 135
|
||||
cam.elevation = -25
|
||||
renderer.update_scene(data, camera=cam)
|
||||
```
|
||||
|
||||
#### Camera properties reference
|
||||
|
||||
| Property | Type | Description |
|
||||
|----------|------|-------------|
|
||||
| `pos` | `real(3)` | Position in parent body frame |
|
||||
| `quat` | `real(4)` | Orientation quaternion (w, x, y, z) |
|
||||
| `xyaxes` | `real(6)` | Alternative orientation: `[x_axis(3), y_axis(3)]` |
|
||||
| `fovy` | `real` | Vertical field of view (degrees, default 45) |
|
||||
| `resolution` | `int(2)` | Sensor resolution — only for camera-based sensors |
|
||||
| `targetbody` | `str` | Track this body (camera always looks at it) |
|
||||
| `mode` | `str` | `"fixed"`, `"track"`, `"trackcom"`, `"targetbody"`, `"targetbodycom"` |
|
||||
|
||||
### Scene options
|
||||
|
||||
Control what is visualized via `MjvOption`:
|
||||
|
||||
```python
|
||||
scene_option = mujoco.MjvOption()
|
||||
# geomgroup is a bool array indexed by group number (0–5).
|
||||
# Each geom's `group` attribute (default 0) assigns it to a group.
|
||||
# Toggle visibility of each group:
|
||||
scene_option.geomgroup[:] = False # hide all groups
|
||||
scene_option.geomgroup[0] = True # show group 0 (e.g. ground plane)
|
||||
scene_option.geomgroup[3] = True # show group 3 (e.g. visualization geoms)
|
||||
|
||||
# Toggle rendering flags
|
||||
scene_option.flags[mujoco.mjtVisFlag.mjVIS_CONTACTFORCE] = True
|
||||
scene_option.flags[mujoco.mjtVisFlag.mjVIS_JOINT] = True
|
||||
|
||||
renderer.update_scene(data, camera=overhead_cam.name, scene_option=scene_option)
|
||||
```
|
||||
|
||||
### Filament backend (experimental)
|
||||
|
||||
MuJoCo's default renderer uses OpenGL. An alternative **Filament** backend
|
||||
(Vulkan-based) is available experimentally and provides higher-quality
|
||||
rendering. Filament does **not** require vertical flip (`np.flipud` is a
|
||||
no-op). It is selected via build flags — see the MuJoCo Filament
|
||||
[source](../src/experimental/filament) for details.
|
||||
|
||||
---
|
||||
|
||||
## Key References
|
||||
|
||||
### Documentation
|
||||
|
||||
| Document | Description |
|
||||
|----------|-------------|
|
||||
| [XMLreference.rst](XMLreference.rst) | Complete MJCF XML element and attribute reference |
|
||||
| [python.rst](python.rst) | Python bindings API: named access, bind, enums, callbacks |
|
||||
| [modeling.rst](modeling.rst) | MJCF modeling guide: coordinate frames, defaults, attachments |
|
||||
| [simulation.rst](programming/simulation.rst) | Simulation loop, state, forward/inverse dynamics |
|
||||
| [modeledit.rst](programming/modeledit.rst) | Procedural model editing with MjSpec |
|
||||
| [visualization.rst](programming/visualization.rst) | Rendering, cameras, scene management |
|
||||
| [APIfunctions.rst](APIreference/APIfunctions.rst) | C API function reference (mj_, mju_, mjv_, mjr_) |
|
||||
| [APItypes.rst](APIreference/APItypes.rst) | All MuJoCo structs and enums |
|
||||
|
||||
### Test Files (Executable Examples)
|
||||
|
||||
| Test file | Key patterns demonstrated |
|
||||
|-----------|--------------------------|
|
||||
| [specs_test.py](../../py/mujoco/specs_test.py) | MjSpec API: compile, recompile, attach, bind, defaults, delete, actuator shortcuts |
|
||||
| [bindings_test.py](../../py/mujoco/bindings_test.py) | Named indexing, mju_ functions, copy/pickle, contacts, mj_step |
|
||||
| [support_test.py](../../py/mujoco/mjx/_src/support_test.py) | MJX bind `.set()` pattern, JAX functional updates |
|
||||
|
||||
### Source Code
|
||||
|
||||
| File | Description |
|
||||
|------|-------------|
|
||||
| [mujoco.h](../include/mujoco.h) | Main C API header with all mju_ function signatures |
|
||||
| [XMLschema.rst](XMLschema.rst) | Schema-level XML structure documentation |
|
||||
@@ -0,0 +1,43 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Adapts environment action into suitable commands format accepted by REAF."""
|
||||
|
||||
import abc
|
||||
from collections.abc import Mapping
|
||||
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
|
||||
|
||||
class ActionSpaceAdapter(abc.ABC):
|
||||
"""Adapts environment action into suitable commands format accepted by REAF.
|
||||
|
||||
Implementations of this interface are responsible for converting the more
|
||||
generic action accepted by the environment (e.g. a flat numpy array) into the
|
||||
more constraining format accepted as commands by REAF, i.e. a dictionary of
|
||||
string to tensors.
|
||||
"""
|
||||
|
||||
@abc.abstractmethod
|
||||
def commands_from_environment_action(
|
||||
self, environment_action: gdmr_types.ActionType
|
||||
) -> Mapping[str, gdmr_types.ArrayType]:
|
||||
"""Converts the environment action into commands accepted by REAF."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def action_spec(self) -> gdmr_types.ActionSpec:
|
||||
"""Returns the action spec exposed by the environment."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def task_commands_keys(self) -> set[str]:
|
||||
"""Returns the keys for the commands exposed to the task layer."""
|
||||
@@ -0,0 +1,91 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Abstract class for commands manipulation in the task logic layer."""
|
||||
|
||||
import abc
|
||||
from collections.abc import Mapping
|
||||
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
|
||||
|
||||
class CommandsProcessor(abc.ABC):
|
||||
"""Perform commands manipulation.
|
||||
|
||||
The following describes the processing pipeline starting from the top (closer
|
||||
to the policy) to the bottom (interfacing with the DACL commands spec).
|
||||
|
||||
Assume that we have two processing units:
|
||||
Processor 1) has a consumed_commands_spec for two keys: "p1/c1" and "p1/c2".
|
||||
Its produced_commands_keys are "p2/c1".
|
||||
Processor 2) has a consumed_commands_spec for "p2/c1". Its
|
||||
produced_commands_keys are "p3/c1" and "p3/c2".
|
||||
|
||||
Specs are propagated starting from the bottom:
|
||||
1) In this example assume that the DACL exposes "p3/c1", "p3/c2" and "p3/c3".
|
||||
2) Processor 2) returns ("p3/c1", "p3/c2") from input "p2/c1". This means that
|
||||
the global commands spec exposed at this level is "p2/c1" and the
|
||||
unprocessed "p3/c3".
|
||||
3) Processor 1) returns "p2/c1" from input ("p1/c1", "p1/c2"). By applying the
|
||||
same transformation rule, we can obtain the final commands spec exposed by
|
||||
the full processing pipeline: "p1/c1", "p1/c2" and "p3/c3".
|
||||
|
||||
"p1/c1" "p1/c2" "p3/c3"
|
||||
| | |
|
||||
----------------- |
|
||||
| P1 | |
|
||||
----------------- |
|
||||
| "p2/c1" |
|
||||
----------------- |
|
||||
| P2 | |
|
||||
----------------- |
|
||||
| "p3/c1" | "p3/c2" |
|
||||
| | |
|
||||
------------------------------------
|
||||
| DACL |
|
||||
------------------------------------
|
||||
"""
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Returns a unique string identifier for this object."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def process_commands(
|
||||
self, consumed_commands: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> Mapping[str, gdmr_types.ArrayType]:
|
||||
"""Processes the commands and returns a new modified version of it.
|
||||
|
||||
Args:
|
||||
consumed_commands: the commands up in the processing chain (or provided by
|
||||
the Environment) that are required by this processor, i.e. with keys
|
||||
specified by `consumed_commands_spec`.
|
||||
|
||||
Returns the new commands. Note that the data in consumed_commands is removed
|
||||
from the global commands dictionary. If users want to keep some of the
|
||||
elements it is their responsibility to retain them in the output
|
||||
dictionary.
|
||||
"""
|
||||
|
||||
@abc.abstractmethod
|
||||
def consumed_commands_spec(self) -> Mapping[str, gdmr_types.AnyArraySpec]:
|
||||
"""Spec of the commands consumed by this processor."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def produced_commands_keys(self) -> set[str]:
|
||||
"""Keys of the commands produced by this processor."""
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Resets the internal state of the command processor."""
|
||||
...
|
||||
@@ -0,0 +1,170 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""REAF data acquisition and control layer to interface with the robotic setup."""
|
||||
|
||||
from collections.abc import Iterable, Mapping
|
||||
|
||||
from dm_env import specs
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
from reaf.core import device as reaf_device
|
||||
from reaf.core import device_coordinator as reaf_coordinator
|
||||
from reaf.core import trigger
|
||||
|
||||
|
||||
class DataAcquisitionAndControlLayer:
|
||||
"""REAF data acquisition and control layer.
|
||||
|
||||
The DACL is responsible to provide an interface for the robotic setup.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
device_coordinator: reaf_coordinator.DeviceCoordinator,
|
||||
commands_trigger: trigger.Trigger | None,
|
||||
measurements_trigger: trigger.Trigger | None,
|
||||
):
|
||||
"""Initializes the DataAcquisitionAndControlLayer.
|
||||
|
||||
Args:
|
||||
device_coordinator: The coordinator representing a specific robotic setup.
|
||||
Note that callers need to explicitly initialize and finalise the
|
||||
coordinator.
|
||||
commands_trigger: A trigger to unblock processing commands during a call
|
||||
to `step`.
|
||||
measurements_trigger: A trigger to unblock processing measurements during
|
||||
a call to `step`.
|
||||
"""
|
||||
self._coordinator = device_coordinator
|
||||
self._devices = self._coordinator.get_devices()
|
||||
# The following checks that names of the devices are unique and their keys
|
||||
# are "mergeable".
|
||||
self._check_device_names_and_keys(self._devices)
|
||||
|
||||
self._commands_trigger = commands_trigger
|
||||
self._measurements_trigger = measurements_trigger
|
||||
|
||||
# Create a map of supported commands keys for each Device.
|
||||
self._commands_for_device = {
|
||||
device.name: device.commands_spec().keys() for device in self._devices
|
||||
}
|
||||
|
||||
def begin_stepping(self) -> Mapping[str, gdmr_types.ArrayType]:
|
||||
"""Begins stepping the DACL and returns the current measurements."""
|
||||
self._coordinator.on_begin_stepping()
|
||||
|
||||
# Wait for the first trigger to happen before collecting the measurements.
|
||||
if self._measurements_trigger is not None:
|
||||
self._measurements_trigger.wait_for_event()
|
||||
return self._get_measurements()
|
||||
|
||||
def end_stepping(self) -> None:
|
||||
"""Ends stepping the data acquisition and control layer."""
|
||||
self._coordinator.on_end_stepping()
|
||||
|
||||
def _set_commands(self, commands: Mapping[str, gdmr_types.ArrayType]) -> None:
|
||||
"""Sets the commands of the data acquisition and control layer."""
|
||||
self._coordinator.before_set_commands()
|
||||
for device in self._devices:
|
||||
device_commands = {
|
||||
k: v
|
||||
for k, v in commands.items()
|
||||
if k in self._commands_for_device[device.name]
|
||||
}
|
||||
device.set_commands(device_commands)
|
||||
self._coordinator.after_set_commands()
|
||||
|
||||
def _get_measurements(self) -> Mapping[str, gdmr_types.ArrayType]:
|
||||
"""Gets the measurements of the data acquisition and control layer."""
|
||||
measurements = {}
|
||||
self._coordinator.before_get_measurements()
|
||||
for device in self._devices:
|
||||
measurements.update(device.get_measurements())
|
||||
|
||||
return measurements
|
||||
|
||||
def step(
|
||||
self, commands: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> Mapping[str, gdmr_types.ArrayType]:
|
||||
"""Steps the data acquisition and control layer."""
|
||||
if self._commands_trigger is not None:
|
||||
self._commands_trigger.wait_for_event()
|
||||
self._set_commands(commands)
|
||||
|
||||
if self._measurements_trigger is not None:
|
||||
self._measurements_trigger.wait_for_event()
|
||||
return self._get_measurements()
|
||||
|
||||
def commands_spec(self) -> Mapping[str, gdmr_types.AnyArraySpec]:
|
||||
"""Returns the specs for the commands."""
|
||||
spec = {}
|
||||
for device in self._devices:
|
||||
spec.update(device.commands_spec())
|
||||
return spec
|
||||
|
||||
def measurements_spec(self) -> Mapping[str, specs.Array]:
|
||||
"""Returns the specs for the measurements."""
|
||||
spec = {}
|
||||
for device in self._devices:
|
||||
spec.update(device.measurements_spec())
|
||||
return spec
|
||||
|
||||
@property
|
||||
def device_coordinator(self) -> reaf_coordinator.DeviceCoordinator:
|
||||
return self._coordinator
|
||||
|
||||
def _check_keys_have_been_formatted_correctly(
|
||||
self, current_key_set: Iterable[str]
|
||||
) -> None:
|
||||
"""Check that keys haven't been left unformatted."""
|
||||
for key in current_key_set:
|
||||
if key.find("{}") != -1:
|
||||
raise ValueError(
|
||||
"Keys should not contain '{}'. Did you mean to use format()?"
|
||||
)
|
||||
|
||||
def _check_device_names_and_keys(
|
||||
self, devices: Iterable[reaf_device.Device]
|
||||
) -> None:
|
||||
"""Raises error if device names are not unique or keys are not exclusive."""
|
||||
# Check names first.
|
||||
all_names = [device.name for device in devices]
|
||||
unique_names = set(all_names)
|
||||
if len(unique_names) != len(all_names):
|
||||
raise RuntimeError(f"Duplicate names when checking devices: {all_names}")
|
||||
|
||||
# Check commands.
|
||||
devices = tuple(devices)
|
||||
current_specs = set()
|
||||
for device in devices:
|
||||
device_keys = device.commands_spec().keys()
|
||||
self._check_keys_have_been_formatted_correctly(device_keys)
|
||||
if not current_specs.isdisjoint(device_keys):
|
||||
raise RuntimeError(
|
||||
f"Duplicate keys when checking device {device.name}:"
|
||||
f" {current_specs.intersection(device_keys)}"
|
||||
)
|
||||
current_specs.update(device_keys)
|
||||
|
||||
# Check measurements.
|
||||
current_specs = set()
|
||||
for device in devices:
|
||||
device_keys = device.measurements_spec().keys()
|
||||
self._check_keys_have_been_formatted_correctly(device_keys)
|
||||
if not current_specs.isdisjoint(device_keys):
|
||||
raise RuntimeError(
|
||||
f"Duplicate keys when checking device {device.name}:"
|
||||
f" {current_specs.intersection(device_keys)}"
|
||||
)
|
||||
current_specs.update(device_keys)
|
||||
@@ -0,0 +1,79 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Computes a constant discount given the termination state.
|
||||
|
||||
This provider returns a discount of 0.0 in case of termination and 1.0
|
||||
otherwise (i.e. for truncation and not termination).
|
||||
|
||||
It is usually safe to use this discount provider for environments that return
|
||||
strictly positive rewards.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
|
||||
from dm_env import specs
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
import numpy as np
|
||||
from reaf.core import discount_provider
|
||||
from reaf.core import termination_checker
|
||||
import tree
|
||||
|
||||
|
||||
class DefaultDiscountProvider(discount_provider.DiscountProvider):
|
||||
"""Computes a constant discount given the termination state.
|
||||
|
||||
This provider returns a discount of 0.0 in case of termination and 1.0
|
||||
otherwise (i.e. for truncation and not termination).
|
||||
|
||||
It is usually safe to use this discount provider for environments that return
|
||||
strictly positive rewards.
|
||||
"""
|
||||
|
||||
def __init__(self, name: str = "default_discount_provider"):
|
||||
self._name = name
|
||||
self._spec = specs.BoundedArray(
|
||||
shape=(), dtype=np.float64, minimum=0.0, maximum=1.0, name="discount"
|
||||
)
|
||||
|
||||
def name(self) -> str:
|
||||
"""Returns a unique string identifier for this object."""
|
||||
return self._name
|
||||
|
||||
def compute_discount(
|
||||
self,
|
||||
unused_required_features: Mapping[str, gdmr_types.ArrayType],
|
||||
termination_state: termination_checker.TerminationResult,
|
||||
) -> tree.Structure[gdmr_types.ArrayType]:
|
||||
"""Computes the discount.
|
||||
|
||||
Args:
|
||||
unused_required_features: Unused
|
||||
termination_state: The termination state as computed by the termination
|
||||
checkers. Returns the discount.
|
||||
|
||||
Returns:
|
||||
The discount.
|
||||
"""
|
||||
if termination_state == termination_state.TERMINATE:
|
||||
return np.asarray(0).astype(self._spec.dtype)
|
||||
else: # TRUNCATION or DO_NOT_TERMINATE
|
||||
return np.asarray(1.0).astype(self._spec.dtype)
|
||||
|
||||
def discount_spec(self) -> tree.Structure[specs.Array]:
|
||||
"""Returns the spec of the discount."""
|
||||
return self._spec
|
||||
|
||||
def required_features_keys(self) -> set[str]:
|
||||
"""Returns the feature keys that are required to compute the discount."""
|
||||
return set()
|
||||
@@ -0,0 +1,231 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""ObservationSpaceAdapter supporting filtering, renaming and type conversion."""
|
||||
|
||||
import abc
|
||||
from collections.abc import Iterable, Mapping
|
||||
import dataclasses
|
||||
|
||||
from dm_env import specs
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
import numpy as np
|
||||
import numpy.typing as npt
|
||||
from reaf.core import observation_space_adapter
|
||||
import tree
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True, kw_only=True)
|
||||
class RenameInfo:
|
||||
original_key: str
|
||||
renamed_key: str
|
||||
|
||||
|
||||
class ObservationTypeMapper(abc.ABC):
|
||||
"""Maps from REAF features and specs into corresponding environment types."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def to_observation_spec(
|
||||
self, features_spec: Mapping[str, specs.Array]
|
||||
) -> gdmr_types.ObservationSpec:
|
||||
"""Convert the features spec into the environment observation spec."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def to_observations(
|
||||
self, features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> tree.Structure[gdmr_types.ArrayType]:
|
||||
"""Convert the features into the environment observations."""
|
||||
|
||||
|
||||
class _DefaultObservationTypeMapper(ObservationTypeMapper):
|
||||
"""An ObservationTypeMapper that returns the input features specs and dict.
|
||||
|
||||
This `ObservationTypeMapper` maps observations from the more constrained
|
||||
`Mapping[str, ArrayType]` used in the task layer to the more generic
|
||||
`tree.Structure[ArrayType]` exposed by the GDM Environment.
|
||||
"""
|
||||
|
||||
def to_observation_spec(
|
||||
self, features_spec: Mapping[str, specs.Array]
|
||||
) -> gdmr_types.ObservationSpec:
|
||||
"""Returns the features spec, unmodified, as a `gdmr_types.ObservationSpec`."""
|
||||
return features_spec
|
||||
|
||||
def to_observations(
|
||||
self, features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> tree.Structure[gdmr_types.ArrayType]:
|
||||
"""Returns the features, unmodified, as a `tree.Structure`."""
|
||||
return features
|
||||
|
||||
|
||||
class DefaultObservationSpaceAdapter(
|
||||
observation_space_adapter.ObservationSpaceAdapter
|
||||
):
|
||||
"""Observation adapter supporting filtering, renaming and type conversion.
|
||||
|
||||
This adapter supports filtering, renaming, and converting REAF features into
|
||||
environment observations.
|
||||
|
||||
The order of operations is the following:
|
||||
1) Filtering, i.e. feature selection.
|
||||
2) Downcasting floats to max_float_dtype.
|
||||
3) Renaming.
|
||||
4) Type conversion.
|
||||
|
||||
Please refer to the constructor documentation for more information.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
task_features_spec: Mapping[str, specs.Array],
|
||||
selected_features: Iterable[str] | None,
|
||||
renamed_features: Iterable[RenameInfo] | None,
|
||||
observation_type_mapper: ObservationTypeMapper | None,
|
||||
max_float_dtype: npt.DTypeLike = np.float64,
|
||||
):
|
||||
"""Initializes the observation space adapter.
|
||||
|
||||
Args:
|
||||
task_features_spec: The spec of all the features exposed by the task
|
||||
layer.
|
||||
selected_features: The features that will be exposed as observations. If
|
||||
None, all features will be exposed, i.e. no filtering.
|
||||
renamed_features: `RenameInfo` objects specifying which features should be
|
||||
renamed and the corresponding new name. If empty or None, no renaming
|
||||
will occur.
|
||||
observation_type_mapper: An `ObservationTypeMapper` specifying how to
|
||||
convert the task layer features data type (i.e. a Mapping[str,
|
||||
ArrayType]) into the more generic type exposed by the GDM Environment
|
||||
(i.e. a tree.Structure[ArrayType]). If None, an instance of
|
||||
`_DefaultObservationTypeMapper` is used which converts the task logic
|
||||
layer features dictionary to the more generic type (i.e.
|
||||
`tree.Structure[ArrayType])` exposed by the environment.
|
||||
max_float_dtype: The maximum float dtype to use for downcasting floats.
|
||||
"""
|
||||
if not np.issubdtype(max_float_dtype, np.floating):
|
||||
raise ValueError(
|
||||
'max_float_dtype must be a floating point dtype. Got'
|
||||
f' {max_float_dtype}'
|
||||
)
|
||||
self._max_float_dtype = max_float_dtype
|
||||
self._max_bits = np.finfo(self._max_float_dtype).bits
|
||||
self._task_features_spec = task_features_spec
|
||||
self._selected_filter = selected_features
|
||||
self._renamed_features = renamed_features or ()
|
||||
self._observation_type_mapper = (
|
||||
observation_type_mapper or _DefaultObservationTypeMapper()
|
||||
)
|
||||
self._check_specs_consistency()
|
||||
# Compute the observation spec only once.
|
||||
self._observation_spec = self._compute_observation_spec()
|
||||
|
||||
def _check_specs_consistency(self) -> None:
|
||||
# Check that filter keys are present in the spec.
|
||||
if self._selected_filter is not None:
|
||||
all_features = self._task_features_spec.keys()
|
||||
features = set()
|
||||
for feature in self._selected_filter:
|
||||
if feature not in all_features:
|
||||
raise ValueError(f'Feature {feature} is not present in the spec.')
|
||||
features.add(feature)
|
||||
else:
|
||||
# No filter applied. Select all features.
|
||||
features = set(self._task_features_spec.keys())
|
||||
|
||||
# Check renaming.
|
||||
for rename_info in self._renamed_features:
|
||||
if rename_info.original_key not in features:
|
||||
raise ValueError(
|
||||
f'Feature {rename_info.original_key} is not present in the spec.'
|
||||
)
|
||||
|
||||
def observations_from_features(
|
||||
self, features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> tree.Structure[gdmr_types.ArrayType]:
|
||||
"""Converts the features into the final environment observations."""
|
||||
# 1. Filter the observations.
|
||||
if (selected_features := self._selected_filter) is None:
|
||||
# No filter. Expose all observations.
|
||||
filtered_features = dict(features)
|
||||
else:
|
||||
filtered_features = {
|
||||
k: v for k, v in features.items() if k in selected_features # pytype: disable=unsupported-operands
|
||||
}
|
||||
|
||||
# 2. Downcast floats to max_float_dtype.
|
||||
filtered_features = {
|
||||
k: self._downcast_if_necessary(v) for k, v in filtered_features.items()
|
||||
}
|
||||
|
||||
# 3. Rename.
|
||||
for rename_info in self._renamed_features:
|
||||
# Rename the feature.
|
||||
value = filtered_features[rename_info.original_key]
|
||||
del filtered_features[rename_info.original_key]
|
||||
filtered_features[rename_info.renamed_key] = value
|
||||
|
||||
# 4. Convert type.
|
||||
return self._observation_type_mapper.to_observations(filtered_features)
|
||||
|
||||
def _compute_observation_spec(self) -> gdmr_types.ObservationSpec:
|
||||
"""Computes the observation spec."""
|
||||
# 1. Filter the specs
|
||||
if (features_to_filter := self._selected_filter) is None:
|
||||
# The observation spec corresponds to the task features spec.
|
||||
filtered_specs = dict(self._task_features_spec)
|
||||
else:
|
||||
filtered_specs = {
|
||||
k: v
|
||||
for k, v in self._task_features_spec.items()
|
||||
if k in features_to_filter # pytype: disable=unsupported-operands
|
||||
}
|
||||
|
||||
# 2. Downcast floats to max_float_dtype.
|
||||
for k, v in filtered_specs.items():
|
||||
if self._dtype_needs_downcast(v.dtype):
|
||||
filtered_specs[k] = v.replace(dtype=self._max_float_dtype)
|
||||
|
||||
# 3. Rename.
|
||||
for rename_info in self._renamed_features:
|
||||
# Rename the feature.
|
||||
value = filtered_specs[rename_info.original_key]
|
||||
del filtered_specs[rename_info.original_key]
|
||||
filtered_specs[rename_info.renamed_key] = value
|
||||
|
||||
# 4. Convert the type.
|
||||
return self._observation_type_mapper.to_observation_spec(filtered_specs)
|
||||
|
||||
def observation_spec(self) -> gdmr_types.ObservationSpec:
|
||||
"""Returns the observation spec."""
|
||||
return self._observation_spec
|
||||
|
||||
def task_features_keys(self) -> set[str]:
|
||||
"""Returns the task features keys that will be converted by this adapter."""
|
||||
return set(self._task_features_spec.keys())
|
||||
|
||||
def _downcast_if_necessary(
|
||||
self, value: gdmr_types.ArrayType
|
||||
) -> gdmr_types.ArrayType:
|
||||
if (
|
||||
hasattr(value, 'dtype') and self._dtype_needs_downcast(value.dtype)
|
||||
) or self._dtype_needs_downcast(type(value)):
|
||||
return np.asarray(value).astype(self._max_float_dtype)
|
||||
else:
|
||||
return value
|
||||
|
||||
def _dtype_needs_downcast(self, dtype: npt.DTypeLike) -> bool:
|
||||
return (
|
||||
np.issubdtype(dtype, np.floating)
|
||||
and np.finfo(dtype).bits > self._max_bits
|
||||
)
|
||||
@@ -0,0 +1,54 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""REAF basic device to interface with the robotic setup."""
|
||||
|
||||
import abc
|
||||
from collections.abc import Mapping
|
||||
from dm_env import specs
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
|
||||
|
||||
class Device(abc.ABC):
|
||||
"""REAF basic device to interface with the robotic setup.
|
||||
|
||||
A device defines a single piece in the robotic setup. It should be
|
||||
hermetic, that is, not depending on other Devices. The coordination of the
|
||||
devices is responsibility of the DeviceCoordinator.
|
||||
|
||||
Important: a Device should return the commands and measurements specs
|
||||
immediately after initialisation without the need for any explicit
|
||||
initialisation, nor for resource acquisition (e.g. connecting to the
|
||||
hardware).
|
||||
"""
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Returns the name of this device."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def commands_spec(self) -> Mapping[str, gdmr_types.AnyArraySpec]:
|
||||
"""Returns the commands specs for this device."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def measurements_spec(self) -> Mapping[str, specs.Array]:
|
||||
"""Returns the measurements specs for this device."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def set_commands(self, commands: Mapping[str, gdmr_types.ArrayType]) -> None:
|
||||
"""Sets the commands for this device."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def get_measurements(self) -> Mapping[str, gdmr_types.ArrayType]:
|
||||
"""Returns the measurements provided by this device."""
|
||||
@@ -0,0 +1,87 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Coordinates the devices composing a robotic setup."""
|
||||
|
||||
import abc
|
||||
from collections.abc import Iterable
|
||||
from reaf.core import device
|
||||
|
||||
|
||||
class DeviceCoordinator(abc.ABC):
|
||||
"""Coordinates the devices composing a robotic setup.
|
||||
|
||||
The `DeviceCoordinator` object is responsible for coordinating all the
|
||||
devices constituting the robotic setup. Whilst the Device is hermetic,
|
||||
the coordinator is responsible for passing information from one device to
|
||||
the other if required. For example in a bimanual setup the coordinator is
|
||||
charged with passing the position of each robot to the other so we can ensure
|
||||
proper and safe interaction such as for example collision avoidance.
|
||||
|
||||
The `DeviceCoordinator` can be configurable to enable different
|
||||
properties on the robotic setup, e.g. adding or not adding a `Device` or
|
||||
forwarding configuration to each `Device`.
|
||||
|
||||
At the very least, the coordinator must implement `get_devices`
|
||||
to return all the devices. We also provide `on_begin_stepping` and
|
||||
`on_end_stepping` methods that will be called before the start of an episode
|
||||
and after the end of the episode respectively. Note that resource acquisition
|
||||
and subsequent release is completely up to the implementation.
|
||||
|
||||
Finally, `before_set_commands`/`before_get_measurements` can be implemented to
|
||||
coordinate devices behaviour before their corresponding functions are
|
||||
called.
|
||||
"""
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Returns the name of the coordinator."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def get_devices(self) -> Iterable[device.Device]:
|
||||
"""Returns the devices composing the embodiment."""
|
||||
|
||||
# Lifecycle methods.
|
||||
|
||||
def on_begin_stepping(self) -> None:
|
||||
"""Prepares the coordinator for having its devices called repeatedly.
|
||||
|
||||
After `on_begin_stepping` the devices returned by `get_devices` will have
|
||||
their `set_commands` and `get_measurements` called repeatedly until
|
||||
`on_end_stepping` is called on this coordinator.
|
||||
"""
|
||||
|
||||
def on_end_stepping(self) -> None:
|
||||
"""Notifies the coordinator that the devices are no longer called.
|
||||
|
||||
After `on_end_stepping` the devices returned by `get_devices` will not have
|
||||
their `set_commands` and `get_measurements` called anymore until this
|
||||
coordinator `on_begin_stepping` method is notified again.
|
||||
"""
|
||||
|
||||
# Step hooks methods.
|
||||
|
||||
def before_set_commands(self) -> None:
|
||||
"""Prepares the coordinator to have its devices set_commands called."""
|
||||
|
||||
def after_set_commands(self) -> None:
|
||||
"""Notifies the coordinator that its devices got `set_commands` called."""
|
||||
|
||||
def before_get_measurements(self) -> None:
|
||||
"""Prepares the coordinator to have its devices get_measurements called.
|
||||
|
||||
This method gets called immediately before the devices `get_measurements`
|
||||
method is called and can be used to customise the devices state given the
|
||||
whole setup state.
|
||||
"""
|
||||
@@ -0,0 +1,61 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Computes the discount."""
|
||||
|
||||
import abc
|
||||
from collections.abc import Mapping
|
||||
|
||||
from dm_env import specs
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
from reaf.core import termination_checker
|
||||
import tree
|
||||
|
||||
|
||||
class DiscountProvider(abc.ABC):
|
||||
"""Computes the discount."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Returns a unique string identifier for this object."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def compute_discount(
|
||||
self,
|
||||
required_features: Mapping[str, gdmr_types.ArrayType],
|
||||
termination_state: termination_checker.TerminationResult,
|
||||
) -> tree.Structure[gdmr_types.ArrayType]:
|
||||
"""Computes the discount.
|
||||
|
||||
Args:
|
||||
required_features: Measurements and features computed by the task logic
|
||||
that are required by this provider, i.e. that have keys specified by
|
||||
`required_features_keys`.
|
||||
termination_state: The termination state as computed by the termination
|
||||
checkers. Returns the discount.
|
||||
|
||||
Returns:
|
||||
The discount.
|
||||
"""
|
||||
|
||||
@abc.abstractmethod
|
||||
def discount_spec(self) -> tree.Structure[specs.Array]:
|
||||
"""Returns the spec of the discount."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def required_features_keys(self) -> set[str]:
|
||||
"""Returns the feature keys that are required to compute the discount."""
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Resets the internal state of the discount provider."""
|
||||
...
|
||||
@@ -0,0 +1,65 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Basic REAF-sim protocol to interface with the simulation."""
|
||||
|
||||
from collections.abc import Mapping
|
||||
import typing
|
||||
|
||||
from dm_env import specs
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
|
||||
|
||||
class Entity(typing.Protocol):
|
||||
"""Basic REAF component to interface with the simulation.
|
||||
|
||||
An entity defines a single component in the simulation that consumes substep
|
||||
commands and outputs substep measurements at every simulation substep. It
|
||||
should be hermetic, that is, not depending on other Entities.
|
||||
|
||||
Important: an Entity should return the substep commands and substep
|
||||
measurements specs immediately after initialisation without the need for any
|
||||
explicit initialisation.
|
||||
"""
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Instance name."""
|
||||
|
||||
def reset(self):
|
||||
"""Resets the entity."""
|
||||
|
||||
def substep_commands_spec(
|
||||
self,
|
||||
) -> Mapping[str, specs.Array]:
|
||||
"""Spec for the substep commands."""
|
||||
|
||||
def substep_measurements_spec(
|
||||
self,
|
||||
) -> Mapping[str, specs.Array]:
|
||||
"""Spec for the substep measurements."""
|
||||
|
||||
def set_substep_commands(
|
||||
self,
|
||||
model: typing.Any,
|
||||
data: typing.Any,
|
||||
consumed_substep_commands: Mapping[str, gdmr_types.ArrayType],
|
||||
) -> None:
|
||||
"""Sets the substep commands."""
|
||||
|
||||
def get_substep_measurements(
|
||||
self,
|
||||
model: typing.Any,
|
||||
data: typing.Any,
|
||||
) -> Mapping[str, gdmr_types.ArrayType]:
|
||||
"""Returns the substep measurements."""
|
||||
@@ -0,0 +1,490 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""The Robotics Environment Authoring Framework (REAF) Environment class."""
|
||||
|
||||
import abc
|
||||
from collections.abc import Mapping
|
||||
import enum
|
||||
from typing import Generic
|
||||
|
||||
from absl import logging
|
||||
import dm_env
|
||||
from dm_env import specs
|
||||
from gdm_robotics.interfaces import environment as gdmr_env
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
import numpy as np
|
||||
from reaf.core import action_space_adapter as reaf_action_space_adapter
|
||||
from reaf.core import data_acquisition_and_control_layer as reaf_dacl
|
||||
from reaf.core import default_observation_space_adapter
|
||||
from reaf.core import logger as reaf_logger
|
||||
from reaf.core import observation_space_adapter as reaf_observation_space_adapter
|
||||
from reaf.core import pass_through_action_space_adapter
|
||||
from reaf.core import task_logic_layer as reaf_tll
|
||||
import tree
|
||||
|
||||
|
||||
class ActionSpecEnforcementOption(enum.StrEnum):
|
||||
"""Options for action spec enforcement."""
|
||||
|
||||
CLIP_TO_SPEC = "clip_to_spec"
|
||||
IGNORE = "ignore"
|
||||
WARNING = "warning"
|
||||
RAISE_ERROR = "raise_error"
|
||||
|
||||
|
||||
class EnvironmentReset(abc.ABC, Generic[gdmr_env.ResetOptions]):
|
||||
"""Support for general resets adhering to the GDM environment API."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def do_reset(
|
||||
self,
|
||||
config: gdmr_env.ResetOptions,
|
||||
) -> None:
|
||||
"""Resets the environment."""
|
||||
|
||||
def default_reset_configuration(self) -> gdmr_env.ResetOptions:
|
||||
"""Returns the default reset configuration."""
|
||||
return gdmr_env.Options()
|
||||
|
||||
|
||||
class EndOfEpisodeHandler:
|
||||
"""Handler called after the last episode step."""
|
||||
|
||||
def on_end_of_episode_stepping(self, final_timestep: dm_env.TimeStep) -> None:
|
||||
"""Called when the episode has ended stepping.
|
||||
|
||||
This will be called at the end of every episode, after all other triggers
|
||||
have been resolved. Episodes can end either due to truncation or
|
||||
termination, i.e. `timestep.step_type` is `StepType.LAST`, or due to an
|
||||
early call to `Environment.reset()`. To verify whether it has indeed
|
||||
ended due to truncation or termination, the implementer should test
|
||||
`timestep.last()`.
|
||||
|
||||
Note that the first reset after environment construction will not trigger
|
||||
this handler, but it will be triggered before resolving any subsequent
|
||||
environment resets, either implicit or explicit.
|
||||
|
||||
Args:
|
||||
final_timestep: The final timestep of the episode that ended stepping.
|
||||
"""
|
||||
|
||||
|
||||
class EnvironmentCloser(abc.ABC):
|
||||
"""Handler called when the environment is closed."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def close(self) -> None:
|
||||
"""Releases resources when the environment is closed.
|
||||
|
||||
This method is called automatically when exiting the environment's
|
||||
context manager (`with` statement).
|
||||
"""
|
||||
|
||||
|
||||
class Environment(gdmr_env.Environment):
|
||||
"""The Robotics Environment Authoring Framework (REAF) Environment class."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
data_acquisition_and_control_layer: reaf_dacl.DataAcquisitionAndControlLayer,
|
||||
task_logic_layer: reaf_tll.TaskLogicLayer,
|
||||
environment_reset: EnvironmentReset,
|
||||
action_space_adapter: (
|
||||
reaf_action_space_adapter.ActionSpaceAdapter | None
|
||||
) = None,
|
||||
observation_space_adapter: (
|
||||
reaf_observation_space_adapter.ObservationSpaceAdapter | None
|
||||
) = None,
|
||||
end_of_episode_handler: EndOfEpisodeHandler | None = None,
|
||||
environment_closer: EnvironmentCloser | None = None,
|
||||
action_spec_enforcement_option: ActionSpecEnforcementOption = ActionSpecEnforcementOption.RAISE_ERROR,
|
||||
):
|
||||
"""Creates an environment.
|
||||
|
||||
Args:
|
||||
data_acquisition_and_control_layer: The layer for communicating with the
|
||||
specific robotic setup.
|
||||
task_logic_layer: The layer in charge of defining the task.
|
||||
environment_reset: The `EnvironmentReset` specifying the function to be
|
||||
called at environment reset and the default environment reset
|
||||
configuration.
|
||||
action_space_adapter: Adapter from the agent action space to the flattened
|
||||
commands accepted by the task layer. If None the
|
||||
PassThroughActionSpaceAdapter is used, meaning the entirety of the
|
||||
commands dictionary is exposed to the agent.
|
||||
observation_space_adapter: Adapter from the computed features to the
|
||||
observations that are exposed to the agent. If None the
|
||||
DefaultObservationSpaceAdapter is used, meaning all the features are
|
||||
exposed to the agent as observations.
|
||||
end_of_episode_handler: Called at the end of an episode, after the last
|
||||
step.
|
||||
environment_closer: Specifies the handler to be called when the
|
||||
environment is closed. This is called automatically on exit if the
|
||||
environment is used as a context manager. If None, no action is
|
||||
performed at close.
|
||||
action_spec_enforcement_option: How to enforce the action spec. If
|
||||
`CLIP_TO_SPEC`, the action will be clipped to the spec. If `WARNING`, an
|
||||
warning logged if the action is outside the spec. If `RAISE_ERROR`, an
|
||||
error will be raised if the action is outside the spec. If `IGNORE`,
|
||||
the action will be passed through. Default is `RAISE_ERROR`.
|
||||
"""
|
||||
|
||||
self._data_acquisition_and_control_layer = (
|
||||
data_acquisition_and_control_layer
|
||||
)
|
||||
self._task_logic_layer = task_logic_layer
|
||||
self._end_of_episode_handler = (
|
||||
end_of_episode_handler or EndOfEpisodeHandler()
|
||||
)
|
||||
self._environment_reset = environment_reset
|
||||
self._environment_closer = environment_closer
|
||||
self._action_spec_enforcement_option = action_spec_enforcement_option
|
||||
|
||||
# Before assigning the adapters, validate the specs on the task logic layer
|
||||
# and the DACL.
|
||||
self._validate_dacl_and_ttl_specs()
|
||||
|
||||
ttl_commands_spec = self._task_logic_layer.commands_spec(
|
||||
self._data_acquisition_and_control_layer.commands_spec()
|
||||
)
|
||||
ttl_features_spec = self._task_logic_layer.features_spec(
|
||||
self._data_acquisition_and_control_layer.measurements_spec()
|
||||
)
|
||||
|
||||
if action_space_adapter is None:
|
||||
action_space_adapter = (
|
||||
pass_through_action_space_adapter.PassThroughActionSpaceAdapter(
|
||||
commands_spec=ttl_commands_spec
|
||||
)
|
||||
)
|
||||
self._action_space_adapter = action_space_adapter
|
||||
|
||||
if observation_space_adapter is None:
|
||||
observation_space_adapter = (
|
||||
default_observation_space_adapter.DefaultObservationSpaceAdapter(
|
||||
task_features_spec=ttl_features_spec,
|
||||
selected_features=None,
|
||||
renamed_features=None,
|
||||
observation_type_mapper=None,
|
||||
)
|
||||
)
|
||||
self._observation_space_adapter = observation_space_adapter
|
||||
|
||||
# Now we can validate the adapters.
|
||||
self._validate_adapters_specs()
|
||||
|
||||
self._last_timestep: dm_env.TimeStep | None = None
|
||||
self._should_finalize_episode = False
|
||||
self._timestep_spec = gdmr_types.TimeStepSpec(
|
||||
step_type=gdmr_types.STEP_TYPE_SPEC,
|
||||
reward=self._task_logic_layer.reward_spec(),
|
||||
discount=self._task_logic_layer.discount_spec(),
|
||||
# The observation spec corresponds to the one exposed by the adapter.
|
||||
observation=self._observation_space_adapter.observation_spec(),
|
||||
)
|
||||
|
||||
self._zero_reward, self._zero_discount = tree.map_structure(
|
||||
_read_only_zeros_like_spec,
|
||||
(self._timestep_spec.reward, self._timestep_spec.discount),
|
||||
)
|
||||
|
||||
def close(self) -> None:
|
||||
"""Frees any resources used by the environment."""
|
||||
if self._environment_closer is not None:
|
||||
self._environment_closer.close()
|
||||
|
||||
def default_reset_options(self) -> gdmr_env.ResetOptions:
|
||||
return self._environment_reset.default_reset_configuration()
|
||||
|
||||
def reset_with_options(
|
||||
self,
|
||||
*,
|
||||
options: gdmr_env.ResetOptions,
|
||||
) -> dm_env.TimeStep:
|
||||
"""Starts a new sequence and returns the first `TimeStep`."""
|
||||
if self._should_finalize_episode:
|
||||
self._finalize_episode()
|
||||
self._environment_reset.do_reset(options)
|
||||
self._task_logic_layer.perform_reset()
|
||||
measurements = self._data_acquisition_and_control_layer.begin_stepping()
|
||||
features = self._task_logic_layer.compute_all_features(measurements)
|
||||
observations = self._compute_observations_from_features(features)
|
||||
|
||||
self._last_timestep = self._restart(observation=observations)
|
||||
# Make sure any early reset after this one triggers `_finalize_episode`.
|
||||
self._should_finalize_episode = True
|
||||
return self._last_timestep
|
||||
|
||||
def action_spec(self) -> gdmr_types.ActionSpec:
|
||||
"""Defines the actions that should be provided to `step`."""
|
||||
# The action spec corresponds to the one exposed by the adapter.
|
||||
return self._action_space_adapter.action_spec()
|
||||
|
||||
def timestep_spec(self) -> gdmr_types.TimeStepSpec:
|
||||
"""Returns the spec associated to the returned TimeStep."""
|
||||
return self._timestep_spec
|
||||
|
||||
def step(self, action: gdmr_types.ActionType) -> dm_env.TimeStep:
|
||||
"""Updates the environment according to action and returns a `TimeStep`."""
|
||||
|
||||
action = self._enforce_action_spec(action)
|
||||
if self._last_timestep is None or self._last_timestep.last():
|
||||
return self.reset()
|
||||
|
||||
# Process the action to obtain a command.
|
||||
commands = self._compute_commands_from_agent_action(action)
|
||||
commands = self._task_logic_layer.compute_final_commands(commands)
|
||||
measurements = self._data_acquisition_and_control_layer.step(commands)
|
||||
|
||||
# Compute all the features.
|
||||
features = self._task_logic_layer.compute_all_features(measurements)
|
||||
|
||||
# Compute the elements of the timestep.
|
||||
reward = self._task_logic_layer.compute_reward(features)
|
||||
termination_state = self._task_logic_layer.check_for_termination(features)
|
||||
discount = self._task_logic_layer.compute_discount(
|
||||
features, termination_state
|
||||
)
|
||||
|
||||
observations = self._compute_observations_from_features(features)
|
||||
|
||||
if termination_state.is_terminated():
|
||||
self._last_timestep = self._termination(
|
||||
reward=reward, observation=observations
|
||||
)
|
||||
elif termination_state.is_truncated():
|
||||
self._last_timestep = self._truncation(
|
||||
reward=reward, observation=observations, discount=discount
|
||||
)
|
||||
else:
|
||||
self._last_timestep = self._transition(
|
||||
reward=reward, observation=observations, discount=discount
|
||||
)
|
||||
|
||||
if self._last_timestep.last():
|
||||
self._finalize_episode()
|
||||
return self._last_timestep
|
||||
|
||||
def _finalize_episode(self) -> None:
|
||||
self._data_acquisition_and_control_layer.end_stepping()
|
||||
# It's crucial to call `end_stepping` on the dacl before invoking the end
|
||||
# of episode handler. This ensures no further `set_command` or
|
||||
# `get_measurements` calls are made. In contrast, the end of episode
|
||||
# handler might interact with devices, requiring them to be informed
|
||||
# beforehand.
|
||||
self._end_of_episode_handler.on_end_of_episode_stepping(self._last_timestep)
|
||||
self._should_finalize_episode = False
|
||||
|
||||
@property
|
||||
def data_acquisition_and_control_layer(
|
||||
self,
|
||||
) -> reaf_dacl.DataAcquisitionAndControlLayer:
|
||||
return self._data_acquisition_and_control_layer
|
||||
|
||||
@property
|
||||
def task_logic_layer(self) -> reaf_tll.TaskLogicLayer:
|
||||
return self._task_logic_layer
|
||||
|
||||
@property
|
||||
def environment_reset(self) -> EnvironmentReset:
|
||||
return self._environment_reset
|
||||
|
||||
@environment_reset.setter
|
||||
def environment_reset(self, environment_reset: EnvironmentReset) -> None:
|
||||
self._environment_reset = environment_reset
|
||||
|
||||
def add_logger(self, logger: reaf_logger.Logger) -> None:
|
||||
self._task_logic_layer.add_logger(logger)
|
||||
|
||||
def remove_logger(self, logger: reaf_logger.Logger) -> None:
|
||||
self._task_logic_layer.remove_logger(logger)
|
||||
|
||||
def _validate_dacl_and_ttl_specs(self) -> None:
|
||||
"""Validates the specs on the task logic layer."""
|
||||
# Validate the spec on the task logic layer.
|
||||
self._task_logic_layer.validate_spec(
|
||||
dacl_commands_spec=(
|
||||
self._data_acquisition_and_control_layer.commands_spec()
|
||||
),
|
||||
dacl_measurements_spec=(
|
||||
self._data_acquisition_and_control_layer.measurements_spec()
|
||||
),
|
||||
)
|
||||
|
||||
def _validate_adapters_specs(self) -> None:
|
||||
# Collect the full commands and features spec and validate them against
|
||||
# the adapters.
|
||||
commands_spec = set(
|
||||
self._task_logic_layer.commands_spec(
|
||||
self._data_acquisition_and_control_layer.commands_spec()
|
||||
).keys()
|
||||
)
|
||||
features_spec = set(
|
||||
self._task_logic_layer.features_spec(
|
||||
self._data_acquisition_and_control_layer.measurements_spec()
|
||||
)
|
||||
)
|
||||
|
||||
# Check the action space adapter.
|
||||
adapter_keys = self._action_space_adapter.task_commands_keys()
|
||||
|
||||
if adapter_keys != commands_spec:
|
||||
raise ValueError(
|
||||
"Mismatch between commands exposed by the action space adapter:"
|
||||
f" {adapter_keys} and commands spec expected by the task layer:"
|
||||
f" {commands_spec}."
|
||||
)
|
||||
|
||||
# Check the observation spec adapter.
|
||||
adapter_keys = self._observation_space_adapter.task_features_keys()
|
||||
if not adapter_keys.issubset(features_spec):
|
||||
raise ValueError(
|
||||
"Failed to validate observation space adapter specs. Missing keys:"
|
||||
f" {adapter_keys - features_spec}"
|
||||
)
|
||||
|
||||
def _compute_observations_from_features(
|
||||
self, features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> tree.Structure[gdmr_types.ArrayType]:
|
||||
return self._observation_space_adapter.observations_from_features(features)
|
||||
|
||||
def _compute_commands_from_agent_action(
|
||||
self, agent_action: gdmr_types.ActionType
|
||||
) -> Mapping[str, gdmr_types.ArrayType]:
|
||||
return self._action_space_adapter.commands_from_environment_action(
|
||||
agent_action
|
||||
)
|
||||
|
||||
def _restart(
|
||||
self,
|
||||
observation: tree.Structure[gdmr_types.ArrayType],
|
||||
) -> dm_env.TimeStep:
|
||||
"""Returns a `TimeStep` with `step_type` set to `StepType.FIRST`."""
|
||||
return dm_env.TimeStep(
|
||||
step_type=np.asarray(dm_env.StepType.FIRST, dtype=np.uint8),
|
||||
observation=observation,
|
||||
reward=self._zero_reward,
|
||||
discount=self._zero_discount,
|
||||
)
|
||||
|
||||
def _transition(
|
||||
self,
|
||||
reward: tree.Structure[gdmr_types.ArrayType],
|
||||
observation: tree.Structure[gdmr_types.ArrayType],
|
||||
discount: tree.Structure[gdmr_types.ArrayType],
|
||||
) -> dm_env.TimeStep:
|
||||
"""Returns a `TimeStep` with `step_type` set to `StepType.MID`."""
|
||||
return dm_env.TimeStep(
|
||||
step_type=np.asarray(dm_env.StepType.MID, dtype=np.uint8),
|
||||
observation=observation,
|
||||
reward=reward,
|
||||
discount=discount,
|
||||
)
|
||||
|
||||
def _termination(
|
||||
self,
|
||||
reward: tree.Structure[gdmr_types.ArrayType],
|
||||
observation: tree.Structure[gdmr_types.ArrayType],
|
||||
) -> dm_env.TimeStep:
|
||||
"""Returns a `TimeStep` with `step_type` set to `StepType.LAST`."""
|
||||
return dm_env.TimeStep(
|
||||
step_type=np.asarray(dm_env.StepType.LAST, dtype=np.uint8),
|
||||
observation=observation,
|
||||
reward=reward,
|
||||
discount=self._zero_discount,
|
||||
)
|
||||
|
||||
def _truncation(
|
||||
self,
|
||||
reward: tree.Structure[gdmr_types.ArrayType],
|
||||
observation: tree.Structure[gdmr_types.ArrayType],
|
||||
discount: tree.Structure[gdmr_types.ArrayType],
|
||||
) -> dm_env.TimeStep:
|
||||
"""Returns a `TimeStep` with `step_type` set to `StepType.LAST`."""
|
||||
return dm_env.TimeStep(
|
||||
step_type=np.asarray(dm_env.StepType.LAST, dtype=np.uint8),
|
||||
observation=observation,
|
||||
reward=reward,
|
||||
discount=discount,
|
||||
)
|
||||
|
||||
def _enforce_action_spec(
|
||||
self, action: gdmr_types.ActionType
|
||||
) -> gdmr_types.ActionType:
|
||||
"""Enforces the action spec."""
|
||||
match self._action_spec_enforcement_option:
|
||||
case ActionSpecEnforcementOption.IGNORE:
|
||||
pass
|
||||
case ActionSpecEnforcementOption.CLIP_TO_SPEC:
|
||||
try:
|
||||
|
||||
def clip_to_spec(a, s):
|
||||
if isinstance(s, specs.BoundedArray):
|
||||
return np.clip(a, s.minimum, s.maximum)
|
||||
return a
|
||||
|
||||
action = tree.map_structure(
|
||||
clip_to_spec,
|
||||
action,
|
||||
self._action_space_adapter.action_spec(),
|
||||
)
|
||||
except ValueError as e:
|
||||
raise ValueError(
|
||||
"Failed to clip action to spec. Action:"
|
||||
f" {action} and spec: {self._action_space_adapter.action_spec()}"
|
||||
) from e
|
||||
case ActionSpecEnforcementOption.WARNING:
|
||||
|
||||
def _validate_without_raising(a, s):
|
||||
dtype_ok = s.dtype == a.dtype
|
||||
shape_ok = s.shape == a.shape
|
||||
minimum_ok = True
|
||||
maximum_ok = True
|
||||
if isinstance(s, specs.BoundedArray):
|
||||
minimum_ok = (s.minimum <= a).all()
|
||||
maximum_ok = (a <= s.maximum).all()
|
||||
return dtype_ok and shape_ok and minimum_ok and maximum_ok
|
||||
|
||||
if not all(
|
||||
tree.flatten(
|
||||
tree.map_structure(
|
||||
_validate_without_raising,
|
||||
action,
|
||||
self._action_space_adapter.action_spec(),
|
||||
)
|
||||
)
|
||||
):
|
||||
logging.warning(
|
||||
"Failed to validate action against spec. Action: %r and spec: %r",
|
||||
action,
|
||||
self._action_space_adapter.action_spec(),
|
||||
)
|
||||
case ActionSpecEnforcementOption.RAISE_ERROR:
|
||||
action = tree.map_structure(
|
||||
lambda a, spec: spec.validate(a), action, self.action_spec()
|
||||
)
|
||||
case _:
|
||||
raise ValueError(
|
||||
"Unknown action spec enforcement option:"
|
||||
f" {self._action_spec_enforcement_option}"
|
||||
)
|
||||
return action
|
||||
|
||||
|
||||
def _read_only_zeros_like_spec(spec: specs.Array) -> np.ndarray:
|
||||
"""Returns a zero array matching the specified spec."""
|
||||
arr = np.zeros(shape=spec.shape, dtype=spec.dtype)
|
||||
arr.flags.writeable = False
|
||||
return arr
|
||||
@@ -0,0 +1,34 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Observe all the produced features and measurements."""
|
||||
|
||||
import abc
|
||||
from collections.abc import Mapping
|
||||
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
|
||||
|
||||
class FeaturesObserver(abc.ABC):
|
||||
"""Observe all the produced features and measurements."""
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Returns a unique string identifier for this object."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def observe_features(
|
||||
self, features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> None:
|
||||
"""Observes all the features and measurements."""
|
||||
@@ -0,0 +1,56 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Produces additional features to be exposed by the task logic layer."""
|
||||
|
||||
import abc
|
||||
from collections.abc import Mapping
|
||||
|
||||
from dm_env import specs
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
|
||||
|
||||
class FeaturesProducer(abc.ABC):
|
||||
"""Produces additional features to be exposed by the task logic layer."""
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Returns a unique string identifier for this object."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def produce_features(
|
||||
self, required_features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> Mapping[str, gdmr_types.ArrayType]:
|
||||
"""Produces additional features for the environment.
|
||||
|
||||
Args:
|
||||
required_features: Measurements and features generated by previous
|
||||
producers in the processing chain that are required by this processor,
|
||||
i.e. with keys specified by `required_features_keys`.
|
||||
|
||||
Returns additional features that will be added to the global measurements
|
||||
and features dictionary.
|
||||
"""
|
||||
|
||||
@abc.abstractmethod
|
||||
def produced_features_spec(self) -> Mapping[str, specs.Array]:
|
||||
"""Returns the spec of the features produced by this producer."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def required_features_keys(self) -> set[str]:
|
||||
"""Returns the keys that are required to produce the new features."""
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Resets the internal state of the feature producer."""
|
||||
...
|
||||
@@ -0,0 +1,80 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Support logging inside the task logic layer."""
|
||||
|
||||
import abc
|
||||
from collections.abc import Mapping
|
||||
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
|
||||
|
||||
class Logger(abc.ABC):
|
||||
"""Support logging inside the task logic layer.
|
||||
|
||||
Lifecycle
|
||||
For each environment step, these member functions are called in this order:
|
||||
1. `record_measurements` is called with raw measurements from the sensors.
|
||||
2. `record_features` is called with features derived from the measurements.
|
||||
3. `record_commands_processing` is called for each
|
||||
`CommandsProcessor.process_commands` invocation, tracking the
|
||||
transformation of commands.
|
||||
4. `record_final_commands` is called once with the final commands sent to
|
||||
the DACL.
|
||||
|
||||
Notes:
|
||||
An environment is first reset(). This triggers the first two steps above.
|
||||
See reset_with_options in ./environment.py.
|
||||
|
||||
After reset, step is called repeatedly.
|
||||
1. This first triggers steps 3 and 4 (See compute_final_commands in TLL
|
||||
called from step in ./environment.py)
|
||||
2. Features are computed (see compute_all_features in TLL called from
|
||||
step in ./environment.py), triggering steps 1 and 2.
|
||||
"""
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Unique string identifier for this object."""
|
||||
|
||||
def record_measurements(
|
||||
self, measurements: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> None:
|
||||
"""Called once with all the measurements from the DACL."""
|
||||
|
||||
def record_features(
|
||||
self, features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> None:
|
||||
"""Called once with all the features computed in the Task Layer."""
|
||||
|
||||
def record_final_commands(
|
||||
self, commands: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> None:
|
||||
"""Called once with the final commands sent to the DACL."""
|
||||
|
||||
def record_commands_processing(
|
||||
self,
|
||||
name: str,
|
||||
consumed_commands: Mapping[str, gdmr_types.ArrayType],
|
||||
produced_commands: Mapping[str, gdmr_types.ArrayType],
|
||||
) -> None:
|
||||
"""Called once per call to `process_commands` for each CommandsProcessor.
|
||||
|
||||
Args:
|
||||
name: Name of the `CommandsProcessor`.
|
||||
consumed_commands: The commands consumed by the current
|
||||
`CommandsProcessor`.
|
||||
produced_commands: The commands produced by the current
|
||||
`CommandsProcessor`.
|
||||
"""
|
||||
@@ -0,0 +1,98 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Testing functions for asserting on Mock objects with numpy structures."""
|
||||
|
||||
from collections.abc import Sequence
|
||||
from unittest import mock
|
||||
import numpy as np
|
||||
|
||||
|
||||
def assert_called_once_with(mock_obj: mock.Mock, *args, **kwargs) -> None:
|
||||
if mock_obj.call_count != 1:
|
||||
raise AssertionError(
|
||||
f"Expected exactly one call to {mock_obj}, got {mock_obj.call_count}"
|
||||
)
|
||||
|
||||
assert_called_with(mock_obj, *args, **kwargs)
|
||||
|
||||
|
||||
def assert_called_with(mock_obj: mock.Mock, *args, **kwargs) -> None:
|
||||
"""Asserts that the last call to mock_obj had the specified arguments."""
|
||||
if mock_obj.call_args is None:
|
||||
raise AssertionError(
|
||||
f"Mock object {mock_obj} not called. Expected one call."
|
||||
)
|
||||
call_args, call_kwargs = mock_obj.call_args
|
||||
np.testing.assert_equal(call_args, args)
|
||||
np.testing.assert_equal(call_kwargs, kwargs)
|
||||
|
||||
|
||||
def assert_has_calls(
|
||||
mock_obj: mock.Mock, calls: Sequence[mock._Call], any_order: bool = False
|
||||
) -> None:
|
||||
"""Asserts that mock_obj has been called with the specified calls."""
|
||||
mock_calls = mock_obj.mock_calls
|
||||
|
||||
# Check that there are at least enough calls.
|
||||
if mock_obj.call_count < len(calls):
|
||||
raise AssertionError(
|
||||
f"Expected at least {len(calls)} calls to {mock_obj}, got"
|
||||
f" {mock_obj.call_count}"
|
||||
)
|
||||
|
||||
def _calls_are_equal(actual: mock._Call, expected: mock._Call) -> bool:
|
||||
_, actual_args, actual_kwargs = actual
|
||||
_, expected_args, expected_kwargs = expected
|
||||
# Quickest way to transform the assertion into a comparator.
|
||||
try:
|
||||
np.testing.assert_equal(actual_args, expected_args)
|
||||
np.testing.assert_equal(actual_kwargs, expected_kwargs)
|
||||
return True
|
||||
except AssertionError:
|
||||
return False
|
||||
|
||||
if any_order:
|
||||
# We just check for the calls to be contained.
|
||||
for expected_call in calls:
|
||||
for actual_call in mock_calls:
|
||||
if _calls_are_equal(actual_call, expected_call):
|
||||
break
|
||||
raise AssertionError(
|
||||
f"Expected call {expected_call} not found in mock calls {mock_calls}."
|
||||
)
|
||||
return
|
||||
|
||||
# We need to check in order, but first find the first call.
|
||||
starting_index = -1
|
||||
first_expected_call = calls[0]
|
||||
for index, actual_call in enumerate(mock_calls):
|
||||
if _calls_are_equal(actual_call, first_expected_call):
|
||||
starting_index = index
|
||||
break
|
||||
if starting_index == -1:
|
||||
raise AssertionError(f"Calls {calls} not found in mock calls {mock_calls}.")
|
||||
|
||||
non_matching_calls = []
|
||||
|
||||
# We have the first element. Now we need to compare element wise.
|
||||
for index, expected_call in enumerate(calls):
|
||||
actual_call = mock_calls[starting_index + index]
|
||||
if not _calls_are_equal(actual_call, expected_call):
|
||||
non_matching_calls.append((index, expected_call, actual_call))
|
||||
|
||||
if non_matching_calls:
|
||||
raise AssertionError(
|
||||
f"Calls {calls} do not match mock calls {mock_calls}. Mismatch (index,"
|
||||
f" expected, actual): {non_matching_calls}"
|
||||
)
|
||||
@@ -0,0 +1,42 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Adapts REAF features into observations exposed by the environment."""
|
||||
|
||||
import abc
|
||||
from collections.abc import Mapping
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
import tree
|
||||
|
||||
|
||||
class ObservationSpaceAdapter(abc.ABC):
|
||||
"""Adapts REAF features into observations exposed by the environment.
|
||||
|
||||
Implementations of this interface are responsible for converting the features
|
||||
generated by the REAF task layer logic (i.e. dictionary of tensors) into the
|
||||
more generic `observation` structure exposed by the environment.
|
||||
"""
|
||||
|
||||
@abc.abstractmethod
|
||||
def observations_from_features(
|
||||
self, features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> tree.Structure[gdmr_types.ArrayType]:
|
||||
"""Converts the REAF features into the environment observations."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def observation_spec(self) -> gdmr_types.ObservationSpec:
|
||||
"""Returns the observation spec."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def task_features_keys(self) -> set[str]:
|
||||
"""Returns the task features keys that will be converted by this adapter."""
|
||||
@@ -0,0 +1,55 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Adapter that passes the commands spec through."""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
from reaf.core import action_space_adapter
|
||||
|
||||
|
||||
class PassThroughActionSpaceAdapter(action_space_adapter.ActionSpaceAdapter):
|
||||
"""Adapter that passes the commands spec through.
|
||||
|
||||
NB the resulting environment will expose a dictionary as the action spec.
|
||||
"""
|
||||
|
||||
def __init__(self, commands_spec: Mapping[str, gdmr_types.AnyArraySpec]):
|
||||
self._commands_spec = commands_spec
|
||||
|
||||
def commands_from_environment_action(
|
||||
self, environment_action: gdmr_types.ActionType
|
||||
) -> Mapping[str, gdmr_types.ArrayType]:
|
||||
"""Returns commands accepted by REAF.
|
||||
|
||||
commands_from_environment_action usually accepts a gdmr_types.ActionType but
|
||||
since this adapter passes the same action as the commands, it needs to be a
|
||||
dict type in order to pass it through as a dict.
|
||||
|
||||
Args:
|
||||
environment_action: The environment action(s) to pass as REAF commands.
|
||||
"""
|
||||
if not isinstance(environment_action, dict):
|
||||
raise ValueError(
|
||||
'environment_action must be a dict, but got: '
|
||||
f'{type(environment_action)}.'
|
||||
)
|
||||
return environment_action
|
||||
|
||||
def action_spec(self) -> gdmr_types.ActionSpec:
|
||||
"""Returns the action spec exposed by the environment."""
|
||||
return self._commands_spec
|
||||
|
||||
def task_commands_keys(self) -> set[str]:
|
||||
"""Returns the keys for the commands exposed to the task layer."""
|
||||
return set(self._commands_spec.keys())
|
||||
@@ -0,0 +1,292 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Computes the reward."""
|
||||
|
||||
import abc
|
||||
from collections.abc import Mapping
|
||||
import operator
|
||||
from typing import Callable, TypeAlias, TypeVar, Union
|
||||
|
||||
from dm_env import specs
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
import numpy as np
|
||||
import tree
|
||||
|
||||
|
||||
RewardValue: TypeAlias = tree.Structure[gdmr_types.ArrayType]
|
||||
RewardSpec: TypeAlias = tree.Structure[specs.Array]
|
||||
|
||||
|
||||
class _RewardProvider(abc.ABC):
|
||||
"""Computes the reward.
|
||||
|
||||
Defines the interface for a reward provider.
|
||||
|
||||
Important: Users should not inherit from this class directly. Instead, use the
|
||||
RewardProvider class later in this file.
|
||||
"""
|
||||
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Returns a unique string identifier for this object."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def compute_reward(
|
||||
self, required_features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> RewardValue:
|
||||
"""Computes the reward.
|
||||
|
||||
Args:
|
||||
required_features: Measurements and features computed by the task logic
|
||||
that are required by this provider, i.e. that have keys specified by
|
||||
`required_features_keys`.
|
||||
|
||||
Returns the computed reward.
|
||||
"""
|
||||
|
||||
@abc.abstractmethod
|
||||
def reward_spec(self) -> RewardSpec:
|
||||
"""Returns the spec of the reward."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def required_features_keys(self) -> set[str]:
|
||||
"""Returns the feature keys that are required to compute the reward."""
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Resets the internal state of the reward provider."""
|
||||
...
|
||||
|
||||
|
||||
RewardProviderOrValue: TypeAlias = Union['RewardProvider', RewardValue]
|
||||
|
||||
|
||||
S = TypeVar('S')
|
||||
T = TypeVar('T')
|
||||
UnaryOperator: TypeAlias = Callable[[S], S]
|
||||
BinaryOperator: TypeAlias = Callable[[S | T, S | T], S | T]
|
||||
|
||||
|
||||
class RewardProvider(_RewardProvider):
|
||||
"""Computes the reward.
|
||||
|
||||
Important: Users should inherit from this class and implement the abstract
|
||||
methods defined in the interface _RewardProvider.
|
||||
"""
|
||||
|
||||
def __add__(self, other: RewardProviderOrValue):
|
||||
return BinaryOperationRewardProvider(operator.add, self, other)
|
||||
|
||||
def __radd__(self, other: RewardProviderOrValue):
|
||||
return BinaryOperationRewardProvider(operator.add, other, self)
|
||||
|
||||
def __sub__(self, other: RewardProviderOrValue):
|
||||
return BinaryOperationRewardProvider(operator.sub, self, other)
|
||||
|
||||
def __rsub__(self, other: RewardProviderOrValue):
|
||||
return BinaryOperationRewardProvider(operator.sub, other, self)
|
||||
|
||||
def __mul__(self, other: RewardProviderOrValue):
|
||||
return BinaryOperationRewardProvider(operator.mul, self, other)
|
||||
|
||||
def __rmul__(self, other: RewardProviderOrValue):
|
||||
return BinaryOperationRewardProvider(operator.mul, other, self)
|
||||
|
||||
def __truediv__(self, other: RewardProviderOrValue):
|
||||
return BinaryOperationRewardProvider(operator.truediv, self, other)
|
||||
|
||||
def __rtruediv__(self, other: RewardProviderOrValue):
|
||||
return BinaryOperationRewardProvider(operator.truediv, other, self)
|
||||
|
||||
def __floordiv__(self, other: RewardProviderOrValue):
|
||||
return BinaryOperationRewardProvider(operator.floordiv, self, other)
|
||||
|
||||
def __rfloordiv__(self, other: RewardProviderOrValue):
|
||||
return BinaryOperationRewardProvider(operator.floordiv, other, self)
|
||||
|
||||
def __pow__(self, other: RewardProviderOrValue):
|
||||
return BinaryOperationRewardProvider(operator.pow, self, other)
|
||||
|
||||
def __rpow__(self, other: RewardProviderOrValue):
|
||||
return BinaryOperationRewardProvider(operator.pow, other, self)
|
||||
|
||||
def __getitem__(self, index: slice):
|
||||
return GetItemOperationRewardProvider(self, index)
|
||||
|
||||
def __neg__(self):
|
||||
return UnaryOperationRewardProvider(operator.neg, self)
|
||||
|
||||
|
||||
class ConstantRewardProvider(RewardProvider):
|
||||
"""A RewardProvider that always returns the same reward."""
|
||||
|
||||
def __init__(self, reward: RewardValue):
|
||||
super().__init__()
|
||||
self._reward = reward
|
||||
|
||||
def name(self) -> str:
|
||||
return str(self._reward)
|
||||
|
||||
def compute_reward(
|
||||
self, required_features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> RewardValue:
|
||||
return self._reward
|
||||
|
||||
def reward_spec(self) -> RewardSpec:
|
||||
return tree.map_structure(
|
||||
lambda v: specs.Array(v.shape, v.dtype), self._reward
|
||||
)
|
||||
|
||||
def required_features_keys(self) -> set[str]:
|
||||
return set()
|
||||
|
||||
|
||||
class BinaryOperationRewardProvider(RewardProvider):
|
||||
"""Applies a binary operator to the result of two reward providers."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
op: BinaryOperator,
|
||||
first_reward_provider: RewardProviderOrValue,
|
||||
second_reward_provider: RewardProviderOrValue,
|
||||
):
|
||||
super().__init__()
|
||||
if not isinstance(first_reward_provider, RewardProvider):
|
||||
first_reward_provider = ConstantRewardProvider(first_reward_provider)
|
||||
if not isinstance(second_reward_provider, RewardProvider):
|
||||
second_reward_provider = ConstantRewardProvider(second_reward_provider)
|
||||
first_spec = first_reward_provider.reward_spec()
|
||||
second_spec = second_reward_provider.reward_spec()
|
||||
tree.assert_same_structure(first_spec, second_spec)
|
||||
assert all(
|
||||
tree.flatten(
|
||||
tree.map_structure(
|
||||
lambda s1, s2: s1.shape == s2.shape and s1.dtype == s2.dtype,
|
||||
first_spec,
|
||||
second_spec,
|
||||
)
|
||||
)
|
||||
)
|
||||
self._op = op
|
||||
self._first_reward_provider = first_reward_provider
|
||||
self._second_reward_provider = second_reward_provider
|
||||
self._reward_spec = first_reward_provider.reward_spec()
|
||||
self._first_required_features_keys = (
|
||||
first_reward_provider.required_features_keys()
|
||||
)
|
||||
self._second_required_features_keys = (
|
||||
second_reward_provider.required_features_keys()
|
||||
)
|
||||
|
||||
def name(self) -> str:
|
||||
op_name = getattr(self._op, '__name__', str(self._op))
|
||||
return (
|
||||
f'{op_name}({self._first_reward_provider.name()},'
|
||||
f' {self._second_reward_provider.name()})'
|
||||
)
|
||||
|
||||
def compute_reward(
|
||||
self, required_features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> RewardValue:
|
||||
first_required_features = {
|
||||
k: v
|
||||
for k, v in required_features.items()
|
||||
if k in self._first_required_features_keys
|
||||
}
|
||||
second_required_features = {
|
||||
k: v
|
||||
for k, v in required_features.items()
|
||||
if k in self._second_required_features_keys
|
||||
}
|
||||
return tree.map_structure(
|
||||
self._op,
|
||||
self._first_reward_provider.compute_reward(first_required_features),
|
||||
self._second_reward_provider.compute_reward(second_required_features),
|
||||
)
|
||||
|
||||
def reward_spec(self) -> RewardSpec:
|
||||
return self._reward_spec
|
||||
|
||||
def required_features_keys(self) -> set[str]:
|
||||
return (
|
||||
self._first_required_features_keys | self._second_required_features_keys
|
||||
)
|
||||
|
||||
def reset(self) -> None:
|
||||
self._first_reward_provider.reset()
|
||||
self._second_reward_provider.reset()
|
||||
|
||||
|
||||
class GetItemOperationRewardProvider(RewardProvider):
|
||||
"""Extracts a slice from the result of a reward provider."""
|
||||
|
||||
def __init__(self, reward_provider: RewardProviderOrValue, index: slice):
|
||||
super().__init__()
|
||||
if not isinstance(reward_provider, RewardProvider):
|
||||
reward_provider = ConstantRewardProvider(reward_provider)
|
||||
self._reward_provider = reward_provider
|
||||
self._index = index
|
||||
|
||||
def name(self) -> str:
|
||||
return f'{self._reward_provider.name}[{self._index}]'
|
||||
|
||||
def compute_reward(
|
||||
self, required_features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> RewardValue:
|
||||
return tree.map_structure(
|
||||
lambda v: v[self._index],
|
||||
self._reward_provider.compute_reward(required_features),
|
||||
)
|
||||
|
||||
def reward_spec(self) -> RewardSpec:
|
||||
return tree.map_structure(
|
||||
lambda s: specs.Array(np.empty(s.shape)[self._index].shape, s.dtype),
|
||||
self._reward_provider.reward_spec(),
|
||||
)
|
||||
|
||||
def required_features_keys(self) -> set[str]:
|
||||
return self._reward_provider.required_features_keys()
|
||||
|
||||
def reset(self) -> None:
|
||||
self._reward_provider.reset()
|
||||
|
||||
|
||||
class UnaryOperationRewardProvider(RewardProvider):
|
||||
"""Applies a unary operator to the result of a reward provider."""
|
||||
|
||||
def __init__(self, op: UnaryOperator, reward_provider: RewardProviderOrValue):
|
||||
super().__init__()
|
||||
if not isinstance(reward_provider, RewardProvider):
|
||||
reward_provider = ConstantRewardProvider(reward_provider)
|
||||
self._op = op
|
||||
self._reward_provider = reward_provider
|
||||
|
||||
def name(self) -> str:
|
||||
op_name = getattr(self._op, '__name__', str(self._op))
|
||||
return f'{op_name}({self._reward_provider.name()})'
|
||||
|
||||
def compute_reward(
|
||||
self, required_features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> RewardValue:
|
||||
return tree.map_structure(
|
||||
self._op, self._reward_provider.compute_reward(required_features)
|
||||
)
|
||||
|
||||
def reward_spec(self) -> RewardSpec:
|
||||
return self._reward_provider.reward_spec()
|
||||
|
||||
def required_features_keys(self) -> set[str]:
|
||||
return self._reward_provider.required_features_keys()
|
||||
|
||||
def reset(self) -> None:
|
||||
self._reward_provider.reset()
|
||||
@@ -0,0 +1,104 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Protocol for substep commands manipulation in REAF-sim."""
|
||||
|
||||
from collections.abc import Mapping
|
||||
import typing
|
||||
|
||||
from dm_env import specs
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
|
||||
|
||||
class SubstepCommandsProcessor(typing.Protocol):
|
||||
"""Processes substep commands, propagating them through a pipeline.
|
||||
|
||||
This processor manipulates substep commands, acting as a node in a pipeline.
|
||||
It consumes substep commands, performs operations, and produces updated
|
||||
substep commands for the next stage in the processing chain.
|
||||
|
||||
The processing pipeline starts with commands provided to the SimulationDevice
|
||||
and progresses towards the substep commands consumed by the individual
|
||||
entities. Each processor consumes a subset of substep commands and produces
|
||||
new, potentially transformed, substep commands. The order of operations is
|
||||
crucial.
|
||||
|
||||
Example Pipeline (conceptual):
|
||||
|
||||
Simulation Device commands --> Processor (1) --> Processor (2) --> Entities
|
||||
|
||||
Specs are propagated starting from the bottom:
|
||||
1) In this example assume that the set of entities expect "p3/c1", "p3/c2" and
|
||||
"p3/c3".
|
||||
2) Processor (2) returns ("p3/c1", "p3/c2") from "p2/c1". This means
|
||||
that the global substep commands spec exposed at this level is "p2/c1" and
|
||||
the unprocessed "p3/c3".
|
||||
3) Processor (1) returns "p2/c1" from ("p1/c1", "p1/c2"). By applying
|
||||
the same transformation rule, we can obtain the final spec exposed
|
||||
by the SimulationDevice: "p1/c1", "p1/c2" and "p3/c3".
|
||||
|
||||
------------------------------------
|
||||
| SimulationDevice |
|
||||
------------------------------------
|
||||
|
||||
"p1/c1" "p1/c2" "p3/c3"
|
||||
| | |
|
||||
----------------- |
|
||||
| P1 | |
|
||||
----------------- |
|
||||
| "p2/c1" |
|
||||
----------------- |
|
||||
| P2 | |
|
||||
----------------- |
|
||||
| "p3/c1" | "p3/c2" |
|
||||
| | |
|
||||
------------------------------------
|
||||
| Entities |
|
||||
------------------------------------
|
||||
"""
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Returns a unique string identifier for this object."""
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Resets the internal state of this processor."""
|
||||
|
||||
def produced_substep_commands_keys(self) -> set[str]:
|
||||
"""Keys of the substep commands produced by this processor."""
|
||||
|
||||
def consumed_substep_commands_spec(
|
||||
self,
|
||||
) -> Mapping[str, specs.Array]:
|
||||
"""Spec of the substep commands consumed by this processor."""
|
||||
|
||||
def process_substep_commands(
|
||||
self,
|
||||
model: typing.Any,
|
||||
data: typing.Any,
|
||||
consumed_substep_commands: Mapping[str, gdmr_types.ArrayType],
|
||||
) -> Mapping[str, gdmr_types.ArrayType]:
|
||||
"""Processes the substep commands and returns a new modified version of it.
|
||||
|
||||
Args:
|
||||
model: the simulation model.
|
||||
data: the simulation data.
|
||||
consumed_substep_commands: the substep commands up in the processing chain
|
||||
that are required by this processor, i.e. with keys specified by
|
||||
`consumed_substep_commands_spec`.
|
||||
|
||||
Returns the new substep commands. Note that the (key, value) pairs in
|
||||
`consumed_substep_commands` are removed from the running substep commands
|
||||
dictionary. If users want to keep some of the elements it is their
|
||||
responsibility to retain them in the output dictionary.
|
||||
"""
|
||||
@@ -0,0 +1,103 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Protocol for substep measurements manipulation in REAF-sim."""
|
||||
|
||||
from collections.abc import Mapping
|
||||
import typing
|
||||
|
||||
from dm_env import specs
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
|
||||
|
||||
class SubstepMeasurementsProcessor(typing.Protocol):
|
||||
"""Processes substep measurements, propagating them through a pipeline.
|
||||
|
||||
This processor manipulates substep measurements, acting as a node in a
|
||||
pipeline. It consumes substep measurements, performs operations, and produces
|
||||
updated substep measurements for the next stage in the processing chain.
|
||||
|
||||
The processing pipeline starts with substep measurements produced by Entities
|
||||
and progresses towards the measurements exposed by the SimulationDevice. Each
|
||||
processor consumes a subset of substep measurements and produces new,
|
||||
potentially transformed, substep measurements. The order of operations is
|
||||
crucial.
|
||||
|
||||
Example Pipeline (conceptual):
|
||||
|
||||
Entities --> Processor (1) --> Processor (2) -> Simulation Device Measurements
|
||||
|
||||
Specs are propagated starting from the bottom:
|
||||
1) In this example assume that the set of entities produce "p1/c1", "p1/c2"
|
||||
and "p1/c3".
|
||||
2) Processor (1) returns "p2/c1" from ("p1/c1", "p1/c2").
|
||||
3) Processor (2) returns ("p3/c1", "p3/c2") from "p2/c1".
|
||||
|
||||
This resulting spec exposed by the SimulationDevice: "p3/c1", "p3/c2"
|
||||
and "p1/c3".
|
||||
|
||||
------------------------------------
|
||||
| SimulationDevice |
|
||||
------------------------------------
|
||||
|
||||
"p3/c1" "p3/c2" "p1/c3"
|
||||
| | |
|
||||
----------------- |
|
||||
| P2 | |
|
||||
----------------- |
|
||||
| "p2/c1" |
|
||||
----------------- |
|
||||
| P1 | |
|
||||
----------------- |
|
||||
| "p1/c1" | "p1/c2" |
|
||||
| | |
|
||||
------------------------------------
|
||||
| Entities |
|
||||
------------------------------------
|
||||
"""
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Returns a unique string identifier for this object."""
|
||||
|
||||
def reset(self):
|
||||
"""Resets the internal state of this processor."""
|
||||
|
||||
def produced_substep_measurements_spec(
|
||||
self,
|
||||
) -> Mapping[str, specs.Array]:
|
||||
"""Spec of the substep measurements consumed by this processor."""
|
||||
|
||||
def consumed_substep_measurements_keys(self) -> set[str]:
|
||||
"""Keys of the substep measurements consumed by this processor."""
|
||||
|
||||
def process_substep_measurements(
|
||||
self,
|
||||
model: typing.Any,
|
||||
data: typing.Any,
|
||||
consumed_substep_measurements: Mapping[str, gdmr_types.ArrayType],
|
||||
) -> Mapping[str, gdmr_types.ArrayType]:
|
||||
"""Processes the substep measurements and returns a new modified version of it.
|
||||
|
||||
Args:
|
||||
model: the simulation model.
|
||||
data: the simulation data.
|
||||
consumed_substep_measurements: the substep measurements up in the
|
||||
processing chain that are required by this processor, i.e. with keys
|
||||
specified by `consumed_substep_measurements_spec`.
|
||||
|
||||
Returns the new substep measurements. Note that the (key, value) pairs in
|
||||
`consumed_substep_measurements` are removed from the running substep
|
||||
measurements dictionary. If users want to keep some of the elements it is
|
||||
their responsibility to retain them in the output dictionary.
|
||||
"""
|
||||
@@ -0,0 +1,342 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Task logic layer for the Robotics Environment Authoring Framework."""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
import itertools
|
||||
from typing import Protocol
|
||||
|
||||
from absl import logging
|
||||
from dm_env import specs
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
from reaf.core import commands_processor as reaf_commands_processor
|
||||
from reaf.core import default_discount_provider
|
||||
from reaf.core import discount_provider as reaf_discount_provider
|
||||
from reaf.core import features_observer as reaf_features_observers
|
||||
from reaf.core import features_producer as reaf_features_producer
|
||||
from reaf.core import logger as reaf_logger
|
||||
from reaf.core import reward_provider as reaf_reward_provider
|
||||
from reaf.core import termination_checker as reaf_termination_checker
|
||||
from reaf.core import zero_reward_provider
|
||||
import tree
|
||||
|
||||
|
||||
class _ResettableObject(Protocol):
|
||||
"""Protocol for an object that can be reset."""
|
||||
|
||||
def reset(self) -> None:
|
||||
...
|
||||
|
||||
|
||||
class TaskLogicLayer:
|
||||
"""Task logic layer for the Robotics Environment Authoring Framework."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
commands_processors: Sequence[reaf_commands_processor.CommandsProcessor],
|
||||
features_producers: Sequence[reaf_features_producer.FeaturesProducer],
|
||||
termination_checkers: Sequence[
|
||||
reaf_termination_checker.TerminationChecker
|
||||
],
|
||||
reward_provider: reaf_reward_provider.RewardProvider | None = None,
|
||||
discount_provider: reaf_discount_provider.DiscountProvider | None = None,
|
||||
features_observers: Sequence[
|
||||
reaf_features_observers.FeaturesObserver
|
||||
] = (),
|
||||
loggers: Sequence[reaf_logger.Logger] = (),
|
||||
):
|
||||
"""Initializes the task logic layer.
|
||||
|
||||
Args:
|
||||
commands_processors: `CommandsProcessor`s that modify the commands before
|
||||
being sent down to the DACL. They are called sequentially, starting from
|
||||
the commands supplied by the policy and ending with the commands that
|
||||
will be sent to the DACL.
|
||||
features_producers: `FeaturesProducer`s that generate new features.
|
||||
Measurements collected by the DACL and features produced by these
|
||||
`FeaturesProducer`s are then merged into the final feature set that is
|
||||
provided to the `reward_provider`, `termination_checkers`,
|
||||
`discount_provider`, `features_observers`, and `loggers`.
|
||||
termination_checkers: `TerminationChecker`s that check the episode
|
||||
termination based on the final feature set.
|
||||
reward_provider: `RewardProvider` that computes a reward based on the
|
||||
final feature set. If None, the ZeroRewardProvider is used and the
|
||||
reward is set to 0.
|
||||
discount_provider: `DiscountProvider` that compute a discount based on the
|
||||
final feature set and final termination state. If None, the
|
||||
DefaultDiscountProvider is used returning 0 for termination and 1 for
|
||||
truncation and non-termination.
|
||||
features_observers: `FeaturesObserver`s that get a view over the final
|
||||
feature set.
|
||||
loggers: `Logger`s for logging measurements, features, and commands in the
|
||||
task layer.
|
||||
"""
|
||||
self._commands_processors = commands_processors
|
||||
self._features_producers = features_producers
|
||||
self._reward_provider = (
|
||||
reward_provider
|
||||
if reward_provider
|
||||
else zero_reward_provider.ZeroRewardProvider()
|
||||
)
|
||||
self._termination_checkers = termination_checkers
|
||||
self._discount_provider = (
|
||||
discount_provider
|
||||
if discount_provider
|
||||
else default_discount_provider.DefaultDiscountProvider()
|
||||
)
|
||||
self._features_observers = features_observers
|
||||
self._loggers = list(loggers)
|
||||
|
||||
# We make a set of all resettable objects so that these objects only get
|
||||
# their resets called once. This is important for e.g. when having a single
|
||||
# object that derives from two interfaces.
|
||||
self._resettable_objects: list[_ResettableObject] = []
|
||||
unique_ids = set()
|
||||
for resettable_object in itertools.chain(
|
||||
self._commands_processors,
|
||||
self._features_producers,
|
||||
self._termination_checkers,
|
||||
[self._reward_provider],
|
||||
[self._discount_provider],
|
||||
):
|
||||
resettable_object_id = id(resettable_object)
|
||||
if resettable_object_id not in unique_ids:
|
||||
unique_ids.add(resettable_object_id)
|
||||
self._resettable_objects.append(resettable_object)
|
||||
|
||||
def validate_spec(
|
||||
self,
|
||||
*,
|
||||
dacl_commands_spec: Mapping[str, gdmr_types.AnyArraySpec],
|
||||
dacl_measurements_spec: Mapping[str, specs.Array],
|
||||
) -> None:
|
||||
"""Checks that the specs have consistent keys."""
|
||||
logging.vlog(3, "Validate features processing")
|
||||
self._validate_features_spec(dacl_measurements_spec)
|
||||
self._validate_commands_spec(dacl_commands_spec)
|
||||
|
||||
def features_spec(
|
||||
self,
|
||||
dacl_measurements_spec: Mapping[str, specs.Array],
|
||||
) -> Mapping[str, specs.Array]:
|
||||
"""Returns the features spec as exposed by the task layer."""
|
||||
spec = dict(dacl_measurements_spec)
|
||||
for features_producer in self._features_producers:
|
||||
spec.update(features_producer.produced_features_spec())
|
||||
|
||||
return spec
|
||||
|
||||
def commands_spec(
|
||||
self, dacl_commands_spec: Mapping[str, gdmr_types.AnyArraySpec]
|
||||
) -> Mapping[str, gdmr_types.AnyArraySpec]:
|
||||
"""Returns the commands spec exposed by the task layer."""
|
||||
# Each processor consumes commands (as described by its
|
||||
# `consumed_commands_spec`) and outputs a potentially different set of
|
||||
# commands (as described by its `produced_commands_keys`).
|
||||
# Starting with the DACL command spec, we iterate in reverse order (i.e. in
|
||||
# the direction DACL -> Policy) through every processor to remove the
|
||||
# `produced_commands_keys` from the spec, and add their
|
||||
# `consumed_commands_spec` to the spec.
|
||||
spec: Mapping[str, gdmr_types.AnyArraySpec] = dict(dacl_commands_spec)
|
||||
for processor in reversed(self._commands_processors):
|
||||
processor_produced_keys = processor.produced_commands_keys()
|
||||
spec = {
|
||||
key: value
|
||||
for key, value in spec.items()
|
||||
if key not in processor_produced_keys
|
||||
}
|
||||
spec.update(processor.consumed_commands_spec())
|
||||
return spec
|
||||
|
||||
def reward_spec(self) -> tree.Structure[specs.Array]:
|
||||
return self._reward_provider.reward_spec()
|
||||
|
||||
def discount_spec(self) -> tree.Structure[specs.Array]:
|
||||
return self._discount_provider.discount_spec()
|
||||
|
||||
def perform_reset(self) -> None:
|
||||
"""Reset the internal state of the task logic layer."""
|
||||
for resettable_object in self._resettable_objects:
|
||||
resettable_object.reset()
|
||||
|
||||
def compute_all_features(
|
||||
self, measurements: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> Mapping[str, gdmr_types.ArrayType]:
|
||||
"""Computes all the task logic features given the current measurements."""
|
||||
for logger in self._loggers:
|
||||
logger.record_measurements(measurements)
|
||||
|
||||
# Produce all the features.
|
||||
current_features = dict(measurements)
|
||||
for feature_producer in self._features_producers:
|
||||
required_features = {
|
||||
key: current_features[key]
|
||||
for key in feature_producer.required_features_keys()
|
||||
}
|
||||
current_features.update(
|
||||
feature_producer.produce_features(required_features)
|
||||
)
|
||||
|
||||
# Observe the features.
|
||||
for feature_observer in self._features_observers:
|
||||
feature_observer.observe_features(current_features)
|
||||
|
||||
# Log the resulting features.
|
||||
for logger in self._loggers:
|
||||
logger.record_features(current_features)
|
||||
return current_features
|
||||
|
||||
def compute_final_commands(
|
||||
self,
|
||||
policy_commands: Mapping[str, gdmr_types.ArrayType],
|
||||
) -> Mapping[str, gdmr_types.ArrayType]:
|
||||
"""Processes the policy commands and returns the final processed commands."""
|
||||
current_commands = dict(policy_commands)
|
||||
for processor in self._commands_processors:
|
||||
# Get commands to be consumed by the processor and remove the commands
|
||||
# from the current_commands.. They correspond to the
|
||||
# `consumed_command_spec`.
|
||||
consumed_commands = {
|
||||
key: current_commands.pop(key)
|
||||
for key in processor.consumed_commands_spec().keys()
|
||||
}
|
||||
produced_commands = processor.process_commands(consumed_commands)
|
||||
current_commands.update(produced_commands)
|
||||
|
||||
# Log the modification.
|
||||
for logger in self._loggers:
|
||||
logger.record_commands_processing(
|
||||
processor.name, consumed_commands, produced_commands
|
||||
)
|
||||
|
||||
for logger in self._loggers:
|
||||
logger.record_final_commands(current_commands)
|
||||
return current_commands
|
||||
|
||||
def compute_reward(
|
||||
self, features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> tree.Structure[gdmr_types.ArrayType]:
|
||||
"""Computes the reward given the features."""
|
||||
return self._reward_provider.compute_reward({
|
||||
key: features[key]
|
||||
for key in self._reward_provider.required_features_keys()
|
||||
})
|
||||
|
||||
def check_for_termination(
|
||||
self, features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> reaf_termination_checker.TerminationResult:
|
||||
"""Checks for termination."""
|
||||
current_state = reaf_termination_checker.TerminationResult.DO_NOT_TERMINATE
|
||||
for termination_checker in self._termination_checkers:
|
||||
current_state = reaf_termination_checker.TerminationResult.combine(
|
||||
current_state,
|
||||
termination_checker.check_termination({
|
||||
key: features[key]
|
||||
for key in termination_checker.required_features_keys()
|
||||
}),
|
||||
)
|
||||
return current_state
|
||||
|
||||
def compute_discount(
|
||||
self,
|
||||
features: Mapping[str, gdmr_types.ArrayType],
|
||||
termination_state: reaf_termination_checker.TerminationResult,
|
||||
) -> tree.Structure[gdmr_types.ArrayType]:
|
||||
"""Computes the discount given the features and termination state."""
|
||||
return self._discount_provider.compute_discount(
|
||||
{
|
||||
key: features[key]
|
||||
for key in self._discount_provider.required_features_keys()
|
||||
},
|
||||
termination_state,
|
||||
)
|
||||
|
||||
def add_logger(self, logger: reaf_logger.Logger) -> None:
|
||||
self._loggers.append(logger)
|
||||
|
||||
def remove_logger(self, logger: reaf_logger.Logger) -> None:
|
||||
self._loggers.remove(logger)
|
||||
|
||||
def _validate_features_spec(
|
||||
self, dacl_measurements_spec: Mapping[str, specs.Array]
|
||||
) -> None:
|
||||
"""Validates the features spec."""
|
||||
# Check measurements/features path.
|
||||
current_key_set = set(dacl_measurements_spec.keys())
|
||||
logging.vlog(4, "DACL measurements keys: %s", current_key_set)
|
||||
|
||||
for producer in self._features_producers:
|
||||
logging.vlog(
|
||||
4,
|
||||
"Producer %s requires %s.",
|
||||
producer.name,
|
||||
producer.required_features_keys(),
|
||||
)
|
||||
# Check required features are available.
|
||||
if not producer.required_features_keys().issubset(current_key_set):
|
||||
raise ValueError(
|
||||
"Failed to validate feature specs for feature producer"
|
||||
f" {producer.name}. Missing keys:"
|
||||
f" {producer.required_features_keys() - current_key_set}"
|
||||
)
|
||||
# Check that there are not duplicates in the output.
|
||||
if not current_key_set.isdisjoint(
|
||||
producer.produced_features_spec().keys()
|
||||
):
|
||||
raise ValueError(
|
||||
"Failed to validate feature specs for feature producer"
|
||||
f" {producer.name}. Duplicate keys:"
|
||||
f" {current_key_set & producer.produced_features_spec().keys()}"
|
||||
)
|
||||
# Now extend the spec.
|
||||
logging.vlog(
|
||||
4,
|
||||
"Update available keys (from producer %s) with %s.",
|
||||
producer.name,
|
||||
producer.produced_features_spec().keys(),
|
||||
)
|
||||
current_key_set.update(producer.produced_features_spec().keys())
|
||||
logging.vlog(4, "Available features keys %s.", current_key_set)
|
||||
|
||||
def _validate_commands_spec(
|
||||
self, dacl_commands_spec: Mapping[str, gdmr_types.AnyArraySpec]
|
||||
) -> None:
|
||||
"""Validates the commands spec."""
|
||||
# Check commands. Starting from the DACL command specs we propagate up in
|
||||
# the chain.
|
||||
logging.vlog(3, "Validate commands processing from DACL to Policy.")
|
||||
current_key_set = set(dacl_commands_spec.keys())
|
||||
logging.vlog(4, "DACL commands keys: %s", current_key_set)
|
||||
|
||||
for processor in reversed(self._commands_processors):
|
||||
produced_command_keys = processor.produced_commands_keys()
|
||||
|
||||
logging.vlog(
|
||||
4,
|
||||
"Processor %s: specs (accepted keys) %s. Exposes %s.",
|
||||
processor.name,
|
||||
processor.consumed_commands_spec().keys(),
|
||||
produced_command_keys,
|
||||
)
|
||||
if not produced_command_keys.issubset(current_key_set):
|
||||
raise ValueError(
|
||||
"Failed to validate commands specs for commands processor"
|
||||
f" {processor.name}. Missing (consumable) keys:"
|
||||
f" {produced_command_keys - current_key_set}"
|
||||
)
|
||||
# Remove the produced keys and add the consumed commands specs (as the
|
||||
# processor is mutable).
|
||||
current_key_set = current_key_set - produced_command_keys
|
||||
current_key_set.update(processor.consumed_commands_spec().keys())
|
||||
@@ -0,0 +1,94 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Checks if the episode should terminate."""
|
||||
|
||||
import abc
|
||||
from collections.abc import Mapping
|
||||
import enum
|
||||
from typing import Self
|
||||
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
|
||||
|
||||
class TerminationResult(enum.IntFlag):
|
||||
"""The result of an episode termination check.
|
||||
|
||||
The TerminationResult refers to the possibility for an episode to terminate.
|
||||
For more details on the concept of termination we refer the readers to
|
||||
https://github.com/google-deepmind/dm_env/blob/master/docs/index.md#environment-api-and-semantics.
|
||||
|
||||
Note that this enum does not refer to the possible causes of termination but
|
||||
only how the termination impacts the learning process.
|
||||
|
||||
The result can be one of the following options:
|
||||
- DO_NOT_TERMINATE: The episode should not terminate.
|
||||
- TRUNCATE: The epsisode should terminate. Truncation implies a non-failure
|
||||
final state. Usually this is associated with a non-zero discount.
|
||||
- TERMINATE: The episode should terminate as the environment is in some
|
||||
final state. Usually this is associated with a zero discount for e.g.
|
||||
finite-horizon RL.
|
||||
"""
|
||||
|
||||
DO_NOT_TERMINATE = 0
|
||||
TRUNCATE = 2**0
|
||||
TERMINATE = 2**1
|
||||
|
||||
def is_terminated(self) -> bool:
|
||||
return self == TerminationResult.TERMINATE
|
||||
|
||||
def is_truncated(self) -> bool:
|
||||
return self == TerminationResult.TRUNCATE
|
||||
|
||||
def combine(self, other: Self) -> Self:
|
||||
# TERMINATE has precedence over TRUNCATE, which in turn has precedence over
|
||||
# DO_NOT_TERMINATE. Given the definitions above, this can be implemented as
|
||||
# a maximum operator. To also enable tracing with JAX, we implement this in
|
||||
# a branchless manner using bitwise operations that preserve the type.
|
||||
# Note that JAX will trace TerminationResult values as ints.
|
||||
# Approach:
|
||||
# - self ^ (self ^ other) == other
|
||||
# - (-1 * (self < other)) will be bitmask of all 1s iff self < other.
|
||||
# - AND with (self ^ other) will result in either update or no-op bitmask.
|
||||
return self ^ ((self ^ other) & (-1 * (self < other)))
|
||||
|
||||
|
||||
class TerminationChecker(abc.ABC):
|
||||
"""Checks if the episode should terminate."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Returns a unique string identifier for this object."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def check_termination(
|
||||
self, required_features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> TerminationResult:
|
||||
"""Checks if the episode should terminate.
|
||||
|
||||
Args:
|
||||
required_features: Measurements and features computed by the task logic
|
||||
that are required by this checker, i.e. that have keys specified by
|
||||
`required_features_keys`.
|
||||
|
||||
Returns if the episode should terminate (and if so, what kind of
|
||||
termination).
|
||||
"""
|
||||
|
||||
@abc.abstractmethod
|
||||
def required_features_keys(self) -> set[str]:
|
||||
"""Returns the feature keys that are required to check the termination."""
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Resets the internal state of the termination checker."""
|
||||
...
|
||||
@@ -0,0 +1,29 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Defines an event-based waiting behaviour."""
|
||||
|
||||
import abc
|
||||
|
||||
|
||||
class Trigger(abc.ABC):
|
||||
"""Defines an event-based waiting behaviour."""
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Returns the name of the trigger."""
|
||||
|
||||
@abc.abstractmethod
|
||||
def wait_for_event(self) -> None:
|
||||
"""Blocks until the next event."""
|
||||
@@ -0,0 +1,48 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# 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
|
||||
#
|
||||
# https://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.
|
||||
"""Reward provider which provides a zero reward."""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dm_env import specs
|
||||
from gdm_robotics.interfaces import types as gdmr_types
|
||||
import numpy as np
|
||||
from reaf.core import reward_provider
|
||||
import tree
|
||||
|
||||
|
||||
class ZeroRewardProvider(reward_provider.RewardProvider):
|
||||
"""Reward provider which provides a zero reward."""
|
||||
|
||||
def __init__(self, name: str = 'zero_reward_provider'):
|
||||
self._name = name
|
||||
|
||||
def name(self) -> str:
|
||||
return self._name
|
||||
|
||||
def compute_reward(
|
||||
self, required_features: Mapping[str, gdmr_types.ArrayType]
|
||||
) -> tree.Structure[gdmr_types.ArrayType]:
|
||||
"""Returns a zero reward."""
|
||||
return np.zeros(1)
|
||||
|
||||
def reward_spec(self) -> tree.Structure[specs.Array]:
|
||||
"""Returns the spec for a constant zero reward."""
|
||||
return specs.Array(shape=(1,), dtype=float)
|
||||
|
||||
def required_features_keys(self) -> set[str]:
|
||||
"""Returns empty set.
|
||||
|
||||
There are no feature keys that are required to compute the reward.
|
||||
"""
|
||||
return set()
|
||||
Reference in New Issue
Block a user