Skip to content
Merged
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
95 changes: 48 additions & 47 deletions scienceworld/scienceworld.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from typing import List, Dict, Tuple, Set, Any
from typing import List, Dict, Tuple, Set, Any, Optional
from typing import OrderedDict as OrderedDictType
import json
import logging
Expand All @@ -20,7 +20,7 @@ class ScienceWorldEnv:
Please look at that for more information on the internals of the system.
"""

def __init__(self, taskName: str = None, serverPath: str = None, envStepLimit: int = 100):
def __init__(self, taskName: Optional[str] = None, serverPath: Optional[str] = None, envStepLimit: int = 100):
'''Start the simulator. Sets up the interface between python and the JVM.
Also does basic init stuff.
:param taskName: The name of the task. Will be run through the infer_task method.
Expand Down Expand Up @@ -70,7 +70,7 @@ def __init__(self, taskName: str = None, serverPath: str = None, envStepLimit: i

# Load the script
self.taskName = taskName
if self.taskName:
if taskName:
self.load(taskName, 0, "")

# Set the environment step limit
Expand Down Expand Up @@ -157,7 +157,7 @@ def close(self) -> None:
self._gateway.java_process.stdin.write("\n".encode("utf-8"))
self._gateway.java_process.stdin.flush()

def __del__(self):
def __del__(self) -> None:
self.close()

# Simplifications
Expand Down Expand Up @@ -190,7 +190,7 @@ def get_task_names(self) -> List[str]:
''' Get the name for the supported tasks in ScienceWorld. '''
return list(self.server.getTaskNames())

def get_max_variations(self, task_name) -> int:
def get_max_variations(self, task_name: str) -> int:
''' Get the maximum number of variations for the tasks. '''
return self.server.getTaskMaxVariations(infer_task(task_name))

Expand Down Expand Up @@ -289,7 +289,7 @@ def get_task_description(self) -> str:
return self.server.getTaskDescription()

# Get the current game's task description
def getObjectTree(self):
def getObjectTree(self) -> Dict[str, Any]:
msg = self.server.getObjectTree(self._obj_tree_tempdir.name)
if msg:
# Game is not initialized.
Expand Down Expand Up @@ -351,7 +351,7 @@ def get_run_history_size(self) -> int:

def clear_run_histories(self) -> None:
''' Clear the run histories. '''
self.runHistories = {}
self.runHistories: Dict[int, Dict[str, Any]] = {}

# A one-stop function to handle saving.
def save_run_histories_buffer_if_full(self, filename_out_prefix: str,
Expand Down Expand Up @@ -477,176 +477,177 @@ def get_goal_progress(self) -> str:
# All of the wrapper methods for camel case, to avoid breaking projects.

# Simplifications
def getSimplificationsUsed(self):
def getSimplificationsUsed(self) -> str:
snake_case_deprecation_warning()

return self.get_simplifications_used()

def getPossibleSimplifications(self):
def getPossibleSimplifications(self) -> List[str]:
snake_case_deprecation_warning()

return self.get_possible_simplifications()

def getTaskNames(self):
def getTaskNames(self) -> List[str]:
""" Get the name for the supported tasks in ScienceWorld. """
snake_case_deprecation_warning()

return self.get_task_names()

# Get the maximum number of variations for this task
def getMaxVariations(self, taskName):
def getMaxVariations(self, taskName: str) -> int:
snake_case_deprecation_warning()

return self.get_max_variations(taskName)

# Get possible actions
def getPossibleActions(self):
def getPossibleActions(self) -> List[str]:
snake_case_deprecation_warning()

return self.get_possible_actions()

# Get possible actions (and also include the template IDs for those actions)
def getPossibleActionsWithIDs(self):
def getPossibleActionsWithIDs(self) -> List[Dict[str, Any]]:
snake_case_deprecation_warning()

return self.get_possible_actions_with_IDs()

# Get possible objects
def getPossibleObjects(self):
def getPossibleObjects(self) -> List[str]:
snake_case_deprecation_warning()

return self.get_possible_objects()

# Get a list of object_ids to unique referents
def getPossibleObjectReferentLUT(self):
def getPossibleObjectReferentLUT(self) -> Dict[str, str]:
snake_case_deprecation_warning()

return self.get_possible_object_referent_LUT()

# As above, but dictionary is referenced by object type ID
def getPossibleObjectReferentTypesLUT(self):
def getPossibleObjectReferentTypesLUT(self) -> Dict[str, Dict[str, str]]:
snake_case_deprecation_warning()

return self.get_possible_object_referent_types_LUT()

# Get a list of *valid* agent-object combinations
def getValidActionObjectCombinations(self):
def getValidActionObjectCombinations(self) -> List[str]:
snake_case_deprecation_warning()

return self.get_valid_action_object_combinations()

def getValidActionObjectCombinationsWithTemplates(self):
def getValidActionObjectCombinationsWithTemplates(self) -> List[Dict[str, Any]]:
snake_case_deprecation_warning()

return self.get_valid_action_object_combinations_with_templates()

# Get a LUT of object_id to type_id
def getAllObjectTypesLUTJSON(self):
def getAllObjectTypesLUTJSON(self) -> Dict[str, str]:
snake_case_deprecation_warning()

return self.get_all_object_types_LUTJSON()

# Get a LUT of {object_id: {type_id, referent:[]} } tuples
def getAllObjectIdsTypesReferentsLUTJSON(self):
def getAllObjectIdsTypesReferentsLUTJSON(self) -> Dict[str, Dict[str, Any]]:
snake_case_deprecation_warning()

return self.get_all_object_ids_types_referents_LUTJSON()

# Get possible action/object combinations
def getPossibleActionObjectCombinations(self):
def getPossibleActionObjectCombinations(self) -> Tuple[List[Dict[str, Any]], Dict[str, str]]:
snake_case_deprecation_warning()

return self.get_possible_action_object_combinations()

# Get a list of object types and their IDs
def getObjectTypes(self):
def getObjectTypes(self) -> Dict[str, int]:
snake_case_deprecation_warning()

return self.get_object_types()

# Get the vocabulary of the model (at the current state)
def getVocabulary(self):
def getVocabulary(self) -> Set[str]:
snake_case_deprecation_warning()

return self.get_vocabulary()

def getNumMoves(self):
def getNumMoves(self) -> int:
snake_case_deprecation_warning()

return self.get_num_moves()

def getTaskDescription(self):
def getTaskDescription(self) -> str:
snake_case_deprecation_warning()

return self.get_task_description()

def getRunHistory(self):
def getRunHistory(self) -> Dict[str, Any]:
snake_case_deprecation_warning()

return self.get_run_history()

def storeRunHistory(self, episodeIdxKey, notes):
def storeRunHistory(self, episodeIdxKey: int, notes: str) -> None:
snake_case_deprecation_warning()

self.store_run_history(episodeIdxKey, notes)

def saveRunHistories(self, filenameOutPrefix):
def saveRunHistories(self, filenameOutPrefix: str) -> None:
snake_case_deprecation_warning()

self.save_run_histories(filenameOutPrefix)

def getRunHistorySize(self):
def getRunHistorySize(self) -> int:
snake_case_deprecation_warning()

return self.get_run_historySize()
return self.get_run_history_size()

def clearRunHistories(self):
def clearRunHistories(self) -> None:
snake_case_deprecation_warning()

self.clear_run_histories()

# A one-stop function to handle saving.
def saveRunHistoriesBufferIfFull(self, filenameOutPrefix, maxPerFile=1000, forceSave=False):
def saveRunHistoriesBufferIfFull(self, filenameOutPrefix: str, maxPerFile: int = 1000,
forceSave: bool = False) -> None:
snake_case_deprecation_warning()

self.save_run_histories_buffer_if_full(filenameOutPrefix, maxPerFile, forceSave)

def getVariationsTrain(self):
def getVariationsTrain(self) -> List[int]:
snake_case_deprecation_warning()

return self.get_variations_train()

def getVariationsDev(self):
def getVariationsDev(self) -> List[int]:
snake_case_deprecation_warning()

return self.get_variations_dev()

def getVariationsTest(self):
def getVariationsTest(self) -> List[int]:
snake_case_deprecation_warning()

return self.get_variations_test()

def getRandomVariationTrain(self):
def getRandomVariationTrain(self) -> int:
snake_case_deprecation_warning()

return self.get_random_variation_train()

def getRandomVariationDev(self):
def getRandomVariationDev(self) -> int:
snake_case_deprecation_warning()

return self.get_random_variation_dev()

def getRandomVariationTest(self):
def getRandomVariationTest(self) -> int:
snake_case_deprecation_warning()

return self.get_random_variation_test()

def getGoldActionSequence(self):
def getGoldActionSequence(self) -> List[str]:
snake_case_deprecation_warning()

return self.get_gold_action_sequence()

def getGoalProgressStr(self):
def getGoalProgressStr(self) -> str:
snake_case_deprecation_warning()

return self.get_goal_progress()
Expand All @@ -657,7 +658,7 @@ class BufferedHistorySaver:
#
# Constructor
#
def __init__(self, filenameOutPrefix):
def __init__(self, filenameOutPrefix: str):
self.filenameOutPrefix = filenameOutPrefix

# Clear the run histories
Expand All @@ -668,7 +669,7 @@ def __init__(self, filenameOutPrefix):
#

# History saving (provides an API to do this, so it's consistent across agents)
def storeRunHistory(self, runHistory, episodeIdxKey, notes):
def storeRunHistory(self, runHistory: Dict[str, Any], episodeIdxKey: int, notes: str) -> None:
packed = {
'episodeIdx': episodeIdxKey,
'notes': notes,
Expand All @@ -677,7 +678,7 @@ def storeRunHistory(self, runHistory, episodeIdxKey, notes):

self.runHistories[episodeIdxKey] = packed

def saveRunHistories(self):
def saveRunHistories(self) -> None:
# Save history

# Create verbose filename
Expand All @@ -695,14 +696,14 @@ def saveRunHistories(self):
with open(filenameOut, 'w') as outfile:
json.dump(self.runHistories, outfile, sort_keys=True, indent=4)

def getRunHistorySize(self):
def getRunHistorySize(self) -> int:
return len(self.runHistories)

def clearRunHistories(self):
self.runHistories = {}
def clearRunHistories(self) -> None:
self.runHistories: Dict[int, Dict[str, Any]] = {}

# A one-stop function to handle saving.
def saveRunHistoriesBufferIfFull(self, maxPerFile=1000, forceSave=False):
def saveRunHistoriesBufferIfFull(self, maxPerFile: int = 1000, forceSave: bool = False) -> None:
if ((self.getRunHistorySize() >= maxPerFile) or forceSave):
self.saveRunHistories()
self.clearRunHistories()
4 changes: 2 additions & 2 deletions scienceworld/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
from scienceworld.constants import NAME2ID, ID2TASK


def infer_task(name_or_id):
def infer_task(name_or_id: str) -> str:
''' Takes a task name or task ID and processes it to produce a uniform task format. '''

if name_or_id in NAME2ID:
Expand All @@ -22,7 +22,7 @@ def infer_task(name_or_id):
return name_or_id


def snake_case_deprecation_warning():
def snake_case_deprecation_warning() -> None:
message = "You are using the camel case api. This feature is deprecated. Please migrate to the snake_case api."
formatted_message = f"\033[91m {message} \033[00m"
warnings.warn(formatted_message, UserWarning, stacklevel=3)
5 changes: 5 additions & 0 deletions tox.ini
Original file line number Diff line number Diff line change
Expand Up @@ -25,3 +25,8 @@ commands = flake8 --max-line-length 120
[testenv:precommit]
deps = pre-commit
commands = pre-commit run --all-files

# Not part of the default envlist: mypy is opt-in for now, run explicitly with `tox -e mypy`.
[testenv:mypy]
deps = mypy
commands = mypy scienceworld
Loading