first commit asr model
This commit is contained in:
@@ -0,0 +1,183 @@
|
||||
import random
|
||||
|
||||
import librosa
|
||||
import torch
|
||||
import torchaudio
|
||||
from torch import Tensor, no_grad, device
|
||||
from torchaudio.transforms import MelSpectrogram, AmplitudeToDB, Resample, TimeStretch
|
||||
from pathlib import Path
|
||||
import torchaudio.functional as F
|
||||
|
||||
from tokenizer import ASRTokenizer
|
||||
from model import ASRModel
|
||||
|
||||
CONFIG = {
|
||||
# 模型配置
|
||||
'input_dim': 640,
|
||||
'num_heads': 8,
|
||||
'ffn_dim': 2048,
|
||||
'num_layers': 8,
|
||||
'dropout': 0.1,
|
||||
}
|
||||
|
||||
class ASRInference:
|
||||
def __init__(self, model_path: Path, vocab_path: Path, device: device, augment: bool = True, augment_prob: float = 0.5) -> None:
|
||||
self.device = device
|
||||
self.augment: bool = augment
|
||||
self.augment_prob: float = augment_prob
|
||||
self.tokenizer = ASRTokenizer(vocab_path=vocab_path)
|
||||
self.model = ASRModel(
|
||||
vocab_size=self.tokenizer.vocab_size(),
|
||||
input_dim=CONFIG['input_dim'],
|
||||
num_heads=CONFIG['num_heads'],
|
||||
ffn_dim=CONFIG['ffn_dim'],
|
||||
num_layers=CONFIG['num_layers'],
|
||||
dropout=CONFIG['dropout'],
|
||||
).to(device)
|
||||
|
||||
self.model.load_state_dict(torch.load(model_path, map_location=device)['model_state_dict'])
|
||||
self.model.eval()
|
||||
|
||||
print(f"params params: {self.model.get_num_params():,}",)
|
||||
|
||||
self.sample_rate = 16000
|
||||
self.mel_transform = MelSpectrogram(
|
||||
sample_rate=self.sample_rate,
|
||||
n_fft=400,
|
||||
win_length=400,
|
||||
hop_length=160,
|
||||
n_mels=80,
|
||||
f_min=0,
|
||||
f_max=8000,
|
||||
power=2.0,
|
||||
)
|
||||
self.amplitude_to_db = AmplitudeToDB()
|
||||
|
||||
def _load_audio(self, audio_path: Path) -> Tensor:
|
||||
waveform, sample_rate = torchaudio.load_with_torchcodec(audio_path)
|
||||
|
||||
if sample_rate != self.sample_rate:
|
||||
waveform = Resample(sample_rate, self.sample_rate)(waveform)
|
||||
|
||||
if waveform.shape[0] > 1:
|
||||
waveform = waveform.mean(dim=0, keepdim=True)
|
||||
|
||||
waveform = waveform / (waveform.abs().max() + 1e-8)
|
||||
return waveform
|
||||
|
||||
def _extract_features(self, waveform: Tensor) -> Tensor:
|
||||
mel_spec: Tensor = self.mel_transform(waveform)
|
||||
log_mel_spec: Tensor = self.amplitude_to_db(mel_spec)
|
||||
return log_mel_spec.squeeze(0) # [n_mels, time]
|
||||
|
||||
def _augment_waveform(self, waveform: Tensor) -> Tensor:
|
||||
if not self.augment or random.random() > self.augment_prob:
|
||||
return waveform
|
||||
|
||||
# 1. voice Stretch/Compress
|
||||
if random.random() < 0.5:
|
||||
waveform = self._voice_stretch_or_compress(waveform=waveform)
|
||||
|
||||
if random.random() < 0.3:
|
||||
waveform = self._drop_frames(waveform)
|
||||
|
||||
if random.random() < 0.4:
|
||||
waveform = self._add_noise(waveform)
|
||||
|
||||
return waveform
|
||||
|
||||
def _voice_stretch_or_compress(self, waveform: Tensor) -> Tensor:
|
||||
speed_factor = random.uniform(0.6, 1.4) # (Speed Change: 0.6x - 1.4x)
|
||||
spec = torch.stft(
|
||||
waveform.squeeze(0),
|
||||
n_fft=400,
|
||||
hop_length=160,
|
||||
window=torch.hann_window(400).to(waveform.device),
|
||||
return_complex=True
|
||||
)
|
||||
|
||||
# 时间拉伸(不改变音高)
|
||||
stretch = TimeStretch(
|
||||
hop_length=160,
|
||||
n_freq=201,
|
||||
fixed_rate=speed_factor
|
||||
)
|
||||
stretched_spec = stretch(spec)
|
||||
|
||||
# 转回波形
|
||||
waveform_stretched = torch.istft(
|
||||
stretched_spec,
|
||||
n_fft=400,
|
||||
hop_length=160,
|
||||
window=torch.hann_window(400).to(waveform.device)
|
||||
).unsqueeze(0)
|
||||
|
||||
return waveform_stretched
|
||||
|
||||
def _drop_frames(self, waveform: Tensor) -> Tensor:
|
||||
audio_len = waveform.shape[1]
|
||||
drop_ratio = random.uniform(0.05, 0.15)
|
||||
drop_len = int(audio_len * drop_ratio)
|
||||
|
||||
if audio_len > drop_len:
|
||||
start_pos = random.randint(0, audio_len - drop_len)
|
||||
# clean
|
||||
waveform = torch.cat([waveform[:, :start_pos], waveform[:, start_pos + drop_len:]], dim=1)
|
||||
|
||||
return waveform
|
||||
|
||||
def _add_noise(self, waveform: Tensor) -> Tensor:
|
||||
snr_db = random.uniform(10, 20)
|
||||
signal_power = torch.mean(waveform ** 2)
|
||||
|
||||
snr_linear = 10 ** (snr_db / 10)
|
||||
noise_power = signal_power / snr_linear
|
||||
|
||||
noise = torch.randn_like(waveform) * torch.sqrt(noise_power)
|
||||
return waveform + noise
|
||||
|
||||
def transcribe(self, audio_path: Path) -> str:
|
||||
waveform = self._load_audio(audio_path=audio_path)
|
||||
waveform = self._augment_waveform(waveform=waveform)
|
||||
mel_spec = self._extract_features(waveform=waveform)
|
||||
|
||||
mel_spec = mel_spec.unsqueeze(0).to(self.device) # [1, n_mels, time]
|
||||
mel_length = torch.tensor([mel_spec.shape[2]], dtype=torch.long, device=self.device)
|
||||
|
||||
with no_grad():
|
||||
log_probs, _ = self.model(mel_specs=mel_spec, mel_lengths=mel_length) # [1, T, vocab]
|
||||
|
||||
text = self.tokenizer.ctc_greedy_decode(log_probs=log_probs[0])
|
||||
return text
|
||||
|
||||
def transcribe_batch(self, audio_paths: list[Path]) -> list[str]:
|
||||
results = []
|
||||
|
||||
for audio_path in audio_paths:
|
||||
text = self.transcribe(audio_path=audio_path)
|
||||
results.append(text)
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def main():
|
||||
workspace_dir = Path(__file__).parent.parent
|
||||
device = torch.device('cuda:0')
|
||||
|
||||
# checkpoint = sorted(workspace_dir.glob('.checkpoints/checkpoint_epoch_*.pt'), key=lambda p: int(p.stem.split('_')[1]))[-1]
|
||||
checkpoint = workspace_dir / ".checkpoints/checkpoint_step_9500.pt"
|
||||
print(f"Load Checkpoint: {checkpoint}")
|
||||
|
||||
inference = ASRInference(model_path=checkpoint, vocab_path=workspace_dir / 'config/uig_vocab.json' , device=device)
|
||||
|
||||
audio_path = workspace_dir / 'data/test/F001_001.wav'
|
||||
print(f"\n转录音频: {audio_path}")
|
||||
text = inference.transcribe(audio_path=audio_path)
|
||||
print(f"\n识别结果: {text}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
# workspace_dir = Path(__file__).parent.parent
|
||||
# checkpoint = sorted(workspace_dir.glob('.checkpoints/checkpoint_step_*.pt'), key=lambda p: int(p.stem.split('_')[-1]))[-1]
|
||||
|
||||
# print(checkpoint)
|
||||
Reference in New Issue
Block a user