(self, x: torch.Tensor, tgt_mask: torch.Tensor,
memory: torch.Tensor,
memory_mask: torch.Tensor)
| 167 | return x, torch.tensor(0.0), olens |
| 168 | |
| 169 | def forward_layers(self, x: torch.Tensor, tgt_mask: torch.Tensor, |
| 170 | memory: torch.Tensor, |
| 171 | memory_mask: torch.Tensor) -> torch.Tensor: |
| 172 | for layer in self.decoders: |
| 173 | x, tgt_mask, memory, memory_mask = layer(x, tgt_mask, memory, |
| 174 | memory_mask) |
| 175 | return x |
| 176 | |
| 177 | @torch.jit.unused |
| 178 | def forward_layers_checkpointed(self, x: torch.Tensor, |