diff --git a/scienceworld/scienceworld.jar b/scienceworld/scienceworld.jar index d1967fa..3b9d980 100644 Binary files a/scienceworld/scienceworld.jar and b/scienceworld/scienceworld.jar differ diff --git a/scripts/reproduce_issue_82.py b/scripts/reproduce_issue_82.py new file mode 100644 index 0000000..bf8ae58 --- /dev/null +++ b/scripts/reproduce_issue_82.py @@ -0,0 +1,11 @@ +from scienceworld import ScienceWorldEnv + + +env = ScienceWorldEnv("1-1") +try: + _, info = env.reset() + for action in info["valid"]: + if action.startswith("open ") and "door" in action: + print(action) +finally: + env.close() diff --git a/simulator/src/main/scala/scienceworld/input/InputParser.scala b/simulator/src/main/scala/scienceworld/input/InputParser.scala index fa7782c..9ea87da 100644 --- a/simulator/src/main/scala/scienceworld/input/InputParser.scala +++ b/simulator/src/main/scala/scienceworld/input/InputParser.scala @@ -5,6 +5,7 @@ import language.runtime.runners.{ActionRunner, PredicateRunner} import language.struct.{DynamicValue, ScopedVariableLUT} import scienceworld.actions.Action import scienceworld.objects.agent.Agent +import scienceworld.objects.portal.Portal import scienceworld.struct.EnvObject import scienceworld.tasks.goals.{GoalSequence, ObjMonitor} import util.UniqueTypeID @@ -114,7 +115,10 @@ class InputParser(actionRequestDefs:Array[ActionRequestDef]) { // Step 2A: Populate an array of the unique referents (as strings) val out = new ArrayBuffer[(String, EnvObject)]() for (i <- 0 until allObjs.length) { - val referent = uniqueReferents(i) + val referent = allObjs(i) match { + case portal:Portal => portal.getCanonicalReferent(perspectiveContainer) + case _ => uniqueReferents(i) + } out.append( (referent.toLowerCase(), allObjs(i)) ) } diff --git a/simulator/src/main/scala/scienceworld/objects/portal/Portal.scala b/simulator/src/main/scala/scienceworld/objects/portal/Portal.scala index bd7dea8..ab5b915 100644 --- a/simulator/src/main/scala/scienceworld/objects/portal/Portal.scala +++ b/simulator/src/main/scala/scienceworld/objects/portal/Portal.scala @@ -103,6 +103,15 @@ class Portal (val _isOpen:Boolean, val connectsFrom:EnvObject, val connectsTo:En return Set(this.name, this.name + " from " + connectsFrom.name + " to " + connectsTo.name, this.name + " from " + connectsTo.name + " to " + connectsFrom.name) } + def getCanonicalReferent(perspectiveContainer:EnvObject):String = { + val connectsToContainer = this.getConnectsTo(perspectiveContainer) + if (connectsToContainer.isDefined) { + return this.name + " to " + connectsToContainer.get.name + } + + return this.name + " from " + connectsFrom.name + " to " + connectsTo.name + } + override def getDescriptName(overrideName: String): String = { return "door between " + this.connectsFrom.name + " and " + this.connectsTo.name } diff --git a/tests/test_scienceworld.py b/tests/test_scienceworld.py index de5bd92..8acfeef 100644 --- a/tests/test_scienceworld.py +++ b/tests/test_scienceworld.py @@ -121,6 +121,31 @@ def test_multiple_instances(): assert obs1_2 == obs2_2 +def test_door_actions_use_canonical_referents(): + env = ScienceWorldEnv("1-1") + try: + _, info = env.reset() + door_actions = { + action for action in info["valid"] + if action.startswith("open ") and "door" in action + } + assert door_actions == { + "open door to art studio", + "open door to bedroom", + "open door to greenhouse", + "open door to kitchen", + "open door to living room", + "open door to workshop", + } + + for action in ("open bedroom door", "open door to bedroom"): + env.reset() + observation, _, _, _ = env.step(action) + assert observation == "The door is now open." + finally: + env.close() + + def test_closing_env(): env = ScienceWorldEnv() env.task_names # Load task names.