| 28 | ]) |
| 29 | |
| 30 | def forward(self, inputs, hidden): |
| 31 | T, B = inputs.size(0), inputs.size(1) |
| 32 | |
| 33 | if self.training: |
| 34 | x_mask = mask2d(B, inputs.size(2), keep_prob=1.-self.dropoutx) |
| 35 | h_mask = mask2d(B, hidden.size(2), keep_prob=1.-self.dropouth) |
| 36 | else: |
| 37 | x_mask = h_mask = None |
| 38 | |
| 39 | hidden = hidden[0] |
| 40 | hiddens = [] |
| 41 | for t in range(T): |
| 42 | hidden = self.cell(inputs[t], hidden, x_mask, h_mask) |
| 43 | hiddens.append(hidden) |
| 44 | hiddens = torch.stack(hiddens) |
| 45 | return hiddens, hiddens[-1].unsqueeze(0) |
| 46 | |
| 47 | def _compute_init_state(self, x, h_prev, x_mask, h_mask): |
| 48 | if self.training: |