Skip to content

Commit 75856e8

Browse files
committed
Hoist loop-invariant slerp setup
Prepare the validated quaternion pair and stable angle terms once per interpolation batch, then reuse them for each sample. This removes the scale-dependent regression introduced when interp() and interp1() were routed through qslerp(), while keeping the public scalar qslerp behavior byte-equivalent.
1 parent 1e0643d commit 75856e8

3 files changed

Lines changed: 64 additions & 41 deletions

File tree

‎spatialmath/base/quaternions.py‎

Lines changed: 49 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,7 @@
1717
from spatialmath.base.argcheck import getunit
1818
from spatialmath.base.types import *
1919
import scipy.interpolate as interpolate
20-
from typing import Optional
20+
from typing import Callable, Optional
2121
from functools import lru_cache
2222
import warnings
2323

@@ -771,6 +771,53 @@ def r2q(
771771
# return np.r_[qs, (math.sqrt(1.0 - qs**2) / nm) * kv]
772772

773773

774+
def _qslerp(
775+
q0: ArrayLike4,
776+
q1: ArrayLike4,
777+
shortest: Optional[bool] = False,
778+
tol: float = 20,
779+
) -> Callable[[float], UnitQuaternionArray]:
780+
"""Prepare an interpolator for a pair of unit quaternions."""
781+
q0 = smb.getvector(q0, 4)
782+
q1 = smb.getvector(q1, 4)
783+
q0_endpoint = q0
784+
785+
dotprod = np.dot(q0, q1)
786+
787+
# If the dot product is negative, the quaternions
788+
# have opposite handed-ness and slerp won't take
789+
# the shorter path. Fix by reversing one quaternion.
790+
if shortest:
791+
if dotprod < 0:
792+
q0 = -q0 # pylint: disable=invalid-unary-operand-type
793+
dotprod = -dotprod # pylint: disable=invalid-unary-operand-type
794+
795+
dotprod = np.clip(dotprod, -1, 1)
796+
797+
# sin(theta) is the length of the component of q1 orthogonal to q0. Computing
798+
# it this way keeps full relative precision as theta approaches 0 or pi, where
799+
# sin(acos(dotprod)) does not: acos loses the small angle to rounding.
800+
sin_theta = float(np.linalg.norm(q1 - dotprod * q0))
801+
theta = math.atan2(sin_theta, dotprod) # theta is the angle between q0 and q1
802+
803+
def interpolate(s: float) -> UnitQuaternionArray:
804+
if s == 0:
805+
return q0_endpoint
806+
elif s == 1:
807+
return q1
808+
809+
if sin_theta > tol * _eps:
810+
s0 = math.sin((1 - s) * theta)
811+
s1 = math.sin(s * theta)
812+
return ((q0 * s0) + (q1 * s1)) / sin_theta
813+
else:
814+
# theta is 0 or pi: q0 and q1 are the same rotation, so is every
815+
# interpolate between them
816+
return q0
817+
818+
return interpolate
819+
820+
774821
def qslerp(
775822
q0: ArrayLike4,
776823
q1: ArrayLike4,
@@ -822,40 +869,7 @@ def qslerp(
822869
"""
823870
if not 0 <= s <= 1:
824871
raise ValueError("s must be in the interval [0,1]")
825-
q0 = smb.getvector(q0, 4)
826-
q1 = smb.getvector(q1, 4)
827-
828-
if s == 0:
829-
return q0
830-
elif s == 1:
831-
return q1
832-
833-
dotprod = np.dot(q0, q1)
834-
835-
# If the dot product is negative, the quaternions
836-
# have opposite handed-ness and slerp won't take
837-
# the shorter path. Fix by reversing one quaternion.
838-
if shortest:
839-
if dotprod < 0:
840-
q0 = -q0 # pylint: disable=invalid-unary-operand-type
841-
dotprod = -dotprod # pylint: disable=invalid-unary-operand-type
842-
843-
dotprod = np.clip(dotprod, -1, 1) # Clip within domain of acos()
844-
845-
# sin(theta) is the length of the component of q1 orthogonal to q0. Computing
846-
# it this way keeps full relative precision as theta approaches 0 or pi, where
847-
# sin(acos(dotprod)) does not: acos loses the small angle to rounding.
848-
sin_theta = float(np.linalg.norm(q1 - dotprod * q0))
849-
theta = math.atan2(sin_theta, dotprod) # theta is the angle between q0 and q1
850-
851-
if sin_theta > tol * _eps:
852-
s0 = math.sin((1 - s) * theta)
853-
s1 = math.sin(s * theta)
854-
return ((q0 * s0) + (q1 * s1)) / sin_theta
855-
else:
856-
# theta is 0 or pi: q0 and q1 are the same rotation, so is every
857-
# interpolate between them
858-
return q0
872+
return _qslerp(q0, q1, shortest=shortest, tol=tol)(s)
859873

860874

861875
def _compute_cdf_sin_squared(theta: float):

‎spatialmath/quaternion.py‎

Lines changed: 5 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@
1919
import numpy as np
2020
from typing import Any
2121
import spatialmath.base as smb
22+
from spatialmath.base.quaternions import _qslerp
2223
from spatialmath.pose3d import SO3, SE3
2324
from spatialmath.baseposelist import BasePoseList
2425
from spatialmath.base.types import *
@@ -1950,9 +1951,8 @@ def interp(
19501951
if not isinstance(end, UnitQuaternion):
19511952
raise TypeError("end argument must be a UnitQuaternion")
19521953

1953-
return UnitQuaternion(
1954-
[smb.qslerp(self.vec, end.vec, sk, shortest=shortest) for sk in s]
1955-
)
1954+
interpolate = _qslerp(self.vec, end.vec, shortest=shortest)
1955+
return UnitQuaternion([interpolate(sk) for sk in s])
19561956

19571957
def interp1(self, s: float = 0, shortest: Optional[bool] = False) -> UnitQuaternion:
19581958
"""
@@ -1999,9 +1999,8 @@ def interp1(self, s: float = 0, shortest: Optional[bool] = False) -> UnitQuatern
19991999
s = smb.getvector(s)
20002000
s = np.clip(s, 0, 1) # enforce valid values
20012001

2002-
return UnitQuaternion(
2003-
[smb.qslerp(smb.qeye(), self.vec, sk, shortest=shortest) for sk in s]
2004-
)
2002+
interpolate = _qslerp(smb.qeye(), self.vec, shortest=shortest)
2003+
return UnitQuaternion([interpolate(sk) for sk in s])
20052004

20062005
def increment(self, w: ArrayLike3, normalize: Optional[bool] = False) -> None:
20072006
"""

‎tests/test_quaternion.py‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
from math import pi
33
import numpy.testing as nt
44
import unittest
5+
from unittest.mock import patch
56

67
from spatialmath import *
78
from spatialmath.base import *
@@ -732,6 +733,15 @@ def test_interp_same_rotation(self):
732733
for qi in p.interp(m, 5):
733734
nt.assert_array_almost_equal(qi.R, p.R)
734735

736+
def test_interp_prepares_slerp_once(self):
737+
q0 = UnitQuaternion.RPY([0.2, 0.3, 0.4])
738+
q1 = UnitQuaternion.RPY([-0.3, 0.1, 0.2])
739+
740+
for interpolate in (lambda: q0.interp1(5), lambda: q0.interp(q1, 5)):
741+
with patch("spatialmath.base.quaternions.np.dot", wraps=np.dot) as dot:
742+
self.assertEqual(len(interpolate()), 5)
743+
self.assertEqual(dot.call_count, 1)
744+
735745
def test_increment(self):
736746
q = UnitQuaternion()
737747

0 commit comments

Comments
 (0)