Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 3 additions & 2 deletions spatialmath/base/animate.py
Original file line number Diff line number Diff line change
Expand Up @@ -163,8 +163,9 @@ def trplot(
else:
self.start = start

# draw axes at the origin
smb.trplot(self.start, ax=self, **kwargs)
# Store frame geometry at the origin so each absolute pose is applied once.
smb.trplot(np.identity(4), ax=self, **kwargs)
self._draw(self.start)

def set_proj_type(self, proj_type: str):
self.ax.set_proj_type(proj_type)
Expand Down
46 changes: 46 additions & 0 deletions tests/base/test_transforms3d_plot.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,9 +22,55 @@
from spatialmath.base.transformsNd import isR, t2r, r2t, rt2tr

import matplotlib.pyplot as plt
from spatialmath.base.animate import Animate


class Test3D(unittest.TestCase):
def test_animate_nonidentity_start(self):
# the axis artists of an Animate must sit where a static trplot puts them, both straight after
# trplot(end, start=start) and after _draw(end); losing either half of the fix breaks one of these
def axis_geometry(artists):
geometry = []
for h in artists[:3]: # the three axes of the frame
# arrow style: compare the shaft only, the head wings differ
if hasattr(h, "_segments3d"):
geometry.append(np.asarray(h._segments3d[0]).T)
else:
geometry.append(np.asarray(h.get_data_3d()))
return geometry

def static_geometry(pose, style):
fig = plt.figure()
try:
ax = fig.add_subplot(projection="3d")
trplot(pose, ax=ax, style=style, frame="moving")
kinds = ("Line3DCollection",) if style == "arrow" else ("Line3D",)
return axis_geometry(
[a for a in ax.get_children() if type(a).__name__ in kinds]
)
finally:
plt.close(fig)

start = transl(-1, 0, 2) @ trotz(0.4)
end = transl(1, 2, 1) @ trotx(0.3)
for style in ("line", "arrow", "rviz"):
fig = plt.figure()
try:
anim = Animate(dim=[-4, 4])
anim.trplot(end, start=start, style=style, frame="moving")
artists = [x.h for x in anim.displaylist]
for actual, expected in zip(
axis_geometry(artists), static_geometry(start, style)
):
nt.assert_allclose(actual, expected, atol=1e-12)
anim._draw(end)
for actual, expected in zip(
axis_geometry(artists), static_geometry(end, style)
):
nt.assert_allclose(actual, expected, atol=1e-12)
finally:
plt.close(fig)

@pytest.mark.skipif(
os.environ.get("CI") == "true"
or (sys.platform.startswith("darwin") and sys.version_info < (3, 11)),
Expand Down