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
16 changes: 16 additions & 0 deletions agent_context/topics/motion-planning/planner-details.md
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,14 @@ the graph. Validate complete observation-layout metadata/fingerprint before
rollout; model dimensions and waypoint capacity belong to the export contract.
Inspect the `policy-deploy` package extra for runtime dependencies.

`NeuralPlannerCfg.dt` is bookkeeping, not solved timing: the rollout integrates
a joint delta per step, so a saturated step implies `action_scale / dt` rad/s
and overshoots a typical arm's velocity limit several times over. Setting
`constraints` swaps that nominal timing for a TOPPRA parameterization through
the shared entry point below; it resamples the path, so positions and poses
change with the timing, and it cannot reduce the path's own jerk. Default
remains `None`, which preserves the nominal-timing behavior.

Runtime/training frame differences must be expressed through explicit policy
frame and TCP transforms. Target/FK quaternions follow the shared
[simulation convention](../simulation-system/simulation-system.md); do not
Expand All @@ -123,6 +131,14 @@ semantics are owned by [collision worlds](collision-worlds.md).

## Retiming and playback

`toppra_planner.py::retime_joint_paths()` is the shared entry point for
re-parameterizing an already-planned joint path under velocity and acceleration
limits, used by planners that emit geometry without executable timing. It is
the planner's own time parameterization with the same backend selection, so
CUDA paths and gradient requests run on Warp; it resamples along the fitted
spline, so recompute derived poses from its output rather than reusing the
input samples.

`compute/trajectory/timing.py` owns time-domain differentiation/resampling and
`retime_to_control_grid()`. Zero-time position changes are invalid; repeated-time
unchanged samples may represent padding/junctions. Time resampling preserves
Expand Down
6 changes: 6 additions & 0 deletions docs/source/api_reference/public_api.rst
Original file line number Diff line number Diff line change
Expand Up @@ -1176,11 +1176,17 @@ embodichain.lab.sim.motion.planners.toppra_planner

.. currentmodule:: embodichain.lab.sim.motion.planners.toppra_planner

Time-optimal joint-path parameterization. The planner solves timing for
waypoint input; :func:`retime_joint_paths` applies the same parameterization to
paths another planner already produced, such as a closed-loop policy rollout
whose samples carry only nominal timing.

.. autosummary::

ToppraPlanner
ToppraPlannerCfg
ToppraPlanOptions
retime_joint_paths

embodichain.lab.sim.motion.planners.se3
---------------------------------------
Expand Down
191 changes: 187 additions & 4 deletions embodichain/lab/sim/motion/planners/neural_planner.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,12 @@
_infer_batch_size,
validate_plan_options,
)
from embodichain.lab.sim.motion.planners.utils import MoveType, PlanResult, PlanState
from embodichain.lab.sim.motion.planners.utils import (
MoveType,
PlanResult,
PlanState,
TrajectorySampleMethod,
)
from embodichain.utils import configclass, logger
from embodichain.utils.math import quat_error_magnitude, quat_from_matrix

Expand Down Expand Up @@ -316,7 +321,30 @@ class NeuralPlannerCfg(BasePlannerCfg):
"""

dt: float = 0.01
"""Nominal timestep reported in PlanResult."""
"""Output sampling period in seconds.

Without :attr:`constraints` this is bookkeeping rather than executable
timing.

The rollout integrates a joint delta per step and never solves a duration,
so this value is bookkeeping rather than executable timing: at the default
``action_scale`` a saturated step implies ``action_scale / dt`` rad/s,
which exceeds a typical arm's joint velocity limit several times over. Set
:attr:`constraints` to replace it with a solved time parameterization, in
which case this becomes the period that trajectory is sampled on.
"""

constraints: dict[str, float | list[float]] | None = None
"""Optional ``velocity`` and ``acceleration`` limits for output retiming.

``None`` keeps the rollout's nominal timing and current behavior. When set,
the closed-loop samples are treated as a geometric path and re-parameterized
under these limits, so the returned timing and derivatives are executable.
Keys and value shapes match :class:`ToppraPlanOptions`: a scalar or one
value per controlled joint. Retiming resamples the path, so it changes the
reported positions and poses as well as the timing, and it cannot reduce
the path's own jerk -- only the duration over which it is traversed.
"""


@configclass
Expand Down Expand Up @@ -583,10 +611,26 @@ def plan(
dt[:, 0] = 0.0
positions_t = positions_t.permute(1, 0, 2)
xpos_t = xpos_t.permute(1, 0, 2, 3)
success = active_idx >= episode_k
if self.cfg.constraints is not None:
return self._retimed_result(
success,
positions_t,
control_part,
waypoints=(
waypoints_pos,
waypoints_quat,
waypoints_joint,
pos_mask,
rot_mask,
joint_mask,
),
episode_k=episode_k,
qpos_limits=(lower, upper),
)
velocities_t, accelerations_t = self._compute_vel_acc_via_finite_diff(
positions_t, dt
)
success = active_idx >= episode_k
return PlanResult(
success=success,
positions=positions_t,
Expand All @@ -596,6 +640,141 @@ def plan(
dt=dt,
)

def _retimed_result(
self,
success: torch.Tensor,
positions: torch.Tensor,
control_part: str,
*,
waypoints: tuple[torch.Tensor, ...],
episode_k: int,
qpos_limits: tuple[torch.Tensor, torch.Tensor],
) -> PlanResult:
"""Re-parameterize the rollout path under the configured limits.

