mirror of
https://github.com/wassname/Deep-reinforcement-learning-with-pytorch.git
synced 2026-09-09 11:13:45 +08:00
Create DQN_CartPole-v0.py
This commit is contained in:
@@ -0,0 +1,120 @@
|
||||
import argparse
|
||||
import pickle
|
||||
from collections import namedtuple
|
||||
from itertools import count
|
||||
|
||||
import os, time
|
||||
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, Categorical
|
||||
from torch.utils.data.sampler import BatchSampler, SubsetRandomSampler
|
||||
from tensorboardX import SummaryWriter
|
||||
|
||||
# Hyper-parameters
|
||||
seed = 1
|
||||
render = False
|
||||
num_episodes = 2000
|
||||
env = gym.make('CartPole-v0').unwrapped
|
||||
num_state = env.observation_space.shape[0]
|
||||
num_action = env.action_space.n
|
||||
torch.manual_seed(seed)
|
||||
env.seed(seed)
|
||||
|
||||
Transition = namedtuple('Transition', ['state', 'action', 'reward', 'next_state'])
|
||||
|
||||
class Net(nn.Module):
|
||||
def __init__(self):
|
||||
super(Net, self).__init__()
|
||||
self.fc1 = nn.Linear(num_state, 100)
|
||||
self.fc2 = nn.Linear(100, num_action)
|
||||
|
||||
def forward(self, x):
|
||||
x = F.relu(self.fc1(x))
|
||||
action_prob = self.fc2(x)
|
||||
return action_prob
|
||||
|
||||
class DQN():
|
||||
|
||||
capacity = 8000
|
||||
learning_rate = 1e-3
|
||||
memory_count = 0
|
||||
batch_size = 256
|
||||
gamma = 0.995
|
||||
update_count = 0
|
||||
|
||||
def __init__(self):
|
||||
super(DQN, self).__init__()
|
||||
self.target_net, self.act_net = Net(), Net()
|
||||
self.memory = [None]*self.capacity
|
||||
self.optimizer = optim.Adam(self.act_net.parameters(), self.learning_rate)
|
||||
self.loss_func = nn.MSELoss()
|
||||
self.writer = SummaryWriter('./DQN/logs')
|
||||
|
||||
|
||||
def select_action(self,state):
|
||||
state = torch.tensor(state, dtype=torch.float).unsqueeze(0)
|
||||
value = self.act_net(state)
|
||||
action_max_value, index = torch.max(value, 1)
|
||||
action = index.item()
|
||||
if np.random.rand(1) >= 0.9: # epslion greedy
|
||||
action = np.random.choice(range(num_action), 1).item()
|
||||
return action
|
||||
|
||||
def store_transition(self,transition):
|
||||
index = self.memory_count % self.capacity
|
||||
self.memory[index] = transition
|
||||
self.memory_count += 1
|
||||
return self.memory_count >= self.capacity
|
||||
|
||||
def update(self):
|
||||
if self.memory_count >= self.capacity:
|
||||
state = torch.tensor([t.state for t in self.memory]).float()
|
||||
action = torch.LongTensor([t.action for t in self.memory]).view(-1,1).long()
|
||||
reward = torch.tensor([t.reward for t in self.memory]).float()
|
||||
next_state = torch.tensor([t.next_state for t in self.memory]).float()
|
||||
|
||||
reward = (reward - reward.mean()) / (reward.std() + 1e-7)
|
||||
with torch.no_grad():
|
||||
target_v = reward + self.gamma * self.target_net(next_state).max(1)[0]
|
||||
|
||||
#Update...
|
||||
for index in BatchSampler(SubsetRandomSampler(range(len(self.memory))), batch_size=self.batch_size, drop_last=False):
|
||||
v = (self.act_net(state).gather(1, action))[index]
|
||||
loss = self.loss_func(target_v[index].unsqueeze(1), (self.act_net(state).gather(1, action))[index])
|
||||
self.optimizer.zero_grad()
|
||||
loss.backward()
|
||||
self.optimizer.step()
|
||||
self.writer.add_scalar('loss/value_loss', loss, self.update_count)
|
||||
self.update_count +=1
|
||||
if self.update_count % 100 ==0:
|
||||
self.target_net.load_state_dict(self.act_net.state_dict())
|
||||
else:
|
||||
print("Memory Buff is too less")
|
||||
def main():
|
||||
|
||||
agent = DQN()
|
||||
for i_ep in range(num_episodes):
|
||||
state = env.reset()
|
||||
if render: env.render()
|
||||
for t in range(10000):
|
||||
action = agent.select_action(state)
|
||||
next_state, reward, done, info = env.step(action)
|
||||
if render: env.render()
|
||||
transition = Transition(state, action, reward, next_state)
|
||||
agent.store_transition(transition)
|
||||
state = next_state
|
||||
if done or t >=9999:
|
||||
agent.writer.add_scalar('live/finish_step', t+1, global_step=i_ep)
|
||||
agent.update()
|
||||
if i_ep % 10 == 0:
|
||||
print("episodes {}, step is {} ".format(i_ep, t))
|
||||
break
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
Reference in New Issue
Block a user