Improve MJX performance on NVIDIA GPU by enabling Triton GEMM.
PiperOrigin-RevId: 604691509 Change-Id: Iba11a9e1b79848624a1ef6eb4e61e2093d36e662
This commit is contained in:
committed by
Copybara-Service
parent
84f20ffb66
commit
38cbc9e543
+10
@@ -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>`__.
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user