Merge pull request #3 from sisl/refactor-sampling

Refactor sampling
This commit is contained in:
ebuehrle
2021-09-15 09:08:13 +02:00
committed by GitHub
2 changed files with 165 additions and 90 deletions

View File

@@ -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)

View File

@@ -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