MCPcopy Create free account
hub / github.com/OUCMachineLearning/OUCML / cell

Method cell

AutoML/darts-master/rnn/model.py:71–88  ·  view source on GitHub ↗
(self, x, h_prev, x_mask, h_mask)

Source from the content-addressed store, hash-verified

69 return f
70
71 def cell(self, x, h_prev, x_mask, h_mask):
72 s0 = self._compute_init_state(x, h_prev, x_mask, h_mask)
73
74 states = [s0]
75 for i, (name, pred) in enumerate(self.genotype.recurrent):
76 s_prev = states[pred]
77 if self.training:
78 ch = (s_prev * h_mask).mm(self._Ws[i])
79 else:
80 ch = s_prev.mm(self._Ws[i])
81 c, h = torch.split(ch, self.nhid, dim=-1)
82 c = c.sigmoid()
83 fn = self._get_activation(name)
84 h = fn(h)
85 s = s_prev + c * (h-s_prev)
86 states += [s]
87 output = torch.mean(torch.stack([states[i] for i in self.genotype.concat], -1), -1)
88 return output
89
90
91class RNNModel(nn.Module):

Callers 1

forwardMethod · 0.95

Calls 2

_compute_init_stateMethod · 0.95
_get_activationMethod · 0.95

Tested by

no test coverage detected