-
Notifications
You must be signed in to change notification settings - Fork 4
Expand file tree
/
Copy pathmodel.py
More file actions
102 lines (84 loc) · 3.91 KB
/
Copy pathmodel.py
File metadata and controls
102 lines (84 loc) · 3.91 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
import torch
from torch import nn
from torch.nn import functional as F
from torch.distributions import Normal
def init_weight(layer, initializer="he normal"):
if initializer == "xavier uniform":
nn.init.xavier_uniform_(layer.weight)
elif initializer == "he normal":
nn.init.kaiming_normal_(layer.weight)
class ValueNetwork(nn.Module):
def __init__(self, n_states, n_hidden_filters=256):
super(ValueNetwork, self).__init__()
self.n_states = n_states
self.n_hidden_filters = n_hidden_filters
self.hidden1 = nn.Linear(in_features=self.n_states, out_features=self.n_hidden_filters)
init_weight(self.hidden1)
self.hidden1.bias.data.zero_()
self.hidden2 = nn.Linear(in_features=self.n_hidden_filters, out_features=self.n_hidden_filters)
init_weight(self.hidden2)
self.hidden2.bias.data.zero_()
self.value = nn.Linear(in_features=self.n_hidden_filters, out_features=1)
init_weight(self.value, initializer="xavier uniform")
self.value.bias.data.zero_()
def forward(self, states):
x = F.relu(self.hidden1(states))
x = F.relu(self.hidden2(x))
return self.value(x)
class QvalueNetwork(nn.Module):
def __init__(self, n_states, n_actions, n_hidden_filters=256):
super(QvalueNetwork, self).__init__()
self.n_states = n_states
self.n_hidden_filters = n_hidden_filters
self.n_actions = n_actions
self.hidden1 = nn.Linear(in_features=self.n_states + self.n_actions, out_features=self.n_hidden_filters)
init_weight(self.hidden1)
self.hidden1.bias.data.zero_()
self.hidden2 = nn.Linear(in_features=self.n_hidden_filters, out_features=self.n_hidden_filters)
init_weight(self.hidden2)
self.hidden2.bias.data.zero_()
self.q_value = nn.Linear(in_features=self.n_hidden_filters, out_features=1)
init_weight(self.q_value, initializer="xavier uniform")
self.q_value.bias.data.zero_()
def forward(self, states, actions):
x = torch.cat([states, actions], dim=1)
x = F.relu(self.hidden1(x))
x = F.relu(self.hidden2(x))
return self.q_value(x)
class PolicyNetwork(nn.Module):
def __init__(self, n_states, n_actions, action_bounds, n_hidden_filters=256):
super(PolicyNetwork, self).__init__()
self.n_states = n_states
self.n_hidden_filters = n_hidden_filters
self.n_actions = n_actions
self.action_bounds = action_bounds
self.hidden1 = nn.Linear(in_features=self.n_states, out_features=self.n_hidden_filters)
init_weight(self.hidden1)
self.hidden1.bias.data.zero_()
self.hidden2 = nn.Linear(in_features=self.n_hidden_filters, out_features=self.n_hidden_filters)
init_weight(self.hidden2)
self.hidden2.bias.data.zero_()
self.mu = nn.Linear(in_features=self.n_hidden_filters, out_features=self.n_actions)
init_weight(self.mu, initializer="xavier uniform")
self.mu.bias.data.zero_()
self.log_std = nn.Linear(in_features=self.n_hidden_filters, out_features=self.n_actions)
init_weight(self.log_std, initializer="xavier uniform")
self.log_std.bias.data.zero_()
def forward(self, states):
x = F.relu(self.hidden1(states))
x = F.relu(self.hidden2(x))
mu = self.mu(x)
log_std = self.log_std(x)
std = log_std.clamp(min=-20, max=2).exp()
dist = Normal(mu, std)
return dist
def sample_or_likelihood(self, states):
dist = self(states)
# Reparameterization trick
u = dist.rsample()
action = torch.tanh(u)
log_prob = dist.log_prob(value=u)
# Enforcing action bounds
log_prob -= torch.log(1 - action ** 2 + 1e-6)
log_prob = log_prob.sum(-1, keepdim=True)
return (action * self.action_bounds[1]).clamp_(self.action_bounds[0], self.action_bounds[1]), log_prob