# ============================================================================== # model.py — Classical Burmese ASR: Conv2D + 4-layer BLSTMP + CTC # Standalone inference-ready architecture. # ============================================================================== import torch import torch.nn as nn import torch.nn.functional as F from pathlib import Path # ============================================================================== # Architecture constants (matching training) # ============================================================================== N_MELS = 80 HIDDEN = 320 PROJ = 512 LAYERS = 4 DROPOUT = 0.15 # Special token ids BLANK_ID = 0 UNK_ID = 1 # ============================================================================== # BLSTMP — 4-layer Bidirectional LSTM with projection + residual # ============================================================================== class BLSTMP(nn.Module): def __init__(self, in_dim, hidden=320, proj=512, layers=4, dropout=0.15): super().__init__() self.layers = nn.ModuleList() self.projs = nn.ModuleList() self.norms = nn.ModuleList() for i in range(layers): d = in_dim if i == 0 else proj self.layers.append(nn.LSTM(d, hidden, 1, bidirectional=True, batch_first=True)) self.projs.append(nn.Linear(hidden * 2, proj)) self.norms.append(nn.LayerNorm(proj)) self.dropout = nn.Dropout(dropout) def forward(self, xs, ilens_cpu): T = xs.shape[1] for i, (l, p, n) in enumerate(zip(self.layers, self.projs, self.norms)): res = xs pk = nn.utils.rnn.pack_padded_sequence( xs, ilens_cpu, batch_first=True, enforce_sorted=False) out, _ = l(pk) out, _ = nn.utils.rnn.pad_packed_sequence( out, batch_first=True, total_length=T) out = n(p(out)) if i > 0 and res.shape == out.shape: out = out + res if i < len(self.layers) - 1: out = self.dropout(F.relu(out)) xs = out return xs # ============================================================================== # FastASR — full model # ============================================================================== class FastASR(nn.Module): def __init__(self, vocab_size=2566): super().__init__() # Frontend: 2 blocks of (Conv2d x2 + BN + ReLU) + MaxPool self.c1 = nn.Sequential( nn.Conv2d(1, 64, 3, padding=1), nn.BatchNorm2d(64), nn.ReLU(), nn.Conv2d(64, 64, 3, padding=1), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2, 2)) self.c2 = nn.Sequential( nn.Conv2d(64, 128, 3, padding=1), nn.BatchNorm2d(128), nn.ReLU(), nn.Conv2d(128, 128, 3, padding=1), nn.BatchNorm2d(128), nn.ReLU(), nn.MaxPool2d(2, 2)) # Encoder self.rnn = BLSTMP(in_dim=128 * (N_MELS // 4)) # Output head self.fc = nn.Linear(PROJ, vocab_size) self.ctc = nn.CTCLoss(blank=BLANK_ID, zero_infinity=True) def forward(self, xs, ilens=None): """Default forward: full audio-to-log-probs. Args: xs : Tensor (B, T_frames, N_MELS) of log-mel features ilens : Tensor or list of input lengths in MEL FRAMES. Defaults to full T_frames for all items. Returns: lp : Tensor (B, T_enc, V) log-softmax over CTC vocab el_cpu : Tensor (B,) encoder lengths after 4x subsampling (on CPU) """ B, T, _ = xs.shape if ilens is None: ilens = torch.full((B,), T, dtype=torch.long, device="cpu") elif not isinstance(ilens, torch.Tensor): ilens = torch.tensor(ilens, dtype=torch.long, device="cpu") x = self.c1(xs.unsqueeze(1)) # (B, 64, T/2, F/2) x = self.c2(x) # (B, 128, T/4, F/4) b, c, t, f = x.shape x = x.transpose(1, 2).reshape(b, t, c * f) # (B, T/4, 2560) el_cpu = torch.clamp(ilens.cpu() // 4, min=1, max=t) x = self.rnn(x, el_cpu) lg = self.fc(x) lp = F.log_softmax(lg, dim=-1) return lp, el_cpu # ============================================================================== # Loader helper # ============================================================================== def load_model(weights_path=None, vocab_size=2566, device="cpu"): """Loads the FastASR model populated from safetensors.""" from safetensors.torch import load_file model = FastASR(vocab_size=vocab_size) if weights_path is not None: sd = load_file(str(weights_path)) model.load_state_dict(sd, strict=True) model.to(device) model.eval() return model