mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-10 11:40:58 +08:00
Rewrite PPO
This commit is contained in:
@@ -46,7 +46,6 @@ class A2CAgent:
|
||||
rewards = config.reward_shift_fn(rewards)
|
||||
for i, terminal in enumerate(terminals):
|
||||
if terminals[i]:
|
||||
next_states[i] = self.task.reset(i)
|
||||
self.last_episode_rewards[i] = self.episode_rewards[i]
|
||||
self.episode_rewards[i] = 0
|
||||
|
||||
|
||||
@@ -48,7 +48,6 @@ class NStepDQNAgent:
|
||||
rewards = config.reward_shift_fn(rewards)
|
||||
for i, terminal in enumerate(terminals):
|
||||
if terminals[i]:
|
||||
next_states[i] = self.task.reset(i)
|
||||
self.last_episode_rewards[i] = self.episode_rewards[i]
|
||||
self.episode_rewards[i] = 0
|
||||
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
#######################################################################
|
||||
# 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
|
||||
|
||||
class PPOAgent:
|
||||
def __init__(self, config):
|
||||
self.config = config
|
||||
self.task = config.task_fn()
|
||||
self.actor = config.actor_network_fn(self.task.state_dim, self.task.action_dim)
|
||||
self.critic = config.critic_network_fn(self.task.state_dim, self.task.action_dim)
|
||||
self.actor_opt = config.actor_optimizer_fn(self.actor.parameters())
|
||||
self.critic_opt = config.critic_optimizer_fn(self.critic.parameters())
|
||||
self.total_steps = 0
|
||||
self.episode_rewards = np.zeros(config.num_workers)
|
||||
self.last_episode_rewards = np.zeros(config.num_workers)
|
||||
self.state_normalizer = Normalizer(self.task.state_dim)
|
||||
self.states = self.task.reset()
|
||||
self.states = self.state_normalizer(self.states)
|
||||
|
||||
def close(self):
|
||||
self.task.close()
|
||||
|
||||
def save(self, file_name):
|
||||
pass
|
||||
# with open(file_name, 'wb') as f:
|
||||
# torch.save(self.network.state_dict(), f)
|
||||
|
||||
def iteration(self):
|
||||
config = self.config
|
||||
rollout = []
|
||||
states = self.states
|
||||
for i in range(config.rollout_length):
|
||||
mean, std, log_std = self.actor.predict(states)
|
||||
values = self.critic.predict(states)
|
||||
dist = torch.distributions.Normal(mean, std)
|
||||
actions = dist.sample()
|
||||
log_probs = dist.log_prob(actions).detach()
|
||||
log_probs = torch.sum(log_probs, dim=1, keepdim=True)
|
||||
next_states, rewards, terminals, _ = self.task.step(actions.data.cpu().numpy())
|
||||
self.episode_rewards += rewards
|
||||
rewards = config.reward_shift_fn(rewards)
|
||||
for i, terminal in enumerate(terminals):
|
||||
if terminals[i]:
|
||||
self.last_episode_rewards[i] = self.episode_rewards[i]
|
||||
self.episode_rewards[i] = 0
|
||||
next_states = self.state_normalizer(next_states)
|
||||
rollout.append([states, values, actions, log_probs, rewards, 1 - terminals])
|
||||
states = next_states
|
||||
|
||||
self.states = states
|
||||
pending_value = self.critic.predict(states)
|
||||
rollout.append([states, pending_value, None, None, None, None])
|
||||
|
||||
processed_rollout = [None] * (len(rollout) - 1)
|
||||
advantages = self.actor.tensor(np.zeros((config.num_workers, 1)))
|
||||
returns = pending_value.data
|
||||
for i in reversed(range(len(rollout) - 1)):
|
||||
states, value, actions, log_probs, rewards, terminals = rollout[i]
|
||||
terminals = self.actor.tensor(terminals).unsqueeze(1)
|
||||
rewards = self.actor.tensor(rewards).unsqueeze(1)
|
||||
actions = self.actor.variable(actions)
|
||||
states = self.actor.variable(states)
|
||||
next_value = rollout[i + 1][1]
|
||||
returns = rewards + config.discount * terminals * returns
|
||||
if not config.use_gae:
|
||||
advantages = returns - value.data
|
||||
else:
|
||||
td_error = rewards + config.discount * terminals * next_value.data - value.data
|
||||
advantages = advantages * config.gae_tau * config.discount * terminals + td_error
|
||||
processed_rollout[i] = [states, actions, log_probs, returns, advantages]
|
||||
|
||||
states, actions, log_probs_old, returns, advantages = map(lambda x: torch.cat(x, dim=0), zip(*processed_rollout))
|
||||
advantages = (advantages - advantages.mean()) / advantages.std()
|
||||
advantages = Variable(advantages)
|
||||
|
||||
for k in range(config.optimization_epochs):
|
||||
mean, std, log_std = self.actor.predict(states)
|
||||
dist = torch.distributions.Normal(mean, std)
|
||||
log_probs = dist.log_prob(actions)
|
||||
log_probs = torch.sum(log_probs, dim=1, keepdim=True)
|
||||
ratio = (log_probs - log_probs_old).exp()
|
||||
obj = ratio * advantages
|
||||
obj_clipped = ratio.clamp(1.0 - self.config.ppo_ratio_clip, 1.0 + self.config.ppo_ratio_clip) * advantages
|
||||
policy_loss = -torch.min(obj, obj_clipped).mean(0)
|
||||
|
||||
v = self.critic.predict(states)
|
||||
value_loss = 0.5 * (Variable(returns) - v).pow(2).mean()
|
||||
|
||||
self.actor_opt.zero_grad()
|
||||
self.critic_opt.zero_grad()
|
||||
policy_loss.backward()
|
||||
value_loss.backward()
|
||||
nn.utils.clip_grad_norm(self.actor.parameters(), config.gradient_clip)
|
||||
nn.utils.clip_grad_norm(self.critic.parameters(), config.gradient_clip)
|
||||
self.actor_opt.step()
|
||||
self.critic_opt.step()
|
||||
|
||||
steps = config.rollout_length * config.num_workers
|
||||
self.total_steps += steps
|
||||
@@ -5,3 +5,4 @@ from .A2C_agent import *
|
||||
from .CategoricalDQN_agent import *
|
||||
from .NStepDQN_agent import *
|
||||
from .QuantileRegressionDQN_agent import *
|
||||
from .PPO_agent import *
|
||||
|
||||
+4
-1
@@ -151,7 +151,10 @@ def sub_task(parent_pipe, pipe, task_fn, rank, log_dir):
|
||||
while True:
|
||||
op, data = pipe.recv()
|
||||
if op == 'step':
|
||||
pipe.send(task.step(data))
|
||||
ob, reward, done, info = task.step(data)
|
||||
if done:
|
||||
ob = task.reset()
|
||||
pipe.send([ob, reward, done, info])
|
||||
elif op == 'reset':
|
||||
pipe.send(task.reset())
|
||||
elif op == 'exit':
|
||||
|
||||
@@ -203,105 +203,32 @@ def dqn_ram_atari(name):
|
||||
# config.double_q = False
|
||||
run_episodes(DQNAgent(config))
|
||||
|
||||
# def a3c_continuous():
|
||||
# config = Config()
|
||||
# config.task_fn = lambda: Pendulum()
|
||||
# # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2')
|
||||
# # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2')
|
||||
# # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2')
|
||||
# task = config.task_fn()
|
||||
# config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.0001)
|
||||
# config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
# config.network_fn = lambda: DisjointActorCriticNet(
|
||||
# # lambda: GaussianActorNet(task.state_dim, task.action_dim, unit_std=False, action_gate=F.tanh, action_scale=2.0),
|
||||
# lambda: GaussianActorNet(task.state_dim, task.action_dim, unit_std=True),
|
||||
# lambda: GaussianCriticNet(task.state_dim))
|
||||
# config.policy_fn = lambda: GaussianPolicy()
|
||||
# config.worker = ContinuousAdvantageActorCritic
|
||||
# config.discount = 0.99
|
||||
# config.num_workers = 8
|
||||
# config.update_interval = 20
|
||||
# config.test_interval = 1
|
||||
# config.test_repetitions = 1
|
||||
# config.entropy_weight = 0
|
||||
# config.gradient_clip = 40
|
||||
# config.logger = Logger('./log', logger)
|
||||
# agent = AsyncAgent(config)
|
||||
# agent.run()
|
||||
#
|
||||
# def p3o_continuous():
|
||||
# config = Config()
|
||||
# config.task_fn = lambda: Pendulum()
|
||||
# # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2')
|
||||
# # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2')
|
||||
# # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2')
|
||||
# # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1')
|
||||
# # config.task_fn = lambda: Roboschool('RoboschoolAnt-v1')
|
||||
# task = config.task_fn()
|
||||
# config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim,
|
||||
# gpu=-1, unit_std=True)
|
||||
# config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim, gpu=-1)
|
||||
# config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn)
|
||||
# config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
# config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
#
|
||||
# config.policy_fn = lambda: GaussianPolicy()
|
||||
# config.replay_fn = lambda: GeneralReplay(memory_size=2048, batch_size=2048)
|
||||
# config.worker = ProximalPolicyOptimization
|
||||
# config.discount = 0.99
|
||||
# config.gae_tau = 0.97
|
||||
# config.num_workers = 6
|
||||
# config.test_interval = 1
|
||||
# config.test_repetitions = 1
|
||||
# config.entropy_weight = 0
|
||||
# config.gradient_clip = 20
|
||||
# config.rollout_length = 10000
|
||||
# config.optimize_epochs = 1
|
||||
# config.ppo_ratio_clip = 0.2
|
||||
# config.logger = Logger('./log', logger)
|
||||
# agent = AsyncAgent(config)
|
||||
# agent.run()
|
||||
#
|
||||
# def d3pg_continuous():
|
||||
# config = Config()
|
||||
# config.task_fn = lambda: Pendulum()
|
||||
# # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2')
|
||||
# # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2')
|
||||
# # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2')
|
||||
# # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1')
|
||||
# # config.task_fn = lambda: Roboschool('RoboschoolReacher-v1')
|
||||
# task = config.task_fn()
|
||||
# config.actor_network_fn = lambda: DeterministicActorNet(
|
||||
# task.state_dim, task.action_dim, F.tanh, 2, 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-4)
|
||||
# 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.random_process_fn = \
|
||||
# lambda: OrnsteinUhlenbeckProcess(size=task.action_dim, theta=0.15, sigma=0.2,
|
||||
# n_steps_annealing=100000)
|
||||
# config.worker = DeterministicPolicyGradient
|
||||
# config.num_workers = 6
|
||||
# config.min_memory_size = 50
|
||||
# config.target_network_mix = 0.001
|
||||
# config.test_interval = 500
|
||||
# config.test_repetitions = 1
|
||||
# config.gradient_clip = 20
|
||||
# config.logger = Logger('./log', logger)
|
||||
# agent = AsyncAgent(config)
|
||||
# agent.run()
|
||||
## continuous control
|
||||
|
||||
def ppo_continuous():
|
||||
config = Config()
|
||||
config.num_workers = 5
|
||||
task_fn = lambda log_dir: Pendulum(log_dir=log_dir)
|
||||
# task_fn = lambda log_dir: Roboschool('RoboschoolInvertedPendulum-v1', log_dir=log_dir)
|
||||
# task_fn = lambda log_dir: Roboschool('RoboschoolAnt-v1', log_dir=log_dir)
|
||||
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers, log_dir=get_default_log_dir(ppo_continuous.__name__))
|
||||
config.actor_network_fn = lambda state_dim, action_dim: GaussianActorNet(state_dim, action_dim)
|
||||
config.critic_network_fn = lambda state_dim, action_dim: GaussianCriticNet(state_dim)
|
||||
config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
config.discount = 0.99
|
||||
config.use_gae = True
|
||||
config.gae_tau = 0.97
|
||||
config.gradient_clip = 0.5
|
||||
config.rollout_length = 20
|
||||
config.optimize_epochs = 4
|
||||
config.ppo_ratio_clip = 0.2
|
||||
config.logger = Logger('./log', logger)
|
||||
run_iterations(PPOAgent(config))
|
||||
|
||||
def ddpg_continuous():
|
||||
config = Config()
|
||||
config.task_fn = lambda: Pendulum()
|
||||
# config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2')
|
||||
# config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2')
|
||||
# config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2')
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1')
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolReacher-v1')
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolHopper-v1')
|
||||
@@ -346,9 +273,8 @@ if __name__ == '__main__':
|
||||
# n_step_dqn_pixel_atari('BreakoutNoFrameskip-v4')
|
||||
# dqn_ram_atari('Breakout-ramNoFrameskip-v4')
|
||||
|
||||
ddpg_continuous()
|
||||
# ddpg_continuous()
|
||||
ppo_continuous()
|
||||
|
||||
# dqn_pixel_atari('BreakoutNoFrameskip-v4')
|
||||
# dqn_ram_atari('Pong-ramNoFrameskip-v4')
|
||||
# acvp.train('PongNoFrameskip-v4')
|
||||
|
||||
|
||||
@@ -92,7 +92,7 @@ class GaussianActorNet(nn.Module, BasicNet):
|
||||
action_scale=1,
|
||||
action_gate=F.tanh,
|
||||
gpu=-1,
|
||||
unit_std=True,
|
||||
# unit_std=True,
|
||||
hidden_size=64,
|
||||
non_linear=F.tanh):
|
||||
super(GaussianActorNet, self).__init__()
|
||||
@@ -100,12 +100,8 @@ class GaussianActorNet(nn.Module, BasicNet):
|
||||
self.fc2 = nn.Linear(hidden_size, hidden_size)
|
||||
self.action_mean = nn.Linear(hidden_size, action_dim)
|
||||
|
||||
if unit_std:
|
||||
self.action_log_std = nn.Parameter(torch.zeros(1, action_dim))
|
||||
else:
|
||||
self.action_std = nn.Linear(hidden_size, action_dim)
|
||||
self.action_log_std = nn.Parameter(torch.zeros(1, action_dim))
|
||||
|
||||
self.unit_std = unit_std
|
||||
self.action_scale = action_scale
|
||||
self.action_gate = action_gate
|
||||
self.non_linear = non_linear
|
||||
@@ -119,24 +115,20 @@ class GaussianActorNet(nn.Module, BasicNet):
|
||||
mean = self.action_mean(phi)
|
||||
if self.action_gate is not None:
|
||||
mean = self.action_scale * self.action_gate(mean)
|
||||
if self.unit_std:
|
||||
log_std = self.action_log_std.expand_as(mean)
|
||||
std = log_std.exp()
|
||||
else:
|
||||
std = F.softplus(self.action_std(phi)) + 1e-5
|
||||
log_std = std.log()
|
||||
log_std = self.action_log_std.expand_as(mean)
|
||||
std = log_std.exp()
|
||||
return mean, std, log_std
|
||||
|
||||
def predict(self, x):
|
||||
return self.forward(x)
|
||||
|
||||
def log_density(self, x, mean, log_std, std):
|
||||
var = std.pow(2)
|
||||
log_density = -(x - mean).pow(2) / (2 * var + 1e-5) - 0.5 * torch.log(2 * Variable(torch.FloatTensor([np.pi])).expand_as(x)) - log_std
|
||||
return log_density.sum(1)
|
||||
|
||||
def entropy(self, std):
|
||||
return 0.5 * (1 + (2 * std.pow(2) * np.pi + 1e-5).log()).sum(1).mean()
|
||||
# def log_density(self, x, mean, log_std, std):
|
||||
# var = std.pow(2)
|
||||
# log_density = -(x - mean).pow(2) / (2 * var + 1e-5) - 0.5 * torch.log(2 * Variable(torch.FloatTensor([np.pi])).expand_as(x)) - log_std
|
||||
# return log_density.sum(1)
|
||||
#
|
||||
# def entropy(self, std):
|
||||
# return 0.5 * (1 + (2 * std.pow(2) * np.pi + 1e-5).log()).sum(1).mean()
|
||||
|
||||
class GaussianCriticNet(nn.Module, BasicNet):
|
||||
def __init__(self,
|
||||
|
||||
@@ -57,3 +57,4 @@ class Config:
|
||||
self.categorical_n_atoms = 51
|
||||
self.num_quantiles = 10
|
||||
self.gaussian_noise_scale = 0.3
|
||||
self.optimization_epochs = 4
|
||||
|
||||
@@ -13,6 +13,17 @@ class Normalizer:
|
||||
self.n = 1.0
|
||||
|
||||
def __call__(self, x):
|
||||
if np.isscalar(x) or len(x.shape) == 1:
|
||||
return self.nomalize_single(x)
|
||||
elif len(x.shape) == 2:
|
||||
new_x = np.zeros(x.shape)
|
||||
for i in range(x.shape[0]):
|
||||
new_x[i] = self.nomalize_single(x[i])
|
||||
return new_x
|
||||
else:
|
||||
assert 'Unsupported Shape'
|
||||
|
||||
def nomalize_single(self, x):
|
||||
is_scalar = np.isscalar(x)
|
||||
if is_scalar:
|
||||
x = np.asarray([x])
|
||||
|
||||
Reference in New Issue
Block a user