diff --git a/scratch/etienne/intersimple/gail_options_image.py b/scratch/etienne/intersimple/gail_options_image.py index 77de2ef..91d14b0 100644 --- a/scratch/etienne/intersimple/gail_options_image.py +++ b/scratch/etienne/intersimple/gail_options_image.py @@ -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)