Refactor sampling

This commit is contained in:
ebuehrle
2021-09-14 11:54:49 +02:00
parent deaef45943
commit d2932374d9

View File

@@ -45,6 +45,115 @@ class OptionsCnnPolicy(stable_baselines3.common.policies.ActorCriticCnnPolicy):
posterior = Categorical(prior.probs * m)
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, policy):
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 = policy.predict({
'obs': torch.tensor(self.s).unsqueeze(0).to(policy.device),
'mask': torch.tensor(self.m).unsqueeze(0).to(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_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):
self._transition_buffer = []
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(
torch.tensor(self.s).unsqueeze(0).to(self.discriminator.device()),
torch.tensor([[self.a]]).to(self.discriminator.device()),
)
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)
def available_actions(env):
"""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))])
@@ -102,65 +211,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
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):
return {
'obs': np.stack(list(t['obs'] for t in transitions), axis=0),
@@ -170,18 +220,18 @@ def flatten_transitions(transitions):
}
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.policy), num_samples))
generator_samples = flatten_transitions(transitions)
discriminator.train_disc(gen_samples=generator_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.policy, discriminator.discrim_net.discriminator), num_samples+1))
generator.rollout_buffer.reset()
for s in generator_samples[:-1]:
generator.rollout_buffer.add(
obs=s['obs'],
action=s['option'].cpu(),
action=s['action'].cpu(),
reward=s['reward'].cpu(),
episode_start=s['episode_start'],
value=s['value'],
@@ -195,17 +245,6 @@ def train_generator(env, generator, discriminator, num_samples):
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):
env = NRasterized(**env_settings)
env.discount = discount
@@ -239,8 +278,8 @@ def train(expert_data, epochs=10, expert_batch_size=32, generator_steps=2048, di
)
for _ in range(epochs):
train_discriminator(env, generator, discriminator, num_samples=expert_batch_size)
train_generator(env, generator, discriminator, num_samples=generator_steps)
train_discriminator(LLOptions(env), generator, discriminator, num_samples=expert_batch_size)
train_generator(HLOptions(env), generator, discriminator, num_samples=generator_steps)
return generator
@@ -257,14 +296,14 @@ if __name__ == '__main__':
# %%
model = stable_baselines3.PPO.load(model_name)
env = NRasterized(**env_settings)
env = LLOptions(NRasterized(**env_settings))
for transition in sample(env, generator, None, 'low'):
env.render()
if transition['dones']:
for s in env.sample_ll(env, generator.policy):
env.env.render()
if s['dones']:
break
env.close(filestr='render/'+model_name)
env.env.close(filestr='render/'+model_name)
# %% Tests
@@ -273,17 +312,14 @@ def test_ll_transitions_vs_expert_data():
expert_trajectories = pickle.load(f)
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(
env=NRasterized(**env_settings),
generator=stable_baselines3.PPO(
gen_transitions = list(itertools.islice(env.sample_ll(
policy=stable_baselines3.PPO(
OptionsCnnPolicy,
OptionsEnv(env),
verbose=1,
),
discriminator=None,
level='low'
).policy
), 10))
gen_transitions = flatten_transitions(gen_transitions)