(self, obs,flatten=True)
| 69 | return mu + eps * std |
| 70 | |
| 71 | def forward_conv(self, obs,flatten=True): |
| 72 | if obs.max() > 1.: |
| 73 | obs = obs / 255. |
| 74 | |
| 75 | self.outputs['obs'] = obs |
| 76 | |
| 77 | conv = torch.relu(self.convs[0](obs)) |
| 78 | self.outputs['conv1'] = conv |
| 79 | |
| 80 | for i in range(1, self.num_layers): |
| 81 | conv = torch.relu(self.convs[i](conv)) |
| 82 | self.outputs['conv%s' % (i + 1)] = conv |
| 83 | |
| 84 | if flatten: |
| 85 | return conv.view(conv.size(0), -1) |
| 86 | else: |
| 87 | return conv |
| 88 | |
| 89 | def forward(self, obs, detach=False): |
| 90 | h = self.forward_conv(obs) |