mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Merge branch 'master' of https://github.com/ShangtongZhang/DeepRL
This commit is contained in:
@@ -9,6 +9,7 @@ Implemented algorithms:
|
||||
* Synchronous N-Step Q-Learning
|
||||
* Deep Deterministic Policy Gradient (DDPG)
|
||||
* (Continuous/Discrete) Synchronous Proximal Policy Optimization (PPO)
|
||||
* The Option-Critic Architecture (OC)
|
||||
* Action Conditional Video Prediction
|
||||
|
||||
Asynchronous algorithms below are removed in current version but can be found in [v0.1](https://github.com/ShangtongZhang/DeepRL/releases/tag/v0.1).
|
||||
@@ -22,6 +23,29 @@ Asynchronous algorithms below are removed in current version but can be found in
|
||||
|
||||
Support for PyTorch v0.3.x can be found in [v0.2](https://github.com/ShangtongZhang/DeepRL/releases/tag/v0.2). Note all the figures are generated via this version. After the upgrade to PyTorch v0.4.0, I have only tested the classical control tasks.
|
||||
|
||||
# Dependency
|
||||
* MacOS 10.12 or Ubuntu 16.04
|
||||
* PyTorch v0.4.0
|
||||
* Python 3.6, 3.5 or 2.7 (deprecated)
|
||||
* Core dependencies: `pip install -e .`
|
||||
* Optional: [Roboschool](https://github.com/openai/roboschool), [DeepMind Control Suite](https://github.com/deepmind/dm_control)+[DMControl2Gym](dm_control2gym)
|
||||
|
||||
# Usage
|
||||
|
||||
```examples.py``` contains examples for all the implemented algorithms
|
||||
|
||||
Please use this bibtex if you want to cite this repo
|
||||
```
|
||||
@misc{deeprl,
|
||||
author = {Shangtong, Zhang},
|
||||
title = {Modularized Implementation of Deep RL Algorithms in PyTorch},
|
||||
year = {2018},
|
||||
publisher = {GitHub},
|
||||
journal = {GitHub Repository},
|
||||
howpublished = {\url{https://github.com/ShangtongZhang/DeepRL}},
|
||||
}
|
||||
```
|
||||
|
||||
# Curves
|
||||
> Curves for CartPole are trivial so I didn't place it here, and there isn't any fixed random seed. The curves are generated in the same manner as OpenAI baselines (one run and smoothed by recent 100 episodes)
|
||||
## DQN
|
||||
@@ -46,6 +70,11 @@ Support for PyTorch v0.3.x can be found in [v0.2](https://github.com/ShangtongZh
|
||||

|
||||

|
||||
|
||||
## OC
|
||||

|
||||
|
||||
This is my synchronous option-critic implementation, not the original one.
|
||||
|
||||
## Action Conditional Video Prediction
|
||||

|
||||
|
||||
@@ -53,17 +82,6 @@ Support for PyTorch v0.3.x can be found in [v0.2](https://github.com/ShangtongZh
|
||||
|
||||
Prediction is sampled after 110K iterations, and I only implemented one-step training
|
||||
|
||||
# Dependency
|
||||
* MacOS 10.12 or Ubuntu 16.04
|
||||
* PyTorch v0.4.0
|
||||
* Python 3.6, 3.5 or 2.7 (deprecated)
|
||||
* Core dependencies: `pip install -e .`
|
||||
* Optional: [Roboschool](https://github.com/openai/roboschool), [DeepMind Control Suite](https://github.com/deepmind/dm_control)+[DMControl2Gym](dm_control2gym)
|
||||
|
||||
# Usage
|
||||
|
||||
```examples.py``` contains examples for all the implemented algorithms
|
||||
|
||||
# References
|
||||
* [Human Level Control through Deep Reinforcement Learning](https://www.nature.com/nature/journal/v518/n7540/full/nature14236.html)
|
||||
* [Asynchronous Methods for Deep Reinforcement Learning](https://arxiv.org/abs/1602.01783)
|
||||
@@ -81,4 +99,5 @@ Prediction is sampled after 110K iterations, and I only implemented one-step tra
|
||||
* [Action-Conditional Video Prediction using Deep Networks in Atari Games](https://arxiv.org/abs/1507.08750)
|
||||
* [A Distributional Perspective on Reinforcement Learning](https://arxiv.org/abs/1707.06887)
|
||||
* [Distributional Reinforcement Learning with Quantile Regression](https://arxiv.org/abs/1710.10044)
|
||||
* [The Option-Critic Architecture](https://arxiv.org/abs/1609.05140)
|
||||
* Some hyper-parameters are from [DeepMind Control Suite](https://arxiv.org/abs/1801.00690), [OpenAI Baselines](https://github.com/openai/baselines) and [Ilya Kostrikov](https://github.com/ikostrikov/pytorch-a2c-ppo-acktr)
|
||||
@@ -0,0 +1,115 @@
|
||||
#######################################################################
|
||||
# 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 ..network import *
|
||||
from .BaseAgent import *
|
||||
|
||||
class OptionCriticAgent(BaseAgent):
|
||||
def __init__(self, config):
|
||||
BaseAgent.__init__(self, config)
|
||||
self.config = config
|
||||
self.task = config.task_fn()
|
||||
self.network = config.network_fn(self.task.state_dim, self.task.action_dim)
|
||||
self.target_network = config.network_fn(self.task.state_dim, self.task.action_dim)
|
||||
self.optimizer = config.optimizer_fn(self.network.parameters())
|
||||
self.target_network.load_state_dict(self.network.state_dict())
|
||||
self.policy = config.policy_fn()
|
||||
|
||||
self.episode_rewards = np.zeros(config.num_workers)
|
||||
self.last_episode_rewards = np.zeros(config.num_workers)
|
||||
|
||||
self.total_steps = 0
|
||||
states = self.config.state_normalizer(self.task.reset())
|
||||
self.q_options, self.betas, self.log_pi = self.network.predict(states)
|
||||
self.options = np.asarray([self.policy.sample(q) for q in self.q_options.detach().cpu().numpy()])
|
||||
self.is_initial_betas = np.ones(self.config.num_workers)
|
||||
self.prev_options = np.copy(self.options)
|
||||
|
||||
def iteration(self):
|
||||
config = self.config
|
||||
rollout = []
|
||||
|
||||
q_options, betas, options, log_pi = self.q_options, self.betas, self.options, self.log_pi
|
||||
for _ in range(config.rollout_length):
|
||||
var_options = self.network.tensor(options).long()
|
||||
worker_index = self.network.tensor(np.arange(config.num_workers)).long()
|
||||
intra_log_pi = log_pi[worker_index, var_options, :]
|
||||
dist = torch.distributions.Categorical(intra_log_pi.exp())
|
||||
actions = dist.sample()
|
||||
next_states, rewards, terminals, _ = self.task.step(actions.cpu().detach().numpy().flatten())
|
||||
next_states = config.state_normalizer(next_states)
|
||||
self.episode_rewards += rewards
|
||||
rewards = config.reward_normalizer(rewards)
|
||||
q_options_next, betas_next, log_pi_next = self.network.predict(next_states)
|
||||
rollout.append([q_options, betas, options, self.prev_options, rewards, 1 - terminals, np.copy(self.is_initial_betas), intra_log_pi, actions])
|
||||
self.is_initial_betas = np.asarray(terminals, dtype=np.float32)
|
||||
|
||||
np_q_options_next = q_options_next.cpu().detach().numpy()
|
||||
np_betas_next = betas_next.gather(1, var_options.unsqueeze(1)).cpu().detach().numpy().flatten()
|
||||
options_next = np.copy(options)
|
||||
dice = np.random.rand(len(options_next))
|
||||
for j in range(len(dice)):
|
||||
if dice[j] < np_betas_next[j]:
|
||||
options_next[j] = self.policy.sample(np_q_options_next[j])
|
||||
for i, terminal in enumerate(terminals):
|
||||
if terminals[i]:
|
||||
self.last_episode_rewards[i] = self.episode_rewards[i]
|
||||
self.episode_rewards[i] = 0
|
||||
self.prev_options = options
|
||||
options = options_next
|
||||
q_options = q_options_next
|
||||
betas = betas_next
|
||||
log_pi = log_pi_next
|
||||
|
||||
self.policy.update_epsilon()
|
||||
self.total_steps += config.num_workers
|
||||
if self.total_steps / config.num_workers % config.target_network_update_freq == 0:
|
||||
self.target_network.load_state_dict(self.network.state_dict())
|
||||
|
||||
self.options = options
|
||||
self.q_options = q_options
|
||||
self.betas = betas
|
||||
self.log_pi = log_pi
|
||||
|
||||
target_q_options, _, _ = self.target_network.predict(next_states)
|
||||
prev_options = self.network.tensor(self.prev_options).long().unsqueeze(1)
|
||||
betas_prev_options = betas.gather(1, prev_options)
|
||||
|
||||
returns = (1 - betas_prev_options) * target_q_options.gather(1, prev_options) +\
|
||||
betas_prev_options * torch.max(target_q_options, dim=1, keepdim=True)[0]
|
||||
returns = returns.detach()
|
||||
|
||||
processed_rollout = [None] * (len(rollout))
|
||||
for i in reversed(range(len(rollout))):
|
||||
q_options, betas, options, prev_options, rewards, terminals, is_initial_betas, log_pi, actions = rollout[i]
|
||||
options = self.network.tensor(options).unsqueeze(1).long()
|
||||
prev_options = self.network.tensor(prev_options).unsqueeze(1).long()
|
||||
terminals = self.network.tensor(terminals).unsqueeze(1)
|
||||
rewards = self.network.tensor(rewards).unsqueeze(1)
|
||||
is_initial_betas = self.network.tensor(is_initial_betas).unsqueeze(1)
|
||||
returns = rewards + config.discount * terminals * returns
|
||||
|
||||
q_omg = q_options.gather(1, options)
|
||||
log_action_prob = log_pi.gather(1, actions.unsqueeze(1))
|
||||
entropy_loss = (log_pi.exp() * log_pi).sum(-1).unsqueeze(1)
|
||||
|
||||
q_prev_omg = q_options.gather(1, prev_options)
|
||||
v_prev_omg = torch.max(q_options, dim=1, keepdim=True)[0]
|
||||
advantage_omg = q_prev_omg - v_prev_omg
|
||||
advantage_omg.add_(config.termination_regularizer)
|
||||
betas = betas.gather(1, prev_options)
|
||||
betas = betas * (1 - is_initial_betas)
|
||||
processed_rollout[i] = [q_omg, returns, betas, advantage_omg.detach(), log_action_prob, entropy_loss]
|
||||
|
||||
q_omg, returns, beta_omg, advantage_omg, log_action_prob, entropy_loss = map(lambda x: torch.cat(x, dim=0), zip(*processed_rollout))
|
||||
pi_loss = -log_action_prob * (returns - q_omg.detach()) + config.entropy_weight * entropy_loss
|
||||
pi_loss = pi_loss.mean()
|
||||
q_loss = 0.5 * (q_omg - returns).pow(2).mean()
|
||||
beta_loss = (advantage_omg * beta_omg).mean()
|
||||
self.optimizer.zero_grad()
|
||||
(pi_loss + q_loss + beta_loss).backward()
|
||||
nn.utils.clip_grad_norm_(self.network.parameters(), config.gradient_clip)
|
||||
self.optimizer.step()
|
||||
@@ -5,3 +5,4 @@ from .CategoricalDQN_agent import *
|
||||
from .NStepDQN_agent import *
|
||||
from .QuantileRegressionDQN_agent import *
|
||||
from .PPO_agent import *
|
||||
from .OptionCritic_agent import *
|
||||
@@ -89,6 +89,26 @@ class QuantileNet(nn.Module, BaseNet):
|
||||
quantiles = quantiles.cpu().detach().numpy()
|
||||
return quantiles
|
||||
|
||||
class OptionCriticNet(nn.Module, BaseNet):
|
||||
def __init__(self, body, action_dim, num_options, gpu=-1):
|
||||
super(OptionCriticNet, self).__init__()
|
||||
self.fc_q = layer_init(nn.Linear(body.feature_dim, num_options))
|
||||
self.fc_pi = layer_init(nn.Linear(body.feature_dim, num_options * action_dim))
|
||||
self.fc_beta = layer_init(nn.Linear(body.feature_dim, num_options))
|
||||
self.num_options = num_options
|
||||
self.action_dim = action_dim
|
||||
self.body = body
|
||||
self.set_gpu(gpu)
|
||||
|
||||
def predict(self, x):
|
||||
phi = self.body(self.tensor(x))
|
||||
q = self.fc_q(phi)
|
||||
beta = F.sigmoid(self.fc_beta(phi))
|
||||
pi = self.fc_pi(phi)
|
||||
pi = pi.view(-1, self.num_options, self.action_dim)
|
||||
log_pi = F.log_softmax(pi, dim=-1)
|
||||
return q, beta, log_pi
|
||||
|
||||
class GaussianActorNet(nn.Module, BaseNet):
|
||||
def __init__(self, action_dim, body, gpu=-1):
|
||||
super(GaussianActorNet, self).__init__()
|
||||
|
||||
@@ -7,6 +7,7 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
|
||||
class BaseNet:
|
||||
def set_gpu(self, gpu):
|
||||
|
||||
@@ -60,6 +60,7 @@ class Config:
|
||||
self.test_interval = 0
|
||||
self.test_repetitions = 10
|
||||
self.evaluation_env = None
|
||||
self.termination_regularizer = 0
|
||||
|
||||
def add_argument(self, *args, **kwargs):
|
||||
self.parser.add_argument(*args, **kwargs)
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# Adapted from https://github.com/openai/baselines/blob/master/baselines/results_plotter.py
|
||||
|
||||
from ..component.bench import load_monitor_log
|
||||
import numpy as np
|
||||
from ..component import *
|
||||
import os
|
||||
import re
|
||||
|
||||
@@ -44,7 +44,7 @@ class Plotter:
|
||||
def load_results(self, dirs, max_timesteps=1e8, x_axis=X_TIMESTEPS, episode_window=100):
|
||||
tslist = []
|
||||
for dir in dirs:
|
||||
ts = component.load_monitor_log(dir)
|
||||
ts = load_monitor_log(dir)
|
||||
ts = ts[ts.l.cumsum() <= max_timesteps]
|
||||
tslist.append(ts)
|
||||
xy_list = [self.ts2xy(ts, x_axis) for ts in tslist]
|
||||
|
||||
+43
-1
@@ -117,6 +117,24 @@ def ppo_cart_pole():
|
||||
config.iteration_log_interval = 1
|
||||
run_iterations(PPOAgent(config))
|
||||
|
||||
def option_critic_cart_pole():
|
||||
config = Config()
|
||||
game = 'CartPole-v0'
|
||||
task_fn = lambda log_dir: ClassicalControl(game, max_steps=200, log_dir=log_dir)
|
||||
config.num_workers = 5
|
||||
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers)
|
||||
config.optimizer_fn = lambda params: torch.optim.RMSprop(params, 0.001)
|
||||
config.network_fn = lambda state_dim, action_dim: OptionCriticNet(
|
||||
FCBody(state_dim), action_dim, num_options=2)
|
||||
config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=10000, min_epsilon=0.1)
|
||||
config.discount = 0.99
|
||||
config.target_network_update_freq = 200
|
||||
config.rollout_length = 5
|
||||
config.termination_regularizer = 0.01
|
||||
config.entropy_weight = 0.01
|
||||
config.logger = Logger('./log', logger)
|
||||
run_iterations(OptionCriticAgent(config))
|
||||
|
||||
## Atari games
|
||||
|
||||
def dqn_pixel_atari(name):
|
||||
@@ -247,6 +265,28 @@ def ppo_pixel_atari(name):
|
||||
config.iteration_log_interval = 1
|
||||
run_iterations(PPOAgent(config))
|
||||
|
||||
def option_ciritc_pixel_atari(name):
|
||||
config = Config()
|
||||
config.history_length = 4
|
||||
task_fn = lambda log_dir: PixelAtari(name, frame_skip=4, history_length=config.history_length, log_dir=log_dir)
|
||||
config.num_workers = 16
|
||||
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers,
|
||||
log_dir=get_default_log_dir(config.tag))
|
||||
config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=1e-4, alpha=0.99, eps=1e-5)
|
||||
config.network_fn = lambda state_dim, action_dim: OptionCriticNet(NatureConvBody(), action_dim, num_options=4, gpu=0)
|
||||
config.policy_fn = lambda: GreedyPolicy(epsilon=0.1, final_step=1000000, min_epsilon=0.1)
|
||||
config.state_normalizer = ImageNormalizer()
|
||||
config.reward_normalizer = SignNormalizer()
|
||||
config.discount = 0.99
|
||||
config.target_network_update_freq = 10000
|
||||
config.rollout_length = 5
|
||||
config.gradient_clip = 5
|
||||
config.max_steps = 1e8
|
||||
config.entropy_weight = 0.01
|
||||
config.termination_regularizer = 0.01
|
||||
config.logger = Logger('./log', logger)
|
||||
run_iterations(OptionCriticAgent(config))
|
||||
|
||||
def dqn_ram_atari(name):
|
||||
config = Config()
|
||||
config.task_fn = lambda: RamAtari(name, no_op=30, frame_skip=4,
|
||||
@@ -331,7 +371,7 @@ def ddpg_continuous():
|
||||
def plot():
|
||||
import matplotlib.pyplot as plt
|
||||
plotter = Plotter()
|
||||
names = plotter.load_log_dirs('')
|
||||
names = plotter.load_log_dirs(pattern='.*')
|
||||
data = plotter.load_results(names)
|
||||
|
||||
for i, name in enumerate(names):
|
||||
@@ -371,6 +411,7 @@ if __name__ == '__main__':
|
||||
# quantile_regression_dqn_cart_pole()
|
||||
# n_step_dqn_cart_pole()
|
||||
# ppo_cart_pole()
|
||||
# option_critic_cart_pole()
|
||||
|
||||
# dqn_pixel_atari('BreakoutNoFrameskip-v4')
|
||||
# a2c_pixel_atari('BreakoutNoFrameskip-v4')
|
||||
@@ -378,6 +419,7 @@ if __name__ == '__main__':
|
||||
# quantile_regression_dqn_pixel_atari('BreakoutNoFrameskip-v4')
|
||||
# n_step_dqn_pixel_atari('BreakoutNoFrameskip-v4')
|
||||
# ppo_pixel_atari('BreakoutNoFrameskip-v4')
|
||||
# option_ciritc_pixel_atari('BreakoutNoFrameskip-v4')
|
||||
# dqn_ram_atari('Breakout-ramNoFrameskip-v4')
|
||||
|
||||
# ddpg_continuous()
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 25 KiB |
Reference in New Issue
Block a user