freococo's picture
Initial release — classical BLSTMP-CTC Burmese ASR
8545139 verified
Raw History Blame Contribute Delete
4.9 kB
# ==============================================================================
# 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