(self, x: torch.Tensor, tgt_mask: torch.Tensor,
memory: torch.Tensor,
memory_mask: torch.Tensor)
| 201 | return x, torch.tensor(0.0), olens |
| 202 | |
| 203 | def forward_layers(self, x: torch.Tensor, tgt_mask: torch.Tensor, |
| 204 | memory: torch.Tensor, |
| 205 | memory_mask: torch.Tensor) -> torch.Tensor: |
| 206 | for layer in self.decoders: |
| 207 | x, tgt_mask, memory, memory_mask = layer(x, tgt_mask, memory, |
| 208 | memory_mask) |
| 209 | return x |
| 210 | |
| 211 | @torch.jit.unused |
| 212 | def forward_layers_checkpointed(self, x: torch.Tensor, |