From 4ecff40d7c9f6d7f576f99b30f9b2e73d0f3ce15 Mon Sep 17 00:00:00 2001 From: Tobias Fischer Date: Tue, 18 Aug 2026 13:33:25 +1000 Subject: [PATCH] Add trchain transform evaluators --- docs/source/arm_ets.rst | 23 ++- src/roboticstoolbox/__init__.py | 3 + src/roboticstoolbox/tools/__init__.py | 4 + src/roboticstoolbox/tools/trchain.py | 233 ++++++++++++++++++++++++++ tests/test_tools.py | 43 ++++- 5 files changed, 302 insertions(+), 4 deletions(-) create mode 100644 src/roboticstoolbox/tools/trchain.py diff --git a/docs/source/arm_ets.rst b/docs/source/arm_ets.rst index f0fbc5a95..1521b416e 100644 --- a/docs/source/arm_ets.rst +++ b/docs/source/arm_ets.rst @@ -66,6 +66,28 @@ The ETS inherits list-like properties and has methods like ``reverse`` and ``pop - `A simple and systematic approach to assigning Denavit-Hartenberg parameters `_. Peter I. Corke, IEEE Transactions on Robotics, 23(3), pp 590-594, June 2007. +Transform chains +---------------- + +The :func:`~roboticstoolbox.tools.trchain.trchain` and +:func:`~roboticstoolbox.tools.trchain.trchain2` functions provide the compact +string notation used by the MATLAB Toolbox for one-off SE(3) and SE(2) +transform chains. Joint variables are numbered from one, while named +constants are passed explicitly in ``variables``. + +.. runblock:: pycon + + >>> import roboticstoolbox as rtb + >>> rtb.trchain("Rz(q1) Tx(a)", [0.3], variables={"a": 1}) + >>> rtb.trchain2("R(q1) Tx(a)", [0.3], variables={"a": 1}) + +Parsed tokens can also be returned and reused when evaluating the same chain +for different joint values. + +.. autofunction:: roboticstoolbox.tools.trchain.trchain + +.. autofunction:: roboticstoolbox.tools.trchain.trchain2 + ETS - 3D -------- @@ -82,4 +104,3 @@ ETS - 2D :members: __str__, __repr__, __mul__, __getitem__, n, m, structure, joints, jindex_set, split, inv, compile, insert, fkine, jacob0, jacobe :undoc-members: :show-inheritance: - diff --git a/src/roboticstoolbox/__init__.py b/src/roboticstoolbox/__init__.py index d7df4fc21..3fe8f44c7 100644 --- a/src/roboticstoolbox/__init__.py +++ b/src/roboticstoolbox/__init__.py @@ -61,6 +61,9 @@ "rtb_path_to_datafile", "rtb_set_param", "rtb_get_param", + "TrChainToken", + "trchain", + "trchain2", # mobile "VehicleBase", "Bicycle", diff --git a/src/roboticstoolbox/tools/__init__.py b/src/roboticstoolbox/tools/__init__.py index 54e9da371..0d5ee9011 100644 --- a/src/roboticstoolbox/tools/__init__.py +++ b/src/roboticstoolbox/tools/__init__.py @@ -22,6 +22,7 @@ ) from roboticstoolbox.tools.plot import xplot from roboticstoolbox.tools.params import rtb_set_param, rtb_get_param +from roboticstoolbox.tools.trchain import TrChainToken, trchain, trchain2 from roboticstoolbox.tools.types import ArrayLike, NDArray, PyArrayLike __all__ = [ @@ -48,6 +49,9 @@ "rtb_path_to_datafile", "rtb_set_param", "rtb_get_param", + "TrChainToken", + "trchain", + "trchain2", "PyArrayLike", "ArrayLike", "NDArray", diff --git a/src/roboticstoolbox/tools/trchain.py b/src/roboticstoolbox/tools/trchain.py new file mode 100644 index 000000000..b5df6b0d8 --- /dev/null +++ b/src/roboticstoolbox/tools/trchain.py @@ -0,0 +1,233 @@ +"""Evaluate chains of elementary homogeneous transforms.""" + +import ast +import operator +import re +from collections.abc import Callable, Mapping, Sequence +from typing import NamedTuple + +import numpy as np +from spatialmath.base import symbolic as sym +from spatialmath.base import transl, transl2, trot2, trotx, troty, trotz + + +class TrChainToken(NamedTuple): + """A parsed elementary-transform token. + + :param op: elementary-transform operator + :param arg: expression inside the operator's parentheses + :param index: one-based joint index, or zero for a constant expression + """ + + op: str + arg: str + index: int + + +_TOKEN_RE = re.compile(r"(?PR.?|T.)\(") +_BINARY_OPERATORS: dict[type[ast.operator], Callable[[object, object], object]] = { + ast.Add: operator.add, + ast.Sub: operator.sub, + ast.Mult: operator.mul, + ast.Div: operator.truediv, + ast.FloorDiv: operator.floordiv, + ast.Mod: operator.mod, + ast.Pow: operator.pow, +} +_UNARY_OPERATORS: dict[type[ast.unaryop], Callable[[object], object]] = { + ast.UAdd: operator.pos, + ast.USub: operator.neg, +} +_DEFAULT_VALUES = { + "pi": np.pi, + "sin": sym.sin, + "cos": sym.cos, + "tan": sym.tan, + "sqrt": sym.sqrt, +} + +_Transform = Callable[[object, str], np.ndarray] +_TRANSFORMS_3D: dict[str, _Transform] = { + "Rx": lambda value, unit: trotx(value, unit), + "Ry": lambda value, unit: troty(value, unit), + "Rz": lambda value, unit: trotz(value, unit), + "Tx": lambda value, _unit: transl(value, 0, 0), + "Ty": lambda value, _unit: transl(0, value, 0), + "Tz": lambda value, _unit: transl(0, 0, value), +} +_TRANSFORMS_2D: dict[str, _Transform] = { + "R": lambda value, unit: trot2(value, unit), + "Rz": lambda value, unit: trot2(value, unit), + "Tx": lambda value, _unit: transl2(value, 0), + "Ty": lambda value, _unit: transl2(0, value), +} + + +def _parse(chain: str, qvar: str) -> tuple[TrChainToken, ...]: + q_re = re.compile(rf"\b{re.escape(qvar)}([1-9][0-9]*)\b") + tokens = [] + position = 0 + + while match := _TOKEN_RE.search(chain, position): + depth = 1 + for position in range(match.end(), len(chain)): + if chain[position] == "(": + depth += 1 + elif chain[position] == ")": + depth -= 1 + if depth == 0: + break + else: + raise ValueError(f"unclosed transform {match['op']!r}") + + arg = chain[match.end() : position] + indices = q_re.findall(arg) + if len(indices) > 1: + raise ValueError("only one joint variable is allowed in each transform") + tokens.append(TrChainToken(match["op"], arg, int(indices[0]) if indices else 0)) + position += 1 + + return tuple(tokens) + + +def _eval_expression(expression: str, values: Mapping[str, object]) -> object: + try: + tree = ast.parse(expression, mode="eval") + except SyntaxError as exc: + raise ValueError(f"cannot evaluate expression {expression!r}") from exc + + def evaluate(node: ast.AST) -> object: + if isinstance(node, ast.Constant) and isinstance(node.value, (int, float)): + return node.value + if isinstance(node, ast.Name): + try: + return values[node.id] + except KeyError as exc: + raise ValueError( + f"unknown name {node.id!r} in expression {expression!r}" + ) from exc + if isinstance(node, ast.BinOp) and type(node.op) in _BINARY_OPERATORS: + return _BINARY_OPERATORS[type(node.op)]( + evaluate(node.left), evaluate(node.right) + ) + if isinstance(node, ast.UnaryOp) and type(node.op) in _UNARY_OPERATORS: + return _UNARY_OPERATORS[type(node.op)](evaluate(node.operand)) + if ( + isinstance(node, ast.Call) + and isinstance(node.func, ast.Name) + and not node.keywords + ): + function = values.get(node.func.id) + if callable(function): + return function(*(evaluate(arg) for arg in node.args)) + raise ValueError(f"unsupported expression {expression!r}") + + return evaluate(tree.body) + + +def _evaluate_chain( + chain: str | Sequence[TrChainToken], + q: Sequence[object] | np.ndarray | None, + unit: str, + qvar: str, + variables: Mapping[str, object] | None, + transforms: Mapping[str, _Transform], + size: int, +) -> tuple[np.ndarray, tuple[TrChainToken, ...]]: + if unit not in {"rad", "deg"}: + raise ValueError("unit must be 'rad' or 'deg'") + if not qvar.isidentifier(): + raise ValueError("qvar must be a valid identifier") + + if isinstance(chain, str): + tokens = _parse(chain, qvar) + else: + tokens = tuple(chain) + if not all(isinstance(token, TrChainToken) for token in tokens): + raise TypeError("chain must be a string or a sequence of TrChainToken") + + q_values = np.asarray(()) if q is None else np.asarray(q).reshape(-1) + values = {**_DEFAULT_VALUES, **(variables or {})} + T = np.eye(size) + + for token in tokens: + try: + transform = transforms[token.op] + except KeyError as exc: + raise ValueError(f"unknown transform {token.op!r}") from exc + + token_values = values + if token.index: + if token.index > len(q_values): + raise ValueError("q has insufficient values") + token_values = { + **values, + f"{qvar}{token.index}": q_values[token.index - 1], + } + + T = T @ transform(_eval_expression(token.arg, token_values), unit) + + return T, tokens + + +def trchain( + chain: str | Sequence[TrChainToken], + q: Sequence[object] | np.ndarray | None = None, + unit: str = "rad", + *, + qvar: str = "q", + variables: Mapping[str, object] | None = None, + return_tokens: bool = False, +) -> np.ndarray | tuple[np.ndarray, tuple[TrChainToken, ...]]: + """Compound SE(3) elementary transforms from a string. + + :param chain: transform chain or tokens returned by an earlier call + :param q: joint values referenced as ``q1``, ``q2``, and so on + :param unit: angular unit, ``"rad"`` or ``"deg"`` + :param qvar: joint-variable prefix + :param variables: values and functions used by token expressions + :param return_tokens: also return the parsed tokens for reuse + :returns: the SE(3) matrix, optionally followed by its parsed tokens + :rtype: ndarray(4, 4) or tuple + + The chain contains ``Rx``, ``Ry``, ``Rz``, ``Tx``, ``Ty``, and ``Tz`` + tokens. Expressions support arithmetic, named values, ``pi``, and direct + calls to named functions. Names other than joint variables are supplied by + ``variables``. + + :seealso: :func:`trchain2`, :class:`roboticstoolbox.ETS` + """ + result = _evaluate_chain(chain, q, unit, qvar, variables, _TRANSFORMS_3D, 4) + return result if return_tokens else result[0] + + +def trchain2( + chain: str | Sequence[TrChainToken], + q: Sequence[object] | np.ndarray | None = None, + unit: str = "rad", + *, + qvar: str = "q", + variables: Mapping[str, object] | None = None, + return_tokens: bool = False, +) -> np.ndarray | tuple[np.ndarray, tuple[TrChainToken, ...]]: + """Compound SE(2) elementary transforms from a string. + + :param chain: transform chain or tokens returned by an earlier call + :param q: joint values referenced as ``q1``, ``q2``, and so on + :param unit: angular unit, ``"rad"`` or ``"deg"`` + :param qvar: joint-variable prefix + :param variables: values and functions used by token expressions + :param return_tokens: also return the parsed tokens for reuse + :returns: the SE(2) matrix, optionally followed by its parsed tokens + :rtype: ndarray(3, 3) or tuple + + The chain contains ``R`` (or ``Rz``), ``Tx``, and ``Ty`` tokens. + Expressions are evaluated as described by :func:`trchain`. + + :seealso: :func:`trchain`, :class:`roboticstoolbox.ETS2` + """ + result = _evaluate_chain(chain, q, unit, qvar, variables, _TRANSFORMS_2D, 3) + return result if return_tokens else result[0] + + +__all__ = ["TrChainToken", "trchain", "trchain2"] diff --git a/tests/test_tools.py b/tests/test_tools.py index 6d7d6d39b..b175666a3 100644 --- a/tests/test_tools.py +++ b/tests/test_tools.py @@ -3,14 +3,51 @@ @author: Jesse Haviland """ -import numpy.testing as nt +import unittest + import numpy as np -import roboticstoolbox as rtb +import numpy.testing as nt import spatialmath as sm -import unittest +import sympy + +import roboticstoolbox as rtb class Testtools(unittest.TestCase): + def test_trchain(self): + T, tokens = rtb.trchain( + "Tx(a) Rx(q1) Ry(45) Tz(2)", + [90], + "deg", + variables={"a": 1}, + return_tokens=True, + ) + expected = ( + sm.SE3.Tx(1) + * sm.SE3.Rx(90, unit="deg") + * sm.SE3.Ry(45, unit="deg") + * sm.SE3.Tz(2) + ) + nt.assert_allclose(T, expected.A) + nt.assert_allclose(rtb.trchain(tokens, [90], unit="deg", variables={"a": 1}), T) + + q1, a = sympy.symbols("q1 a", real=True) + symbolic = rtb.trchain("Rz(q1) Tx(a)", [q1], variables={"a": a}) + self.assertEqual(sympy.simplify(symbolic[0, 3] - a * sympy.cos(q1)), 0) + + with self.assertRaises(ValueError): + rtb.trchain("Tx(__import__('os'))") + + def test_trchain2(self): + T = rtb.trchain2( + "R(theta1 - pi / 2) Tx(a) Rz(theta2) Ty(2)", + [np.pi / 2, np.pi / 4], + qvar="theta", + variables={"a": 1}, + ) + expected = sm.SE2.Rot(0) * sm.SE2.Tx(1) * sm.SE2.Rot(np.pi / 4) * sm.SE2.Ty(2) + nt.assert_allclose(T, expected.A) + def test_null(self): a0 = np.array([1, 2, 3])