Merge different methods to sample the policy and collect transitions

This commit is contained in:
Johannes Fischer
2021-09-13 15:50:49 +02:00
parent f1ece358d7
commit 826c0fa219

View File

@@ -89,7 +89,6 @@ def check_future_collisions_fast(env, actions):
radius = (env._env._lengths**2 + env._env._widths**2).sqrt() / 2
min_distance = radius[env._agent] + radius
min_distance = min_distance.unsqueeze(0).unsqueeze(0)
print('min_distance', min_distance.shape)
assert min_distance.shape == (1, 1, nv)
return (distance > min_distance).all(-1).all(-1)
@@ -103,38 +102,11 @@ 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_ll(env, generator):
"""Sample low-level (state, action) pairs for discriminator training."""
done = True
while True:
if done:
s = env.reset()
done = False
m = available_actions(env)
ch, _, _ = 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}'
while not done and plan and feasible(env, plan, ch):
a, plan = env._normalize(plan[0]), plan[1:]
nexts, _, done, _ = env.step(a)
yield {
'obs': s,
'next_obs': nexts,
'acts': np.array((a,)),
'dones': np.array(done),
}
s = nexts
def sample_hl(env, generator, discriminator):
"""Sample high-level (state, action, reward) tuples for generator training."""
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
@@ -150,19 +122,36 @@ def sample_hl(env, generator, discriminator):
'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
steps = 0
while not done and plan and feasible(env, plan, ch):
a, plan = env._normalize(plan[0]), plan[1:]
if level == 'high':
r += discriminator.discrim_net.discriminator(
torch.tensor(s).unsqueeze(0).to(discriminator.discrim_net.device()),
torch.tensor([[a]]).to(discriminator.discrim_net.device()),
)
steps += 1
s, _, done, _ = env.step(a)
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':
assert steps > 0
yield {
'obs': obs,
'action': ch,
@@ -182,12 +171,12 @@ def flatten_transitions(transitions):
}
def train_discriminator(env, generator, discriminator, num_samples):
transitions = list(itertools.islice(sample_ll(env, generator), num_samples))
transitions = list(itertools.islice(sample(env, generator, None, 'low'), 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_hl(env, generator, discriminator), num_samples+1))
generator_samples = list(itertools.islice(sample(env, generator, discriminator, 'high'), num_samples+1))
generator.rollout_buffer.reset()
for s in generator_samples[:-1]:
@@ -270,20 +259,7 @@ if __name__ == '__main__':
env = NRasterized(**env_settings)
s = env.reset()
done = False
env.render()
while not done:
m = available_actions(env)
ch, _, _ = 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)))
while not done and plan and feasible(env, plan, ch):
a, plan = env._normalize(plan[0]), plan[1:]
s, _, done, _ = env.step(a)
for _ in sample(env, generator, None, 'low'):
env.render()
env.close(filestr='render/'+model_name)
@@ -297,13 +273,15 @@ def test_ll_transitions_vs_expert_data():
env = NRasterized(agent=51, width=36, height=36, m_per_px=2)
gen_transitions = list(itertools.islice(sample_ll(
gen_transitions = list(itertools.islice(sample(
env=NRasterized(**env_settings),
generator=stable_baselines3.PPO(
OptionsCnnPolicy,
OptionsEnv(env),
verbose=1,
)
),
discriminator=None,
level='low'
), 10))
gen_transitions = flatten_transitions(gen_transitions)