""" Full ASR model: waveform -> log-mel -> SpecAugment -> Zipformer encoder -> {CTC head, attention decoder}, trained with hybrid CR-CTC + attention loss. CR-CTC (arXiv:2410.05101): two SpecAugmented views of each utterance are stacked on the batch dim for one encoder pass, then split; each gets a CTC loss plus a symmetric-KL consistency loss between the two. The decoder only consumes view 1 -- consistency regularization is a CTC-branch technique. """ import random from typing import Optional import torch import torch.nn as nn import torch.nn.functional as F import torchaudio from src.decoder import AttentionDecoder from src.zipformer import ZipformerEncoder, make_pad_mask class LogMelFrontend(nn.Module): """Waveform -> log-mel fbank with masked per-utterance mean/var normalization.""" def __init__( self, sample_rate: int = 16000, n_mels: int = 80, n_fft: int = 400, hop_length: int = 160, win_length: int = 400, ): super().__init__() self.hop_length = hop_length self.mel = torchaudio.transforms.MelSpectrogram( sample_rate=sample_rate, n_fft=n_fft, hop_length=hop_length, win_length=win_length, n_mels=n_mels, center=True, ) def forward(self, waveforms: torch.Tensor, wave_lengths: torch.Tensor): # waveforms: (B, L) power_spec = self.mel(waveforms) # (B, n_mels, T) log_mel = torch.log(power_spec.clamp(min=1e-6)) feats = log_mel.transpose(1, 2) # (B, T, n_mels) feat_lengths = torch.div(wave_lengths, self.hop_length, rounding_mode="floor") + 1 feat_lengths = feat_lengths.clamp(max=feats.size(1)) pad_mask = make_pad_mask(feat_lengths, feats.size(1)) # True = padded valid = (~pad_mask).unsqueeze(-1).to(feats.dtype) # (B, T, 1) count = valid.sum(dim=1, keepdim=True).clamp(min=1.0) mean = (feats * valid).sum(dim=1, keepdim=True) / count var = ((feats - mean) ** 2 * valid).sum(dim=1, keepdim=True) / count feats = (feats - mean) / torch.sqrt(var + 1e-5) feats = feats * valid # re-zero padding after normalization return feats, feat_lengths def spec_augment( feats: torch.Tensor, lengths: torch.Tensor, freq_mask_param: int = 27, time_mask_param: int = 100, num_freq_masks: int = 2, num_time_masks: int = 2, ) -> torch.Tensor: """Zeroes random frequency bands and time spans, per-utterance, on (B, T, F).""" feats = feats.clone() b, t, f = feats.shape for i in range(b): length = int(lengths[i].item()) for _ in range(num_freq_masks): width = random.randint(0, min(freq_mask_param, f)) if width == 0: continue start = random.randint(0, f - width) feats[i, :, start : start + width] = 0.0 for _ in range(num_time_masks): width = random.randint(0, min(time_mask_param, length)) if width == 0: continue start = random.randint(0, length - width) feats[i, start : start + width, :] = 0.0 return feats def masked_mean(values: torch.Tensor, pad_mask: torch.Tensor) -> torch.Tensor: """values: (B, T); pad_mask: (B, T) True at padded positions.""" valid = (~pad_mask).to(values.dtype) return (values * valid).sum() / valid.sum().clamp(min=1.0) class ASRModel(nn.Module): def __init__( self, vocab_size: int, blank_id: int = 0, pad_id: int = 0, n_mels: int = 80, d_model: int = 256, conv_channels: int = 128, encoder_nhead: int = 4, encoder_d_ff: int = 1024, conv_kernel: int = 15, stage_layers=(2, 3, 4, 3, 2), downsample_factors=(2, 2), encoder_dropout: float = 0.1, decoder_nhead: int = 4, decoder_d_ff: int = 1024, decoder_num_layers: int = 4, decoder_dropout: float = 0.1, ctc_weight: float = 0.3, attn_weight: float = 0.7, cr_loss_weight: float = 0.2, sample_rate: int = 16000, specaugment: Optional[dict] = None, ): super().__init__() self.blank_id = blank_id self.pad_id = pad_id self.ctc_weight = ctc_weight self.attn_weight = attn_weight self.cr_loss_weight = cr_loss_weight self.specaugment_cfg = specaugment or {} self.frontend = LogMelFrontend(sample_rate=sample_rate, n_mels=n_mels) self.encoder = ZipformerEncoder( n_mels=n_mels, d_model=d_model, conv_channels=conv_channels, nhead=encoder_nhead, d_ff=encoder_d_ff, conv_kernel=conv_kernel, stage_layers=stage_layers, downsample_factors=downsample_factors, dropout=encoder_dropout, ) self.ctc_head = nn.Linear(d_model, vocab_size) self.decoder = AttentionDecoder( vocab_size=vocab_size, d_model=d_model, nhead=decoder_nhead, d_ff=decoder_d_ff, num_layers=decoder_num_layers, dropout=decoder_dropout, ) def _augmented_view(self, feats, lengths): return spec_augment(feats, lengths, **self.specaugment_cfg) def forward_train( self, waveforms: torch.Tensor, wave_lengths: torch.Tensor, ctc_targets: torch.Tensor, ctc_target_lengths: torch.Tensor, decoder_input_tokens: torch.Tensor, decoder_input_lengths: torch.Tensor, decoder_target_tokens: torch.Tensor, ) -> dict: feats, feat_lengths = self.frontend(waveforms, wave_lengths) view1 = self._augmented_view(feats, feat_lengths) view2 = self._augmented_view(feats, feat_lengths) combined_feats = torch.cat([view1, view2], dim=0) combined_lengths = torch.cat([feat_lengths, feat_lengths], dim=0) enc_out, enc_lengths = self.encoder(combined_feats, combined_lengths) b = feats.size(0) enc_out1, enc_out2 = enc_out[:b], enc_out[b:] enc_lengths1, enc_lengths2 = enc_lengths[:b], enc_lengths[b:] ctc_logits1 = self.ctc_head(enc_out1) ctc_logits2 = self.ctc_head(enc_out2) log_probs1 = F.log_softmax(ctc_logits1, dim=-1) log_probs2 = F.log_softmax(ctc_logits2, dim=-1) ctc_loss1 = F.ctc_loss( log_probs1.transpose(0, 1), ctc_targets, enc_lengths1, ctc_target_lengths, blank=self.blank_id, zero_infinity=True, ) ctc_loss2 = F.ctc_loss( log_probs2.transpose(0, 1), ctc_targets, enc_lengths2, ctc_target_lengths, blank=self.blank_id, zero_infinity=True, ) ctc_loss = 0.5 * (ctc_loss1 + ctc_loss2) kl_1_given_2 = F.kl_div(log_probs1, log_probs2, log_target=True, reduction="none").sum(-1) kl_2_given_1 = F.kl_div(log_probs2, log_probs1, log_target=True, reduction="none").sum(-1) symmetric_kl = 0.5 * (kl_1_given_2 + kl_2_given_1) # (B, T) pad_mask1 = make_pad_mask(enc_lengths1, enc_out1.size(1)) cr_loss = masked_mean(symmetric_kl, pad_mask1) decoder_logits = self.decoder( decoder_input_tokens, decoder_input_lengths, enc_out1, enc_lengths1 ) ce_loss = F.cross_entropy( decoder_logits.reshape(-1, decoder_logits.size(-1)), decoder_target_tokens.reshape(-1), ignore_index=self.pad_id, ) total_loss = ( self.ctc_weight * ctc_loss + self.attn_weight * ce_loss + self.cr_loss_weight * cr_loss ) return { "loss": total_loss, "ctc_loss": ctc_loss.detach(), "ce_loss": ce_loss.detach(), "cr_loss": cr_loss.detach(), } @torch.no_grad() def forward_val_loss( self, waveforms: torch.Tensor, wave_lengths: torch.Tensor, ctc_targets: torch.Tensor, ctc_target_lengths: torch.Tensor, decoder_input: torch.Tensor, decoder_input_lengths: torch.Tensor, decoder_target: torch.Tensor, ) -> dict: """Single clean pass (no SpecAugment, no CR-CTC) for monitoring validation loss.""" feats, feat_lengths = self.frontend(waveforms, wave_lengths) enc_out, enc_lengths = self.encoder(feats, feat_lengths) log_probs = F.log_softmax(self.ctc_head(enc_out), dim=-1) ctc_loss = F.ctc_loss( log_probs.transpose(0, 1), ctc_targets, enc_lengths, ctc_target_lengths, blank=self.blank_id, zero_infinity=True, ) decoder_logits = self.decoder(decoder_input, decoder_input_lengths, enc_out, enc_lengths) ce_loss = F.cross_entropy( decoder_logits.reshape(-1, decoder_logits.size(-1)), decoder_target.reshape(-1), ignore_index=self.pad_id, ) loss = self.ctc_weight * ctc_loss + self.attn_weight * ce_loss return {"loss": loss, "ctc_loss": ctc_loss, "ce_loss": ce_loss} @torch.no_grad() def forward_eval(self, waveforms: torch.Tensor, wave_lengths: torch.Tensor): """No augmentation, single pass. Returns encoder out/lengths and CTC log-probs.""" feats, feat_lengths = self.frontend(waveforms, wave_lengths) enc_out, enc_lengths = self.encoder(feats, feat_lengths) ctc_log_probs = F.log_softmax(self.ctc_head(enc_out), dim=-1) return enc_out, enc_lengths, ctc_log_probs def count_parameters(model: nn.Module) -> int: return sum(p.numel() for p in model.parameters() if p.requires_grad)