MCPcopy Create free account
hub / github.com/WendyShang/flare / PixelEncoder

Class PixelEncoder

encoder.py:35–131  ·  view source on GitHub ↗

Convolutional encoder of pixels observations.

Source from the content-addressed store, hash-verified

33
34
35class PixelEncoder(nn.Module):
36 """Convolutional encoder of pixels observations."""
37 def __init__(self, obs_shape, feature_dim,
38 channels=[16, 32, 32],
39 num_layers=2,
40 num_filters=32,
41 output_logits=False,
42 image_channel=3):
43 super().__init__()
44
45 assert len(obs_shape) == 3
46 self.obs_shape = obs_shape
47 self.feature_dim = feature_dim
48 self.num_layers = num_layers
49 self.image_channel = image_channel
50 self.convs = nn.ModuleList(
51 [nn.Conv2d(obs_shape[0], num_filters, 3, stride=2)]
52 )
53 for i in range(num_layers - 1):
54 self.convs.append(nn.Conv2d(num_filters, num_filters, 3, stride=1))
55 self.outputs = dict()
56
57 x = torch.randn([32]+list(obs_shape))
58 self.out_dim = self.forward_conv(x,flatten=False).shape[-1]
59 print('conv output dim: ' + str(self.out_dim))
60
61 self.fc = nn.Linear(num_filters * self.out_dim * self.out_dim, self.feature_dim)
62 self.ln = nn.LayerNorm(self.feature_dim)
63
64 self.output_logits = output_logits
65
66 def reparameterize(self, mu, logstd):
67 std = torch.exp(logstd)
68 eps = torch.randn_like(std)
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)
91 if detach:
92 h = h.detach()

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected