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