mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Update PPO
This commit is contained in:
+32
-19
@@ -78,28 +78,41 @@ class PPOAgent(BaseAgent):
|
||||
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)
|
||||
returns = Variable(returns)
|
||||
|
||||
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)
|
||||
batcher = Batcher(states.size(0) // config.num_mini_batches, [np.arange(states.size(0))])
|
||||
for _ in range(config.optimization_epochs):
|
||||
batcher.shuffle()
|
||||
while not batcher.end():
|
||||
batch_indices = batcher.next_batch()[0]
|
||||
batch_indices = self.actor.variable(batch_indices, torch.LongTensor)
|
||||
sampled_states = states[batch_indices]
|
||||
sampled_actions = actions[batch_indices]
|
||||
sampled_log_probs_old = log_probs_old[batch_indices]
|
||||
sampled_returns = returns[batch_indices]
|
||||
sampled_advantages = advantages[batch_indices]
|
||||
|
||||
v = self.critic.predict(states)
|
||||
value_loss = 0.5 * (Variable(returns) - v).pow(2).mean()
|
||||
mean, std, log_std = self.actor.predict(sampled_states)
|
||||
dist = torch.distributions.Normal(mean, std)
|
||||
log_probs = dist.log_prob(sampled_actions)
|
||||
log_probs = torch.sum(log_probs, dim=1, keepdim=True)
|
||||
ratio = (log_probs - sampled_log_probs_old).exp()
|
||||
obj = ratio * sampled_advantages
|
||||
obj_clipped = ratio.clamp(1.0 - self.config.ppo_ratio_clip,
|
||||
1.0 + self.config.ppo_ratio_clip) * sampled_advantages
|
||||
policy_loss = -torch.min(obj, obj_clipped).mean(0)
|
||||
|
||||
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()
|
||||
v = self.critic.predict(sampled_states)
|
||||
value_loss = 0.5 * (sampled_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
|
||||
|
||||
@@ -186,14 +186,15 @@ def n_step_dqn_pixel_atari(name):
|
||||
config.num_workers = 16
|
||||
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers,
|
||||
log_dir=get_default_log_dir(n_step_dqn_pixel_atari.__name__))
|
||||
config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01)
|
||||
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: ConvNet(config.history_length, action_dim, gpu=3)
|
||||
config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1)
|
||||
config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.05)
|
||||
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.logger = Logger('./log', logger)
|
||||
run_iterations(NStepDQNAgent(config))
|
||||
|
||||
@@ -220,25 +221,29 @@ def dqn_ram_atari(name):
|
||||
|
||||
def ppo_continuous():
|
||||
config = Config()
|
||||
config.num_workers = 16
|
||||
config.num_workers = 1
|
||||
# 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)
|
||||
# task_fn = lambda log_dir: Roboschool('RoboschoolReacher-v1', log_dir=log_dir)
|
||||
task_fn = lambda log_dir: Roboschool('RoboschoolHopper-v1', log_dir=log_dir)
|
||||
# task_fn = lambda log_dir: DMControl('cartpole', 'balance', log_dir=log_dir)
|
||||
# task_fn = lambda log_dir: DMControl('hopper', 'hop', 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, 1e-4)
|
||||
config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 1e-4)
|
||||
config.state_normalizer = RunningStatsNormalizer()
|
||||
config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 3e-4, eps=1e-5)
|
||||
config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 3e-4, eps=1e-5)
|
||||
# config.state_normalizer = RunningStatsNormalizer()
|
||||
config.discount = 0.99
|
||||
config.use_gae = True
|
||||
config.gae_tau = 0.97
|
||||
config.gae_tau = 0.95
|
||||
config.gradient_clip = 0.5
|
||||
config.rollout_length = 20
|
||||
config.optimize_epochs = 5
|
||||
config.rollout_length = 2048
|
||||
config.optimization_epochs = 10
|
||||
config.num_mini_batches = 32
|
||||
config.ppo_ratio_clip = 0.2
|
||||
config.iteration_log_interval = 1
|
||||
config.logger = Logger('./log', logger)
|
||||
run_iterations(PPOAgent(config))
|
||||
|
||||
@@ -247,17 +252,19 @@ def ddpg_continuous():
|
||||
log_dir = get_default_log_dir(ddpg_continuous.__name__)
|
||||
# config.task_fn = lambda: Pendulum(log_dir=log_dir)
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1', log_dir=log_dir)
|
||||
config.task_fn = lambda: Roboschool('RoboschoolReacher-v1', log_dir=log_dir)
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolHopper-v1', log_dir=log_dir)
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolReacher-v1', log_dir=log_dir)
|
||||
config.task_fn = lambda: Roboschool('RoboschoolHopper-v1', log_dir=log_dir)
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolAnt-v1', log_dir=log_dir)
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolWalker2d-v1', log_dir=log_dir)
|
||||
config.task_fn = lambda: DMControl('cartpole', 'balance', log_dir=log_dir)
|
||||
# config.task_fn = lambda: DMControl('cartpole', 'balance', log_dir=log_dir)
|
||||
# config.task_fn = lambda: DMControl('finger', 'spin', log_dir=log_dir)
|
||||
config.actor_network_fn = lambda state_dim, action_dim: DeterministicActorNet(state_dim, action_dim)
|
||||
config.critic_network_fn = lambda state_dim, action_dim: DeterministicCriticNet(state_dim, action_dim)
|
||||
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: HighDimActionReplay(memory_size=1000000, batch_size=64)
|
||||
config.discount = 0.99
|
||||
config.state_normalizer = RunningStatsNormalizer()
|
||||
config.random_process_fn = \
|
||||
lambda action_dim: OrnsteinUhlenbeckProcess(size=action_dim, theta=0.15, sigma=0.3,
|
||||
n_steps_annealing=1000000)
|
||||
@@ -270,12 +277,17 @@ def ddpg_continuous():
|
||||
def plot():
|
||||
import matplotlib.pyplot as plt
|
||||
plotter = Plotter()
|
||||
# name = 'log/ppo_continuous-180408-002056'
|
||||
# plotter.plot_results([name])
|
||||
# plt.show()
|
||||
names = ['a2c_pixel_atari-180407-92711',
|
||||
'categorical_dqn_pixel_atari-180407-094006',
|
||||
'dqn_pixel_atari-180407-01414',
|
||||
'quantile_regression_dqn_pixel_atari-180407-01604',
|
||||
'n_step_dqn_pixel_atari-180407-163421',
|
||||
'ppo_continuous-180407-111715']
|
||||
'n_step_dqn_pixel_atari-180408-001104',
|
||||
'ppo_continuous-180408-002056',
|
||||
'ddpg_continuous-180407-234141'
|
||||
]
|
||||
for name in names:
|
||||
plotter.plot_results(['to_plot/%s' % (name)])
|
||||
plt.savefig('images/%s.png' % (name))
|
||||
|
||||
@@ -87,31 +87,36 @@ class GaussianActorNet(nn.Module, BasicNet):
|
||||
def __init__(self,
|
||||
state_dim,
|
||||
action_dim,
|
||||
action_scale=1,
|
||||
action_gate=F.tanh,
|
||||
gpu=-1,
|
||||
hidden_size=64,
|
||||
non_linear=F.tanh):
|
||||
super(GaussianActorNet, self).__init__()
|
||||
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)
|
||||
self.fc_action = nn.Linear(hidden_size, action_dim)
|
||||
|
||||
self.action_log_std = nn.Parameter(torch.zeros(1, action_dim))
|
||||
|
||||
self.action_scale = action_scale
|
||||
self.action_gate = action_gate
|
||||
self.non_linear = non_linear
|
||||
|
||||
self.init_weights()
|
||||
BasicNet.__init__(self, gpu)
|
||||
|
||||
def init_weights(self):
|
||||
bound = 3e-3
|
||||
nn.init.uniform(self.fc_action.weight.data, -bound, bound)
|
||||
nn.init.constant(self.fc_action.bias.data, 0)
|
||||
|
||||
nn.init.orthogonal(self.fc1.weight.data)
|
||||
nn.init.constant(self.fc1.bias.data, 0)
|
||||
nn.init.orthogonal(self.fc2.weight.data)
|
||||
nn.init.constant(self.fc2.bias.data, 0)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.variable(x)
|
||||
phi = self.non_linear(self.fc1(x))
|
||||
phi = self.non_linear(self.fc2(phi))
|
||||
mean = self.action_mean(phi)
|
||||
if self.action_gate is not None:
|
||||
mean = self.action_scale * self.action_gate(mean)
|
||||
mean = F.tanh(self.fc_action(phi))
|
||||
log_std = self.action_log_std.expand_as(mean)
|
||||
std = log_std.exp()
|
||||
return mean, std, log_std
|
||||
@@ -130,8 +135,19 @@ class GaussianCriticNet(nn.Module, BasicNet):
|
||||
self.fc2 = nn.Linear(hidden_size, hidden_size)
|
||||
self.fc_value = nn.Linear(hidden_size, 1)
|
||||
self.non_linear = non_linear
|
||||
self.init_weights()
|
||||
BasicNet.__init__(self, gpu)
|
||||
|
||||
def init_weights(self):
|
||||
bound = 3e-3
|
||||
nn.init.uniform(self.fc_value.weight.data, -bound, bound)
|
||||
nn.init.constant(self.fc_value.bias.data, 0)
|
||||
|
||||
nn.init.orthogonal(self.fc1.weight.data)
|
||||
nn.init.constant(self.fc1.bias.data, 0)
|
||||
nn.init.orthogonal(self.fc2.weight.data)
|
||||
nn.init.constant(self.fc2.bias.data, 0)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.variable(x)
|
||||
phi = self.non_linear(self.fc1(x))
|
||||
|
||||
@@ -55,3 +55,4 @@ class Config:
|
||||
self.num_quantiles = 10
|
||||
self.gaussian_noise_scale = 0.3
|
||||
self.optimization_epochs = 4
|
||||
self.num_mini_batches = 32
|
||||
|
||||
Reference in New Issue
Block a user