Source code for embodichain.lab.sim.atomic_actions.affordance

# ----------------------------------------------------------------------------
# Copyright (c) 2021-2026 DexForce Technology Co., Ltd.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ----------------------------------------------------------------------------

from __future__ import annotations

import torch
from dataclasses import dataclass, field
from typing import Any

from embodichain.toolkits.graspkit.pg_grasp import (
    GraspGenerator,
    GraspGeneratorCfg,
)
from embodichain.toolkits.graspkit.pg_grasp.gripper_collision_checker import (
    GripperCollisionCfg,
)
from embodichain.utils import logger


[docs] @dataclass class Affordance: """Base class for affordance data. Represents an object's interaction possibilities. Subclasses carry whatever typed fields they need (mesh tensors, interaction points, etc.); the base class only carries an object label and a free-form custom_config dict. """ object_label: str = "" """Label of the object this affordance belongs to.""" custom_config: dict[str, Any] = field(default_factory=dict) """User-defined configuration payload."""
[docs] def set_custom_config(self, key: str, value: Any) -> None: """Set a custom affordance configuration value.""" self.custom_config[key] = value
[docs] def get_custom_config(self, key: str, default: Any = None) -> Any: """Get a custom affordance configuration value.""" return self.custom_config.get(key, default)
[docs] def get_batch_size(self) -> int: """Return the batch size of this affordance data.""" return 1
[docs] @dataclass class AntipodalAffordance(Affordance): """Antipodal grasp affordance for parallel-jaw grippers.""" mesh_vertices: torch.Tensor | None = None """Object mesh vertices, shape [N, 3].""" mesh_triangles: torch.Tensor | None = None """Object mesh triangle indices, shape [M, 3].""" generator_cfg: GraspGeneratorCfg | None = None """Optional grasp-generator configuration.""" gripper_collision_cfg: GripperCollisionCfg | None = None """Optional gripper-collision configuration.""" force_reannotate: bool = False """If True, recompute the grasp annotation on each access.""" _generator: GraspGenerator | None = field(default=None, init=False, repr=False) def _init_generator(self) -> None: if self.mesh_vertices is None or self.mesh_triangles is None: logger.log_error( "mesh_vertices and mesh_triangles must be provided to initialize " "AntipodalAffordance.", ValueError, ) self._generator = GraspGenerator( vertices=self.mesh_vertices, triangles=self.mesh_triangles, cfg=self.generator_cfg, gripper_collision_cfg=self.gripper_collision_cfg, ) if self.force_reannotate or self._generator._hit_point_pairs is None: self._generator.annotate() def _resolve_approach_direction( self, approach_direction: torch.Tensor ) -> torch.Tensor: """Move the approach direction to the grasp generator device.""" return approach_direction.to( device=self._generator.device, dtype=torch.float32, ) def get_valid_grasp_poses( self, obj_poses: torch.Tensor, approach_direction: torch.Tensor = torch.tensor( [0, 0, -1], dtype=torch.float32 ), ) -> list[tuple[torch.Tensor, torch.Tensor]]: if self._generator is None: self._init_generator() approach_direction = self._resolve_approach_direction(approach_direction) results = [] for i, obj_pose in enumerate(obj_poses): is_success, grasp_poses, _, costs = self._generator.get_valid_grasp_poses( obj_pose, approach_direction ) if grasp_poses.shape == (4, 4): grasp_poses = grasp_poses.unsqueeze(0) if costs.dim() == 0: costs = costs.unsqueeze(0) if not is_success: logger.log_warning( f"Failed to find valid grasp poses for {i}-th object." ) costs = torch.full( (grasp_poses.shape[0],), torch.inf, dtype=torch.float32, device=grasp_poses.device, ) results.append((grasp_poses, costs)) return results def get_best_grasp_poses( self, obj_poses: torch.Tensor, approach_direction: torch.Tensor = torch.tensor( [0, 0, -1], dtype=torch.float32 ), ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: if self._generator is None: self._init_generator() approach_direction = self._resolve_approach_direction(approach_direction) grasp_xpos_list: list[torch.Tensor] = [] is_success_list: list[bool] = [] open_length_list: list[float] = [] for i, obj_pose in enumerate(obj_poses): is_success, grasp_xpos, open_length = self._generator.get_grasp_poses( obj_pose, approach_direction ) if is_success: grasp_xpos_list.append(grasp_xpos.unsqueeze(0)) else: logger.log_warning(f"No valid grasp pose found for {i}-th object.") grasp_xpos_list.append( torch.eye( 4, dtype=torch.float32, device=self._generator.device ).unsqueeze(0) ) is_success_list.append(is_success) open_length_list.append(open_length) is_success_t = torch.tensor( is_success_list, dtype=torch.bool, device=self._generator.device ) grasp_xpos = torch.concatenate(grasp_xpos_list, dim=0) open_length_t = torch.tensor( open_length_list, dtype=torch.float32, device=self._generator.device ) return is_success_t, grasp_xpos, open_length_t
[docs] @dataclass class InteractionPoints(Affordance): """Batch of 3D interaction points on an object surface.""" points: torch.Tensor = field(default_factory=lambda: torch.zeros(1, 3)) """Batch of 3D interaction points with shape [B, 3].""" normals: torch.Tensor | None = None """Optional surface normals at each interaction point with shape [B, 3].""" point_types: list[str] = field(default_factory=list) """Optional labels for each point's interaction type."""
[docs] def get_points_by_type(self, point_type: str) -> torch.Tensor | None: """Get points by their interaction type.""" if point_type in self.point_types: indices = [i for i, t in enumerate(self.point_types) if t == point_type] return self.points[indices] return None
[docs] def get_batch_size(self) -> int: """Return the number of interaction points in this affordance.""" return self.points.shape[0]
[docs] def get_approach_direction(self, point_idx: int) -> torch.Tensor: """Get recommended approach direction for a given point.""" if self.normals is not None: return -self.normals[point_idx] return torch.tensor( [0, 0, 1], dtype=self.points.dtype, device=self.points.device )
__all__ = ["Affordance", "AntipodalAffordance", "InteractionPoints"]