diff --git a/DQN/DQN_CartPole-v0.py b/DQN/DQN_CartPole-v0.py new file mode 100644 index 0000000..a06e23d --- /dev/null +++ b/DQN/DQN_CartPole-v0.py @@ -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()