-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathmain.py
More file actions
executable file
·145 lines (110 loc) · 4.7 KB
/
Copy pathmain.py
File metadata and controls
executable file
·145 lines (110 loc) · 4.7 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
# main function that sets up environments
# perform training loop
import torch
import numpy as np
from tensorboardX import SummaryWriter
import os
from src import envs
from src.maddpg import MADDPG
from src.replay_buffer import ReplayBuffer
from src.utils import transpose_to_tensor
def seeding(seed=1):
np.random.seed(seed)
torch.manual_seed(seed)
def pre_process(entity, batchsize):
processed_entity = []
for j in range(3):
list = []
for i in range(batchsize):
b = entity[i][j]
list.append(b)
c = torch.Tensor(list)
processed_entity.append(c)
return processed_entity
def main():
seeding()
# number of parallel agents
parallel_envs = 8
# number of training episodes.
number_of_episodes = 30000
episode_length = 100
batchsize = 2000
# how many episodes to save policy and gif
save_interval = 1000
t = 0
# amplitude of OU noise this slowly decreases to 0
noise = 2
noise_reduction = 0.9999
# how many episodes before update
episode_per_update = 2 * parallel_envs
log_path = os.getcwd()+"/log"
model_dir= os.getcwd()+"/model_dir"
os.makedirs(model_dir, exist_ok=True)
torch.set_num_threads(parallel_envs)
env = envs.make_parallel_env(parallel_envs)
# keep 5000 episodes worth of replay
buffer = ReplayBuffer(int(5000*episode_length))
# initialize policy and critic
maddpg = MADDPG()
logger = SummaryWriter(log_dir=log_path)
agent0_reward = []
agent1_reward = []
agent2_reward = []
# training loop
# show progressbar
import progressbar as pb
widget = ['episode: ', pb.Counter(),'/',str(number_of_episodes),' ',
pb.Percentage(), ' ', pb.ETA(), ' ', pb.Bar(marker=pb.RotatingMarker()), ' ' ]
timer = pb.ProgressBar(widgets=widget, maxval=number_of_episodes).start()
# Modified training loop to remove rendering
for episode in range(0, number_of_episodes, parallel_envs):
timer.update(episode)
reward_this_episode = np.zeros((parallel_envs, 3))
all_obs = env.reset()
obs, obs_full = all_obs
# save info or not
save_info = ((episode) % save_interval < parallel_envs or episode==number_of_episodes-parallel_envs)
for episode_t in range(episode_length):
t += parallel_envs
actions = maddpg.act(transpose_to_tensor(obs), noise=noise)
noise *= noise_reduction
actions_array = torch.stack(actions).detach().numpy()
actions_for_env = np.rollaxis(actions_array,1)
next_obs, next_obs_full, rewards, dones, info = env.step(actions_for_env)
transition = (obs, obs_full, actions_for_env, rewards, next_obs, next_obs_full, dones)
buffer.push(transition)
reward_this_episode += rewards
obs, obs_full = next_obs, next_obs_full
# update once after every episode_per_update
if len(buffer) > batchsize and episode % episode_per_update < parallel_envs:
for a_i in range(3):
samples = buffer.sample(batchsize)
maddpg.update(samples, a_i, logger)
maddpg.update_targets() #soft update the target network towards the actual networks
for i in range(parallel_envs):
agent0_reward.append(reward_this_episode[i,0])
agent1_reward.append(reward_this_episode[i,1])
agent2_reward.append(reward_this_episode[i,2])
if episode % 100 == 0 or episode == number_of_episodes-1:
avg_rewards = [np.mean(agent0_reward), np.mean(agent1_reward), np.mean(agent2_reward)]
agent0_reward = []
agent1_reward = []
agent2_reward = []
for a_i, avg_rew in enumerate(avg_rewards):
logger.add_scalar('agent%i/mean_episode_rewards' % a_i, avg_rew, episode)
#saving model
if save_info:
save_dict_list =[]
for i in range(3):
save_dict = {'actor_params' : maddpg.maddpg_agent[i].actor.state_dict(),
'actor_optim_params': maddpg.maddpg_agent[i].actor_optimizer.state_dict(),
'critic_params' : maddpg.maddpg_agent[i].critic.state_dict(),
'critic_optim_params' : maddpg.maddpg_agent[i].critic_optimizer.state_dict()}
save_dict_list.append(save_dict)
torch.save(save_dict_list,
os.path.join(model_dir, 'episode-{}.pt'.format(episode)))
env.close()
logger.close()
timer.finish()
if __name__=='__main__':
main()