From a897805683fef38173020e544d4ff2f98b282b2b Mon Sep 17 00:00:00 2001 From: Yuval Tassa Date: Sun, 17 May 2026 12:02:52 -0700 Subject: [PATCH] Improve test comparing forward and inverse dynamics mismatch PiperOrigin-RevId: 916864072 Change-Id: I9952c7dd1bde2033dfe1380552f2c4b6b7344b26 --- test/engine/engine_inverse_test.cc | 41 +++++++++++++++++++++--------- 1 file changed, 29 insertions(+), 12 deletions(-) diff --git a/test/engine/engine_inverse_test.cc b/test/engine/engine_inverse_test.cc index 988e73e6..97a48e2d 100644 --- a/test/engine/engine_inverse_test.cc +++ b/test/engine/engine_inverse_test.cc @@ -41,19 +41,36 @@ TEST_F(InverseTest, ForwardInverseMatch) { ASSERT_THAT(model, NotNull()) << error; mjData* data = mj_makeData(model); - // simulate, call mj_forward - for (int i = 0; i < kSteps; ++i) { - mj_step(model, data); + // set small tolerance and enough iterations for all solvers to converge + model->opt.iterations = 500; + model->opt.tolerance = 0; + + // solver names for diagnostics + const char* solver_name[] = {"PGS", "CG", "Newton"}; + + for (mjtSolver solver : {mjSOL_PGS, mjSOL_CG, mjSOL_NEWTON}) { + model->opt.solver = solver; + mj_resetData(model, data); + + // simulate, call mj_forward + for (int i = 0; i < kSteps; ++i) { + mj_step(model, data); + } + mj_forward(model, data); + + // call built-in testing function + mj_compareFwdInv(model, data); + + // per-solver tolerances + mjtNum epsilon; + switch (solver) { + case mjSOL_PGS: epsilon = MjTol(1e-6, 1e-2); break; + case mjSOL_CG: epsilon = MjTol(1e-3, 1e-1); break; + case mjSOL_NEWTON: epsilon = MjTol(1e-10, 5e-3); break; + } + EXPECT_LT(data->solver_fwdinv[0], epsilon) << solver_name[solver]; + EXPECT_LT(data->solver_fwdinv[1], epsilon) << solver_name[solver]; } - mj_forward(model, data); - - // call built-in testing function - mj_compareFwdInv(model, data); - - // expect mismatch to be small - mjtNum epsilon = MjTol(1e-10, 0.05); - EXPECT_LT(data->solver_fwdinv[0], epsilon); - EXPECT_LT(data->solver_fwdinv[1], epsilon); mj_deleteData(data); mj_deleteModel(model);