85ab6183bc
PiperOrigin-RevId: 597316447 Change-Id: If9619320f0024c486bb7abe2c63e9a1ba1a631f0
353 lines
12 KiB
Python
353 lines
12 KiB
Python
# Copyright 2023 DeepMind Technologies Limited
|
|
#
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
# you may not use this file except in compliance with the License.
|
|
# You may obtain a copy of the License at
|
|
#
|
|
# http://www.apache.org/licenses/LICENSE-2.0
|
|
#
|
|
# Unless required by applicable law or agreed to in writing, software
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
# See the License for the specific language governing permissions and
|
|
# limitations under the License.
|
|
# ==============================================================================
|
|
"""Functions to initialize, load, or save data."""
|
|
|
|
import copy
|
|
from typing import List, Union
|
|
|
|
import jax
|
|
from jax import numpy as jp
|
|
import mujoco
|
|
from mujoco.mjx._src import collision_driver
|
|
from mujoco.mjx._src import constraint
|
|
from mujoco.mjx._src import mesh
|
|
from mujoco.mjx._src import types
|
|
import numpy as np
|
|
|
|
|
|
def _put_option(o: mujoco.MjOption, device=None) -> types.Option:
|
|
"""Puts mujoco.MjOption onto a device, resulting in mjx.Option."""
|
|
if o.integrator not in set(types.IntegratorType):
|
|
raise NotImplementedError(f'{mujoco.mjtIntegrator(o.integrator)}')
|
|
|
|
if o.cone not in set(types.ConeType):
|
|
raise NotImplementedError(f'{mujoco.mjtCone(o.cone)}')
|
|
|
|
if o.solver not in set(types.SolverType):
|
|
raise NotImplementedError(f'{mujoco.mjtSolver(o.solver)}')
|
|
|
|
for i in range(mujoco.mjtEnableBit.mjNENABLE):
|
|
if o.enableflags & 2**i:
|
|
raise NotImplementedError(f'{mujoco.mjtEnableBit(2 ** i)}')
|
|
|
|
static_fields = {
|
|
f.name: copy.copy(getattr(o, f.name))
|
|
for f in types.Option.fields()
|
|
if f.type in (int, bytes, np.ndarray)
|
|
}
|
|
static_fields['integrator'] = types.IntegratorType(o.integrator)
|
|
static_fields['cone'] = types.ConeType(o.cone)
|
|
static_fields['solver'] = types.SolverType(o.solver)
|
|
static_fields['disableflags'] = types.DisableBit(o.disableflags)
|
|
|
|
device_fields = {
|
|
f.name: copy.copy(getattr(o, f.name))
|
|
for f in types.Option.fields()
|
|
if f.type is jax.Array
|
|
}
|
|
device_fields = jax.device_put(device_fields, device=device)
|
|
|
|
has_fluid_params = o.density > 0 or o.viscosity > 0 or o.wind.any()
|
|
|
|
return types.Option(
|
|
has_fluid_params=has_fluid_params,
|
|
**static_fields,
|
|
**device_fields,
|
|
)
|
|
|
|
|
|
def _put_statistic(s: mujoco.MjStatistic, device=None) -> types.Statistic:
|
|
"""Puts mujoco.MjStatistic onto a device, resulting in mjx.Statistic."""
|
|
return types.Statistic(
|
|
meaninertia=jax.device_put(s.meaninertia, device=device)
|
|
)
|
|
|
|
|
|
def put_model(m: mujoco.MjModel, device=None) -> types.Model:
|
|
"""Puts mujoco.MjModel onto a device, resulting in mjx.Model."""
|
|
|
|
if m.ntendon:
|
|
raise NotImplementedError('tendons are not supported')
|
|
|
|
if (m.geom_condim != 3).any() or (m.pair_dim != 3).any():
|
|
raise NotImplementedError('only condim=3 is supported')
|
|
|
|
# check collision geom types
|
|
for g1, g2, *_ in collision_driver.collision_candidates(m):
|
|
if collision_driver.get_collision_fn((g1, g2)) is None:
|
|
g1, g2 = mujoco.mjtGeom(g1), mujoco.mjtGeom(g2)
|
|
raise NotImplementedError(f'({g1}, {g2}) has no collision function')
|
|
|
|
for enum_field, enum_type, mj_type in (
|
|
(m.actuator_biastype, types.BiasType, mujoco.mjtBias),
|
|
(m.actuator_dyntype, types.DynType, mujoco.mjtDyn),
|
|
(m.actuator_gaintype, types.GainType, mujoco.mjtGain),
|
|
(m.actuator_trntype, types.TrnType, mujoco.mjtTrn),
|
|
(m.eq_type, types.EqType, mujoco.mjtEq),
|
|
):
|
|
missing = set(enum_field) - set(enum_type)
|
|
if missing:
|
|
raise NotImplementedError(
|
|
f'{[mj_type(m) for m in missing]} not supported'
|
|
)
|
|
|
|
opt = _put_option(m.opt, device=device)
|
|
stat = _put_statistic(m.stat, device=device)
|
|
|
|
static_fields = {
|
|
f.name: getattr(m, f.name)
|
|
for f in types.Model.fields()
|
|
if f.type in (int, bytes, np.ndarray)
|
|
}
|
|
static_fields['geom_rgba'] = static_fields['geom_rgba'].reshape((-1, 4))
|
|
static_fields['mat_rgba'] = static_fields['mat_rgba'].reshape((-1, 4))
|
|
|
|
device_fields = {
|
|
f.name: copy.copy(getattr(m, f.name)) # copy because device_put is async
|
|
for f in types.Model.fields()
|
|
if f.type is jax.Array
|
|
}
|
|
device_fields.update(mesh.get(m))
|
|
device_fields = jax.device_put(device_fields, device=device)
|
|
|
|
return types.Model(
|
|
opt=opt,
|
|
stat=stat,
|
|
**static_fields,
|
|
**device_fields,
|
|
)
|
|
|
|
|
|
def make_data(m: Union[types.Model, mujoco.MjModel]) -> types.Data:
|
|
"""Allocate and initialize Data."""
|
|
|
|
ncon = collision_driver.ncon(m)
|
|
ne, nf, nl, nc = constraint.count_constraints(m)
|
|
nefc = ne + nf + nl + nc
|
|
|
|
zero_nv = jp.zeros(m.nv, dtype=jp.float32)
|
|
zero_nv_6 = jp.zeros((m.nv, 6), dtype=jp.float32)
|
|
zero_nbody_3 = jp.zeros((m.nbody, 3), dtype=jp.float32)
|
|
zero_nbody_6 = jp.zeros((m.nbody, 6), dtype=jp.float32)
|
|
zero_nbody_10 = jp.zeros((m.nbody, 10), dtype=jp.float32)
|
|
zero_nbody_3_3 = jp.zeros((m.nbody, 3, 3), dtype=jp.float32)
|
|
zero_nefc = jp.zeros(nefc, dtype=jp.float32)
|
|
zero_na = jp.zeros(m.na, dtype=jp.float32)
|
|
zero_nu = jp.zeros(m.nu, dtype=jp.float32)
|
|
zero_njnt_3 = jp.zeros((m.njnt, 3), dtype=jp.float32)
|
|
zero_nm = jp.zeros(m.nM, dtype=jp.float32)
|
|
|
|
# create first d to get num contacts and nc
|
|
d = types.Data(
|
|
solver_niter=jp.array(0, dtype=jp.int32),
|
|
time=jp.array(0.0),
|
|
qpos=jp.array(m.qpos0),
|
|
qvel=zero_nv,
|
|
act=zero_na,
|
|
qacc_warmstart=zero_nv,
|
|
ctrl=zero_nu,
|
|
qfrc_applied=zero_nv,
|
|
xfrc_applied=zero_nbody_6,
|
|
eq_active=jp.zeros(m.neq, dtype=jp.int32),
|
|
qacc=zero_nv,
|
|
act_dot=zero_na,
|
|
xpos=zero_nbody_3,
|
|
xquat=jp.zeros((m.nbody, 4), dtype=jp.float32),
|
|
xmat=zero_nbody_3_3,
|
|
xipos=zero_nbody_3,
|
|
ximat=zero_nbody_3_3,
|
|
xanchor=zero_njnt_3,
|
|
xaxis=zero_njnt_3,
|
|
geom_xpos=jp.zeros((m.ngeom, 3), dtype=jp.float32),
|
|
geom_xmat=jp.zeros((m.ngeom, 3, 3), dtype=jp.float32),
|
|
site_xpos=jp.zeros((m.nsite, 3), dtype=jp.float32),
|
|
site_xmat=jp.zeros((m.nsite, 3, 3), dtype=jp.float32),
|
|
subtree_com=zero_nbody_3,
|
|
cdof=zero_nv_6,
|
|
cinert=zero_nbody_10,
|
|
actuator_length=zero_nu,
|
|
actuator_moment=jp.zeros((m.nu, m.nv), dtype=jp.float32),
|
|
crb=zero_nbody_10,
|
|
qM=zero_nm,
|
|
qLD=zero_nm,
|
|
qLDiagInv=zero_nv,
|
|
qLDiagSqrtInv=zero_nv,
|
|
contact=types.Contact.zero(ncon),
|
|
efc_J=jp.zeros((nefc, m.nv), dtype=jp.float32),
|
|
efc_frictionloss=zero_nefc,
|
|
efc_D=zero_nefc,
|
|
actuator_velocity=zero_nu,
|
|
cvel=zero_nbody_6,
|
|
cdof_dot=zero_nv_6,
|
|
qfrc_bias=zero_nv,
|
|
qfrc_passive=zero_nv,
|
|
efc_aref=zero_nefc,
|
|
qfrc_actuator=zero_nv,
|
|
qfrc_smooth=zero_nv,
|
|
qacc_smooth=zero_nv,
|
|
qfrc_constraint=zero_nv,
|
|
qfrc_inverse=zero_nv,
|
|
efc_force=zero_nefc,
|
|
)
|
|
|
|
return d
|
|
|
|
|
|
def _get_contact(
|
|
c: mujoco._structs._MjContactList,
|
|
cx: types.Contact,
|
|
efc_start: int,
|
|
):
|
|
"""Converts mjx.Contact to mujoco._structs._MjContactList."""
|
|
con_id = np.nonzero(cx.dist <= 0)[0]
|
|
for field in types.Contact.fields():
|
|
value = getattr(cx, field.name)[con_id]
|
|
if field.name == 'frame':
|
|
value = value.reshape((-1, 9))
|
|
getattr(c, field.name)[:] = value
|
|
|
|
ncon = cx.dist.shape[0]
|
|
c.efc_address[:] = np.arange(efc_start, efc_start + ncon * 4, 4)[con_id]
|
|
|
|
|
|
def get_data(
|
|
m: mujoco.MjModel, d: types.Data
|
|
) -> Union[mujoco.MjData, List[mujoco.MjData]]:
|
|
"""Gets mjx.Data from a device, resulting in mujoco.MjData or List[MjData]."""
|
|
dx = jax.device_get(d)
|
|
batched = len(d.qpos.shape) > 1
|
|
batch_size = d.qpos.shape[0] if batched else 1
|
|
ne, nf, nl, nc = constraint.count_constraints(m)
|
|
efc_type = np.array([
|
|
mujoco.mjtConstraint.mjCNSTR_EQUALITY,
|
|
mujoco.mjtConstraint.mjCNSTR_FRICTION_DOF,
|
|
mujoco.mjtConstraint.mjCNSTR_LIMIT_JOINT,
|
|
mujoco.mjtConstraint.mjCNSTR_CONTACT_PYRAMIDAL,
|
|
]).repeat([ne, nf, nl, nc])
|
|
|
|
ds = []
|
|
for i in range(batch_size):
|
|
dx_i = jax.tree_map(lambda x, i=i: x[i], dx) if batched else d
|
|
ncon = (dx_i.contact.dist <= 0).sum()
|
|
efc_active = (dx_i.efc_J != 0).any(axis=1)
|
|
efc_con = efc_type == mujoco.mjtConstraint.mjCNSTR_CONTACT_PYRAMIDAL
|
|
nefc, nc = efc_active.sum(), (efc_active & efc_con).sum()
|
|
d_i = mujoco.MjData(m)
|
|
d_i.nnzJ = nefc * m.nv
|
|
mujoco._functions._realloc_con_efc(d_i, ncon=ncon, nefc=nefc) # pylint: disable=protected-access
|
|
d_i.efc_J_rownnz[:] = np.repeat(m.nv, nefc)
|
|
d_i.efc_J_rowadr[:] = np.arange(0, nefc * m.nv, m.nv)
|
|
d_i.efc_J_colind[:] = np.tile(np.arange(m.nv), nefc)
|
|
|
|
for field in types.Data.fields():
|
|
if field.name == 'contact':
|
|
_get_contact(d_i.contact, dx_i.contact, nefc - nc)
|
|
continue
|
|
|
|
value = getattr(dx_i, field.name)
|
|
|
|
if field.name in ('xmat', 'ximat', 'geom_xmat', 'site_xmat'):
|
|
value = value.reshape((-1, 9))
|
|
|
|
if field.name in ('efc_frictionloss', 'efc_D', 'efc_aref', 'efc_force'):
|
|
value = value[efc_active]
|
|
|
|
if field.name == 'efc_J':
|
|
value = value[efc_active].reshape(-1)
|
|
|
|
if value.shape:
|
|
getattr(d_i, field.name)[:] = value
|
|
else:
|
|
setattr(d_i, field.name, value)
|
|
|
|
d_i.efc_type[:] = efc_type[efc_active]
|
|
ds.append(d_i)
|
|
|
|
return ds if batched else ds[0]
|
|
|
|
|
|
def _put_contact(
|
|
c: mujoco._structs._MjContactList, ncon: int, device=None
|
|
) -> types.Contact:
|
|
"""Puts mujoco.structs._MjContactList onto a device, resulting in mjx.Contact."""
|
|
fields = {
|
|
f.name: copy.copy(getattr(c, f.name)) for f in types.Contact.fields()
|
|
}
|
|
fields['frame'] = fields['frame'].reshape((-1, 3, 3))
|
|
pad_size = ncon - c.dist.shape[0]
|
|
pad_fn = lambda x: np.concatenate(
|
|
(x, np.zeros((pad_size,) + x.shape[1:], dtype=x.dtype))
|
|
)
|
|
fields = jax.tree_map(pad_fn, fields)
|
|
fields['dist'][-pad_size:] = np.inf
|
|
fields = jax.device_put(fields, device=device)
|
|
|
|
return types.Contact(**fields)
|
|
|
|
|
|
def put_data(m: mujoco.MjModel, d: mujoco.MjData, device=None) -> types.Data:
|
|
"""Puts mujoco.MjData onto a device, resulting in mjx.Data."""
|
|
ncon = collision_driver.ncon(m)
|
|
ne, nf, nl, nc = constraint.count_constraints(m)
|
|
nefc = ne + nf + nl + nc
|
|
|
|
for d_val, val, name in (
|
|
(d.ncon, ncon, 'ncon'),
|
|
(d.ne, ne, 'ne'),
|
|
(d.nf, nf, 'nf'),
|
|
(d.nl, nl, 'nl'),
|
|
(d.nefc, nefc, 'nefc'),
|
|
):
|
|
if d_val > val:
|
|
raise ValueError(f'd.{name} too high, d.{name} = {d_val}, model = {val}')
|
|
|
|
fields = {
|
|
f.name: copy.copy(getattr(d, f.name)) # copy because device_put is async
|
|
for f in types.Data.fields()
|
|
if f.type is jax.Array
|
|
}
|
|
|
|
for fname in ('xmat', 'ximat', 'geom_xmat', 'site_xmat'):
|
|
fields[fname] = fields[fname].reshape((-1, 3, 3))
|
|
|
|
# pad efc fields: MuJoCo efc arrays are sparse for inactive constraints.
|
|
# efc_J is also optionally column-sparse (typically for large nv). MJX is
|
|
# neither: it contains zeros for inactive constraints, and efc_J is always
|
|
# (nefc, nv). this may change in the future.
|
|
if mujoco.mj_isSparse(m):
|
|
nr = d.efc_J_rownnz.shape[0]
|
|
efc_j = np.zeros((nr, m.nv))
|
|
for i in range(nr):
|
|
rowadr = d.efc_J_rowadr[i]
|
|
for j in range(d.efc_J_rownnz[i]):
|
|
efc_j[i, d.efc_J_colind[rowadr + j]] = fields['efc_J'][rowadr + j]
|
|
fields['efc_J'] = efc_j
|
|
else:
|
|
fields['efc_J'] = fields['efc_J'].reshape((-1 if m.nv else 0, m.nv))
|
|
|
|
for fname in ('efc_J', 'efc_frictionloss', 'efc_D', 'efc_aref', 'efc_force'):
|
|
value = np.zeros((nefc, m.nv)) if fname == 'efc_J' else np.zeros(nefc)
|
|
for i in range(4):
|
|
value_beg = sum([ne, nf, nl][:i])
|
|
d_beg = sum([d.ne, d.nf, d.nl][:i])
|
|
size = [d.ne, d.nf, d.nl, d.nefc - d.nl - d.nf - d.ne][i]
|
|
value[value_beg:value_beg+size] = fields[fname][d_beg:d_beg+size]
|
|
fields[fname] = value
|
|
|
|
fields = jax.device_put(fields, device=device)
|
|
fields['contact'] = _put_contact(d.contact, ncon, device=device)
|
|
|
|
return types.Data(**fields)
|