Refactor sampling
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user