336 lines
13 KiB
Python
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 |