(self, x, h_prev, x_mask, h_mask)
| 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 | |
| 91 | class RNNModel(nn.Module): |
no test coverage detected