#mjx Update bind type checks.

PiperOrigin-RevId: 729566363
Change-Id: Ifb4fd3e83988140f781ac1d5db690cf6fbe70fbf
This commit is contained in:
Google DeepMind
2025-02-21 10:01:04 -08:00
committed by Copybara-Service
parent f9bb6cb3db
commit 1532cffdf0
+11 -11
View File
@@ -13,7 +13,7 @@
# limitations under the License.
# ==============================================================================
"""Engine support functions."""
from collections.abc import Sequence
from collections.abc import Iterable, Sequence
from typing import Optional, Tuple, Union
import jax
@@ -293,10 +293,6 @@ class BindModel(object):
def __init__(self, model: Model, specs: Sequence[mujoco.MjStruct]):
self.model = model
self.prefix = ''
try:
iter(specs)
except TypeError:
specs = [specs]
ids = []
for spec in specs:
if isinstance(spec, mujoco.MjsBody):
@@ -383,9 +379,13 @@ class BindModel(object):
def _bind_model(
self: Model, obj: Sequence[mujoco.MjStruct]
self: Model, obj: mujoco.MjStruct | Iterable[mujoco.MjStruct]
) -> BindModel:
"""Bind a Mujoco spec to an MJX Model."""
if isinstance(obj, mujoco.MjStruct):
obj = (obj,)
else:
obj = tuple(obj)
return BindModel(self, obj)
@@ -398,10 +398,6 @@ class BindData(object):
self.data = data
self.model = model
self.prefix = ''
try:
iter(specs)
except TypeError:
specs = [specs]
ids = []
for spec in specs:
if isinstance(spec, mujoco.MjsBody):
@@ -530,9 +526,13 @@ class BindData(object):
def _bind_data(
self: Data, model: Model, obj: Sequence[mujoco.MjStruct]
self: Data, model: Model, obj: mujoco.MjStruct | Iterable[mujoco.MjStruct]
) -> BindData:
"""Bind a Mujoco spec to an MJX Data."""
if isinstance(obj, mujoco.MjStruct):
obj = (obj,)
else:
obj = tuple(obj)
return BindData(self, model, obj)