diff --git a/rlbench/tasks/place_cups.py b/rlbench/tasks/place_cups.py index fd8c53838..d74c26239 100644 --- a/rlbench/tasks/place_cups.py +++ b/rlbench/tasks/place_cups.py @@ -1,10 +1,11 @@ +from itertools import combinations from typing import List, Tuple import numpy as np from pyrep.objects.dummy import Dummy from pyrep.objects.proximity_sensor import ProximitySensor from pyrep.objects.shape import Shape from rlbench.backend.conditions import DetectedCondition, NothingGrasped, \ - OrConditions + OrConditions, ConditionSet from rlbench.backend.spawn_boundary import SpawnBoundary from rlbench.backend.task import Task @@ -32,8 +33,12 @@ def init_episode(self, index: int) -> List[str]: self._index = index b = SpawnBoundary([self._cups_boundary]) [b.sample(c, min_distance=0.10) for c in self._cups] - success_conditions = [NothingGrasped(self.robot.gripper) - ] + self._on_peg_conditions[:index + 1] + # Descriptions specify a count, not particular cup identities. + success_conditions = [ + NothingGrasped(self.robot.gripper), + OrConditions([ConditionSet(list(cups)) for cups in + combinations(self._on_peg_conditions, index + 1)]) + ] self.register_success_conditions(success_conditions) self.register_waypoint_ability_start( 0, self._move_above_next_target) diff --git a/tests/unit/test_place_cups_conditions.py b/tests/unit/test_place_cups_conditions.py new file mode 100644 index 000000000..2f8e8c59e --- /dev/null +++ b/tests/unit/test_place_cups_conditions.py @@ -0,0 +1,38 @@ +"""Count distinct cups rather than requiring an unmentioned cup-ID prefix.""" +import itertools +import unittest +from types import SimpleNamespace +from unittest.mock import Mock, patch + +from rlbench.backend.conditions import DetectedCondition, OrConditions +from rlbench.tasks.place_cups import PlaceCups + + +class TestPlaceCupsConditions(unittest.TestCase): + def test_all_sensor_states_and_variations(self): + cups = [object() for _ in range(3)] + gripper = Mock() + task = PlaceCups(None, SimpleNamespace(gripper=gripper)) + task._cups = cups + task._cups_boundary = object() + sensors = [Mock() for _ in range(3)] + task._on_peg_conditions = [OrConditions([ + DetectedCondition(cup, sensor) for sensor in sensors + ]) for cup in cups] + with patch('rlbench.tasks.place_cups.SpawnBoundary'): + for variation in range(3): + task.init_episode(variation) + for bits in itertools.product((False, True), repeat=9): + matrix = [bits[i * 3:(i + 1) * 3] for i in range(3)] + for j, sensor in enumerate(sensors): + sensor.is_detected.side_effect = ( + lambda cup, j=j: matrix[cups.index(cup)][j]) + for held in (False, True): + gripper.get_grasped_objects.return_value = [cups[0]] if held else [] + expected = not held and sum(map(any, matrix)) >= variation + 1 + with self.subTest(variation=variation, bits=bits, held=held): + self.assertEqual(bool(task.success()[0]), expected) + + +if __name__ == '__main__': + unittest.main()