Move files to src

This commit is contained in:
ebuehrle
2021-10-26 17:41:05 +02:00
parent 24b91d4eec
commit bcddf422f0
10 changed files with 35 additions and 343 deletions

View File

@@ -1,5 +1,8 @@
# %%
from gail.discriminator import CnnDiscriminatorFlatAction
import sys
sys.path.append('../../../')
from src.discriminator import CnnDiscriminatorFlatAction
from imitation.algorithms import adversarial
import stable_baselines3
import torch.utils.data
@@ -16,16 +19,16 @@ import pathlib
from imitation.util import logger
from stable_baselines3.common.env_util import make_vec_env
from tqdm import tqdm
from gail.policy import OptionsCnnPolicy
from gail.options import OptionsEnv, LLOptions, HLOptions, RenderOptions
from gail.train import train_discriminator, train_generator
from src.policies.options import OptionsCnnPolicy
from src.gail.options import OptionsEnv, LLOptions, HLOptions, RenderOptions
from src.gail.train import train_discriminator, train_generator
model_name = 'gail_options_image_random'
env_settings = {'width': 70, 'height': 70, 'm_per_px': 1}
ALL_OPTIONS = [(v,t) for v in [0,2,4,6,8] for t in [5, 10, 20]] # option 0 is safe fallback
def train(expert_data, epochs=20, expert_batch_size=32, generator_steps=1024, discount=0.99):
def train(expert_data, epochs=100, expert_batch_size=64, generator_steps=1024, discount=0.99):
env = NRasterizedRouteRandomAgent(**env_settings)
env.discount = discount
@@ -61,27 +64,28 @@ def train(expert_data, epochs=20, expert_batch_size=32, generator_steps=1024, di
for _ in tqdm(range(epochs)):
train_discriminator(LLOptions(env, options=ALL_OPTIONS), generator, discriminator, num_samples=expert_batch_size)
train_generator(HLOptions(env, options=ALL_OPTIONS), generator, discriminator, num_samples=generator_steps)
generator.save(model_name)
return generator
# %%
if __name__ == '__main__':
# %%
with open("data/NormalizedIntersimpleExpertMu.001N10000_NRasterizedRouteRandomAgentw70h70mppx1.pkl", "rb") as f:
trajectories = pickle.load(f)
transitions = rollout.flatten_trajectories(trajectories)
generator = train(transitions, epochs=100)
generator.save(model_name)
# %%
def video(model_name, env):
model = stable_baselines3.PPO.load(model_name)
env = RenderOptions(NRasterizedRouteRandomAgent(**env_settings))
env = RenderOptions(env, options=ALL_OPTIONS)
for s in env.sample_ll(model):
if s['dones']:
break
env.close(filestr='render/'+model_name)
def evaluate():
video(
model_name=model_name,
env=NRasterizedRouteRandomAgent(**env_settings)
)
# %%
if __name__ == '__main__':
with open("data/NormalizedIntersimpleExpertMu.001N10000_NRasterizedRouteRandomAgentw70h70mppx1.pkl", "rb") as f:
trajectories = pickle.load(f)
transitions = rollout.flatten_trajectories(trajectories)
train(transitions)