| 75 | 105x80 grayscale frames (4 channels).""" |
| 76 | |
| 77 | def __init__(self, n_actions, gru_dim=GRU_DIM): |
| 78 | super().__init__() |
| 79 | self.conv = nn.Sequential( |
| 80 | _ortho(nn.Conv2d(4, 32, 8, stride=4), 2 ** 0.5), nn.ReLU(), |
| 81 | _ortho(nn.Conv2d(32, 64, 4, stride=2), 2 ** 0.5), nn.ReLU(), |
| 82 | _ortho(nn.Conv2d(64, 64, 3, stride=1), 2 ** 0.5), nn.ReLU(), |
| 83 | nn.Flatten(), |
| 84 | ) |
| 85 | with torch.no_grad(): |
| 86 | n_flat = self.conv(torch.zeros(1, 4, 105, 80)).shape[1] |
| 87 | self.fc = _ortho(nn.Linear(n_flat, gru_dim), 2 ** 0.5) |
| 88 | self.ln = nn.LayerNorm(gru_dim) |
| 89 | self.gru = nn.GRUCell(gru_dim, gru_dim) |
| 90 | self.pi = _ortho(nn.Linear(gru_dim, n_actions), 0.01) |
| 91 | self.v = _ortho(nn.Linear(gru_dim, 1), 1.0) |
| 92 | self.gru_dim = gru_dim |
| 93 | |
| 94 | def features(self, obs): |
| 95 | h = self.ln(torch.relu(self.fc(self.conv(obs / 255.0)))) |