Improve MJX performance on NVIDIA GPU by enabling Triton GEMM.

PiperOrigin-RevId: 604691509
Change-Id: Iba11a9e1b79848624a1ef6eb4e61e2093d36e662
This commit is contained in:
Erik Frey
2024-02-06 10:29:29 -08:00
committed by Copybara-Service
parent 84f20ffb66
commit 38cbc9e543
3 changed files with 20 additions and 0 deletions
+10
View File
@@ -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 <https://jax.readthedocs.io/en/latest/gpu_performance_tips.html>`__.
+5
View File
@@ -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 = {
+5
View File
@@ -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",