mirror of
https://github.com/wassname/Deep-reinforcement-learning-with-pytorch.git
synced 2026-09-09 11:13:45 +08:00
104 lines
3.9 KiB
Python
104 lines
3.9 KiB
Python
import argparse
|
|
import pickle
|
|
from collections import namedtuple
|
|
|
|
import os
|
|
import numpy as np
|
|
import matplotlib.pyplot as plt
|
|
|
|
import gym
|
|
import torch
|
|
import torch.nn as nn
|
|
import torch.nn.functional as F
|
|
import torch.optim as optim
|
|
from torch.distributions import Normal
|
|
from torch.utils.data.sampler import BatchSampler, SubsetRandomSampler
|
|
from resnet import ResNet
|
|
|
|
class PPO():
|
|
clip_param = 0.2
|
|
max_grad_norm = 0.5
|
|
ppo_epoch = 10
|
|
buffer_capacity = 1000
|
|
batch_size = 8
|
|
|
|
def __init__(self):
|
|
super(PPO, self).__init__()
|
|
self.resnet = ResNet()
|
|
self.buffer = []
|
|
self.counter = 0
|
|
self.training_step = 0
|
|
|
|
self.actor_optimizer = optim.Adam(self.actor_net.parameters(), 1e-3)
|
|
self.critic_net_optimizer = optim.Adam(self.critic_net.parameters(), 4e-3)
|
|
if not os.path.exists('../param'):
|
|
os.makedirs('../param/net_param')
|
|
os.makedirs('../param/img')
|
|
|
|
def select_action(self, state):
|
|
state = torch.from_numpy(state).float().unsqueeze(0)
|
|
with torch.no_grad():
|
|
mu, sigma = self.actor_net(state)
|
|
dist = Normal(mu, sigma)
|
|
action = dist.sample()
|
|
action_log_prob = dist.log_prob(action)
|
|
action = action.clamp(-2, 2)
|
|
return action.item(), action_log_prob.item()
|
|
|
|
def get_value(self, state):
|
|
state = torch.from_numpy(state)
|
|
with torch.no_grad():
|
|
value = self.critic_net(state)
|
|
return value.item()
|
|
|
|
def save_param(self):
|
|
torch.save(self.actor_net.state_dict(), '../param/net_param/actor_net' + str(time.time())[:10], +'.pkl')
|
|
torch.save(self.critic_net.state_dict(), '../param/net_param/critic_net' + str(time.time())[:10], +'.pkl')
|
|
|
|
def store_transition(self, transition):
|
|
self.buffer.append(transition)
|
|
self.counter += 1
|
|
return counter % self.buffer_capacity == 0
|
|
|
|
def update(self):
|
|
self.training_step += 1
|
|
|
|
state = torch.tensor([t.state for t in self.buffer], dtype=torch.float)
|
|
action = torch.tensor([t.action for t in self.buffer], dtype=torch.float).view(-1, 1)
|
|
reward = torch.tensor([t.reward for t in self.buffer], dtype=torch.float).view(-1, 1)
|
|
next_state = torch.tensor([t.next_state for t in self.buffer], dtype=torch.float)
|
|
old_action_log_prob = torch.tensor([t.a_log_prob for t in self.buffer], dtype=torch.float).view(-1, 1)
|
|
|
|
reward = (reward - reward.mean()) / (reward.std() + 1e-10)
|
|
with torch.no_grad():
|
|
target_v = reward + args.gamma * self.critic_net(next_state)
|
|
|
|
advantage = (target_v - self.critic_net(state)).detach()
|
|
for _ in range(self.ppo_epoch): # iteration ppo_epoch
|
|
for index in BatchSampler(SubsetRandomSampler(range(self.buffer_capacity), self.batch_size, True)):
|
|
# epoch iteration, PPO core!!!
|
|
mu, sigma = self.actor_net(state[index])
|
|
n = Normal(mu, sigma)
|
|
action_log_prob = n.log_prob(action[index])
|
|
ratio = torch.exp(action_log_prob - old_action_log_prob)
|
|
|
|
L1 = ratio * advantage[index]
|
|
L2 = torch.clamp(ratio, 1 - self.clip_param, 1 + self.clip_param) * advantage[index]
|
|
action_loss = -torch.min(L1, L2).mean() # MAX->MIN desent
|
|
self.actor_optimizer.zero_grad()
|
|
action_loss.backward()
|
|
nn.utils.clip_grad_norm_(self.actor_net.parameters(), self.max_grad_norm)
|
|
self.actor_optimizer.step()
|
|
|
|
value_loss = F.smooth_l1_loss(self.critic_net(state[index]), target_v[index])
|
|
self.critic_net_optimizer.zero_grad()
|
|
value_loss.backward()
|
|
nn.utils.clip_grad_norm_(self.critic_net.parameters(), self.max_grad_norm)
|
|
self.critic_net_optimizer.step()
|
|
|
|
del self.buffer[:]
|
|
|
|
|
|
if __name__ == '__main__':
|
|
pass
|