Add cuda support, normalize actions
This commit is contained in:
@@ -1,3 +1,4 @@
|
|||||||
|
# %%
|
||||||
from gail.discriminator import CnnDiscriminator
|
from gail.discriminator import CnnDiscriminator
|
||||||
from imitation.algorithms import adversarial
|
from imitation.algorithms import adversarial
|
||||||
import stable_baselines3
|
import stable_baselines3
|
||||||
@@ -82,11 +83,14 @@ def sample_ll(env, generator):
|
|||||||
done = False
|
done = False
|
||||||
|
|
||||||
m = available_actions(env)
|
m = available_actions(env)
|
||||||
ch, _, _ = generator.policy.predict({'obs': torch.tensor(s).unsqueeze(0), 'mask': torch.tensor(m).unsqueeze(0)})
|
ch, _, _ = generator.policy.predict({
|
||||||
plan = list(generate_plan(env, ch))
|
'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):
|
while not done and plan and feasible(env, plan):
|
||||||
a, plan = plan[0], plan[1:]
|
a, plan = env._normalize(plan[0]), plan[1:]
|
||||||
nexts, _, done, _ = env.step(a)
|
nexts, _, done, _ = env.step(a)
|
||||||
yield {
|
yield {
|
||||||
'obs': s,
|
'obs': s,
|
||||||
@@ -108,14 +112,20 @@ def sample_hl(env, generator, discriminator):
|
|||||||
episode_start = True
|
episode_start = True
|
||||||
|
|
||||||
obs = {'obs': s, 'mask': m}
|
obs = {'obs': s, 'mask': m}
|
||||||
ch, value, log_prob = generator.policy.predict({'obs': torch.tensor(s).unsqueeze(0), 'mask': torch.tensor(m).unsqueeze(0)})
|
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)))
|
plan = list(map(float, generate_plan(env, ch)))
|
||||||
r = 0
|
r = 0
|
||||||
steps = 0
|
steps = 0
|
||||||
|
|
||||||
while not done and plan and feasible(env, plan):
|
while not done and plan and feasible(env, plan):
|
||||||
a, plan = plan[0], plan[1:]
|
a, plan = env._normalize(plan[0]), plan[1:]
|
||||||
r += discriminator.discrim_net.discriminator(torch.tensor(s).unsqueeze(0), torch.tensor([[a]]))
|
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
|
steps += 1
|
||||||
s, _, done, _ = env.step(a)
|
s, _, done, _ = env.step(a)
|
||||||
m = available_actions(env)
|
m = available_actions(env)
|
||||||
@@ -150,8 +160,8 @@ def train_generator(env, generator, discriminator, num_samples):
|
|||||||
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['action'],
|
action=s['action'].cpu(),
|
||||||
reward=s['reward'],
|
reward=s['reward'].cpu(),
|
||||||
episode_start=s['episode_start'],
|
episode_start=s['episode_start'],
|
||||||
value=s['value'],
|
value=s['value'],
|
||||||
log_prob=s['log_prob'],
|
log_prob=s['log_prob'],
|
||||||
@@ -213,38 +223,62 @@ def train(expert_data, epochs=10, expert_batch_size=32, generator_steps=2048):
|
|||||||
return generator
|
return generator
|
||||||
|
|
||||||
# %%
|
# %%
|
||||||
|
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, epochs=2, expert_batch_size=2, generator_steps=2)
|
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 = NRasterized(**env_settings)
|
||||||
|
|
||||||
s = env.reset()
|
s = env.reset()
|
||||||
done = False
|
done = False
|
||||||
env.render()
|
env.render()
|
||||||
while not done:
|
while not done:
|
||||||
m = available_actions(env)
|
m = available_actions(env)
|
||||||
ch, _, _ = generator.policy.predict({'obs': torch.tensor(s).unsqueeze(0), 'mask': torch.tensor(m).unsqueeze(0)})
|
ch, _, _ = generator.policy.predict({
|
||||||
plan = list(generate_plan(env, ch))
|
'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):
|
while not done and plan and feasible(env, plan):
|
||||||
a, plan = plan[0], plan[1:]
|
a, plan = env._normalize(plan[0]), plan[1:]
|
||||||
s, _, done, _ = env.step(a)
|
s, _, done, _ = env.step(a)
|
||||||
env.render()
|
env.render()
|
||||||
|
|
||||||
env.close(filestr='render/'+model_name)
|
env.close(filestr='render/'+model_name)
|
||||||
|
|
||||||
# %% Tests
|
# %% Tests
|
||||||
|
|
||||||
def test_ll_transitions_vs_expert_data():
|
def test_ll_transitions_vs_expert_data():
|
||||||
pass
|
with open("data/NormalizedIntersimpleExpertMu.001_NRasterizedAgent51w36h36mppx2.pkl", "rb") as f:
|
||||||
|
expert_trajectories = pickle.load(f)
|
||||||
|
expert_transitions = rollout.flatten_trajectories(expert_trajectories)
|
||||||
|
|
||||||
|
env = NRasterized(agent=51, width=36, height=36, m_per_px=2)
|
||||||
|
|
||||||
|
gen_transitions = list(itertools.islice(sample_ll(
|
||||||
|
env=NRasterized(**env_settings),
|
||||||
|
generator=stable_baselines3.PPO(
|
||||||
|
OptionsCnnPolicy,
|
||||||
|
OptionsEnv(env),
|
||||||
|
verbose=1,
|
||||||
|
)
|
||||||
|
), 10))
|
||||||
|
gen_transitions = flatten_transitions(gen_transitions)
|
||||||
|
|
||||||
|
assert expert_transitions[:10].obs.shape == gen_transitions['obs'].shape
|
||||||
|
assert expert_transitions[:10].next_obs.shape == gen_transitions['next_obs'].shape
|
||||||
|
assert expert_transitions[:10].acts.shape == gen_transitions['acts'].shape
|
||||||
|
assert expert_transitions[:10].dones.shape == gen_transitions['dones'].shape
|
||||||
|
|
||||||
|
|
||||||
def test_hl_transitions():
|
def test_hl_transitions():
|
||||||
pass
|
pass
|
||||||
|
|||||||
Reference in New Issue
Block a user