-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
105 lines (82 loc) · 2.65 KB
/
Copy pathmain.py
File metadata and controls
105 lines (82 loc) · 2.65 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
import gym
import torch
import numpy as np
import argparse
from parameters import *
from PPO import Ppo
from collections import deque
parser = argparse.ArgumentParser()
parser.add_argument('--env_name', type=str, default="Ant-v2",
help='name of Mujoco environement')
args = parser.parse_args()
env = gym.make(args.env_name)
N_S = env.observation_space.shape[0]
N_A = env.action_space.shape[0]
#初始化随机种子
env.seed(500)
torch.manual_seed(500)
np.random.seed(500)
##状态的归一化
class Nomalize:
def __init__(self, N_S):
self.mean = np.zeros((N_S,))
self.std = np.zeros((N_S, ))
self.stdd = np.zeros((N_S, ))
self.n = 0
#可以像函数一样调用类
def __call__(self, x):
x = np.asarray(x)
self.n += 1
if self.n == 1:
self.mean = x
else:
#更新样本均值和方差
old_mean = self.mean.copy()
self.mean = old_mean + (x - old_mean) / self.n
self.stdd = self.stdd + (x - old_mean) * (x - self.mean)
#状态归一化
if self.n > 1:
self.std = np.sqrt(self.stdd / (self.n - 1))
else:
self.std = self.mean
x = x - self.mean
x = x / (self.std + 1e-8)
x = np.clip(x, -5, +5)
return x
ppo = Ppo(N_S,N_A)
nomalize = Nomalize(N_S)
episodes = 0
eva_episodes = 0
for iter in range(Iter):
memory = deque()
scores = []
steps = 0
while steps <2048: #Horizen ,超过这个步数,才开始一次训练。否则继续添加MEMORY
episodes += 1
#归一化s
s = nomalize(env.reset())
score = 0
for _ in range(MAX_STEP):
steps += 1
#选择行为
a=ppo.actor_net.choose_action(torch.from_numpy(np.array(s).astype(np.float32)).unsqueeze(0))[0]
s_ , r ,done,info = env.step(a)
#env.render()
# 归一化s
s_ = nomalize(s_)
mask = (1-done)*1
memory.append([s,a,r,mask])
score += r
s = s_
if done:
break
with open('log_' + args.env_name + '.txt', 'a') as outfile:
outfile.write('\t' + str(episodes) + '\t' + str(score) + '\n')
scores.append(score)
score_avg = np.mean(scores)
print('{} episode score is {:.2f}'.format(episodes, score_avg))
if iter % 200 == 0:
torch.save(ppo.actor_net.state_dict(), './model/ppo_actor_{}'.format(iter))
torch.save(ppo.critic_net.state_dict(), './model/ppo_critic_{}'.format(iter))
#每隔一定的timesteps 进行参数更新
ppo.train(memory)