Forward decoder. Args: memory: encoded memory, float32 (batch, maxlen_in, feat) memory_mask: encoder memory mask, (batch, 1, maxlen_in) ys_in_pad: padded input token ids, int64 (batch, maxlen_out) ys_in_lens: input lengths of this batch (batch
(
self,
memory: torch.Tensor,
memory_mask: torch.Tensor,
ys_in_pad: torch.Tensor,
ys_in_lens: torch.Tensor,
r_ys_in_pad: torch.Tensor = torch.empty(0),
reverse_weight: float = 0.0,
)
| 144 | self.use_sdpa = use_sdpa |
| 145 | |
| 146 | def forward( |
| 147 | self, |
| 148 | memory: torch.Tensor, |
| 149 | memory_mask: torch.Tensor, |
| 150 | ys_in_pad: torch.Tensor, |
| 151 | ys_in_lens: torch.Tensor, |
| 152 | r_ys_in_pad: torch.Tensor = torch.empty(0), |
| 153 | reverse_weight: float = 0.0, |
| 154 | ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: |
| 155 | """Forward decoder. |
| 156 | Args: |
| 157 | memory: encoded memory, float32 (batch, maxlen_in, feat) |
| 158 | memory_mask: encoder memory mask, (batch, 1, maxlen_in) |
| 159 | ys_in_pad: padded input token ids, int64 (batch, maxlen_out) |
| 160 | ys_in_lens: input lengths of this batch (batch) |
| 161 | r_ys_in_pad: not used in transformer decoder, in order to unify api |
| 162 | with bidirectional decoder |
| 163 | reverse_weight: not used in transformer decoder, in order to unify |
| 164 | api with bidirectional decode |
| 165 | Returns: |
| 166 | (tuple): tuple containing: |
| 167 | x: decoded token score before softmax (batch, maxlen_out, |
| 168 | vocab_size) if use_output_layer is True, |
| 169 | torch.tensor(0.0), in order to unify api with bidirectional decoder |
| 170 | olens: (batch, ) |
| 171 | NOTE(xcsong): |
| 172 | We pass the `__call__` method of the modules instead of `forward` to the |
| 173 | checkpointing API because `__call__` attaches all the hooks of the module. |
| 174 | https://discuss.pytorch.org/t/any-different-between-model-input-and-model-forward-input/3690/2 |
| 175 | """ |
| 176 | tgt = ys_in_pad |
| 177 | maxlen = tgt.size(1) |
| 178 | # tgt_mask: (B, 1, L) |
| 179 | tgt_mask = ~make_pad_mask(ys_in_lens, maxlen).unsqueeze(1) |
| 180 | tgt_mask = tgt_mask.to(tgt.device) |
| 181 | # m: (1, L, L) |
| 182 | m = subsequent_mask(tgt_mask.size(-1), |
| 183 | device=tgt_mask.device).unsqueeze(0) |
| 184 | # tgt_mask: (B, L, L) |
| 185 | tgt_mask = tgt_mask & m |
| 186 | if self.use_sdpa: |
| 187 | tgt_mask = mask_to_bias(tgt_mask, memory.dtype) |
| 188 | memory_mask = mask_to_bias(memory_mask, memory.dtype) |
| 189 | |
| 190 | x, _ = self.embed(tgt) |
| 191 | if self.gradient_checkpointing and self.training: |
| 192 | x = self.forward_layers_checkpointed(x, tgt_mask, memory, |
| 193 | memory_mask) |
| 194 | else: |
| 195 | x = self.forward_layers(x, tgt_mask, memory, memory_mask) |
| 196 | if self.normalize_before: |
| 197 | x = self.after_norm(x) |
| 198 | if self.use_output_layer: |
| 199 | x = self.output_layer(x) |
| 200 | olens = tgt_mask.sum(1) |
| 201 | return x, torch.tensor(0.0), olens |
| 202 | |
| 203 | def forward_layers(self, x: torch.Tensor, tgt_mask: torch.Tensor, |
nothing calls this directly
no test coverage detected