Cal decoder with predictor. Args: encoder_out: Encoder output tensor. encoder_out_lens: Encoder output lengths. sematic_embeds: TODO. ys_pad_lens: Lengths of ys_pad.
(
self, encoder_out, encoder_out_lens, sematic_embeds, ys_pad_lens
)
| 329 | return pre_acoustic_embeds, pre_token_length, alphas, pre_peak_index |
| 330 | |
| 331 | def cal_decoder_with_predictor( |
| 332 | self, encoder_out, encoder_out_lens, sematic_embeds, ys_pad_lens |
| 333 | ): |
| 334 | |
| 335 | """Cal decoder with predictor. |
| 336 | |
| 337 | Args: |
| 338 | encoder_out: Encoder output tensor. |
| 339 | encoder_out_lens: Encoder output lengths. |
| 340 | sematic_embeds: TODO. |
| 341 | ys_pad_lens: Lengths of ys_pad. |
| 342 | """ |
| 343 | decoder_outs = self.decoder(encoder_out, encoder_out_lens, sematic_embeds, ys_pad_lens) |
| 344 | decoder_out = decoder_outs[0] |
| 345 | decoder_out = torch.log_softmax(decoder_out, dim=-1) |
| 346 | return decoder_out, ys_pad_lens |
| 347 | |
| 348 | def _calc_att_loss( |
| 349 | self, |
no test coverage detected