@@ -36,7 +36,8 @@ class CnnDiscriminator(torch.nn.Module):
|
|||||||
|
|
||||||
def forward(self, state, action):
|
def forward(self, state, action):
|
||||||
sa = self._concatenate(state, action)
|
sa = self._concatenate(state, action)
|
||||||
return self.cnn(sa).squeeze()
|
assert sa.ndim == 4
|
||||||
|
return self.cnn(sa).squeeze(1)
|
||||||
|
|
||||||
class MlpDiscriminator(torch.nn.Module):
|
class MlpDiscriminator(torch.nn.Module):
|
||||||
"""MLP similar to stable_baselines3.common.policies.ActorCriticPolicy."""
|
"""MLP similar to stable_baselines3.common.policies.ActorCriticPolicy."""
|
||||||
@@ -55,4 +56,5 @@ class MlpDiscriminator(torch.nn.Module):
|
|||||||
def forward(self, state, action):
|
def forward(self, state, action):
|
||||||
flat = self.flatten(state)
|
flat = self.flatten(state)
|
||||||
sa = torch.cat((action, flat), -1)
|
sa = torch.cat((action, flat), -1)
|
||||||
return self.mlp(sa).squeeze()
|
assert sa.ndim == 2
|
||||||
|
return self.mlp(sa).squeeze(1)
|
||||||
|
|||||||
@@ -45,6 +45,128 @@ class OptionsCnnPolicy(stable_baselines3.common.policies.ActorCriticCnnPolicy):
|
|||||||
posterior = Categorical(prior.probs * m)
|
posterior = Categorical(prior.probs * m)
|
||||||
return values, posterior.log_prob(ch), posterior.entropy() # additional values used by PPO.train
|
return values, posterior.log_prob(ch), posterior.entropy() # additional values used by PPO.train
|
||||||
|
|
||||||
|
class OptionsEnv(gym.Wrapper):
|
||||||
|
|
||||||
|
def __init__(self, env, *args, **kwargs):
|
||||||
|
super().__init__(env, *args, *kwargs)
|
||||||
|
num_hl_options = len(ALL_OPTIONS)
|
||||||
|
self.action_space = gym.spaces.Discrete(num_hl_options)
|
||||||
|
self.observation_space = gym.spaces.Dict({
|
||||||
|
'obs': env.observation_space,
|
||||||
|
'mask': gym.spaces.Box(low=0, high=1, shape=(num_hl_options,)),
|
||||||
|
})
|
||||||
|
|
||||||
|
def _after_choice(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def _after_step(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def _transitions(self):
|
||||||
|
raise NotImplementedError('Use `LLOptions` or `HLOptions` for sampling.')
|
||||||
|
|
||||||
|
def sample(self, generator):
|
||||||
|
self.done = True
|
||||||
|
while True:
|
||||||
|
self.episode_start = False
|
||||||
|
if self.done:
|
||||||
|
self.s = self.env.reset()
|
||||||
|
self.m = available_actions(self.env)
|
||||||
|
self.done = False
|
||||||
|
self.episode_start = True
|
||||||
|
|
||||||
|
self.ch, self.value, self.log_prob = generator.policy.predict({
|
||||||
|
'obs': torch.tensor(self.s).unsqueeze(0).to(generator.policy.device),
|
||||||
|
'mask': torch.tensor(self.m).unsqueeze(0).to(generator.policy.device),
|
||||||
|
})
|
||||||
|
self.plan = list(map(float, generate_plan(self.env, self.ch)))
|
||||||
|
|
||||||
|
self._after_choice()
|
||||||
|
|
||||||
|
assert not self.done
|
||||||
|
assert self.plan
|
||||||
|
assert feasible(self.env, self.plan, self.ch)
|
||||||
|
|
||||||
|
while not self.done and self.plan and feasible(self.env, self.plan, self.ch):
|
||||||
|
self.a, self.plan = self.plan[0], self.plan[1:]
|
||||||
|
self.a = self.env._normalize(self.a)
|
||||||
|
self.nexts, _, self.done, _ = self.env.step(self.a)
|
||||||
|
self.nextm = available_actions(self.env)
|
||||||
|
|
||||||
|
self._after_step()
|
||||||
|
|
||||||
|
self.s = self.nexts
|
||||||
|
self.m = self.nextm
|
||||||
|
|
||||||
|
yield from self._transitions()
|
||||||
|
|
||||||
|
class LLOptions(OptionsEnv):
|
||||||
|
"""Sample low-level (state, action) tuples for discriminator training."""
|
||||||
|
|
||||||
|
def __init__(self, *args, **kwargs):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
self.observation_space = self.observation_space['obs']
|
||||||
|
|
||||||
|
def _after_choice(self):
|
||||||
|
self._transition_buffer = []
|
||||||
|
|
||||||
|
def _after_step(self):
|
||||||
|
self._transition_buffer.append({
|
||||||
|
'obs': self.s,
|
||||||
|
'next_obs': self.nexts,
|
||||||
|
'acts': np.array((self.a,)),
|
||||||
|
'dones': np.array(self.done),
|
||||||
|
})
|
||||||
|
|
||||||
|
def _transitions(self):
|
||||||
|
yield from self._transition_buffer
|
||||||
|
|
||||||
|
def sample_ll(self, policy):
|
||||||
|
return self.sample(policy)
|
||||||
|
|
||||||
|
class HLOptions(OptionsEnv):
|
||||||
|
"""Sample high-level (state, action, reward) tuples for generator training."""
|
||||||
|
|
||||||
|
def __init__(self, *args, **kwargs):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
|
||||||
|
def _after_choice(self):
|
||||||
|
self.r = 0
|
||||||
|
self.steps = 0
|
||||||
|
|
||||||
|
def _after_step(self):
|
||||||
|
self.r += self.discount**self.steps * self.discriminator.discrim_net.reward_train(
|
||||||
|
state=torch.tensor(self.s).unsqueeze(0).to(self.discriminator.discrim_net.device()),
|
||||||
|
action=torch.tensor([[self.a]]).to(self.discriminator.discrim_net.device()),
|
||||||
|
next_state=torch.tensor(self.s).unsqueeze(0).to(self.discriminator.discrim_net.device()), # unused
|
||||||
|
done=torch.tensor(self.done).unsqueeze(0).to(self.discriminator.discrim_net.device()), # unused
|
||||||
|
)
|
||||||
|
self.steps += 1
|
||||||
|
|
||||||
|
def _transitions(self):
|
||||||
|
yield {
|
||||||
|
'obs': {'obs': self.s, 'mask': self.m},
|
||||||
|
'action': self.ch,
|
||||||
|
'reward': self.r.detach(),
|
||||||
|
'episode_start': self.episode_start,
|
||||||
|
'value': self.value.detach(),
|
||||||
|
'log_prob': self.log_prob.detach(),
|
||||||
|
'done': self.done,
|
||||||
|
}
|
||||||
|
|
||||||
|
def sample_hl(self, policy, discriminator):
|
||||||
|
self.discriminator = discriminator
|
||||||
|
return self.sample(policy)
|
||||||
|
|
||||||
|
class RenderOptions(LLOptions):
|
||||||
|
|
||||||
|
def _after_step(self):
|
||||||
|
super()._after_step()
|
||||||
|
self.env.render()
|
||||||
|
|
||||||
|
def close(self, *args, **kwargs):
|
||||||
|
self.env.close(*args, **kwargs)
|
||||||
|
|
||||||
def available_actions(env):
|
def available_actions(env):
|
||||||
"""Return mask of available actions given current `env` state."""
|
"""Return mask of available actions given current `env` state."""
|
||||||
valid = np.array([feasible(env, generate_plan(env, i), i) for i in range(len(ALL_OPTIONS))])
|
valid = np.array([feasible(env, generate_plan(env, i), i) for i in range(len(ALL_OPTIONS))])
|
||||||
@@ -102,65 +224,6 @@ def feasible(env, plan, ch):
|
|||||||
valid = check_future_collisions_fast(env, [full_plan]) # check_future_collisions_fast takes in B-list and outputs (B,) bool tensor
|
valid = check_future_collisions_fast(env, [full_plan]) # check_future_collisions_fast takes in B-list and outputs (B,) bool tensor
|
||||||
return ch == 0 or valid.item()
|
return ch == 0 or valid.item()
|
||||||
|
|
||||||
def sample(env, generator, discriminator, level: str):
|
|
||||||
"""
|
|
||||||
Sample low-level (state, action, next_state) tuples for discriminator training or
|
|
||||||
high-level (state, action, reward) tuples for generator training.
|
|
||||||
"""
|
|
||||||
done = True
|
|
||||||
while True:
|
|
||||||
episode_start = False
|
|
||||||
if done:
|
|
||||||
s = env.reset()
|
|
||||||
m = available_actions(env)
|
|
||||||
done = False
|
|
||||||
episode_start = True
|
|
||||||
|
|
||||||
obs = {'obs': s, 'mask': m}
|
|
||||||
ch, value, log_prob = generator.policy.predict({
|
|
||||||
'obs': torch.tensor(s).unsqueeze(0).to(generator.policy.device),
|
|
||||||
'mask': torch.tensor(m).unsqueeze(0).to(generator.policy.device),
|
|
||||||
})
|
|
||||||
plan = list(map(float, generate_plan(env, ch)))
|
|
||||||
|
|
||||||
assert not done
|
|
||||||
assert plan
|
|
||||||
assert feasible(env, plan, ch), f'Infeasible hl action {ch}'
|
|
||||||
|
|
||||||
r = 0
|
|
||||||
discount = 1
|
|
||||||
while not done and plan and feasible(env, plan, ch):
|
|
||||||
a, plan = env._normalize(plan[0]), plan[1:]
|
|
||||||
if level == 'high':
|
|
||||||
r += discount * discriminator.discrim_net.discriminator(
|
|
||||||
torch.tensor(s).unsqueeze(0).to(discriminator.discrim_net.device()),
|
|
||||||
torch.tensor([[a]]).to(discriminator.discrim_net.device()),
|
|
||||||
)
|
|
||||||
discount *= env.discount
|
|
||||||
|
|
||||||
nexts, _, done, _ = env.step(a)
|
|
||||||
m = available_actions(env)
|
|
||||||
|
|
||||||
if level == 'low':
|
|
||||||
yield {
|
|
||||||
'obs': s,
|
|
||||||
'next_obs': nexts,
|
|
||||||
'acts': np.array((a,)),
|
|
||||||
'dones': np.array(done),
|
|
||||||
}
|
|
||||||
s = nexts
|
|
||||||
|
|
||||||
if level == 'high':
|
|
||||||
yield {
|
|
||||||
'obs': obs,
|
|
||||||
'option': ch,
|
|
||||||
'reward': r.detach(),
|
|
||||||
'episode_start': episode_start,
|
|
||||||
'value': value.detach(),
|
|
||||||
'log_prob': log_prob.detach(),
|
|
||||||
'done': done,
|
|
||||||
}
|
|
||||||
|
|
||||||
def flatten_transitions(transitions):
|
def flatten_transitions(transitions):
|
||||||
return {
|
return {
|
||||||
'obs': np.stack(list(t['obs'] for t in transitions), axis=0),
|
'obs': np.stack(list(t['obs'] for t in transitions), axis=0),
|
||||||
@@ -170,18 +233,18 @@ def flatten_transitions(transitions):
|
|||||||
}
|
}
|
||||||
|
|
||||||
def train_discriminator(env, generator, discriminator, num_samples):
|
def train_discriminator(env, generator, discriminator, num_samples):
|
||||||
transitions = list(itertools.islice(sample(env, generator, None, 'low'), num_samples))
|
transitions = list(itertools.islice(env.sample_ll(generator), num_samples))
|
||||||
generator_samples = flatten_transitions(transitions)
|
generator_samples = flatten_transitions(transitions)
|
||||||
discriminator.train_disc(gen_samples=generator_samples)
|
discriminator.train_disc(gen_samples=generator_samples)
|
||||||
|
|
||||||
def train_generator(env, generator, discriminator, num_samples):
|
def train_generator(env, generator, discriminator, num_samples):
|
||||||
generator_samples = list(itertools.islice(sample(env, generator, discriminator, 'high'), num_samples+1))
|
generator_samples = list(itertools.islice(env.sample_hl(generator, discriminator), num_samples+1))
|
||||||
|
|
||||||
generator.rollout_buffer.reset()
|
generator.rollout_buffer.reset()
|
||||||
for s in generator_samples[:-1]:
|
for s in generator_samples[:-1]:
|
||||||
generator.rollout_buffer.add(
|
generator.rollout_buffer.add(
|
||||||
obs=s['obs'],
|
obs=s['obs'],
|
||||||
action=s['option'].cpu(),
|
action=s['action'].cpu(),
|
||||||
reward=s['reward'].cpu(),
|
reward=s['reward'].cpu(),
|
||||||
episode_start=s['episode_start'],
|
episode_start=s['episode_start'],
|
||||||
value=s['value'],
|
value=s['value'],
|
||||||
@@ -195,17 +258,6 @@ def train_generator(env, generator, discriminator, num_samples):
|
|||||||
|
|
||||||
generator.train()
|
generator.train()
|
||||||
|
|
||||||
class OptionsEnv(gym.Wrapper):
|
|
||||||
|
|
||||||
def __init__(self, env):
|
|
||||||
super().__init__(env)
|
|
||||||
num_hl_options = len(ALL_OPTIONS)
|
|
||||||
self.action_space = gym.spaces.Discrete(num_hl_options)
|
|
||||||
self.observation_space = gym.spaces.Dict({
|
|
||||||
'obs': env.observation_space,
|
|
||||||
'mask': gym.spaces.Box(low=0, high=1, shape=(num_hl_options,)),
|
|
||||||
})
|
|
||||||
|
|
||||||
def train(expert_data, epochs=10, expert_batch_size=32, generator_steps=2048, discount=0.99):
|
def train(expert_data, epochs=10, expert_batch_size=32, generator_steps=2048, discount=0.99):
|
||||||
env = NRasterized(**env_settings)
|
env = NRasterized(**env_settings)
|
||||||
env.discount = discount
|
env.discount = discount
|
||||||
@@ -239,8 +291,8 @@ def train(expert_data, epochs=10, expert_batch_size=32, generator_steps=2048, di
|
|||||||
)
|
)
|
||||||
|
|
||||||
for _ in range(epochs):
|
for _ in range(epochs):
|
||||||
train_discriminator(env, generator, discriminator, num_samples=expert_batch_size)
|
train_discriminator(LLOptions(env), generator, discriminator, num_samples=expert_batch_size)
|
||||||
train_generator(env, generator, discriminator, num_samples=generator_steps)
|
train_generator(HLOptions(env), generator, discriminator, num_samples=generator_steps)
|
||||||
|
|
||||||
return generator
|
return generator
|
||||||
|
|
||||||
@@ -250,40 +302,36 @@ if __name__ == '__main__':
|
|||||||
with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedAgent51w36h36mppx2.pkl", "rb") as f:
|
with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedAgent51w36h36mppx2.pkl", "rb") as f:
|
||||||
trajectories = pickle.load(f)
|
trajectories = pickle.load(f)
|
||||||
transitions = rollout.flatten_trajectories(trajectories)
|
transitions = rollout.flatten_trajectories(trajectories)
|
||||||
generator = train(transitions, generator_steps=200)
|
generator = train(transitions, epochs=2, expert_batch_size=2, generator_steps=2)
|
||||||
|
|
||||||
generator.save(model_name)
|
generator.save(model_name)
|
||||||
|
|
||||||
# %%
|
# %%
|
||||||
model = stable_baselines3.PPO.load(model_name)
|
model = stable_baselines3.PPO.load(model_name)
|
||||||
|
|
||||||
env = NRasterized(**env_settings)
|
env = RenderOptions(NRasterized(**env_settings))
|
||||||
|
|
||||||
for transition in sample(env, generator, None, 'low'):
|
for s in env.sample_ll(generator):
|
||||||
env.render()
|
if s['dones']:
|
||||||
if transition['dones']:
|
|
||||||
break
|
break
|
||||||
|
|
||||||
env.close(filestr='render/'+model_name)
|
env.close(filestr='render/'+model_name)
|
||||||
|
|
||||||
# %% Tests
|
# %% Tests
|
||||||
|
|
||||||
def test_ll_transitions_vs_expert_data():
|
def test_ll_expert_data():
|
||||||
with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedAgent51w36h36mppx2.pkl", "rb") as f:
|
with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedAgent51w36h36mppx2.pkl", "rb") as f:
|
||||||
expert_trajectories = pickle.load(f)
|
expert_trajectories = pickle.load(f)
|
||||||
expert_transitions = rollout.flatten_trajectories(expert_trajectories)
|
expert_transitions = rollout.flatten_trajectories(expert_trajectories)
|
||||||
|
|
||||||
env = NRasterized(agent=51, width=36, height=36, m_per_px=2)
|
env = LLOptions(NRasterized(agent=51, width=36, height=36, m_per_px=2))
|
||||||
|
|
||||||
gen_transitions = list(itertools.islice(sample(
|
gen_transitions = list(itertools.islice(env.sample_ll(
|
||||||
env=NRasterized(**env_settings),
|
policy=stable_baselines3.PPO(
|
||||||
generator=stable_baselines3.PPO(
|
|
||||||
OptionsCnnPolicy,
|
OptionsCnnPolicy,
|
||||||
OptionsEnv(env),
|
OptionsEnv(env),
|
||||||
verbose=1,
|
verbose=1,
|
||||||
),
|
)
|
||||||
discriminator=None,
|
|
||||||
level='low'
|
|
||||||
), 10))
|
), 10))
|
||||||
gen_transitions = flatten_transitions(gen_transitions)
|
gen_transitions = flatten_transitions(gen_transitions)
|
||||||
|
|
||||||
@@ -292,6 +340,31 @@ def test_ll_transitions_vs_expert_data():
|
|||||||
assert expert_transitions[:10].acts.shape == gen_transitions['acts'].shape
|
assert expert_transitions[:10].acts.shape == gen_transitions['acts'].shape
|
||||||
assert expert_transitions[:10].dones.shape == gen_transitions['dones'].shape
|
assert expert_transitions[:10].dones.shape == gen_transitions['dones'].shape
|
||||||
|
|
||||||
|
def test_ll_states():
|
||||||
|
env = NRasterized()
|
||||||
|
policy = stable_baselines3.PPO(
|
||||||
|
OptionsCnnPolicy,
|
||||||
|
OptionsEnv(env),
|
||||||
|
verbose=1,
|
||||||
|
)
|
||||||
|
llenv = LLOptions(env)
|
||||||
|
transitions = list(itertools.islice(llenv.sample_ll(policy=policy), 100))
|
||||||
|
|
||||||
|
env2 = NRasterized()
|
||||||
|
s2 = env2.reset()
|
||||||
|
for i, t in enumerate(transitions):
|
||||||
|
assert i == 0 or np.array_equal(t['obs'], transitions[i-1]['next_obs'])
|
||||||
|
assert np.array_equal(t['obs'], s2)
|
||||||
|
assert t['acts'].shape == (1,)
|
||||||
|
|
||||||
|
nexts2, _, done2, _ = env2.step(t['acts'])
|
||||||
|
assert np.array_equal(t['next_obs'], nexts2)
|
||||||
|
assert np.array_equal(t['dones'], done2)
|
||||||
|
|
||||||
|
if done2:
|
||||||
|
break
|
||||||
|
|
||||||
|
s2 = nexts2
|
||||||
|
|
||||||
def test_hl_transitions():
|
def test_hl_transitions():
|
||||||
pass
|
pass
|
||||||
|
|||||||
Reference in New Issue
Block a user