Update scratch

This commit is contained in:
Johannes Fischer
2021-10-28 18:31:52 +02:00
parent 1a1f6d8836
commit 673b565e11
4 changed files with 24 additions and 109 deletions

View File

@@ -22,13 +22,16 @@ from tqdm import tqdm
from src.policies.options import OptionsCnnPolicy
from src.gail.options import OptionsEnv, LLOptions, HLOptions, RenderOptions
from src.gail.train import train_discriminator, train_generator
from src.evaluation.evaluation import Evaluation
from torch.utils.tensorboard import SummaryWriter
import os
model_name = 'gail_options_image_random'
env_settings = {'width': 36, 'height': 36, 'm_per_px': 2}
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=100, expert_batch_size=64, generator_steps=1024, discount=0.99):
def train(expert_data, epochs=100, expert_batch_size=16, generator_steps=16, discount=0.99):
env = NRasterizedRandomAgent(**env_settings)
env.discount = discount
@@ -60,12 +63,19 @@ def train(expert_data, epochs=100, expert_batch_size=64, generator_steps=1024, d
generator.verbose,
generator.tensorboard_log,
)
for _ in tqdm(range(epochs)):
filestr = os.path.join('out', model_name)
writer = SummaryWriter(filestr)
ev = Evaluation(filestr, env, expert_data, n_eval_episodes=100)
for epoch 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)
metrics = ev.evaluate(epoch, generator, discriminator)
for metric, value in metrics.items():
writer.add_scalar(metric, value, epoch)
return generator
def video(model_name, env):
@@ -85,7 +95,7 @@ def evaluate():
# %%
if __name__ == '__main__':
with open("../../../scratch/etienne/intersimple/data/NormalizedIntersimpleExpertMu.001N10000_NRasterizedRandomAgentw36h36mppx2.pkl", "rb") as f:
with open("scratch/etienne/intersimple/data/NormalizedIntersimpleExpertMu.001N10000_NRasterizedRandomAgentInfow36h36mppx2.pkl", "rb") as f:
trajectories = pickle.load(f)
transitions = rollout.flatten_trajectories(trajectories)
train(transitions)