From 38cbc9e54321f30e504d898657228ed208aa2617 Mon Sep 17 00:00:00 2001 From: Erik Frey Date: Tue, 6 Feb 2024 10:29:29 -0800 Subject: [PATCH] Improve MJX performance on NVIDIA GPU by enabling Triton GEMM. PiperOrigin-RevId: 604691509 Change-Id: Iba11a9e1b79848624a1ef6eb4e61e2093d36e662 --- doc/mjx.rst | 10 ++++++++++ mjx/mujoco/mjx/testspeed.py | 5 +++++ mjx/tutorial.ipynb | 5 +++++ 3 files changed, 20 insertions(+) diff --git a/doc/mjx.rst b/doc/mjx.rst index 50c47389..8dc4bf66 100644 --- a/doc/mjx.rst +++ b/doc/mjx.rst @@ -357,3 +357,13 @@ For MJX to perform well, some configuration parameters should be adjusted from t greater, or if MJX detects a TPU as the default backend, otherwise "dense". For TPU, using "sparse" with the Newton solver can speed up simulation by 2x to 3x. For GPU, choosing "dense" may impart a more modest speedup of 10% to 20%, as long as the dense matrices can fit on the device. + +GPU performance tuning +---------------------- + +The following environment variables should be set: + +``XLA_FLAGS=--xla_gpu_triton_gemm_any=true`` + This enables the Triton-based GEMM (matmul) emitter for any GEMM that it supports. This can yield a 30% speedup on + NVIDIA GPUs. If you have multiple GPUs, you may also benefit from enabling flags related to + `communciation between GPUs `__. diff --git a/mjx/mujoco/mjx/testspeed.py b/mjx/mujoco/mjx/testspeed.py index 66937a8b..d66d475f 100644 --- a/mjx/mujoco/mjx/testspeed.py +++ b/mjx/mujoco/mjx/testspeed.py @@ -14,6 +14,7 @@ # ============================================================================== """Run benchmarks on various devices.""" +import os import time from typing import Sequence, Tuple @@ -56,6 +57,10 @@ def _measure(fn, *args) -> Tuple[float, float]: def _main(argv: Sequence[str]): """Benchmark a model.""" + xla_flags = os.environ.get('XLA_FLAGS', '') + xla_flags += ' --xla_gpu_triton_gemm_any=True' + os.environ['XLA_FLAGS'] = xla_flags + f = epath.resource_path('mujoco.mjx') / 'test_data' / FLAGS.mjcf m = mujoco.MjModel.from_xml_path(f.as_posix()) m.opt.solver = { diff --git a/mjx/tutorial.ipynb b/mjx/tutorial.ipynb index 8644de41..ade2a83b 100644 --- a/mjx/tutorial.ipynb +++ b/mjx/tutorial.ipynb @@ -102,6 +102,11 @@ "}\n", "\"\"\")\n", "\n", + "# Tell XLA to use Triton GEMM, this improves steps/sec by ~30% on some GPUs\n", + "xla_flags = os.environ.get('XLA_FLAGS', '')\n", + "xla_flags += ' --xla_gpu_triton_gemm_any=True'\n", + "os.environ['XLA_FLAGS'] = xla_flags\n", + "\n", "# Configure MuJoCo to use the EGL rendering backend (requires GPU)\n", "print('Setting environment variable to use GPU rendering:')\n", "%env MUJOCO_GL=egl\n",