asr / src /model.py
shubhexists's picture
Add Zipformer-inspired ASR model: weights, tokenizer, config, and training code
ce3c8df verified
Raw History Blame Contribute Delete
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(),
}
@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)