0a60d555c6
PiperOrigin-RevId: 627773755 Change-Id: Iebbcaa41cfdde126312e388d368843a4754abfb8
87 lines
2.7 KiB
Python
87 lines
2.7 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.
|
|
# ==============================================================================
|
|
"""Run benchmarks on various devices."""
|
|
|
|
from typing import Sequence
|
|
|
|
from absl import app
|
|
from absl import flags
|
|
from etils import epath
|
|
import mujoco
|
|
from mujoco import mjx
|
|
|
|
_MJCF = flags.DEFINE_string('mjcf', None, 'path to model', required=True)
|
|
_BASE_PATH = flags.DEFINE_string(
|
|
'base_path', None, 'base path, defaults to mujoco.mjx resource path'
|
|
)
|
|
_NSTEP = flags.DEFINE_integer('nstep', 1000, 'number of steps per rollout')
|
|
_BATCH_SIZE = flags.DEFINE_integer(
|
|
'batch_size', 1024, 'number of parallel rollouts'
|
|
)
|
|
_UNROLL = flags.DEFINE_integer('unroll', 1, 'loop unroll length')
|
|
_SOLVER = flags.DEFINE_enum(
|
|
'solver', 'cg', ['cg', 'newton'], 'constraint solver'
|
|
)
|
|
_ITERATIONS = flags.DEFINE_integer(
|
|
'iterations', 1, 'number of solver iterations'
|
|
)
|
|
_LS_ITERATIONS = flags.DEFINE_integer(
|
|
'ls_iterations', 4, 'number of linesearch iterations'
|
|
)
|
|
_OUTPUT = flags.DEFINE_enum(
|
|
'output', 'text', ['text', 'tsv'], 'format to print results'
|
|
)
|
|
|
|
|
|
def _main(argv: Sequence[str]):
|
|
"""Runs testpeed function."""
|
|
path = epath.resource_path('mujoco.mjx') / 'test_data'
|
|
path = _BASE_PATH.value or path
|
|
f = epath.Path(path) / _MJCF.value
|
|
m = mujoco.MjModel.from_xml_path(f.as_posix())
|
|
|
|
print(f'Rolling out {_NSTEP.value} steps at dt = {m.opt.timestep:.3f}...')
|
|
jit_time, run_time, steps = mjx.benchmark(
|
|
m,
|
|
_NSTEP.value,
|
|
_BATCH_SIZE.value,
|
|
_UNROLL.value,
|
|
_SOLVER.value,
|
|
_ITERATIONS.value,
|
|
_LS_ITERATIONS.value,
|
|
)
|
|
|
|
name = argv[0]
|
|
if _OUTPUT.value == 'text':
|
|
print(f"""
|
|
Summary for {_BATCH_SIZE.value} parallel rollouts
|
|
|
|
Total JIT time: {jit_time:.2f} s
|
|
Total simulation time: {run_time:.2f} s
|
|
Total steps per second: { steps / run_time:.0f}
|
|
Total realtime factor: { steps * m.opt.timestep / run_time:.2f} x
|
|
Total time per step: { 1e6 * run_time / steps:.2f} µs""")
|
|
elif _OUTPUT.value == 'tsv':
|
|
name = name.split('/')[-1].replace('testspeed_', '')
|
|
print(f'{name}\tjit: {jit_time:.2f}s\tsteps/second: {steps / run_time:.0f}')
|
|
|
|
|
|
def main():
|
|
app.run(_main)
|
|
|
|
|
|
if __name__ == '__main__':
|
|
main()
|