From 881544c0c58dc2e95fbd132a5ec90b99e012b6f7 Mon Sep 17 00:00:00 2001 From: Taylor Howell Date: Thu, 12 Feb 2026 11:38:34 -0800 Subject: [PATCH] Modify MuJoCo Warp dev version check PiperOrigin-RevId: 869311325 Change-Id: I2f34bd932d7998d37b96b36e5d561f734567b8d5 --- .../mjx/third_party/mujoco_warp/_src/io.py | 20 +++++++------------ 1 file changed, 7 insertions(+), 13 deletions(-) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py index 2febee1f..d2bd14d9 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py @@ -15,12 +15,12 @@ import dataclasses import importlib.metadata -import re import warnings from typing import Any, Optional, Sequence import mujoco import numpy as np +import packaging.version import warp as wp from mujoco.mjx.third_party.mujoco_warp._src import bvh @@ -31,18 +31,12 @@ from mujoco.mjx.third_party.mujoco_warp._src import warp_util def _is_mujoco_dev() -> bool: - _DEV_VERSION_PATTERN = re.compile(r"^\d+\.\d+\.\d+.+") # anything after x.y.z - - version = getattr(__import__("mujoco"), "__version__", None) - if version and _DEV_VERSION_PATTERN.match(version): - return True - - # fall back to metadata - dist_version = importlib.metadata.version("mujoco") - if _DEV_VERSION_PATTERN.match(dist_version): - return True - - return False + """Checks if mujoco version is > 3.4.0.""" + version_str = getattr(mujoco, "__version__", None) + if not version_str: + version_str = importlib.metadata.version("mujoco") + version_str = version_str.split("-")[0].split(".dev")[0] + return packaging.version.parse(version_str) > packaging.version.parse("3.4.0") BLEEDING_EDGE_MUJOCO = _is_mujoco_dev()