Download src/model.py from shubhexists/asr: direct link, hf CLI and curl.
- Browser
- Download file 9.86 kB
-
https://huggingface.co/shubhexists/asr/resolve/main/src/model.py
- Command line
-
hf download hf://shubhexists/asr/src/model.py
-
curl -L -o model.py https://huggingface.co/shubhexists/asr/resolve/main/src/model.py
9.86 kB
| """ | |
| 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(), | |
| } | |
| 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} | |
| 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) | |