| | |
| | | import torch.nn.functional as F |
| | | |
| | | from funasr.layers.abs_normalize import AbsNormalize |
| | | from funasr.losses.label_smoothing_loss import ( |
| | | LabelSmoothingLoss, NllLoss # noqa: H301 |
| | | ) |
| | | from funasr.losses.label_smoothing_loss import LabelSmoothingLoss, NllLoss # noqa: H301 |
| | | from funasr.models.ctc import CTC |
| | | from funasr.models.decoder.abs_decoder import AbsDecoder |
| | | from funasr.models.encoder.abs_encoder import AbsEncoder |
| | |
| | | |
| | | self.error_calculator = None |
| | | |
| | | |
| | | # we set self.decoder = None in the CTC mode since |
| | | # self.decoder parameters were never used and PyTorch complained |
| | | # and threw an Exception in the multi-GPU experiment. |
| | |
| | | profile: torch.Tensor, |
| | | profile_lengths: torch.Tensor, |
| | | text_id: torch.Tensor, |
| | | text_id_lengths: torch.Tensor |
| | | text_id_lengths: torch.Tensor, |
| | | ) -> Tuple[torch.Tensor, Dict[str, torch.Tensor], torch.Tensor]: |
| | | """Frontend + Encoder + Decoder + Calc loss |
| | | |
| | |
| | | assert text_lengths.dim() == 1, text_lengths.shape |
| | | # Check that batch_size is unified |
| | | assert ( |
| | | speech.shape[0] |
| | | == speech_lengths.shape[0] |
| | | == text.shape[0] |
| | | == text_lengths.shape[0] |
| | | speech.shape[0] == speech_lengths.shape[0] == text.shape[0] == text_lengths.shape[0] |
| | | ), (speech.shape, speech_lengths.shape, text.shape, text_lengths.shape) |
| | | batch_size = speech.shape[0] |
| | | |
| | |
| | | asr_encoder_out, encoder_out_lens, text, text_lengths |
| | | ) |
| | | |
| | | |
| | | # Intermediate CTC (optional) |
| | | loss_interctc = 0.0 |
| | | if self.interctc_weight != 0.0 and intermediate_outs is not None: |
| | |
| | | loss_interctc = loss_interctc / len(intermediate_outs) |
| | | |
| | | # calculate whole encoder loss |
| | | loss_ctc = ( |
| | | 1 - self.interctc_weight |
| | | ) * loss_ctc + self.interctc_weight * loss_interctc |
| | | |
| | | loss_ctc = (1 - self.interctc_weight) * loss_ctc + self.interctc_weight * loss_interctc |
| | | |
| | | # 2b. Attention decoder branch |
| | | if self.ctc_weight != 1.0: |
| | | loss_att, loss_spk, acc_att, acc_spk, cer_att, wer_att = self._calc_att_loss( |
| | | asr_encoder_out, spk_encoder_out, encoder_out_lens, text, text_lengths, profile, profile_lengths, text_id, text_id_lengths |
| | | asr_encoder_out, |
| | | spk_encoder_out, |
| | | encoder_out_lens, |
| | | text, |
| | | text_lengths, |
| | | profile, |
| | | profile_lengths, |
| | | text_id, |
| | | text_id_lengths, |
| | | ) |
| | | |
| | | # 3. CTC-Att loss definition |
| | |
| | | loss = loss_asr |
| | | else: |
| | | loss = self.spk_weight * loss_spk + (1 - self.spk_weight) * loss_asr |
| | | |
| | | |
| | | stats = dict( |
| | | loss=loss.detach(), |
| | |
| | | # feats: (Batch, Length, Dim) |
| | | # -> encoder_out: (Batch, Length2, Dim2) |
| | | if self.asr_encoder.interctc_use_conditioning: |
| | | encoder_out, encoder_out_lens, _ = self.asr_encoder( |
| | | feats, feats_lengths, ctc=self.ctc |
| | | ) |
| | | encoder_out, encoder_out_lens, _ = self.asr_encoder(feats, feats_lengths, ctc=self.ctc) |
| | | else: |
| | | encoder_out, encoder_out_lens, _ = self.asr_encoder(feats, feats_lengths) |
| | | intermediate_outs = None |
| | |
| | | encoder_out_spk_ori = self.spk_encoder(feats_raw, feats_lengths)[0] |
| | | # import ipdb;ipdb.set_trace() |
| | | if encoder_out_spk_ori.size(1)!=encoder_out.size(1): |
| | | encoder_out_spk=F.interpolate(encoder_out_spk_ori.transpose(-2,-1), size=(encoder_out.size(1)), mode='nearest').transpose(-2,-1) |
| | | encoder_out_spk = F.interpolate( |
| | | encoder_out_spk_ori.transpose(-2, -1), size=(encoder_out.size(1)), mode="nearest" |
| | | ).transpose(-2, -1) |
| | | else: |
| | | encoder_out_spk=encoder_out_spk_ori |
| | | |
| | |
| | | profile: torch.Tensor, |
| | | profile_lens: torch.Tensor, |
| | | text_id: torch.Tensor, |
| | | text_id_lengths: torch.Tensor |
| | | text_id_lengths: torch.Tensor, |
| | | ): |
| | | ys_in_pad, ys_out_pad = add_sos_eos(ys_pad, self.sos, self.eos, self.ignore_id) |
| | | ys_in_lens = ys_pad_lens + 1 |
| | | |
| | | # 1. Forward decoder |
| | | decoder_out, weights_no_pad, _ = self.decoder( |
| | | asr_encoder_out, spk_encoder_out, encoder_out_lens, ys_in_pad, ys_in_lens, profile, profile_lens |
| | | asr_encoder_out, |
| | | spk_encoder_out, |
| | | encoder_out_lens, |
| | | ys_in_pad, |
| | | ys_in_lens, |
| | | profile, |
| | | profile_lens, |
| | | ) |
| | | |
| | | spk_num_no_pad=weights_no_pad.size(-1) |
| | | pad=(0,self.max_spk_num-spk_num_no_pad) |
| | | weights=F.pad(weights_no_pad, pad, mode='constant', value=0) |
| | | weights = F.pad(weights_no_pad, pad, mode="constant", value=0) |
| | | |
| | | # pre_id=weights.argmax(-1) |
| | | # pre_text=decoder_out.argmax(-1) |