google-deepmind/mujoco

[MJX] jax.lax.while_loop in solver.py prevents computation of backward gradients

Open

#2,259 opened on Nov 29, 2024

 (14 comments) (5 reactions) (2 assignees)C++ (1,689 forks)github user discovery
MJXenhancementgood first issue

Repository metrics

Stars
 (14,570 stars)
PR merge metrics
 (Avg merge 15d 1h) (13 merged PRs in 30d)

Description

The feature, motivation and pitch

Problem

The solver's jax.lax.while_loop implementation prevents gradient computation through the environment step during gradient based trajectory optimization. This occurs in the solver implementation when iterations > 1.

Error encountered with jax.jit compiled grad function:

ValueError: Reverse-mode differentiation does not work for lax.while_loop or lax.fori_loop with dynamic start/stop values.

Current workaround of using opt.iteration=1 leads to potentially inaccurate simulation and gradients.

Proposed Solution

Add an option to set a fixed iteration count (e.g., 4) that would be compatible with reverse-mode differentiation using either lax.scan or lax.fori_loop with static bounds.

Alternatives

No response

Additional context

No response

Contributor guide