Skip to content

open_state

Open

Bases: TensorizedAbsoluteState, BooleanStateMixin

Tensorized Open state.

VALUES shape: (S, O) — bool True = open, False = closed.

Source code in OmniGibson/omnigibson/object_states/open_state.py
class Open(TensorizedAbsoluteState, BooleanStateMixin):
    """
    Tensorized Open state.

    VALUES shape: (S, O) — bool True = open, False = closed.
    """

    # wp.array3d (S, O, max_dof) uint8 — 1 for DOF columns corresponding to openable joints.
    OPENABLE_MASK = None

    # wp.array3d (S, O, max_dof) float32 — per-DOF threshold/direction for side=+1.
    THRESHOLDS_S1 = None
    DIRECTIONS_S1 = None

    # wp.array3d (S, O, max_dof) float32 — per-DOF threshold/direction for side=-1 (0 for non-both_sides).
    THRESHOLDS_S2 = None
    DIRECTIONS_S2 = None

    # wp.array2d (S, O) uint8 — whether each object uses both-sides open logic.
    BOTH_SIDES = None

    # wp.array2d (S, O) int32 — row indices into ArticulatedObjectViewAPI._JOINT_POSITIONS.
    OBJ_IDXES_IN_ARTICULATION_VIEW = None

    # Cached max_dof int (kernel input). Updated in initialize_view.
    _MAX_DOF = 0

    @classproperty
    def value_name(cls):
        return "open"

    @classproperty
    def value_type(cls):
        return th.bool

    @classmethod
    def is_compatible(cls, obj, **kwargs):
        # Run super first
        compatible, reason = super().is_compatible(obj, **kwargs)
        if not compatible:
            return compatible, reason

        # Check whether this object has any openable joints
        return (
            (True, None)
            if obj.n_joints > 0
            else (False, f"No relevant joints for Open state found for object {obj.name}")
        )

    @classmethod
    def is_compatible_asset(cls, prim, **kwargs):
        # Run super first
        compatible, reason = super().is_compatible_asset(prim, **kwargs)
        if not compatible:
            return compatible, reason

        def _find_articulated_joints(prim):
            for child in prim.GetChildren():
                child_type = child.GetTypeName().lower()
                if "joint" in child_type and "fixed" not in child_type:
                    return True
                for gchild in child.GetChildren():
                    gchild_type = gchild.GetTypeName().lower()
                    if "joint" in gchild_type and "fixed" not in gchild_type:
                        return True
            return False

        # Check whether this object has any openable joints
        return (
            (True, None)
            if _find_articulated_joints(prim=prim)
            else (False, f"No relevant joints for Open state found for asset prim {prim.GetName()}")
        )

    @classmethod
    def initialize_view(cls):
        super().initialize_view()  # builds OBJ_IDXS, IDX_OBJS, VALUES (S, O)

        S = len(cls.IDX_OBJS)
        O = len(cls.OBJ_IDXS)

        if O == 0:
            cls._MAX_DOF = 0
            cls.OPENABLE_MASK = None
            cls.THRESHOLDS_S1 = None
            cls.DIRECTIONS_S1 = None
            cls.THRESHOLDS_S2 = None
            cls.DIRECTIONS_S2 = None
            cls.BOTH_SIDES = None
            cls.OBJ_IDXES_IN_ARTICULATION_VIEW = None
            return

        max_dof = ArticulatedObjectViewAPI.get_max_dof()
        cls._MAX_DOF = int(max_dof)

        # Build per-DOF tables on CPU first, then upload once. Per-cell GPU writes via torch
        # indexing would round-trip through CUDA per cell — CPU scratch + one bulk upload is faster.
        openable_mask_cpu = th.zeros((S, O, max_dof), dtype=th.uint8)
        thresholds_s1_cpu = th.zeros((S, O, max_dof), dtype=th.float32)
        directions_s1_cpu = th.zeros((S, O, max_dof), dtype=th.float32)
        thresholds_s2_cpu = th.zeros((S, O, max_dof), dtype=th.float32)
        directions_s2_cpu = th.zeros((S, O, max_dof), dtype=th.float32)
        both_sides_cpu = th.zeros((S, O), dtype=th.uint8)

        for scene_idx, scene in enumerate(cls.IDX_OBJS):
            for obj_idx in range(O):
                obj = scene[obj_idx]
                if obj is None or obj.joints is None:
                    continue  # obj not initialized yet
                both_sides, relevant_joints, joint_directions = _get_relevant_joints(obj)
                both_sides_cpu[scene_idx, obj_idx] = 1 if both_sides else 0
                for joint, direction in zip(relevant_joints, joint_directions):
                    for dof_col in joint.dof_indices:
                        openable_mask_cpu[scene_idx, obj_idx, dof_col] = 1
                        for side, threshold_attr, direction_attr in [
                            (1, thresholds_s1_cpu, directions_s1_cpu),
                            (-1, thresholds_s2_cpu, directions_s2_cpu),
                        ]:
                            threshold, open_end, _ = _compute_joint_threshold(joint, direction * side)
                            threshold_attr[scene_idx, obj_idx, dof_col] = threshold
                            direction_attr[scene_idx, obj_idx, dof_col] = 1.0 if open_end > threshold else -1.0

        # (S, O) row index into ArticulatedObjectViewAPI._JOINT_POSITIONS.
        obj_view_rows_cpu = th.zeros((S, O), dtype=th.int32)
        for _, obj_idx in cls.OBJ_IDXS.items():
            for scene_idx, scene in enumerate(cls.IDX_OBJS):
                obj = scene[obj_idx]
                if obj is None:
                    continue
                row = ArticulatedObjectViewAPI.get_view_row(obj.articulation_root_path)
                obj_view_rows_cpu[scene_idx, obj_idx] = row

        cls.OPENABLE_MASK = lazy.isaacsim.core.utils.warp.tensor.create_tensor_from_list(
            openable_mask_cpu, "uint8", device="cuda"
        )
        cls.THRESHOLDS_S1 = lazy.isaacsim.core.utils.warp.tensor.create_tensor_from_list(
            thresholds_s1_cpu, "float32", device="cuda"
        )
        cls.DIRECTIONS_S1 = lazy.isaacsim.core.utils.warp.tensor.create_tensor_from_list(
            directions_s1_cpu, "float32", device="cuda"
        )
        cls.THRESHOLDS_S2 = lazy.isaacsim.core.utils.warp.tensor.create_tensor_from_list(
            thresholds_s2_cpu, "float32", device="cuda"
        )
        cls.DIRECTIONS_S2 = lazy.isaacsim.core.utils.warp.tensor.create_tensor_from_list(
            directions_s2_cpu, "float32", device="cuda"
        )
        cls.BOTH_SIDES = lazy.isaacsim.core.utils.warp.tensor.create_tensor_from_list(
            both_sides_cpu, "uint8", device="cuda"
        )
        cls.OBJ_IDXES_IN_ARTICULATION_VIEW = lazy.isaacsim.core.utils.warp.tensor.create_tensor_from_list(
            obj_view_rows_cpu, "int32", device="cuda"
        )

    @classmethod
    def _update_values(cls, values):
        if cls.OPENABLE_MASK is None or cls._MAX_DOF == 0:
            return
        if ArticulatedObjectViewAPI._JOINT_POSITIONS is None:
            return
        S, O = values.shape[:2]
        wp.launch(
            kernel=_open_update_kernel,
            dim=(S, O),
            inputs=[
                ArticulatedObjectViewAPI._JOINT_POSITIONS,
                cls.OBJ_IDXES_IN_ARTICULATION_VIEW,
                cls.OPENABLE_MASK,
                cls.THRESHOLDS_S1,
                cls.DIRECTIONS_S1,
                cls.THRESHOLDS_S2,
                cls.DIRECTIONS_S2,
                cls.BOTH_SIDES,
                wp.int32(cls._MAX_DOF),
                cls.VALUES_WP,
            ],
            device="cuda",
        )

    def _get_value(self):
        return bool(super()._get_value())

    def _set_value(self, new_value, fully=False):
        """
        Set the openness state, either to a random joint position satisfying the new value, or fully open/closed.

        Args:
            new_value (bool): The new value for the openness state of the object.
            fully (bool): Whether the object should be fully opened/closed (e.g. all relevant joints to 0/1).

        Returns:
            bool: A boolean indicating the success of the setter. Failure may happen due to unannotated objects.
        """
        both_sides, relevant_joints, joint_directions = _get_relevant_joints(self.obj)
        if not relevant_joints:
            return False

        # The "sides" variable is used to check open/closed state for objects whose joints can switch
        # positions. These objects are annotated with the both_sides annotation and the idea is that switching
        # the directions of *all* of the joints results in a similarly valid checkable state. We want our object to be
        # open from *both* of the two sides, and I was too lazy to implement the logic for this without rejection
        # sampling, so that's what we do.
        # TODO: Implement a sampling method that's guaranteed to be correct, ditch the rejection method.
        sides = [1, -1] if both_sides else [1]

        for _ in range(m.OPEN_SAMPLING_ATTEMPTS):
            side = random.choice(sides)

            joints_to_set = list(zip(relevant_joints, joint_directions))

            # All joints are relevant if we are closing, but if we are opening let's sample a subset.
            if new_value and not fully:
                num_to_open = th.randint(1, len(relevant_joints) + 1, (1,)).item()
                random_indices = th.randperm(len(relevant_joints))[:num_to_open]
                joints_to_set = [joints_to_set[i] for i in random_indices]

            # Go through the relevant joints & set random positions.
            for joint, joint_direction in joints_to_set:
                threshold, open_end, closed_end = _compute_joint_threshold(joint, joint_direction * side)

                # Get the range
                if new_value:
                    joint_range = (threshold, open_end)
                else:
                    joint_range = (threshold, closed_end)

                if fully:
                    joint_pos = joint_range[1]
                else:
                    # Convert the range to the format numpy accepts.
                    low = min(joint_range)
                    high = max(joint_range)

                    # Sample a position.
                    joint_pos = (th.rand(1) * (high - low) + low).item()

                # Save sampled position. JointPrim.set_pos flips TensorizedState.caches_dirty,
                # so the next get_value() below will trigger a refresh and observe the new pose.
                joint.set_pos(joint_pos)

            if self.get_value() == new_value:
                return True

        # We exhausted our attempts and could not find a working sample.
        return False

    # We don't need to load / save anything since the joints are saved elsewhere.
    # Overriding the inherited TensorizedAbsoluteState._load_state which would
    # call self._set_value(stored_value)
    def _load_state(self, state):
        return