summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/babyai_shared_run.py120
-rw-r--r--experiments/babyai_shared_smoke.py26
2 files changed, 113 insertions, 33 deletions
diff --git a/experiments/babyai_shared_run.py b/experiments/babyai_shared_run.py
index 6b6ddd2..0b7c72a 100644
--- a/experiments/babyai_shared_run.py
+++ b/experiments/babyai_shared_run.py
@@ -16,8 +16,9 @@ import torch
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sdil.babyai_shared import (
- BabyAISharedConfig, BabyAISharedNet, CONDITIONS, encode_visual,
- evaluate_actions, manual_step, missions_to_bow, visual_input_dim,
+ BabyAISharedConfig, BabyAISharedNet, CONDITIONS, build_history_index,
+ encode_history_visual, encode_visual, evaluate_actions, history_input_dim,
+ manual_step, missions_to_bow, visual_input_dim,
)
@@ -35,16 +36,16 @@ def load_data(path):
archive = np.load(path, allow_pickle=False)
metadata = json.loads(str(archive["metadata_json"]))
train = {name: archive[f"train_{name}"] for name in (
- "image", "direction", "mission_bow", "action")}
+ "image", "direction", "mission_bow", "action", "episode_offset")}
validation = {name: archive[f"validation_{name}"] for name in (
- "image", "direction", "mission_bow", "action")}
+ "image", "direction", "mission_bow", "action", "episode_offset")}
rollout_seeds = archive["rollout_seed"].copy()
return metadata, train, validation, rollout_seeds
@torch.no_grad()
def rollout_policy(net, env_id, seeds, vocabulary, cardinalities,
- context_enabled=True, chunk_size=64):
+ context_enabled=True, chunk_size=64, history_steps=1):
successes = 0
returns = []
lengths = []
@@ -52,6 +53,8 @@ def rollout_policy(net, env_id, seeds, vocabulary, cardinalities,
chunk = seeds[chunk_start:chunk_start + chunk_size]
envs = [gym.make(env_id) for _ in chunk]
observations = []
+ observation_histories = []
+ previous_action_histories = []
active = []
chunk_returns = [0.0] * len(envs)
chunk_lengths = [0] * len(envs)
@@ -59,20 +62,51 @@ def rollout_policy(net, env_id, seeds, vocabulary, cardinalities,
for index, (env, seed) in enumerate(zip(envs, chunk)):
observation, _ = env.reset(seed=int(seed))
observations.append(observation)
+ observation_histories.append([observation])
+ previous_action_histories.append([-1])
active.append(index)
while active:
- images = np.asarray([
- observations[index]["image"] for index in active],
- dtype=np.uint8)
- directions = np.asarray([
- observations[index]["direction"] for index in active],
- dtype=np.uint8)
+ if history_steps == 1:
+ images = np.asarray([
+ observations[index]["image"] for index in active],
+ dtype=np.uint8)
+ directions = np.asarray([
+ observations[index]["direction"] for index in active],
+ dtype=np.uint8)
+ else:
+ image_shape = observations[active[0]]["image"].shape
+ images = np.zeros(
+ (len(active), history_steps, *image_shape),
+ dtype=np.uint8)
+ directions = np.zeros(
+ (len(active), history_steps), dtype=np.uint8)
+ previous_actions = np.full(
+ (len(active), history_steps), -1, dtype=np.int64)
+ history_mask = np.zeros(
+ (len(active), history_steps), dtype=bool)
+ for row, index in enumerate(active):
+ obs_history = observation_histories[index][-history_steps:]
+ action_history = previous_action_histories[index][
+ -history_steps:]
+ offset = history_steps - len(obs_history)
+ images[row, offset:] = np.asarray([
+ value["image"] for value in obs_history])
+ directions[row, offset:] = np.asarray([
+ value["direction"] for value in obs_history])
+ previous_actions[row, offset:] = action_history
+ history_mask[row, offset:] = True
mission_text = [
observations[index]["mission"] for index in active]
mission_bow = missions_to_bow(mission_text, vocabulary)
- features = encode_visual(
- images, directions, *cardinalities,
- device=net.device, dtype=net.dtype)
+ if history_steps == 1:
+ features = encode_visual(
+ images, directions, *cardinalities,
+ device=net.device, dtype=net.dtype)
+ else:
+ features = encode_history_visual(
+ images, directions, previous_actions, history_mask,
+ *cardinalities, net.config.action_dim,
+ device=net.device, dtype=net.dtype)
mission_tensor = torch.as_tensor(
mission_bow, device=net.device, dtype=net.dtype)
actions = net.forward_features(
@@ -84,6 +118,9 @@ def rollout_policy(net, env_id, seeds, vocabulary, cardinalities,
observation, reward, terminated, truncated, _ = (
envs[index].step(int(actions[position])))
observations[index] = observation
+ observation_histories[index].append(observation)
+ previous_action_histories[index].append(
+ int(actions[position]))
chunk_returns[index] += float(reward)
chunk_lengths[index] += 1
if terminated or truncated:
@@ -104,6 +141,21 @@ def rollout_policy(net, env_id, seeds, vocabulary, cardinalities,
}
+def encode_compact_batch(split, indices, cardinalities, net, history):
+ if history is None:
+ return encode_visual(
+ split["image"][indices], split["direction"][indices],
+ *cardinalities, device=net.device, dtype=net.dtype)
+ history_indices, previous_actions, mask = history
+ selected_history = history_indices[indices]
+ return encode_history_visual(
+ split["image"][selected_history],
+ split["direction"][selected_history],
+ previous_actions[indices], mask[indices],
+ *cardinalities, net.config.action_dim,
+ device=net.device, dtype=net.dtype)
+
+
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--data", type=Path, default=DEFAULT_DATA)
@@ -115,8 +167,10 @@ def main():
parser.add_argument("--epochs", type=int, default=15)
parser.add_argument("--batch-size", type=int, default=256)
parser.add_argument("--neutral-examples", type=int, default=4096)
+ parser.add_argument("--history-steps", type=int, default=1)
parser.add_argument("--model-seed", type=int, default=4101)
parser.add_argument("--shuffle-seed", type=int, default=4101)
+ parser.add_argument("--stage")
parser.add_argument("--device", default="cuda")
parser.add_argument("--out", type=Path, required=True)
args = parser.parse_args()
@@ -138,13 +192,24 @@ def main():
cardinalities = (
metadata["object_cardinality"], metadata["color_cardinality"],
metadata["state_cardinality"])
+ if args.history_steps < 1:
+ raise ValueError("history steps must be positive")
+ input_dim = (visual_input_dim(*cardinalities)
+ if args.history_steps == 1 else history_input_dim(
+ *cardinalities, args.history_steps))
config = BabyAISharedConfig(
- input_dim=visual_input_dim(*cardinalities),
+ input_dim=input_dim,
mission_dim=len(metadata["vocabulary"]), width=args.width,
hidden_layers=args.hidden_layers, learning_rate=args.learning_rate,
reciprocal_learning_rate=args.learning_rate,
context_gain=args.context_gain)
net = BabyAISharedNet(config, seed=args.model_seed, device=device)
+ train_history = (None if args.history_steps == 1 else build_history_index(
+ train["episode_offset"], train["action"], args.history_steps))
+ validation_history = (
+ None if args.history_steps == 1 else build_history_index(
+ validation["episode_offset"], validation["action"],
+ args.history_steps))
shuffle = torch.Generator(device="cpu").manual_seed(args.shuffle_seed)
neutral_count = min(args.neutral_examples, len(train["action"]))
neutral_indices = np.arange(neutral_count)
@@ -155,10 +220,8 @@ def main():
for epoch in range(args.epochs):
if args.condition == "sdil":
- neutral_features = encode_visual(
- train["image"][neutral_indices],
- train["direction"][neutral_indices], *cardinalities,
- device=device, dtype=net.dtype)
+ neutral_features = encode_compact_batch(
+ train, neutral_indices, cardinalities, net, train_history)
neutral_missions = torch.as_tensor(
train["mission_bow"][neutral_indices],
device=device, dtype=net.dtype)
@@ -169,9 +232,8 @@ def main():
losses = []
for start in range(0, len(permutation), args.batch_size):
indices = permutation[start:start + args.batch_size]
- features = encode_visual(
- train["image"][indices], train["direction"][indices],
- *cardinalities, device=device, dtype=net.dtype)
+ features = encode_compact_batch(
+ train, indices, cardinalities, net, train_history)
missions = torch.as_tensor(
train["mission_bow"][indices],
device=device, dtype=net.dtype)
@@ -189,23 +251,26 @@ def main():
training_seconds = time.time() - started
validation_metrics = evaluate_actions(
net, validation["image"], validation["direction"],
- validation["mission_bow"], validation["action"], cardinalities)
+ validation["mission_bow"], validation["action"], cardinalities,
+ history=validation_history)
lesion_validation_metrics = evaluate_actions(
net, validation["image"], validation["direction"],
validation["mission_bow"], validation["action"], cardinalities,
- context_enabled=False)
+ context_enabled=False, history=validation_history)
rollout = rollout_policy(
net, metadata["env_id"], rollout_seeds, metadata["vocabulary"],
- cardinalities)
+ cardinalities, history_steps=args.history_steps)
lesion_rollout = rollout_policy(
net, metadata["env_id"], rollout_seeds, metadata["vocabulary"],
- cardinalities, context_enabled=False)
+ cardinalities, context_enabled=False,
+ history_steps=args.history_steps)
total_seconds = time.time() - started
parameter_count = sum(value.numel() for values in (
net.W, net.b, net.Q[1:], net.C, net.P, net.P_bias)
for value in values)
result = {
- "stage": "babyai_shared_b0" if args.epochs == 15 else "babyai_shared_b1",
+ "stage": (args.stage or (
+ "babyai_shared_b0" if args.epochs == 15 else "babyai_shared_b1")),
"condition": args.condition,
"finite": first_nonfinite_epoch is None,
"first_nonfinite_epoch": first_nonfinite_epoch,
@@ -218,6 +283,7 @@ def main():
"shuffle_seed": args.shuffle_seed,
"neutral_examples_per_epoch": (
neutral_count if args.condition == "sdil" else 0),
+ "history_steps": args.history_steps,
"parameter_count_including_reciprocal_context_predictor": (
parameter_count),
"training_wall_seconds": training_seconds,
diff --git a/experiments/babyai_shared_smoke.py b/experiments/babyai_shared_smoke.py
index e5bbb55..16d3e90 100644
--- a/experiments/babyai_shared_smoke.py
+++ b/experiments/babyai_shared_smoke.py
@@ -10,10 +10,12 @@ from minigrid.core.constants import COLOR_TO_IDX, OBJECT_TO_IDX, STATE_TO_IDX
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sdil.babyai_shared import (
- BabyAISharedConfig, BabyAISharedNet, build_vocabulary, encode_visual,
- manual_step, missions_to_bow, select_teaching_signal, visual_input_dim,
+ BabyAISharedConfig, BabyAISharedNet, build_history_index, build_vocabulary,
+ encode_history_visual, history_input_dim, manual_step, missions_to_bow,
+ select_teaching_signal,
)
from prepare_babyai_shared import generate_split
+from babyai_shared_run import rollout_policy
def maximum_difference(left, right):
@@ -31,12 +33,18 @@ def main():
max(COLOR_TO_IDX.values()) + 1,
max(STATE_TO_IDX.values()) + 1,
)
- features = encode_visual(
- data["image"], data["direction"], *cardinalities,
- device="cpu", dtype=torch.float64)
actions = torch.from_numpy(data["action"].astype("int64"))
+ history_steps = 4
+ history_indices, previous_actions, history_mask = build_history_index(
+ data["episode_offset"], data["action"], history_steps)
+ features = encode_history_visual(
+ data["image"][history_indices], data["direction"][history_indices],
+ previous_actions, history_mask, *cardinalities, 7,
+ device="cpu", dtype=torch.float64)
+ assert history_mask[:, -1].all()
+ assert (previous_actions[data["episode_offset"][:-1], -1] == -1).all()
config = BabyAISharedConfig(
- input_dim=visual_input_dim(*cardinalities),
+ input_dim=history_input_dim(*cardinalities, history_steps),
mission_dim=len(vocabulary), width=16, hidden_layers=2,
learning_rate=0.01, reciprocal_learning_rate=0.01,
momentum=0.0, weight_decay=0.0)
@@ -129,6 +137,11 @@ def main():
assert all(not value.requires_grad for values in (
kp.W, kp.b, kp.Q[1:], kp.C, kp.P, kp.P_bias) for value in values)
+ rollout = rollout_policy(
+ base, "BabyAI-GoToObjS6-v1", [50_000, 50_001], vocabulary,
+ cardinalities, history_steps=history_steps)
+ assert rollout["episodes"] == 2
+
print({
"expert_episodes": 32,
"expert_steps": len(actions),
@@ -143,6 +156,7 @@ def main():
row["action_observations"] for row in predictor_reports),
"predictor_teaching_observations": max(
row["teaching_observations"] for row in predictor_reports),
+ "history_rollout_episodes": rollout["episodes"],
})