Retiming fits a spline through the rollout samples and resamples it, so
the returned trajectory is not the one the rollout verified. Success is
therefore re-derived from the samples actually returned rather than
carried over: the resampled grid can step past a waypoint the rollout
stopped on, and a spline through samples clamped at a joint limit can
overshoot that limit between them. A row is reported as successful only
when its returned trajectory parameterizes, moves in nonzero time,
stays inside the joint limits, and still reaches every waypoint.

Args:
success: Per-env rollout convergence of shape ``(B,)``.
positions: Rollout joint samples of shape ``(B, N, DOF)``.
control_part: Robot control part used for forward kinematics.
waypoints: Parsed waypoint targets and masks from
:meth:`_parse_waypoints`.
episode_k: Number of waypoints this episode must reach.
qpos_limits: Lower and upper joint-position limits.

Returns:
PlanResult with solved timing, derivatives and recomputed poses.
"""
from .toppra_planner import retime_joint_paths

# Sample the solved trajectory on the configured control period rather
# than on the rollout's step count. The rollout emits however many
# steps it happened to take, which is far too coarse once the duration
# stretches: a waypoint the rollout stopped on then falls between two
# output samples and is reported as missed.
retimed = retime_joint_paths(
positions,
constraints=dict(self.cfg.constraints),
sample_method=TrajectorySampleMethod.TIME,
sample_interval=float(self.cfg.dt),
Comment thread
greptile-apps[bot] marked this conversation as resolved.
device=self.device,
)
Comment thread
greptile-apps[bot] marked this conversation as resolved.
Comment thread
greptile-apps[bot] marked this conversation as resolved.
retimed_positions = retimed.positions
# One batched FK over every sample; the waypoint re-check reuses it.
poses = self.robot.compute_batch_fk(
qpos=retimed_positions, name=control_part, to_matrix=True
)
success = success & retimed.success.to(success.device)
success = success & self._retimed_is_executable(
retimed_positions, retimed.dt, qpos_limits
)
success = success & self._retimed_reaches_waypoints(
retimed_positions, poses, waypoints, episode_k
)
return PlanResult(
success=success,
positions=retimed_positions,
velocities=retimed.velocities,
accelerations=retimed.accelerations,
xpos_list=poses,
dt=retimed.dt,
Comment thread
greptile-apps[bot] marked this conversation as resolved.
)

def _retimed_is_executable(
self,
positions: torch.Tensor,
dt: torch.Tensor,
qpos_limits: tuple[torch.Tensor, torch.Tensor],
) -> torch.Tensor:
"""Return whether each retimed row has real timing and legal positions.

The shared retiming kernel returns a zero-duration result for a path
whose endpoints nearly coincide, which would claim a move with no time
to execute it. Spline interpolation through samples the rollout clamped
at a joint limit can also overshoot that limit between them, since the
fit sees only the sampled values.
"""
lower, upper = qpos_limits
arm = positions[..., : self._action_dim]
moves = (arm - arm[:, :1]).abs().amax(dim=(1, 2)) > 1.0e-9
timed = dt.sum(dim=-1) > 0.0
within = ((arm >= lower) & (arm <= upper)).all(dim=-1).all(dim=-1)
return within & (timed | ~moves)

def _retimed_reaches_waypoints(
self,
positions: torch.Tensor,
poses: torch.Tensor,
waypoints: tuple[torch.Tensor, ...],
episode_k: int,
) -> torch.Tensor:
"""Return whether each retimed row still reaches every waypoint in order.

Applies the rollout's own arrival test to the resampled grid, so a
trajectory that passes a waypoint only between output samples is not
reported as reaching it.
"""
(
waypoints_pos,
waypoints_quat,
waypoints_joint,
pos_mask,
rot_mask,
joint_mask,
) = waypoints
batch, samples = positions.shape[:2]
policy_poses = self._policy_pose_xyzw(poses.flatten(0, 1)).view(
batch, samples, -1
)
active_idx = torch.zeros(batch, dtype=torch.long, device=self.device)
for index in range(samples):
qpos = positions[:, index]
reached = self._is_active_reached(
qpos[:, : self._action_dim],
policy_poses[:, index],
waypoints_pos,
waypoints_quat,
waypoints_joint,
pos_mask,
rot_mask,
joint_mask,
active_idx,
)
active_idx = torch.where(
reached & (active_idx < episode_k), active_idx + 1, active_idx
)
if bool((active_idx >= episode_k).all()):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Check completion on the CPU at every sample spend a lot of time.

break
return active_idx >= episode_k

def _parse_waypoints(self, target_states: list[PlanState]) -> tuple[
torch.Tensor,
torch.Tensor,
Expand Down Expand Up @@ -702,7 +881,11 @@ def _fk_matrix(self, qpos: torch.Tensor, control_part: str) -> torch.Tensor:

def _fk_pose_xyzw(self, qpos: torch.Tensor, control_part: str) -> torch.Tensor:
"""Return the policy-frame FK pose as ``xyz + xyzw``."""
fk = self._to_policy_frame(self._fk_matrix(qpos, control_part))
return self._policy_pose_xyzw(self._fk_matrix(qpos, control_part))

def _policy_pose_xyzw(self, fk_matrix: torch.Tensor) -> torch.Tensor:
"""Map ``(M, 4, 4)`` runtime TCP poses to policy-frame ``xyz + xyzw``."""
fk = self._to_policy_frame(fk_matrix)
pos = fk[:, :3, 3]
# ``quat_from_matrix`` is an EmbodiChain ``xyzw`` producer; converting
# it again would turn a valid pose into a different rotation.
Expand Down
Loading