MCPcopy Create free account
hub / github.com/drinkingcoder/NeuralMarker / ConvGRU

Class ConvGRU

core/update.py:26–41  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

24 return self.conv2(self.relu(self.conv1(x)))
25
26class ConvGRU(nn.Module):
27 def __init__(self, hidden_dim=128, input_dim=192+128):
28 super(ConvGRU, self).__init__()
29 self.convz = nn.Conv2d(hidden_dim+input_dim, hidden_dim, 3, padding=1)
30 self.convr = nn.Conv2d(hidden_dim+input_dim, hidden_dim, 3, padding=1)
31 self.convq = nn.Conv2d(hidden_dim+input_dim, hidden_dim, 3, padding=1)
32
33 def forward(self, h, x):
34 hx = torch.cat([h, x], dim=1)
35
36 z = torch.sigmoid(self.convz(hx))
37 r = torch.sigmoid(self.convr(hx))
38 q = torch.tanh(self.convq(torch.cat([r*h, x], dim=1)))
39
40 h = (1-z) * h + z * q
41 return h
42
43class SepConvGRU(nn.Module):
44 def __init__(self, hidden_dim=128, input_dim=192+128):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected