diff --git a/.gitignore b/.gitignore index fd63ecc..803df77 100644 --- a/.gitignore +++ b/.gitignore @@ -8,6 +8,7 @@ exp_* upload.py *.sh data +dataset draw_* log evaluation_log diff --git a/README.md b/README.md index 8f85185..f1e9cf3 100644 --- a/README.md +++ b/README.md @@ -15,6 +15,7 @@ Implemented algorithms: * Distributed Deep Deterministic Policy Gradient (Distributed DDPG, aka D3PG) * Hybrid Reward Architecture (HRA) * Parallelized Proximal Policy Optimization (P3O, similar to DPPO) +* Action Conditional Video Prediction # Curves > Curves for CartPole are trivial so I didn't place it here. There isn't any fixed random seed. @@ -79,22 +80,27 @@ but is wrong with high-dimensional action. And its computation of entropy is wro I use 8 threads and a two tanh hidden layer network, each hidden layer has 64 hidden units. +## Action Conditional Video Prediction + +![Loading...](https://raw.githubusercontent.com/ShangtongZhang/DeepRL/master/images/ACVP.png) + +**Left**: One-step prediction **Right**: Ground truth + +Prediction is sampled after 110K iterations and I only implemented one-step training + # Dependency +> Tested in macOS 10.12 and CentO/S 6.8 * Open AI gym * [Roboschool](https://github.com/openai/roboschool) (Optional) -* PyTorch v0.2.0 +* PyTorch v0.3.0 * Python 2.7 or Python 3.6 -* Tensorflow (Optional, but tensorboard is awesome) -> If you want to use Roboschool, you have to use Python3. And don't try to use Roboschool with parallelized algorithms, -> there is a known [critical bug](https://github.com/openai/roboschool/issues/86). +* [TensorboardX](https://github.com/lanpa/tensorboard-pytorch) + # Usage -Detailed usage and all training parameters can be found in ```main.py```. -And you need to create following directories before running the program: -``` -cd DeepRL -mkdir data log -``` +```dataset.py```: generate dataset for action conditional video prediction + +```main.py```: all other algorithms # References * [Human Level Control through Deep Reinforcement Learning](https://www.nature.com/nature/journal/v518/n7540/full/nature14236.html) @@ -110,3 +116,4 @@ mkdir data log * [Trust Region Policy Optimization](https://arxiv.org/abs/1502.05477) * [Proximal Policy Optimization Algorithms](https://arxiv.org/abs/1707.06347) * [Emergence of Locomotion Behaviours in Rich Environments](https://arxiv.org/abs/1707.02286) +* [Action-Conditional Video Prediction using Deep Networks in Atari Games](https://arxiv.org/abs/1507.08750) diff --git a/agent/DDPG_agent.py b/agent/DDPG_agent.py new file mode 100644 index 0000000..19f26d0 --- /dev/null +++ b/agent/DDPG_agent.py @@ -0,0 +1,109 @@ +####################################################################### +# Copyright (C) 2017 Shangtong Zhang(zhangshangtong.cpp@gmail.com) # +# Permission given to modify the code as long as you keep this # +# declaration at the top # +####################################################################### + +import numpy as np +import torch.multiprocessing as mp +from network import * +from utils import * +from component import * +import pickle +import os +import time +import gym.monitoring + +class DDPGAgent: + def __init__(self, config): + self.config = config + self.task = config.task_fn() + self.worker_network = config.network_fn() + self.target_network = config.network_fn() + self.target_network.load_state_dict(self.worker_network.state_dict()) + self.actor_opt = config.actor_optimizer_fn(self.worker_network.actor.parameters()) + self.critic_opt = config.critic_optimizer_fn(self.worker_network.critic.parameters()) + self.replay = config.replay_fn() + self.random_process = config.random_process_fn() + self.criterion = nn.MSELoss() + self.total_steps = 0 + + self.state_normalizer = Normalizer(self.task.state_dim) + self.reward_normalizer = Normalizer(1) + + def soft_update(self, target, src): + for target_param, param in zip(target.parameters(), src.parameters()): + target_param.data.copy_(target_param.data * (1.0 - self.config.target_network_mix) + + param.data * self.config.target_network_mix) + + def save(self, file_name): + with open(file_name, 'wb') as f: + torch.save(self.worker_network.state_dict(), f) + + def episode(self, deterministic=False, video_recorder=None): + self.random_process.reset_states() + state = self.task.reset() + state = self.state_normalizer(state) + + config = self.config + actor = self.worker_network.actor + critic = self.worker_network.critic + target_actor = self.target_network.actor + target_critic = self.target_network.critic + + steps = 0 + total_reward = 0.0 + while True: + actor.eval() + action = actor.predict(np.stack([state])).flatten() + if not deterministic: + action += self.random_process.sample() + next_state, reward, done, info = self.task.step(action) + if video_recorder is not None: + video_recorder.capture_frame() + done = (done or (config.max_episode_length and steps >= config.max_episode_length)) + next_state = self.state_normalizer(next_state) + total_reward += reward + reward = self.reward_normalizer(reward) + + if not deterministic: + self.replay.feed([state, action, reward, next_state, int(done)]) + self.total_steps += 1 + + steps += 1 + state = next_state + + if done: + break + + if not deterministic and self.replay.size() >= config.min_memory_size: + self.worker_network.train() + experiences = self.replay.sample() + states, actions, rewards, next_states, terminals = experiences + q_next = target_critic.predict(next_states, target_actor.predict(next_states)) + terminals = critic.to_torch_variable(terminals).unsqueeze(1) + rewards = critic.to_torch_variable(rewards).unsqueeze(1) + q_next = config.discount * q_next * (1 - terminals) + q_next.add_(rewards) + q_next = q_next.detach() + q = critic.predict(states, actions) + critic_loss = self.criterion(q, q_next) + + critic.zero_grad() + self.critic_opt.zero_grad() + critic_loss.backward() + self.critic_opt.step() + + actions = actor.predict(states, False) + var_actions = Variable(actions.data, requires_grad=True) + q = critic.predict(states, var_actions) + q.backward(torch.ones(q.size())) + + actor.zero_grad() + self.actor_opt.zero_grad() + actions.backward(-var_actions.grad.data) + self.actor_opt.step() + + self.soft_update(self.target_network, self.worker_network) + + return total_reward, steps diff --git a/agent/DQN_agent.py b/agent/DQN_agent.py index 3f36503..80c3a16 100644 --- a/agent/DQN_agent.py +++ b/agent/DQN_agent.py @@ -16,29 +16,25 @@ import torch class DQNAgent: def __init__(self, config): self.config = config - self.learning_network = config.network_fn(config.optimizer_fn) - self.target_network = config.network_fn(config.optimizer_fn) + self.learning_network = config.network_fn() + self.target_network = config.network_fn() + self.optimizer = config.optimizer_fn(self.learning_network.parameters()) + self.criterion = nn.MSELoss() self.target_network.load_state_dict(self.learning_network.state_dict()) self.task = config.task_fn() self.replay = config.replay_fn() self.policy = config.policy_fn() self.total_steps = 0 - self.history_buffer = None def episode(self, deterministic=False): episode_start_time = time.time() state = self.task.reset() - if self.history_buffer is None: - self.history_buffer = [np.zeros_like(state)] * self.config.history_length - else: - self.history_buffer.pop(0) - self.history_buffer.append(state) + self.history_buffer = [state] * self.config.history_length state = np.vstack(self.history_buffer) total_reward = 0.0 steps = 0 while True: - value = self.learning_network.predict(np.stack([self.task.normalize_state(state)]), False) - value = value.cpu().data.numpy().flatten() + value = self.learning_network.predict(np.stack([self.task.normalize_state(state)]), True).flatten() if deterministic: action = np.argmax(value) elif self.total_steps < self.config.exploration_steps: @@ -50,10 +46,11 @@ class DQNAgent: self.history_buffer.pop(0) self.history_buffer.append(next_state) next_state = np.vstack(self.history_buffer) + total_reward += np.sum(reward * self.config.reward_weight) + reward = self.config.reward_shift_fn(reward) if not deterministic: self.replay.feed([state, action, reward, next_state, int(done)]) self.total_steps += 1 - total_reward += np.sum(reward * self.config.reward_weight) steps += 1 state = next_state if done: @@ -97,10 +94,10 @@ class DQNAgent: actions = self.learning_network.to_torch_variable(actions, 'int64').unsqueeze(1) q = self.learning_network.predict(states, False) q = q.gather(1, actions).squeeze(1) - loss = self.learning_network.criterion(q, q_next) - self.learning_network.zero_grad() + loss = self.criterion(q, q_next) + self.optimizer.zero_grad() loss.backward() - self.learning_network.optimizer.step() + self.optimizer.step() if not deterministic and self.total_steps % self.config.target_network_update_freq == 0: self.target_network.load_state_dict(self.learning_network.state_dict()) if not deterministic and self.total_steps > self.config.exploration_steps: @@ -112,4 +109,4 @@ class DQNAgent: def save(self, file_name): with open(file_name, 'wb') as f: - pickle.dump(self.learning_network.state_dict(), f) + torch.save(self.learning_network.state_dict(), f) diff --git a/agent/__init__.py b/agent/__init__.py index 8462c94..ab0c05d 100644 --- a/agent/__init__.py +++ b/agent/__init__.py @@ -1,2 +1,3 @@ from .async_agent import * from .DQN_agent import * +from .DDPG_agent import * \ No newline at end of file diff --git a/agent/async_agent.py b/agent/async_agent.py index b5e50d7..7dc26b0 100644 --- a/agent/async_agent.py +++ b/agent/async_agent.py @@ -13,8 +13,11 @@ from async_worker import * import pickle import os import time +import sys def train(id, config, learning_network, extra): + np.random.seed() + torch.manual_seed(np.random.randint(sys.maxsize)) worker = config.worker(config, learning_network, extra) episode = 0 rewards = [] @@ -27,6 +30,8 @@ def train(id, config, learning_network, extra): episode += 1 def evaluate(config, task, learning_network, extra): + np.random.seed() + torch.manual_seed(np.random.randint(sys.maxsize)) test_rewards = [] test_points = [] test_wall_times = [] diff --git a/async_worker/actor_critic.py b/async_worker/actor_critic.py index 9ad9417..e821da6 100644 --- a/async_worker/actor_critic.py +++ b/async_worker/actor_critic.py @@ -33,6 +33,7 @@ class AdvantageActorCritic: steps += 1 total_reward += reward + reward = config.reward_shift_fn(reward) if deterministic: if terminal: diff --git a/async_worker/n_step_q.py b/async_worker/n_step_q.py index 3a5899b..47c4cfb 100644 --- a/async_worker/n_step_q.py +++ b/async_worker/n_step_q.py @@ -34,6 +34,7 @@ class NStepQLearning: steps += 1 total_reward += reward + reward = config.reward_shift_fn(reward) if deterministic: if terminal: diff --git a/async_worker/one_step_q.py b/async_worker/one_step_q.py index d190753..a37bc4c 100644 --- a/async_worker/one_step_q.py +++ b/async_worker/one_step_q.py @@ -34,6 +34,7 @@ class OneStepQLearning: steps += 1 total_reward += reward + reward = config.reward_shift_fn(reward) if deterministic: if terminal: diff --git a/async_worker/one_step_sarsa.py b/async_worker/one_step_sarsa.py index d763867..1434108 100644 --- a/async_worker/one_step_sarsa.py +++ b/async_worker/one_step_sarsa.py @@ -37,6 +37,7 @@ class OneStepSarsa: steps += 1 total_reward += reward + reward = config.reward_shift_fn(reward) if deterministic: if terminal: diff --git a/async_worker/ppo.py b/async_worker/ppo.py index 3d7b6f6..aee1398 100644 --- a/async_worker/ppo.py +++ b/async_worker/ppo.py @@ -104,6 +104,7 @@ class ProximalPolicyOptimization: advs.append(cum_adv) advantages = advs[::-1] returns = list(returns) + replay.feed([states, actions, returns, advantages]) batched_rewards /= batched_episode diff --git a/component/atari_wrapper.py b/component/atari_wrapper.py index 936eec0..acd8e90 100644 --- a/component/atari_wrapper.py +++ b/component/atari_wrapper.py @@ -105,6 +105,39 @@ class MaxAndSkipEnv(gym.Wrapper): self._obs_buffer.append(obs) return obs +def _process_frame84_rgb(frame): + img = np.reshape(frame, [210, 160, 3]).astype(np.float32) + img = Image.fromarray(img) + resized_screen = img.resize((84, 110, 3), Image.BILINEAR) + resized_screen = np.array(resized_screen) + x_t = resized_screen[18:102, :, :] + x_t = x_t.reshape((84, 84, 3)) + return x_t + +class DatasetEnv(gym.Wrapper): + def __init__(self, env=None): + super(DatasetEnv, self).__init__(env) + self.saved_obs = [] + self.saved_actions = [] + + def get_saved(self): + return self.saved_obs, self.saved_actions + + def clear_saved(self): + self.saved_obs = [] + self.saved_actions = [] + + def _step(self, action): + obs, reward, done, info = self.env.step(action) + self.saved_actions.append(action) + self.saved_obs.append(obs) + return obs, reward, done, info + + def _reset(self): + obs = self.env.reset() + self.saved_obs.append(obs) + return obs + def _process_frame84(frame): img = np.reshape(frame, [210, 160, 3]).astype(np.float32) img = img[:, :, 0] * 0.299 + img[:, :, 1] * 0.587 + img[:, :, 2] * 0.114 @@ -143,7 +176,16 @@ class ProcessFrame(gym.Wrapper): def _reset(self): return self.process_fn(self.env.reset()) -class ClippedRewardsWrapper(gym.Wrapper): +class NormalizeFrame(gym.Wrapper): + def __init__(self, env=None): + super(NormalizeFrame, self).__init__(env) + + def _normalize(self, obs): + return np.asarray(obs, dtype=np.float32) / 255.0 + def _step(self, action): obs, reward, done, info = self.env.step(action) - return obs, np.sign(reward), done, info + return self._normalize(obs), reward, done, info + + def _reset(self): + return self._normalize(self.env.reset()) diff --git a/component/task.py b/component/task.py index d52d009..8079d34 100644 --- a/component/task.py +++ b/component/task.py @@ -70,8 +70,7 @@ class PixelAtari(BasicTask): env = MaxAndSkipEnv(env, skip=frame_skip) if 'FIRE' in env.unwrapped.get_action_meanings(): env = FireResetEnv(env) - env = ProcessFrame(env, frame_size) - self.env = ClippedRewardsWrapper(env) + self.env = ProcessFrame(env, frame_size) self.action_dim = self.env.action_space.n def normalize_state(self, state): diff --git a/dataset.py b/dataset.py new file mode 100644 index 0000000..6193f8d --- /dev/null +++ b/dataset.py @@ -0,0 +1,107 @@ +####################################################################### +# Copyright (C) 2017 Shangtong Zhang(zhangshangtong.cpp@gmail.com) # +# Permission given to modify the code as long as you keep this # +# declaration at the top # +####################################################################### + +from agent import * +from component import * +from utils import * +import torchvision +import torch + +# PREFIX = '.' +PREFIX = '/local/data' + +def dqn_pixel_atari(name): + config = Config() + config.history_length = 4 + config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False) + action_dim = config.task_fn().action_dim + config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01) + config.network_fn = lambda optimizer_fn: NatureConvNet(config.history_length, action_dim, optimizer_fn) + config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1) + config.replay_fn = lambda: Replay(memory_size=1000000, batch_size=32, dtype=np.uint8) + config.discount = 0.99 + config.target_network_update_freq = 10000 + config.max_episode_length = 0 + config.exploration_steps= 50000 + config.logger = Logger('./log', logger) + config.test_interval = 10 + config.test_repetitions = 1 + config.double_q = False + return DQNAgent(config) + +def train_dqn(game): + agent = dqn_pixel_atari(game) + run_episodes(agent) + +def episode(env, agent): + config = agent.config + policy = GreedyPolicy(epsilon=0.3, final_step=1, min_epsilon=0.3) + state = env.reset() + history_buffer = [state] * config.history_length + state = np.vstack(history_buffer) + total_reward = 0.0 + steps = 0 + while True: + value = agent.learning_network.predict(np.stack([state]), False) + value = value.cpu().data.numpy().flatten() + action = policy.sample(value) + next_state, reward, done, info = env.step(action) + history_buffer.pop(0) + history_buffer.append(next_state) + state = np.vstack(history_buffer) + done = (done or (config.max_episode_length and steps > config.max_episode_length)) + steps += 1 + total_reward += reward + if done: + break + return total_reward, steps + +def generate_dateset(game): + agent = dqn_pixel_atari(game) + model_file = 'data/%s-%s-model-%s.bin' % (agent.__class__.__name__, agent.config.tag, agent.task.name) + with open(model_file, 'rb') as f: + saved_state = torch.load(model_file, map_location=lambda storage, loc: storage) + agent.learning_network.load_state_dict(saved_state) + + env = gym.make(game) + env = EpisodicLifeEnv(env) + env = MaxAndSkipEnv(env, skip=4) + dataset_env = DatasetEnv(env) + env = ProcessFrame(dataset_env, 84) + env = NormalizeFrame(env) + env = ClippedRewardsWrapper(env) + + ep = 0 + max_ep = 200 + mkdir('%s/dataset/%s' % (PREFIX, game)) + obs_sum = 0.0 + obs_count = 0 + while True: + rewards, steps = episode(env, agent) + path = '%s/dataset/%s/%05d' % (PREFIX, game, ep) + mkdir(path) + logger.info('Episode %d, reward %f, steps %d' % (ep, rewards, steps)) + with open('%s/action.bin' % (path), 'wb') as f: + pickle.dump(dataset_env.saved_actions, f) + obs_sum += np.asarray(dataset_env.saved_obs).sum(0) + obs_count += len(dataset_env.saved_obs) + for ind, obs in enumerate(dataset_env.saved_obs): + obs = torch.from_numpy(np.transpose(obs, (2, 0, 1))) + torchvision.utils.save_image(obs, '%s/%05d.png' % (path, ind)) + dataset_env.clear_saved() + ep += 1 + if ep >= max_ep: + break + obs_mean = np.transpose(obs_sum, (2, 0, 1)) / obs_count + with open('%s/dataset/%s/meta.bin' % (PREFIX, game), 'wb') as f: + pickle.dump({'episodes': ep, + 'mean_obs': obs_mean}, f) + +if __name__ == '__main__': + mkdir('dataset') + game = 'PongNoFrameskip-v4' + # train_dqn(game) + generate_dateset(game) \ No newline at end of file diff --git a/images/ACVP.png b/images/ACVP.png new file mode 100644 index 0000000..8ce8312 Binary files /dev/null and b/images/ACVP.png differ diff --git a/main.py b/main.py index 6e2ca81..598cc82 100644 --- a/main.py +++ b/main.py @@ -1,13 +1,20 @@ +####################################################################### +# Copyright (C) 2017 Shangtong Zhang(zhangshangtong.cpp@gmail.com) # +# Permission given to modify the code as long as you keep this # +# declaration at the top # +####################################################################### + import logging from agent import * from component import * from utils import * +import model.action_conditional_video_prediction as acvp def dqn_cart_pole(): config = Config() config.task_fn = lambda: CartPole() config.optimizer_fn = lambda params: torch.optim.RMSprop(params, 0.001) - config.network_fn = lambda optimizer_fn: FCNet([8, 50, 200, 2], optimizer_fn) + config.network_fn = lambda: FCNet([8, 50, 200, 2]) # config.network_fn = lambda optimizer_fn: DuelingFCNet([8, 50, 200, 2], optimizer_fn) config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=10000, min_epsilon=0.1) config.replay_fn = lambda: Replay(memory_size=10000, batch_size=10) @@ -15,7 +22,7 @@ def dqn_cart_pole(): config.target_network_update_freq = 200 config.max_episode_length = 200 config.exploration_steps = 1000 - config.logger = Logger('./log', gym.logger) + config.logger = Logger('./log', logger) config.history_length = 2 config.test_interval = 100 config.test_repetitions = 50 @@ -29,8 +36,8 @@ def async_cart_pole(): config.optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) config.network_fn = lambda: FCNet([4, 50, 200, 2]) config.policy_fn = lambda: GreedyPolicy(epsilon=0.5, final_step=5000, min_epsilon=0.1) - config.worker = OneStepQLearning - # config.worker = NStepQLearning + # config.worker = OneStepQLearning + config.worker = NStepQLearning # config.worker = OneStepSarsa config.discount = 0.99 config.target_network_update_freq = 200 @@ -39,7 +46,7 @@ def async_cart_pole(): config.update_interval = 6 config.test_interval = 1 config.test_repetitions = 50 - config.logger = Logger('./log', gym.logger) + config.logger = Logger('./log', logger) agent = AsyncAgent(config) agent.run() @@ -56,7 +63,7 @@ def a3c_cart_pole(): config.update_interval = 6 config.test_interval = 1 config.test_repetitions = 30 - config.logger = Logger('./log', gym.logger) + config.logger = Logger('./log', logger) config.gae_tau = 1.0 config.entropy_weight = 0.01 agent = AsyncAgent(config) @@ -68,15 +75,16 @@ def dqn_pixel_atari(name): config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False) action_dim = config.task_fn().action_dim config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01) - config.network_fn = lambda optimizer_fn: NatureConvNet(config.history_length, action_dim, optimizer_fn) + config.network_fn = lambda: NatureConvNet(config.history_length, action_dim) # config.network_fn = lambda optimizer_fn: DuelingNatureConvNet(config.history_length, n_actions, optimizer_fn) config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1) config.replay_fn = lambda: Replay(memory_size=1000000, batch_size=32, dtype=np.uint8) + config.reward_shift_fn = lambda r: np.sign(r) config.discount = 0.99 config.target_network_update_freq = 10000 config.max_episode_length = 0 config.exploration_steps= 50000 - config.logger = Logger('./log', gym.logger) + config.logger = Logger('./log', logger) config.test_interval = 10 config.test_repetitions = 1 # config.double_q = True @@ -97,14 +105,15 @@ def async_pixel_atari(name): # config.worker = OneStepSarsa # config.worker = NStepQLearning config.worker = OneStepQLearning + config.reward_shift_fn = lambda r: np.sign(r) config.discount = 0.99 config.target_network_update_freq = 10000 config.max_episode_length = 10000 - config.num_workers = 10 + config.num_workers = 6 config.update_interval = 20 config.test_interval = 50000 config.test_repetitions = 1 - config.logger = Logger('./log', gym.logger) + config.logger = Logger('./log', logger) agent = AsyncAgent(config) agent.run() @@ -116,15 +125,16 @@ def a3c_pixel_atari(name): config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.0001) config.network_fn = lambda: OpenAIActorCriticConvNet( config.history_length, task.env.action_space.n, LSTM=True) + config.reward_shift_fn = lambda r: np.sign(r) config.policy_fn = SamplePolicy config.worker = AdvantageActorCritic config.discount = 0.99 config.max_episode_length = 10000 - config.num_workers = 10 + config.num_workers = 6 config.update_interval = 20 config.test_interval = 50000 config.test_repetitions = 1 - config.logger = Logger('./log', gym.logger) + config.logger = Logger('./log', logger) agent = AsyncAgent(config) agent.run() @@ -134,15 +144,14 @@ def dqn_fruit(): config.optimizer_fn = lambda params: torch.optim.SGD(params, 0.01, momentum=0.9) config.reward_weight = np.ones(10) / 10 config.hybrid_reward = False - config.network_fn = lambda optimizer_fn: FruitHRFCNet( - 98, 4, config.reward_weight, optimizer_fn) + config.network_fn = lambda: FruitHRFCNet(98, 4, config.reward_weight) config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=10000, min_epsilon=0.1) config.replay_fn = lambda: Replay(memory_size=10000, batch_size=15) config.discount = 0.95 config.target_network_update_freq = 200 config.max_episode_length = 100 config.exploration_steps = 200 - config.logger = Logger('./log', gym.logger) + config.logger = Logger('./log', logger) config.history_length = 1 config.test_interval = 0 config.test_repetitions = 10 @@ -156,15 +165,14 @@ def hrdqn_fruit(): config.hybrid_reward = True config.reward_weight = np.ones(10) / 10 config.optimizer_fn = lambda params: torch.optim.SGD(params, 0.01, momentum=0.9) - config.network_fn = lambda optimizer_fn: FruitHRFCNet( - 98, 4, config.reward_weight, optimizer_fn) + config.network_fn = lambda optimizer_fn: FruitHRFCNet(98, 4, config.reward_weight) config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=10000, min_epsilon=0.1) config.replay_fn = lambda: HybridRewardReplay(memory_size=10000, batch_size=15) config.discount = 0.95 config.target_network_update_freq = 200 config.max_episode_length = 100 config.exploration_steps = 200 - config.logger = Logger('./log', gym.logger) + config.logger = Logger('./log', logger) config.history_length = 1 config.test_interval = 0 config.test_repetitions = 10 @@ -195,7 +203,7 @@ def a3c_continuous(): config.test_repetitions = 1 config.entropy_weight = 0 config.gradient_clip = 40 - config.logger = Logger('./log', gym.logger) + config.logger = Logger('./log', logger) agent = AsyncAgent(config) agent.run() @@ -228,16 +236,16 @@ def p3o_continuous(): config.rollout_length = 10000 config.optimize_epochs = 1 config.ppo_ratio_clip = 0.2 - config.logger = Logger('./log', gym.logger) + config.logger = Logger('./log', logger) agent = AsyncAgent(config) agent.run() def d3pg_continuous(): config = Config() - # config.task_fn = lambda: Pendulum() + config.task_fn = lambda: Pendulum() # config.task_fn = lambda: ContinuousLunarLander() # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1') - config.task_fn = lambda: Roboschool('RoboschoolReacher-v1') + # config.task_fn = lambda: Roboschool('RoboschoolReacher-v1') # config.task_fn = lambda: BipedalWalker() task = config.task_fn() config.actor_network_fn = lambda: DeterministicActorNet( @@ -262,20 +270,61 @@ def d3pg_continuous(): config.test_interval = 500 config.test_repetitions = 1 config.gradient_clip = 20 - config.logger = Logger('./log', gym.logger) + config.logger = Logger('./log', logger) agent = AsyncAgent(config) agent.run() +def ddpg_continuous(): + config = Config() + # config.task_fn = lambda: Pendulum() + # config.task_fn = lambda: ContinuousLunarLander() + # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1') + config.task_fn = lambda: Roboschool('RoboschoolReacher-v1') + # config.task_fn = lambda: Roboschool('RoboschoolHopper-v1') + # config.task_fn = lambda: Roboschool('RoboschoolAnt-v1') + # config.task_fn = lambda: Roboschool('RoboschoolWalker2d-v1') + # config.task_fn = lambda: BipedalWalker() + task = config.task_fn() + config.actor_network_fn = lambda: DeterministicActorNet( + task.state_dim, task.action_dim, F.tanh, 1, non_linear=F.relu, batch_norm=False) + config.critic_network_fn = lambda: DeterministicCriticNet( + task.state_dim, task.action_dim, non_linear=F.relu, batch_norm=False) + config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn) + config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, lr=1e-4) + config.critic_optimizer_fn =\ + lambda params: torch.optim.Adam(params, lr=1e-3, weight_decay=0.01) + config.replay_fn = lambda: SharedReplay(memory_size=1000000, batch_size=64, + state_shape=(task.state_dim, ), action_shape=(task.action_dim, )) + config.discount = 0.99 + config.max_episode_length = task.max_episode_steps + config.random_process_fn = \ + lambda: OrnsteinUhlenbeckProcess(size=task.action_dim, theta=0.15, sigma=0.2, + n_steps_annealing=100000) + config.worker = DeterministicPolicyGradient + config.min_memory_size = 50 + config.target_network_mix = 0.001 + config.test_interval = 0 + config.test_repetitions = 1 + config.gradient_clip = 40 + config.render_episode_freq = 0 + config.logger = Logger('./log', logger) + run_episodes(DDPGAgent(config)) + if __name__ == '__main__': - # gym.logger.setLevel(logging.DEBUG) - gym.logger.setLevel(logging.INFO) + mkdir('data') + mkdir('data/video') + mkdir('log') + os.system('export OMP_NUM_THREADS=1') + # logger.setLevel(logging.DEBUG) + logger.setLevel(logging.INFO) # dqn_cart_pole() # async_cart_pole() # a3c_cart_pole() # a3c_continuous() # p3o_continuous() - d3pg_continuous() + # d3pg_continuous() + ddpg_continuous() # dqn_fruit() # hrdqn_fruit() @@ -288,3 +337,5 @@ if __name__ == '__main__': # async_pixel_atari('BreakoutNoFrameskip-v4') # a3c_pixel_atari('BreakoutNoFrameskip-v4') + # acvp.train('PongNoFrameskip-v4') + diff --git a/model/__init__.py b/model/__init__.py new file mode 100644 index 0000000..d577c8d --- /dev/null +++ b/model/__init__.py @@ -0,0 +1 @@ +from .action_conditional_video_prediction import train \ No newline at end of file diff --git a/model/action_conditional_video_prediction.py b/model/action_conditional_video_prediction.py new file mode 100644 index 0000000..a16fa44 --- /dev/null +++ b/model/action_conditional_video_prediction.py @@ -0,0 +1,212 @@ +####################################################################### +# Copyright (C) 2017 Shangtong Zhang(zhangshangtong.cpp@gmail.com) # +# Permission given to modify the code as long as you keep this # +# declaration at the top # +####################################################################### + +import torch +from torch.autograd import Variable +import torch.nn as nn +import torch.nn.functional as F +import numpy as np +import pickle +import torchvision +from skimage import io +from collections import deque +import gym +import torch.optim +from utils import * +from tqdm import tqdm + +PREFIX = '.' +# PREFIX = '/local/data' + +class Network(nn.Module): + def __init__(self, num_actions, gpu=True): + super(Network, self).__init__() + + self.conv1 = nn.Conv2d(12, 64, 8, 2, (0, 1)) + self.conv2 = nn.Conv2d(64, 128, 6, 2, (1, 1)) + self.conv3 = nn.Conv2d(128, 128, 6, 2, (1, 1)) + self.conv4 = nn.Conv2d(128, 128, 4, 2, (0, 0)) + + self.hidden_units = 128 * 11 * 8 + + self.fc5 = nn.Linear(self.hidden_units, 2048) + self.fc_encode = nn.Linear(2048, 2048) + self.fc_action = nn.Linear(num_actions, 2048) + self.fc_decode = nn.Linear(2048, 2048) + self.fc8 = nn.Linear(2048, self.hidden_units) + + self.deconv9 = nn.ConvTranspose2d(128, 128, 4, 2) + self.deconv10 = nn.ConvTranspose2d(128, 128, 6, 2, (1, 1)) + self.deconv11 = nn.ConvTranspose2d(128, 128, 6, 2, (1, 1)) + self.deconv12 = nn.ConvTranspose2d(128, 3, 8, 2, (0, 1)) + + self.gpu = gpu and torch.cuda.is_available() + if self.gpu: + self.cuda() + self.FloatTensor = torch.cuda.FloatTensor + else: + self.FloatTensor = torch.FloatTensor + + self.init_weights() + self.criterion = nn.MSELoss() + self.opt = torch.optim.Adam(self.parameters(), 1e-4) + + def init_weights(self): + for layer in self.children(): + if isinstance(layer, nn.Conv2d) or isinstance(layer, nn.ConvTranspose2d): + nn.init.xavier_uniform(layer.weight.data) + nn.init.constant(layer.bias.data, 0) + nn.init.uniform(self.fc_encode.weight.data, -1, 1) + nn.init.uniform(self.fc_decode.weight.data, -1, 1) + nn.init.uniform(self.fc_action.weight.data, -0.1, 0.1) + + def to_torch_variable(self, x, dtype='float32'): + if isinstance(x, Variable): + return x + if not isinstance(x, torch.FloatTensor): + x = torch.from_numpy(np.asarray(x, dtype=dtype)) + if self.gpu: + x = x.cuda() + return Variable(x) + + def forward(self, obs, action): + x = F.relu(self.conv1(obs)) + x = F.relu(self.conv2(x)) + x = F.relu(self.conv3(x)) + x = F.relu(self.conv4(x)) + x = x.view((-1, self.hidden_units)) + x = F.relu(self.fc5(x)) + x = self.fc_encode(x) + action = self.fc_action(action) + x = torch.mul(x, action) + x = self.fc_decode(x) + x = F.relu(self.fc8(x)) + x = x.view((-1, 128, 11, 8)) + x = F.relu(self.deconv9(x)) + x = F.relu(self.deconv10(x)) + x = F.relu(self.deconv11(x)) + x = self.deconv12(x) + return x + + def fit(self, x, a, y): + x = self.to_torch_variable(x) + a = self.to_torch_variable(a) + y = self.to_torch_variable(y) + y_ = self.forward(x, a) + loss = self.criterion(y_, y) + self.opt.zero_grad() + loss.backward() + for param in self.parameters(): + param.grad.data.clamp_(-0.1, 0.1) + self.opt.step() + return np.asscalar(loss.cpu().data.numpy()) + + def evaluate(self, x, a, y): + x = self.to_torch_variable(x) + a = self.to_torch_variable(a) + y = self.to_torch_variable(y) + y_ = self.forward(x, a) + loss = self.criterion(y_, y) + return np.asscalar(loss.cpu().data.numpy()) + + def predict(self, x, a): + x = self.to_torch_variable(x) + a = self.to_torch_variable(a) + return self.forward(x, a).cpu().data.numpy() + +def load_episode(game, ep, num_actions): + path = '%s/dataset/%s/%05d' % (PREFIX, game, ep) + with open('%s/action.bin' % (path), 'rb') as f: + actions = pickle.load(f) + num_frames = len(actions) + 1 + frames = [] + + for i in range(1, num_frames): + frame = io.imread('%s/%05d.png' % (path, i)) + frame = np.transpose(frame, (2, 0, 1)) + frames.append(frame.astype(np.uint8)) + + actions = actions[1:] + encoded_actions = np.zeros((len(actions), num_actions)) + encoded_actions[np.arange(len(actions)), actions] = 1 + + return frames, encoded_actions + +def extend_frames(frames, actions): + buffer = deque(maxlen=4) + extended_frames = [] + targets = [] + + for i in range(len(frames) - 1): + buffer.append(frames[i]) + if len(buffer) >= 4: + extended_frames.append(np.vstack(buffer)) + targets.append(frames[i + 1]) + actions = actions[3:, :] + + return np.stack(extended_frames), actions, np.stack(targets) + +def train(game): + env = gym.make(game) + num_actions = env.action_space.n + + net = Network(num_actions) + + with open('%s/dataset/%s/meta.bin' % (PREFIX, game), 'rb') as f: + meta = pickle.load(f) + episodes = meta['episodes'] + mean_obs = meta['mean_obs'] + + def pre_process(x): + if x.shape[1] == 12: + return (x - np.vstack([mean_obs] * 4)) / 255.0 + elif x.shape[1] == 3: + return (x - mean_obs) / 255.0 + else: + assert False + + def post_process(y): + return (y * 255 + mean_obs).astype(np.uint8) + + train_episodes = int(episodes * 0.95) + indices_train = np.arange(train_episodes) + iteration = 0 + while True: + np.random.shuffle(indices_train) + for ep in indices_train: + frames, actions = load_episode(game, ep, num_actions) + frames, actions, targets = extend_frames(frames, actions) + batcher = Batcher(32, [frames, actions, targets]) + batcher.shuffle() + while not batcher.end(): + if iteration % 10000 == 0: + mkdir('data/acvp-sample') + losses = [] + test_indices = range(train_episodes, episodes) + ep_to_print = np.random.choice(test_indices) + for test_ep in tqdm(test_indices): + frames, actions = load_episode(game, test_ep, num_actions) + frames, actions, targets = extend_frames(frames, actions) + test_batcher = Batcher(32, [frames, actions, targets]) + while not test_batcher.end(): + x, a, y = test_batcher.next_batch() + losses.append(net.evaluate(pre_process(x), a, pre_process(y))) + if test_ep == ep_to_print: + test_batcher.reset() + x, a, y = test_batcher.next_batch() + y_ = post_process(net.predict(pre_process(x), a)) + torchvision.utils.save_image(torch.from_numpy(y_), 'data/acvp-sample/%s-%09d.png' % (game, iteration)) + torchvision.utils.save_image(torch.from_numpy(y), 'data/acvp-sample/%s-%09d-truth.png' % (game, iteration)) + + logger.info('Iteration %d, test loss %f' % (iteration, np.mean(losses))) + torch.save(net.state_dict(), 'data/acvp-%s.bin' % (game)) + + x, a, y = batcher.next_batch() + loss = net.fit(pre_process(x), a, pre_process(y)) + if iteration % 100 == 0: + logger.info('Iteration %d, loss %f' % (iteration, loss)) + + iteration += 1 diff --git a/network/base_network.py b/network/base_network.py index 08beffd..21086a7 100644 --- a/network/base_network.py +++ b/network/base_network.py @@ -13,8 +13,6 @@ import numpy as np # Base class for all kinds of network class BasicNet: def __init__(self, optimizer_fn, gpu, LSTM=False): - if optimizer_fn is not None: - self.optimizer = optimizer_fn(self.parameters()) self.gpu = gpu and torch.cuda.is_available() self.LSTM = LSTM if self.gpu: @@ -57,8 +55,8 @@ class ActorCriticNet(BasicNet): def predict(self, x): phi = self.forward(x, True) pre_prob = self.fc_actor(phi) - prob = F.softmax(pre_prob) - log_prob = F.log_softmax(pre_prob) + prob = F.softmax(pre_prob, dim=1) + log_prob = F.log_softmax(pre_prob, dim=1) value = self.fc_critic(phi) return prob, log_prob, value diff --git a/network/continuous_action_network.py b/network/continuous_action_network.py index d900094..1ad91ff 100644 --- a/network/continuous_action_network.py +++ b/network/continuous_action_network.py @@ -14,9 +14,9 @@ class DeterministicActorNet(nn.Module, BasicNet): action_scale, gpu=False, batch_norm=False, - non_linear=F.relu): + non_linear=F.relu, + hidden_size=64): super(DeterministicActorNet, self).__init__() - hidden_size = 64 self.layer1 = nn.Linear(state_dim, hidden_size) self.layer3 = nn.Linear(hidden_size, action_dim) self.action_gate = action_gate @@ -70,9 +70,9 @@ class DeterministicCriticNet(nn.Module, BasicNet): action_dim, gpu=False, batch_norm=False, - non_linear=F.relu): + non_linear=F.relu, + hidden_size=64): super(DeterministicCriticNet, self).__init__() - hidden_size = 64 self.layer1 = nn.Linear(state_dim, hidden_size) self.layer2 = nn.Linear(hidden_size + action_dim, hidden_size) self.layer3 = nn.Linear(hidden_size, 1) @@ -116,9 +116,15 @@ class DeterministicCriticNet(nn.Module, BasicNet): return self.forward(x, action) class GaussianActorNet(nn.Module, BasicNet): - def __init__(self, state_dim, action_dim, action_scale=1.0, action_gate=None, gpu=False, unit_std=True): + def __init__(self, + state_dim, + action_dim, + action_scale=1.0, + action_gate=None, + gpu=False, + unit_std=True, + hidden_size=64): super(GaussianActorNet, self).__init__() - hidden_size = 64 self.fc1 = nn.Linear(state_dim, hidden_size) self.fc2 = nn.Linear(hidden_size, hidden_size) self.action_mean = nn.Linear(hidden_size, action_dim) @@ -161,9 +167,11 @@ class GaussianActorNet(nn.Module, BasicNet): return 0.5 * (1 + (2 * std.pow(2) * np.pi + 1e-5).log()).sum(1).mean() class GaussianCriticNet(nn.Module, BasicNet): - def __init__(self, state_dim, gpu=False): + def __init__(self, + state_dim, + gpu=False, + hidden_size=64): super(GaussianCriticNet, self).__init__() - hidden_size = 64 self.fc1 = nn.Linear(state_dim, hidden_size) self.fc2 = nn.Linear(hidden_size, hidden_size) self.fc_value = nn.Linear(hidden_size, 1) diff --git a/network/conv_network.py b/network/conv_network.py index 4951b4f..331b734 100644 --- a/network/conv_network.py +++ b/network/conv_network.py @@ -15,8 +15,7 @@ class NatureConvNet(nn.Module, VanillaNet): self.conv3 = nn.Conv2d(64, 64, kernel_size=3, stride=1) self.fc4 = nn.Linear(7 * 7 * 64, 512) self.fc5 = nn.Linear(512, n_actions) - self.criterion = nn.MSELoss() - BasicNet.__init__(self, optimizer_fn, gpu) + BasicNet.__init__(self, None, gpu) def forward(self, x): x = self.to_torch_variable(x) @@ -37,8 +36,7 @@ class DuelingNatureConvNet(nn.Module, DuelingNet): self.fc4 = nn.Linear(7 * 7 * 64, 512) self.fc_advantage = nn.Linear(512, n_actions) self.fc_value = nn.Linear(512, 1) - self.criterion = nn.MSELoss() - BasicNet.__init__(self, optimizer_fn, gpu) + BasicNet.__init__(self, None, gpu) def forward(self, x): x = self.to_torch_variable(x) diff --git a/network/shallow_network.py b/network/shallow_network.py index 148a2bb..265a9cb 100644 --- a/network/shallow_network.py +++ b/network/shallow_network.py @@ -13,7 +13,6 @@ class FCNet(nn.Module, VanillaNet): self.fc1 = nn.Linear(dims[0], dims[1]) self.fc2 = nn.Linear(dims[1], dims[2]) self.fc3 = nn.Linear(dims[2], dims[3]) - self.criterion = nn.MSELoss() BasicNet.__init__(self, optimizer_fn, gpu) def forward(self, x): @@ -32,7 +31,6 @@ class DuelingFCNet(nn.Module, DuelingNet): self.fc2 = nn.Linear(dims[1], dims[2]) self.fc_value = nn.Linear(dims[2], 1) self.fc_advantage = nn.Linear(dims[2], dims[3]) - self.criterion = nn.MSELoss() BasicNet.__init__(self, optimizer_fn, gpu) def forward(self, x): @@ -67,7 +65,6 @@ class FruitHRFCNet(nn.Module, VanillaNet): hidden_size = 250 self.fc1 = nn.Linear(state_dim, hidden_size) self.fc2 = nn.ModuleList([nn.Linear(hidden_size, action_dim) for _ in head_weights]) - self.criterion = nn.MSELoss() self.head_weights = head_weights BasicNet.__init__(self, optimizer_fn, gpu) @@ -93,7 +90,6 @@ class FruitMultiStatesFCNet(nn.Module, BasicNet): hidden_size = 250 self.fc1 = nn.ModuleList([nn.Linear(state_dim, hidden_size) for _ in head_weights]) self.fc2 = nn.ModuleList([nn.Linear(hidden_size, action_dim) for _ in head_weights]) - self.criterion = nn.MSELoss() self.head_weights = head_weights self.state_dim = state_dim self.n_heads = head_weights.shape[0] diff --git a/utils/__init__.py b/utils/__init__.py index c9b2200..e71c201 100644 --- a/utils/__init__.py +++ b/utils/__init__.py @@ -1,8 +1,8 @@ from .config import * from .normalizer import * from .misc import * - -try: - from .tf_logger import Logger -except: - from .vanilla_logger import Logger \ No newline at end of file +from .tf_logger import Logger +import logging +logging.basicConfig(format='%(asctime)s - %(name)s - %(levelname)s: %(message)s') +logger = logging.getLogger('MAIN') +logger.setLevel(logging.INFO) \ No newline at end of file diff --git a/utils/config.py b/utils/config.py index 132a866..9df817d 100644 --- a/utils/config.py +++ b/utils/config.py @@ -37,6 +37,7 @@ class Config: self.noise_decay_interval = 0 self.target_network_mix = 0.001 self.action_shift_fn = lambda a: a + self.reward_shift_fn = lambda r: r self.reward_weight = 1 self.hybrid_reward = False self.target_type = self.q_target @@ -49,3 +50,4 @@ class Config: self.save_interval = 0 self.max_steps = 0 self.success_threshold = float('inf') + self.render_episode_freq = 0 diff --git a/utils/misc.py b/utils/misc.py index 8f4ba1a..8510196 100644 --- a/utils/misc.py +++ b/utils/misc.py @@ -6,7 +6,8 @@ import numpy as np import pickle - +import os +import gym.monitoring def run_episodes(agent): config = agent.config @@ -30,9 +31,18 @@ def run_episodes(agent): agent_type, config.tag, agent.task.name), 'wb') as f: pickle.dump([steps, rewards], f) + if config.render_episode_freq and ep % config.render_episode_freq == 0: + video_recoder = gym.monitoring.VideoRecorder( + env=agent.task.env, base_path='./data/video/%s-%s-%s-%d' % (agent_type, config.tag, agent.task.name, ep)) + agent.episode(True, video_recoder) + video_recoder.close() + if config.episode_limit and ep > config.episode_limit: break + if config.max_steps and agent.total_steps > config.max_steps: + break + if config.test_interval and ep % config.test_interval == 0: config.logger.info('Testing...') agent.save('data/%s-%s-model-%s.bin' % (agent_type, config.tag, agent.task.name)) @@ -47,11 +57,42 @@ def run_episodes(agent): pickle.dump({'rewards': rewards, 'steps': steps, 'test_rewards': avg_test_rewards}, f) - if avg_reward > agent.task.success_threshold: + if avg_reward > config.success_threshold: break return steps, rewards, avg_test_rewards def sync_grad(target_network, src_network): for param, src_param in zip(target_network.parameters(), src_network.parameters()): - param._grad = src_param.grad.clone() \ No newline at end of file + param._grad = src_param.grad.clone() + +def mkdir(path): + if not os.path.exists(path): + os.mkdir(path) + +class Batcher: + def __init__(self, batch_size, data): + self.batch_size = batch_size + self.data = data + self.num_entries = len(data[0]) + self.reset() + + def reset(self): + self.batch_start = 0 + self.batch_end = self.batch_start + self.batch_size + + def end(self): + return self.batch_start >= self.num_entries + + def next_batch(self): + batch = [] + for d in self.data: + batch.append(d[self.batch_start: self.batch_end]) + self.batch_start = self.batch_end + self.batch_end = min(self.batch_start + self.batch_size, self.num_entries) + return batch + + def shuffle(self): + indices = np.arange(self.num_entries) + np.random.shuffle(indices) + self.data = [d[indices] for d in self.data] diff --git a/utils/tf_logger.py b/utils/tf_logger.py index 2b53f71..72e2905 100644 --- a/utils/tf_logger.py +++ b/utils/tf_logger.py @@ -1,84 +1,25 @@ -# Adapted from https://github.com/yunjey/pytorch-tutorial/blob/master/tutorials/04-utils/tensorboard/logger.py -# Code referenced from https://gist.github.com/gyglim/1f8dfb1b5c82627ae3efcfbbadb9f514 -import tensorflow as tf -import numpy as np -import scipy.misc -import logging - -try: - from StringIO import StringIO # Python 2.7 -except ImportError: - from io import BytesIO # Python 3.x +####################################################################### +# Copyright (C) 2017 Shangtong Zhang(zhangshangtong.cpp@gmail.com) # +# Permission given to modify the code as long as you keep this # +# declaration at the top # +####################################################################### +from tensorboardX import SummaryWriter class Logger(object): def __init__(self, log_dir, vanilla_logger, skip=False): - """Create a summary writer logging to log_dir.""" - self.writer = tf.summary.FileWriter(log_dir) + self.writer = SummaryWriter(log_dir) self.info = vanilla_logger.info self.debug = vanilla_logger.debug self.warning = vanilla_logger.warning self.skip = skip - logging.info('') def scalar_summary(self, tag, value, step): if self.skip: return - """Log a scalar variable.""" - summary = tf.Summary(value=[tf.Summary.Value(tag=tag, simple_value=value)]) - self.writer.add_summary(summary, step) + self.writer.add_scalar(tag, value, step) - def image_summary(self, tag, images, step): + def histo_summary(self, tag, values, step): if self.skip: return - """Log a list of images.""" - - img_summaries = [] - for i, img in enumerate(images): - # Write the image to a string - try: - s = StringIO() - except: - s = BytesIO() - scipy.misc.toimage(img).save(s, format="png") - - # Create an Image object - img_sum = tf.Summary.Image(encoded_image_string=s.getvalue(), - height=img.shape[0], - width=img.shape[1]) - # Create a Summary value - img_summaries.append(tf.Summary.Value(tag='%s/%d' % (tag, i), image=img_sum)) - - # Create and write Summary - summary = tf.Summary(value=img_summaries) - self.writer.add_summary(summary, step) - - def histo_summary(self, tag, values, step, bins=1000): - if self.skip: - return - """Log a histogram of the tensor of values.""" - - # Create a histogram using numpy - counts, bin_edges = np.histogram(values, bins=bins) - - # Fill the fields of the histogram proto - hist = tf.HistogramProto() - hist.min = float(np.min(values)) - hist.max = float(np.max(values)) - hist.num = int(np.prod(values.shape)) - hist.sum = float(np.sum(values)) - hist.sum_squares = float(np.sum(values ** 2)) - - # Drop the start of the first bin - bin_edges = bin_edges[1:] - - # Add bin edges and counts - for edge in bin_edges: - hist.bucket_limit.append(edge) - for c in counts: - hist.bucket.append(c) - - # Create and write Summary - summary = tf.Summary(value=[tf.Summary.Value(tag=tag, histo=hist)]) - self.writer.add_summary(summary, step) - self.writer.flush() \ No newline at end of file + self.writer.add_histogram(tag, values, step, bins=1000) \ No newline at end of file diff --git a/utils/vanilla_logger.py b/utils/vanilla_logger.py deleted file mode 100644 index ba50148..0000000 --- a/utils/vanilla_logger.py +++ /dev/null @@ -1,23 +0,0 @@ -import numpy as np -import logging - -class Logger(object): - def __init__(self, log_dir, vanilla_logger, skip=False): - """Create a summary writer logging to log_dir.""" - self.info = vanilla_logger.info - self.debug = vanilla_logger.debug - self.warning = vanilla_logger.warning - self.skip = skip - logging.info('') - - def scalar_summary(self, tag, value, step): - if self.skip: - return - - def image_summary(self, tag, images, step): - if self.skip: - return - - def histo_summary(self, tag, values, step, bins=1000): - if self.skip: - return