diff --git a/doc/changelog.rst b/doc/changelog.rst
index 25d866a4..471e7105 100644
--- a/doc/changelog.rst
+++ b/doc/changelog.rst
@@ -2,6 +2,16 @@
Changelog
=========
+Upcoming version (not yet released)
+-----------------------------------
+
+General
+^^^^^^^
+
+- Improved convergence of Signed Distance Function (SDF) collisions by using line search and a new objective function
+ for the optimization. This allows to decrease the number of initial points needed for finding the contacts and is more
+ robust for very small or large geom sizes.
+
Version 3.0.1 (November 15, 2023)
---------------------------------
diff --git a/model/plugin/sdf/nutbolt.xml b/model/plugin/sdf/nutbolt.xml
index d3d80ff7..51e2b554 100644
--- a/model/plugin/sdf/nutbolt.xml
+++ b/model/plugin/sdf/nutbolt.xml
@@ -8,7 +8,7 @@
-
+
@@ -30,7 +30,7 @@
-
+
diff --git a/plugin/sdf/bolt.cc b/plugin/sdf/bolt.cc
index 99264d23..4030a164 100644
--- a/plugin/sdf/bolt.cc
+++ b/plugin/sdf/bolt.cc
@@ -166,12 +166,12 @@ void Bolt::RegisterPlugin() {
plugin.sdf_distance =
+[](const mjtNum point[3], const mjData* d, int instance) {
auto* sdf = reinterpret_cast(d->plugin_data[instance]);
- sdf->visualizer_.AddPoint(point);
return sdf->Distance(point);
};
plugin.sdf_gradient = +[](mjtNum gradient[3], const mjtNum point[3],
const mjData* d, int instance) {
auto* sdf = reinterpret_cast(d->plugin_data[instance]);
+ sdf->visualizer_.AddPoint(point);
sdf->Gradient(gradient, point);
};
plugin.sdf_staticdistance =
diff --git a/plugin/sdf/bowl.cc b/plugin/sdf/bowl.cc
index 451b2b4c..55b6515f 100644
--- a/plugin/sdf/bowl.cc
+++ b/plugin/sdf/bowl.cc
@@ -164,12 +164,12 @@ void Bowl::RegisterPlugin() {
plugin.sdf_distance =
+[](const mjtNum point[3], const mjData* d, int instance) {
auto* sdf = reinterpret_cast(d->plugin_data[instance]);
- sdf->visualizer_.AddPoint(point);
return sdf->Distance(point);
};
plugin.sdf_gradient = +[](mjtNum gradient[3], const mjtNum point[3],
const mjData* d, int instance) {
auto* sdf = reinterpret_cast(d->plugin_data[instance]);
+ sdf->visualizer_.AddPoint(point);
sdf->Gradient(gradient, point);
};
plugin.sdf_staticdistance =
diff --git a/plugin/sdf/gear.cc b/plugin/sdf/gear.cc
index 67f4a8be..d7b49da3 100644
--- a/plugin/sdf/gear.cc
+++ b/plugin/sdf/gear.cc
@@ -250,12 +250,12 @@ void Gear::RegisterPlugin() {
plugin.sdf_distance =
+[](const mjtNum point[3], const mjData* d, int instance) {
auto* sdf = reinterpret_cast(d->plugin_data[instance]);
- sdf->visualizer_.AddPoint(point);
return sdf->Distance(point);
};
plugin.sdf_gradient = +[](mjtNum gradient[3], const mjtNum point[3],
const mjData* d, int instance) {
auto* sdf = reinterpret_cast(d->plugin_data[instance]);
+ sdf->visualizer_.AddPoint(point);
sdf->Gradient(gradient, point);
};
plugin.sdf_staticdistance =
diff --git a/plugin/sdf/nut.cc b/plugin/sdf/nut.cc
index 9e0fb53b..448fc4a8 100644
--- a/plugin/sdf/nut.cc
+++ b/plugin/sdf/nut.cc
@@ -166,12 +166,12 @@ void Nut::RegisterPlugin() {
plugin.sdf_distance =
+[](const mjtNum point[3], const mjData* d, int instance) {
auto* sdf = reinterpret_cast(d->plugin_data[instance]);
- sdf->visualizer_.AddPoint(point);
return sdf->Distance(point);
};
plugin.sdf_gradient = +[](mjtNum gradient[3], const mjtNum point[3],
const mjData* d, int instance) {
auto* sdf = reinterpret_cast(d->plugin_data[instance]);
+ sdf->visualizer_.AddPoint(point);
sdf->Gradient(gradient, point);
};
plugin.sdf_staticdistance =
diff --git a/plugin/sdf/sdflib.cc b/plugin/sdf/sdflib.cc
index 62628d13..392119ae 100644
--- a/plugin/sdf/sdflib.cc
+++ b/plugin/sdf/sdflib.cc
@@ -146,12 +146,12 @@ void SdfLib::RegisterPlugin() {
plugin.sdf_distance =
+[](const mjtNum point[3], const mjData* d, int instance) {
auto* sdf = reinterpret_cast(d->plugin_data[instance]);
- sdf->visualizer_.AddPoint(point);
return sdf->Distance(point);
};
plugin.sdf_gradient = +[](mjtNum gradient[3], const mjtNum point[3],
const mjData* d, int instance) {
auto* sdf = reinterpret_cast(d->plugin_data[instance]);
+ sdf->visualizer_.AddPoint(point);
sdf->Gradient(gradient, point);
};
diff --git a/src/engine/engine_collision_sdf.c b/src/engine/engine_collision_sdf.c
index ed5b8d06..ec46002a 100644
--- a/src/engine/engine_collision_sdf.c
+++ b/src/engine/engine_collision_sdf.c
@@ -249,7 +249,6 @@ void mjc_gradient(const mjModel* m, const mjData* d, const mjSDF* s,
gradient[1] = - grad1[1] * B - grad2[1] * A;
gradient[2] = - grad1[2] * B - grad2[2] * A;
}
- mju_normalize3(gradient);
break;
case mjSDFTYPE_SINGLE:
geomGradient(gradient, m, d, s->plugin[0], s->id[0], point[0], s->geomtype[0]);
@@ -388,10 +387,15 @@ static mjtNum stepFrankWolfe(mjtNum x[3], const mjtNum* corners, int ncorners,
// finds minimum using gradient descent
static mjtNum stepGradient(mjtNum x[3], const mjModel* m, const mjSDF* s,
mjData* d) {
- mjtNum alpha = 0.2; // step along the gradient direction
+ const mjtNum c = .1; // reduction factor for the target decrease in the objective function
+ const mjtNum rho = .5; // reduction factor for the gradient scaling (alpha)
+ const mjtNum amin = 1e-4; // minimum value for alpha
+ mjtNum dist = mjMAXVAL;
for (int step=0; step < m->opt.sdf_iterations; step++) {
mjtNum grad[3];
+ mjtNum alpha = 2.; // initial line search factor scaling the gradient
+ // the units of the gradient depend on s->type
// evaluate gradient
mjc_gradient(m, d, s, grad, x);
@@ -403,12 +407,29 @@ static mjtNum stepGradient(mjtNum x[3], const mjModel* m, const mjSDF* s,
return mjMAXVAL;
}
- // update solution
- mju_addToScl3(x, grad, -alpha/(mjtNum)(step+1));
+ // save current solution
+ mjtNum x0[] = {x[0], x[1], x[2]};
+
+ // evaluate distance
+ mjtNum dist0 = mjc_distance(m, d, s, x0);
+ mjtNum wolfe = - c * alpha * mju_dot3(grad, grad);
+
+ // backtracking line search
+ do {
+ alpha *= rho;
+ wolfe *= rho;
+ mju_addScl3(x, x0, grad, -alpha);
+ dist = mjc_distance(m, d, s, x);
+ } while (alpha > amin && dist - dist0 > wolfe);
+
+ // if no improvement, early stop
+ if (dist0 < dist) {
+ return dist;
+ }
}
- // compute distance
- return mjc_distance(m, d, s, x);
+ // the distance will be used for the contact creation
+ return dist;
}
//---------------------------- bounding box vs sdf -------------------------------------------------