Files
audio_model/src/dataset.py
T

336 lines
13 KiB
Python

import random
import torch
from torch import Tensor
from torch.utils.data import Dataset, DataLoader
import torchaudio
from torch_audiomentations import ApplyImpulseResponse, Gain, PitchShift, LowPassFilter, HighPassFilter, PolarityInversion
from torchaudio.transforms import Resample, TimeStretch
from pathlib import Path
import pandas as pd
from typing import List, TypedDict
from handle.text_normalizer import collapse_spaces, normalize_extended_uyghur_characters
from tokenizer import ASRTokenizer
# 单个样本的数据结构(Dataset.__getitem__ 返回)
class BatchItem(TypedDict):
waveform: Tensor # [time]
target_ids: Tensor # [seq_len] 目标文本的token IDs
target_text: str # 原始文本
audio_path: str # 音频文件路径
# 批量数据的数据结构(collate_fn 返回,DataLoader 输出)
class Batch(TypedDict):
waveforms: Tensor # [batch, time]
targets: Tensor # [batch, max_len] padding后的目标IDs
waveform_lengths: Tensor # [batch] 每个样本的实际Waveform长度
target_lengths: Tensor # [batch] 每个样本的实际目标长度
target_texts: List[str] # [batch] 原始文本列表
audio_paths: List[str] # [batch] 音频路径列表
class TsvFormat(TypedDict):
client_id: str
path: str
sentence: str
up_votes: int
down_votes: int
age: str
gender: str
locale: str
class NoiseAugmentor:
def __init__(self, noise_root: Path, sample_rate: int=16000):
self.sample_rate = sample_rate
self.noise_files = list(Path(noise_root).rglob("*.wav"))
def apply_real_noise(self, waveform: Tensor):
# 1. 随机选一个噪音文件
noise_path = random.choice(self.noise_files)
noise_waveform, sr = torchaudio.load_with_torchcodec(noise_path)
# Resample to target sample rate.
if sr != self.sample_rate:
noise_waveform = Resample(sr, self.sample_rate)(noise_waveform)
# Convert to mono if it is setro.
if waveform.shape[0] > 1:
waveform = waveform.mean(dim=0, keepdim=True)
# 3. 截取或填充,使其长度与语音一致
sig_len = waveform.shape[1]
noise_len = noise_waveform.shape[1]
if noise_len >= sig_len:
# 随机截取一段
start = random.randint(0, noise_len - sig_len)
noise_waveform = noise_waveform[:, start:start + sig_len]
else:
full_noise = torch.zeros_like(waveform)
start = random.randint(0, sig_len - noise_len)
full_noise[:, start : start + noise_len] = noise_waveform
noise_waveform = full_noise
# 4. 设定随机信噪比 SNR (5dB 到 20dB)
snr_db = random.uniform(5, 20)
# 5. 混合
return self._mix_at_snr(waveform, noise_waveform, snr_db)
def _mix_at_snr(self, signal: Tensor, noise: Tensor, snr_db: float):
s_p = signal.pow(2).mean()
n_p = noise.pow(2).mean()
snr_linear = 10**(snr_db / 10)
scale = torch.sqrt(s_p / (n_p * snr_linear + 1e-8))
noisy = signal + scale * noise
# 归一化,防止溢出
return noisy / (noisy.abs().max() + 1e-8)
class CommonVoiceDataset(Dataset[BatchItem]):
def __init__(
self,
tsv_path: Path,
audio_dir: Path,
noise_dir: Path,
corridor_noise_dir: Path,
tokenizer: ASRTokenizer,
sample_rate: int = 16000,
max_audio_len: int = 480000, # 30秒 @ 16kHz
augment: bool = True,
augment_prob: float = 0.5,
) -> None:
super().__init__()
self.noise_augmentor = NoiseAugmentor(noise_root=noise_dir, sample_rate=sample_rate)
self.audio_dir = Path(audio_dir)
self.tokenizer = tokenizer
self.sample_rate = sample_rate
self.max_audio_len = max_audio_len
self.augment = augment
self.augment_prob = augment_prob
self.data: pd.DataFrame = pd.read_csv(tsv_path, sep='\t')
valid_indices = []
for index, row in self.data.iterrows():
audio_path: Path = self.audio_dir / row['path']
if audio_path.exists():
valid_indices.append(index)
self.data = self.data.loc[valid_indices].reset_index(drop=True)
self.gain_up = Gain(min_gain_in_db=4, max_gain_in_db=8, p=1.0, output_type='tensor')
self.gain_down = Gain(min_gain_in_db=-15, max_gain_in_db=-8, p=1.0, output_type='tensor')
self.pitch_up = PitchShift(min_transpose_semitones=1, max_transpose_semitones=4, p=1.0, sample_rate=self.sample_rate, output_type='tensor')
self.pitch_down = PitchShift(min_transpose_semitones=-4, max_transpose_semitones=-1, p=1.0, sample_rate=self.sample_rate, output_type='tensor')
self.lowpass = LowPassFilter(min_cutoff_freq=100, max_cutoff_freq=2000, p=1.0, output_type='tensor')
self.highpass = HighPassFilter(min_cutoff_freq=1000, max_cutoff_freq=2000, p=1.0, output_type='tensor')
self.apply_ir = ApplyImpulseResponse(ir_paths=corridor_noise_dir, convolve_mode='same', p=1, output_type="tensor")
self.polarity_inversion = PolarityInversion(p=1.0, output_type="tensor")
def __len__(self):
return len(self.data)
def _load_audio(self, audio_path: Path) -> Tensor:
waveform, sample_rate = torchaudio.load_with_torchcodec(audio_path)
# Resample to target sample rate.
if sample_rate != self.sample_rate:
waveform = Resample(sample_rate, self.sample_rate)(waveform)
# Convert to mono if it is setro.
if waveform.shape[0] > 1:
waveform = waveform.mean(dim=0, keepdim=True)
# Normalization
waveform = waveform / waveform.abs().max()
# Clip waveform exceeds from max length.
if waveform.shape[1] > self.max_audio_len:
waveform = waveform[:, :self.max_audio_len]
return waveform
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.6:
waveform = self._stretch_or_compress(waveform=waveform)
if random.random() < 0.7:
waveform = self.noise_augmentor.apply_real_noise(waveform)
if random.random() < 0.3:
waveform = self._time_mask_waveform(waveform=waveform)
# torch_audiomentations: [1, time] -> [1, 1, time]
if waveform.dim() == 2:
waveform_3d = waveform.unsqueeze(0)
# 随机选择一种物理特性增强 (互斥区)
choice = random.random()
if choice < 0.25: # [0.00 - 0.25] 25% 概率:增益
if random.random() < 0.5:
waveform_3d = self.gain_up(waveform_3d, sample_rate=self.sample_rate)
else:
waveform_3d = self.gain_down(waveform_3d, sample_rate=self.sample_rate)
elif choice < 0.50: # [0.25 - 0.50] 25% 概率:音高
if random.random() < 0.5:
waveform_3d = self.pitch_up(waveform_3d, sample_rate=self.sample_rate)
else:
waveform_3d = self.pitch_down(waveform_3d, sample_rate=self.sample_rate)
elif choice < 0.70: # [0.50 - 0.70] 20% 概率:低通
waveform_3d = self.lowpass(waveform_3d, sample_rate=self.sample_rate)
elif choice < 0.85: # [0.70 - 0.85] 15% 概率:高通
waveform_3d = self.highpass(waveform_3d, sample_rate=self.sample_rate)
elif choice < 0.95: # [0.85 - 0.95] 10% 概率:走廊混响 (IR)
# 使用你测试过最好的 0.8/0.2 比例
dry = waveform_3d.clone()
wet = self.apply_ir(waveform_3d, sample_rate=self.sample_rate)
waveform_3d = 0.8 * dry + 0.2 * wet
else: # [0.95 - 1.00] 5% 概率:极性翻转
waveform_3d = self.polarity_inversion(waveform_3d, sample_rate=self.sample_rate)
# [1, 1, time] -> [1, time]
waveform = waveform_3d.squeeze(0)
# 防止多次 augment 后振幅溢出,最后归一化
max_amp = waveform.abs().max()
if max_amp > 1.0:
waveform = waveform / max_amp
return waveform
def _stretch_or_compress(self, waveform: Tensor) -> Tensor:
speed_factor = random.uniform(0.80, 1.6) # (Speed Change: 0.85x - 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
)
# 时间拉伸(不改变音高)
stretched_spec = TimeStretch(hop_length=160, n_freq=spec.shape[-2], fixed_rate=speed_factor)(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 _time_mask_waveform(self, waveform: Tensor) -> Tensor:
audio_len = waveform.shape[1]
sr = self.sample_rate # 16000
# 设置参数:单次遮盖最长 0.4 秒 (6400个点)
max_mask_time = 0.4
max_mask_samples = int(sr * max_mask_time)
# 根据音频长度决定遮盖次数:
# 比如每 3 秒钟允许遮盖 1 次
num_masks = max(1, audio_len // (sr * 3))
for _ in range(num_masks):
# 每次随机遮盖 0.1s 到 0.4s
current_mask_len = random.randint(int(sr * 0.1), max_mask_samples)
if audio_len > current_mask_len:
start_pos = random.randint(0, audio_len - current_mask_len)
# 填充微小噪音(模拟环境底噪)
noise = torch.randn(1, current_mask_len).to(waveform.device) * 0.002
waveform[:, start_pos : start_pos + current_mask_len] = noise
return waveform
def __getitem__(self, index) -> BatchItem:
row: TsvFormat = self.data.iloc[index]
audio_path: Path = self.audio_dir / row['path']
text: str = normalize_extended_uyghur_characters(collapse_spaces(row['sentence'].strip()))
waveform = self._load_audio(audio_path=audio_path)
waveform = self._augment_waveform(waveform)
waveform = waveform.squeeze(0)
return BatchItem(
waveform=waveform,
target_ids=torch.tensor(self.tokenizer.encode(text), dtype=torch.long),
target_text=text,
audio_path=str(audio_path)
)
def collate_fn(items: List[BatchItem]) -> Batch:
max_waveform_len = max(item['waveform'].shape[0] for item in items)
max_target_len = max(len(item['target_ids']) for item in items)
batch_size = len(items)
waveforms = torch.zeros(batch_size, max_waveform_len)
targets = torch.zeros(batch_size, max_target_len, dtype=torch.long)
waveform_lengths = torch.zeros(batch_size, dtype=torch.long)
target_lengths = torch.zeros(batch_size, dtype=torch.long)
target_texts = []
audio_paths = []
for i, item in enumerate(items):
waveform_len = item['waveform'].shape[0]
target_len = len(item['target_ids'])
waveforms[i, :waveform_len] = item['waveform']
targets[i, :target_len] = item['target_ids']
waveform_lengths[i] = waveform_len
target_lengths[i] = target_len
target_texts.append(item['target_text'])
audio_paths.append(item['audio_path'])
return Batch(
waveforms=waveforms,
targets=targets,
waveform_lengths=waveform_lengths,
target_lengths=target_lengths,
target_texts=target_texts,
audio_paths=audio_paths
)
def create_dataloader(tsv_path: Path, audio_dir: Path, noise_dir: Path, corridor_noise_dir: Path, tokenizer: ASRTokenizer, batch_size: int = 8, shuffle: bool = True, augment: bool = True, augment_prob: int = 0.5) -> DataLoader:
dataset = CommonVoiceDataset(tsv_path=tsv_path, audio_dir=audio_dir, noise_dir=noise_dir, corridor_noise_dir=corridor_noise_dir, tokenizer=tokenizer, augment=augment, augment_prob=augment_prob)
return DataLoader(dataset=dataset, batch_size=batch_size, shuffle=shuffle, collate_fn=collate_fn, pin_memory=True, num_workers=8, prefetch_factor=8, persistent_workers=True)
# ============ 测试代码 ============
if __name__ == "__main__":
from pathlib import Path
workspace_dir = Path(__file__).parent.parent
# 初始化tokenizer
tokenizer = ASRTokenizer(workspace_dir / 'config' / 'asr_vocab.json')
# 创建数据加载器
dataloader = create_dataloader(
tsv_path=workspace_dir / 'data' / 'ug' / 'train.tsv',
audio_dir=workspace_dir / 'data' / 'ug' / 'clips',
tokenizer=tokenizer,
batch_size=2,
shuffle=True,
)
# 测试加载一个batch
print("测试数据加载:")
for batch in dataloader:
batch: Batch
print(f"Mel specs shape: {batch['mel_specs'].shape}")
print(f"Targets shape: {batch['targets'].shape}")
print(f"Mel lengths: {batch['mel_lengths']}")
print(f"Target lengths: {batch['target_lengths']}")
print(f"Target texts: {batch['target_texts']}")
print(f"Audio paths: {batch['audio_paths']}")
break