Skip to content

Commit 9e2f89a

Browse files
committed
General env wrapper for textworld, textworld express and alfworld added
1 parent 5d81e14 commit 9e2f89a

2 files changed

Lines changed: 126 additions & 2 deletions

File tree

‎tales/get_env_splits.py‎

Lines changed: 124 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,124 @@
1+
# This is literally just a wrapper to get the train and test-time splits. Is 99.99% just building on Marc's existing code.
2+
import glob
3+
from os.path import join as pjoin
4+
from tales.textworld import textworld_data, textworld_env
5+
from tales.textworld_express import twx_data, twx_env
6+
from tales.alfworld import alfworld_data, alfworld_env
7+
8+
9+
def get_textworld_env_splits(difficulties = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10], games_per_difficulty=1):
10+
# Returns a list of envs for training and test splits for Textworld-Cookingworld:
11+
# For training, we let the user specify difficulties and how many games per difficulty to include.
12+
# For testing, we use all difficulties from 1 to 10, and use one game each, similar to the evaluation in the original paper.
13+
textworld_data.prepare_twcooking_data() # make sure the data is ready
14+
15+
# Training split:
16+
# Get the game files:
17+
train_games_files = []
18+
for diff in difficulties:
19+
all_games = sorted(textworld_data.get_cooking_game(diff, split="train"))
20+
train_games_files.extend(all_games[:games_per_difficulty])
21+
22+
# Testing split:
23+
test_games_files = []
24+
for i in range(1, 11):
25+
# Just get one game per difficulty for testing.
26+
# This is similar to the evaluation in the original paper.
27+
all_games = sorted(textworld_data.get_cooking_game(i, split="test"))
28+
test_games_files.append(all_games[0])
29+
30+
return train_games_files, test_games_files
31+
32+
def get_alfworld_env_splits(games_per_task = 2):
33+
# For alfworld, we just generate the test split first and then condition the train split to not have the same files as the text split.
34+
alfworld_data.prepare_alfworld_data() # make sure the data is ready
35+
test_games_files = []
36+
for task in alfworld_data.TASK_TYPES:
37+
game_files_seen = sorted(glob.glob(pjoin(alfworld_data.TALES_CACHE_ALFWORLD_VALID_SEEN, f"{task}*", "**", "*.tw-pddl")))
38+
game_files_unseen = sorted(glob.glob(pjoin(alfworld_data.TALES_CACHE_ALFWORLD_VALID_UNSEEN, f"{task}*", "**", "*.tw-pddl")))
39+
# The test split always only takes the first game file in the split.
40+
test_games_files.extend(game_files_seen[[0]])
41+
test_games_files.extend(game_files_unseen[[0]])
42+
43+
# Assert we have the right number of files.
44+
assert len(test_games_files) == 2 * len(alfworld_data.TASK_TYPES)
45+
46+
# Now, get the training split.
47+
# We want to make sure that the training split does not have any files that are in the test split.
48+
train_games_files = []
49+
for task in alfworld_data.TASK_TYPES:
50+
game_files_seen = sorted(glob.glob(pjoin(alfworld_data.TALES_CACHE_ALFWORLD_VALID_SEEN, f"{task}*", "**", "*.tw-pddl")))
51+
game_files_unseen = sorted(glob.glob(pjoin(alfworld_data.TALES_CACHE_ALFWORLD_VALID_UNSEEN, f"{task}*", "**", "*.tw-pddl")))
52+
# Remove any files that are in the test split.
53+
filtered_game_files_seen = [f for f in game_files_seen if not any(s in f for s in test_games_files)]
54+
filtered_game_files_unseen = [f for f in game_files_unseen if not any(s in f for s in test_games_files)]
55+
56+
# Now get the requested number of games per task type
57+
train_games_files.extend(filtered_game_files_seen[:games_per_task])
58+
train_games_files.extend(filtered_game_files_unseen[:games_per_task])
59+
60+
return train_games_files, test_games_files
61+
62+
class GeneralTALESEnv:
63+
# A general env wrapper such that the train/test files gotten from the above functions can easily just be plugged into an env and ran.
64+
# This returns a 'fake' batch env that will always deterministically cycle through the provided env file/seeds unless explicitly told to randomize (for training)
65+
# TODO: implement for Scienceworld and Jericho
66+
def __init__(self, env_name, split, *args, **kwargs):
67+
self.env_name = env_name
68+
self.split = split
69+
self.env_idx = 0
70+
self.kwargs = kwargs
71+
self.args = args
72+
self.game_files = None
73+
if env_name == "textworld":
74+
self.train_envs, self.test_envs = get_textworld_env_splits(**kwargs)
75+
if split == "train":
76+
self.game_files = self.train_envs
77+
else:
78+
self.game_files = self.test_envs
79+
self.env = textworld_env.TextWorldEnv(self.game_files[self.env_idx],
80+
*args, **kwargs)
81+
elif env_name == "twx":
82+
# Train/test in twx are just seed based.
83+
self.game_files = twx_data.TASKS
84+
self.env = twx_env.TextWorldExpressEnv(game_name = self.game_files[self.env_idx][1],
85+
game_params = self.game_files[self.env_idx][2],
86+
admissible_commands=False,
87+
split=split,
88+
*args, **kwargs)
89+
elif env_name == "alfworld":
90+
self.train_envs, self.test_envs = get_alfworld_env_splits(**kwargs)
91+
if split == "train":
92+
self.game_files = self.train_envs
93+
else:
94+
self.game_files = self.test_envs
95+
self.env = alfworld_env.ALFWorldEnv(self.game_files[self.env_idx],
96+
*args, **kwargs)
97+
else:
98+
raise ValueError(f"Unknown environment name: {env_name}, please choose from textworld, twx, or alfworld.")
99+
100+
# Not sure if this is right, need to double check w/ Marc
101+
def reset(self, *, seed=None, options=None):
102+
return self.env.reset(seed=seed, options=options)
103+
104+
def get_next_task(self, seed = None, options=None):
105+
# Move to the next env in the list.
106+
self.env_idx = (self.env_idx + 1) % len(self.game_files)
107+
if self.env is not None:
108+
self.env.close()
109+
if self.env_name == "textworld":
110+
self.env = textworld_env.TextWorldEnv(self.game_files[self.env_idx], *self.args, **self.kwargs)
111+
elif self.env_name == "twx":
112+
self.env = twx_env.TextWorldExpressEnv(game_name = self.game_files[self.env_idx][1],
113+
game_params = self.game_files[self.env_idx][2],
114+
*self.args, **self.kwargs)
115+
elif self.env_name == "alfworld":
116+
self.env = alfworld_env.ALFWorldEnv(self.game_files[self.env_idx], *self.args, **self.kwargs)
117+
else:
118+
raise ValueError(f"next_task not implemented for env {self.env_name}, only for textworld and alfworld.")
119+
return self.reset(seed = seed, options = options)
120+
121+
def step(self, action):
122+
return self.env.step(action)
123+
124+

‎tales/textworld_express/twx_env.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -10,13 +10,13 @@
1010
class TextWorldExpressEnv(gym.Env):
1111

1212
def __init__(
13-
self, game_name, game_params, admissible_commands=False, *args, **kwargs
13+
self, game_name, game_params, admissible_commands=False, split="test", *args, **kwargs
1414
):
1515
self.game_name = game_name
1616
self.game_params = game_params
1717
self.admissible_commands = admissible_commands
1818
self.env = twx.TextWorldExpressEnv(envStepLimit=np.inf)
19-
self.seeds = twx_data.get_seeds(split="test", env=self.env)
19+
self.seeds = twx_data.get_seeds(split=split, env=self.env)
2020
self.seed = self.seeds[0]
2121

2222
def reset(self, *, seed=None, options=None):

0 commit comments

Comments
 (0)