From 669c366bc0aea4308aba2c7941c7d6717e70ae21 Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Tue, 10 Apr 2018 11:32:59 -0600 Subject: [PATCH] Update ACVP --- main.py | 6 +++--- model/action_conditional_video_prediction.py | 2 ++ model/dataset.py | 1 + 3 files changed, 6 insertions(+), 3 deletions(-) diff --git a/main.py b/main.py index 12ba7d3..9ca98d0 100644 --- a/main.py +++ b/main.py @@ -302,8 +302,8 @@ def action_conditional_video_prediction(): # a2c_pixel_atari(game) # Generate a dataset with the trained model - a2c_model_file = './data/A2CAgent-vanilla-model-%s.bin' % (game) - generate_dataset(game, a2c_model_file, prefix) + # a2c_model_file = './data/A2CAgent-vanilla-model-%s.bin' % (game) + # generate_dataset(game, a2c_model_file, prefix) # Train the action conditional video prediction model acvp_train(game, prefix) @@ -334,7 +334,7 @@ if __name__ == '__main__': # ddpg_continuous() # ppo_continuous() - action_conditional_video_prediction() + # action_conditional_video_prediction() # plot() diff --git a/model/action_conditional_video_prediction.py b/model/action_conditional_video_prediction.py index 3e5002b..e0f8ec0 100644 --- a/model/action_conditional_video_prediction.py +++ b/model/action_conditional_video_prediction.py @@ -4,6 +4,8 @@ # declaration at the top # ####################################################################### +__all__ = ['acvp_train'] + import torch from torch.autograd import Variable import torch.nn as nn diff --git a/model/dataset.py b/model/dataset.py index 5b50a71..7f601e4 100644 --- a/model/dataset.py +++ b/model/dataset.py @@ -50,6 +50,7 @@ def generate_dataset(game, a2c_model, prefix): config.gradient_clip = 0.5 config.logger = Logger('./log', logger, skip=True) agent = A2CAgent(config) + agent.close() agent.load(a2c_model) task = PixelAtari(game, frame_skip=4, history_length=4, log_dir=None, dataset=True)