zhifu gao
2024-04-24 861147c7308b91068ffa02724fdf74ee623a909e
funasr/models/sa_asr/e2e_sa_asr.py
@@ -14,9 +14,7 @@
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
@@ -97,7 +95,6 @@
        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.
@@ -141,7 +138,7 @@
            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
@@ -156,10 +153,7 @@
        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]
@@ -183,7 +177,6 @@
                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:
@@ -204,15 +197,20 @@
            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
@@ -227,7 +225,6 @@
            loss = loss_asr
        else:
            loss = self.spk_weight * loss_spk + (1 - self.spk_weight) * loss_asr
        stats = dict(
            loss=loss.detach(),
@@ -291,9 +288,7 @@
        # 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
@@ -304,7 +299,9 @@
        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
@@ -440,19 +437,25 @@
            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)