diff --git a/python/LQR.ipynb b/python/LQR.ipynb
index c185b044..2b40aa98 100644
--- a/python/LQR.ipynb
+++ b/python/LQR.ipynb
@@ -244,7 +244,7 @@
" if len(frames) \u003c data.time * FRAMERATE:\n",
" renderer.update_scene(data)\n",
" pixels = renderer.render()\n",
- " frames.append(pixels.copy())\n",
+ " frames.append(pixels)\n",
"\n",
"# Display video.\n",
"media.show_video(frames, fps=FRAMERATE)"
@@ -294,7 +294,7 @@
"\n",
" renderer.update_scene(data, camera)\n",
" pixels = renderer.render()\n",
- " frames.append(pixels.copy())\n",
+ " frames.append(pixels)\n",
"\n",
"media.show_video(frames, fps=FRAMERATE)"
]
@@ -565,7 +565,7 @@
" camera.lookat = data.body('torso').subtree_com\n",
" renderer.update_scene(data, camera)\n",
" pixels = renderer.render()\n",
- " frames.append(pixels.copy())\n",
+ " frames.append(pixels)\n",
"\n",
"media.show_video(frames, fps=FRAMERATE)"
]
@@ -827,7 +827,7 @@
" if len(frames) \u003c data.time * FRAMERATE:\n",
" renderer.update_scene(data)\n",
" pixels = renderer.render()\n",
- " frames.append(pixels.copy())\n",
+ " frames.append(pixels)\n",
"\n",
"media.show_video(frames, fps=FRAMERATE)"
]
@@ -921,7 +921,7 @@
" camera.azimuth = azimuth(data.time)\n",
" renderer.update_scene(data, camera, scene_option)\n",
" pixels = renderer.render()\n",
- " frames.append(pixels.copy())\n",
+ " frames.append(pixels)\n",
"\n",
"media.show_video(frames, fps=FRAMERATE)"
]
diff --git a/python/mujoco/renderer.py b/python/mujoco/renderer.py
index 36c2f233..d541c22f 100644
--- a/python/mujoco/renderer.py
+++ b/python/mujoco/renderer.py
@@ -76,10 +76,6 @@ the clause:
self._rect = _render.MjrRect(0, 0, self._width, self._height)
- # Internal buffers.
- self._rgb_buffer = np.empty((self._height, self._width, 3), dtype=np.uint8)
- self._depth_buffer = np.empty((self._height, self._width), dtype=np.float32)
-
# Create render contexts.
self._gl_context = gl_context.GLContext(width, height)
self._gl_context.make_current()
@@ -124,12 +120,20 @@ the clause:
def disable_segmentation_rendering(self):
self._segmentation_rendering = False
- def render(self) -> np.ndarray:
+ def render(self, *, out: Optional[np.ndarray] = None) -> np.ndarray:
"""Renders the scene as a numpy array of pixel values.
+ Args:
+ out: Alternative output array in which to place the resulting pixels. It
+ must have the same shape as the expected output but the type will be
+ cast if necessary. The expted shape depends on the value of
+ `self._depth_rendering`: when `True`, we expect `out.shape == (width,
+ height)`, and `out.shape == (width, height, 3)` when `False`.
+
Returns:
- A numpy array of pixels with dimensions (H, W, 3). The array will be
- mutated by future calls to `render`.
+ A new numpy array holding the pixels with shape `(H, W)` or `(H, W, 3)`,
+ depending on the value of `self._depth_rendering` unless
+ `out is None`, in which case a reference to `out` is returned.
"""
original_flags = self._scene.flags.copy()
@@ -139,12 +143,30 @@ the clause:
self._gl_context.make_current()
+ if self._depth_rendering:
+ out_shape = (self._height, self._width)
+ out_dtype = np.float32
+ else:
+ out_shape = (self._height, self._width, 3)
+ out_dtype = np.uint8
+
+ if out is None:
+ out = np.empty(out_shape, dtype=out_dtype)
+ else:
+ if out.shape != out_shape:
+ raise ValueError(
+ f'Expected `out.shape == {out_shape}`. Got `out.shape={out.shape}`'
+ ' instead. When using depth rendering, the out array should be of'
+ ' shape `(width, height)` and otherwise (width, height, 3).'
+ f' Got `(self.height, self.width)={(self.height, self.width)}` and'
+ f' `self._depth_rendering={self._depth_rendering}`.'
+ )
+
# Render scene and read contents of RGB and depth buffers.
_render.mjr_render(self._rect, self._scene, self._mjr_context)
- _render.mjr_readPixels(self._rgb_buffer, self._depth_buffer, self._rect,
- self._mjr_context)
-
if self._depth_rendering:
+ _render.mjr_readPixels(None, out, self._rect, self._mjr_context)
+
# Get the distances to the near and far clipping planes.
extent = self._model.stat.extent
near = self._model.vis.map.znear * extent
@@ -153,32 +175,40 @@ the clause:
# Convert from [0 1] to depth in units of length, see links below:
# http://stackoverflow.com/a/6657284/1461210
# https://www.khronos.org/opengl/wiki/Depth_Buffer_Precision
- pixels = near / (1 - self._depth_buffer * (1 - near / far))
+ out = near / (1 - out * (1 - near / far))
elif self._segmentation_rendering:
+ _render.mjr_readPixels(out, None, self._rect, self._mjr_context)
+
# Convert 3-channel uint8 to 1-channel uint32.
- image3 = self._rgb_buffer.astype(np.uint32)
- segimage = (image3[:, :, 0] +
- image3[:, :, 1] * (2**8) +
- image3[:, :, 2] * (2**16))
+ image3 = out.astype(np.uint32)
+ segimage = (
+ image3[:, :, 0]
+ + image3[:, :, 1] * (2**8)
+ + image3[:, :, 2] * (2**16)
+ )
# Remap segid to 2-channel (object ID, object type) pair.
# Seg ID 0 is background -- will be remapped to (-1, -1).
ngeoms = self._scene.ngeom
- segid2output = np.full((ngeoms + 1, 2), fill_value=-1,
- dtype=np.int32) # Seg id cannot be > ngeom + 1.
+ segid2output = np.full(
+ (ngeoms + 1, 2), fill_value=-1, dtype=np.int32
+ ) # Seg id cannot be > ngeom + 1.
visible_geoms = [g for g in self._scene.geoms[:ngeoms] if g.segid != -1]
visible_segids = np.array([g.segid + 1 for g in visible_geoms], np.int32)
visible_objid = np.array([g.objid for g in visible_geoms], np.int32)
visible_objtype = np.array([g.objtype for g in visible_geoms], np.int32)
segid2output[visible_segids, 0] = visible_objid
segid2output[visible_segids, 1] = visible_objtype
- pixels = segid2output[segimage]
+ out = segid2output[segimage]
# Reset scene flags.
np.copyto(self._scene.flags, original_flags)
else:
- pixels = self._rgb_buffer
- return np.flipud(pixels)
+ _render.mjr_readPixels(out, None, self._rect, self._mjr_context)
+
+ out[:] = np.flipud(out)
+
+ return out
def update_scene(
self,
diff --git a/python/mujoco/renderer_test.py b/python/mujoco/renderer_test.py
index 295de178..f8ec0907 100644
--- a/python/mujoco/renderer_test.py
+++ b/python/mujoco/renderer_test.py
@@ -17,6 +17,7 @@
from absl.testing import absltest
from absl.testing import parameterized
import mujoco
+import numpy as np
@absltest.skipUnless(hasattr(mujoco, 'GLContext'),
@@ -46,5 +47,71 @@ class MuJoCoRendererTest(parameterized.TestCase):
not_all_black = True
break
self.assertTrue(not_all_black)
+
+ def test_renderer_output_without_out(self):
+ xml = """
+
+
+
+
+
+
+"""
+ model = mujoco.MjModel.from_xml_string(xml)
+ data = mujoco.MjData(model)
+ mujoco.mj_forward(model, data)
+ renderer = mujoco.Renderer(model, 50, 50)
+ renderer.update_scene(data, 'closeup')
+ pixels = [renderer.render()]
+
+ colors = (
+ (1.0, 0.0, 0.0, 1.0),
+ (0.0, 1.0, 0.0, 1.0),
+ (0.0, 0.0, 1.0, 1.0),
+ )
+
+ for i, color in enumerate(colors):
+ model.geom_rgba[0, :] = color
+ mujoco.mj_forward(model, data)
+ renderer.update_scene(data, 'closeup')
+ pixels.append(renderer.render())
+ self.assertIsNot(pixels[-2], pixels[-1])
+
+ # Pixels should change over steps.
+ self.assertFalse((pixels[i + 1] == pixels[i]).all())
+
+ def test_renderer_output_with_out(self):
+ xml = """
+
+
+
+
+
+
+"""
+ render_size = (50, 50)
+ render_out = np.zeros((*render_size, 3), np.uint8)
+ model = mujoco.MjModel.from_xml_string(xml)
+ data = mujoco.MjData(model)
+ mujoco.mj_forward(model, data)
+ renderer = mujoco.Renderer(model, *render_size)
+ renderer.update_scene(data, 'closeup')
+
+ self.assertTrue(np.all(render_out == 0))
+
+ pixels = renderer.render(out=render_out)
+
+ # Pixels should always refer to the same `render_out` array.
+ self.assertIs(pixels, render_out)
+ self.assertFalse(np.all(render_out == 0))
+
+ failing_render_size = (10, 10)
+ self.assertNotEqual(failing_render_size, render_size)
+ with self.assertRaises(ValueError):
+ pixels = renderer.render(
+ out=np.zeros((*failing_render_size, 3), np.uint8)
+ )
+
+
if __name__ == '__main__':
absltest.main()
diff --git a/python/tutorial.ipynb b/python/tutorial.ipynb
index cb74b4a0..8c0bbb69 100644
--- a/python/tutorial.ipynb
+++ b/python/tutorial.ipynb
@@ -561,7 +561,7 @@
" mujoco.mj_step(model, data)\n",
" if len(frames) \u003c data.time * framerate:\n",
" renderer.update_scene(data)\n",
- " pixels = renderer.render().copy()\n",
+ " pixels = renderer.render()\n",
" frames.append(pixels)\n",
"media.show_video(frames, fps=framerate)"
]
@@ -614,7 +614,7 @@
" mujoco.mj_step(model, data)\n",
" if len(frames) \u003c data.time * framerate:\n",
" renderer.update_scene(data, scene_option=scene_option)\n",
- " pixels = renderer.render().copy()\n",
+ " pixels = renderer.render()\n",
" frames.append(pixels)\n",
"\n",
"# Simulate and display video.\n",
@@ -670,7 +670,7 @@
" mujoco.mj_step(model, data)\n",
" if len(frames) \u003c data.time * framerate:\n",
" renderer.update_scene(data, scene_option=scene_option)\n",
- " pixels = renderer.render().copy()\n",
+ " pixels = renderer.render()\n",
" frames.append(pixels)\n",
"\n",
"media.show_video(frames, fps=60)"
@@ -838,7 +838,7 @@
" mujoco.mj_step(model, data)\n",
" if len(frames) \u003c data.time * framerate:\n",
" renderer.update_scene(data, \"closeup\")\n",
- " pixels = renderer.render().copy()\n",
+ " pixels = renderer.render()\n",
" frames.append(pixels)\n",
"\n",
"media.show_video(frames, fps=framerate)"
@@ -1004,7 +1004,7 @@
" renderer.update_scene(data, \"fixed\")\n",
" frame = renderer.render()\n",
" render_time += time.time() - tic\n",
- " frames.append(frame.copy())\n",
+ " frames.append(frame)\n",
"\n",
"# print timing and play video\n",
"step_time = 1e6*sim_time/n_steps\n",
@@ -1299,7 +1299,7 @@
" mujoco.mj_step(model, data)\n",
" renderer.update_scene(data, \"track\", options)\n",
" frame = renderer.render()\n",
- " frames.append(frame.copy())\n",
+ " frames.append(frame)\n",
"\n",
"# show video\n",
"media.show_video(frames, fps=30)"
@@ -1462,7 +1462,7 @@
" mujoco.mj_step(model, data)\n",
" renderer.update_scene(data, \"y\")\n",
" frame = renderer.render()\n",
- " frames.append(frame.copy())\n",
+ " frames.append(frame)\n",
"media.show_video(frames, fps=30)"
]
},
@@ -1577,7 +1577,7 @@
" sensordata.append(data.sensor('accelerometer').data.copy())\n",
" renderer.update_scene(data, \"fixed\")\n",
" frame = renderer.render()\n",
- " frames.append(frame.copy())\n",
+ " frames.append(frame)\n",
"\n",
"media.show_video(frames, fps=fps)"
]
@@ -1921,7 +1921,7 @@
" if len(frames) \u003c data.time * framerate:\n",
" renderer.update_scene(data)\n",
" modify_scene(renderer.scene)\n",
- " pixels = renderer.render().copy()\n",
+ " pixels = renderer.render()\n",
" frames.append(pixels)\n",
"media.show_video(frames, fps=framerate)"
]