from core.data.audio_dataset import * import torch from torch.nn.utils.rnn import pad_sequence from torch.utils.data import Dataset, DataLoader, Sampler from datasets import load_dataset import torchaudio import random ### 主测试函数 def test_data_pipeline(): print("Loading dataset...") ds = load_dataset("hf-internal-testing/librispeech_asr_dummy", split="validation").with_format("python") dataset = AudioDataset(ds) sampler = DynamicBatchSampler(dataset, max_total_length=48000) collator = AudioCollator() dataloader = DataLoader(dataset, batch_sampler=sampler, collate_fn=collator) for i, batch in enumerate(dataloader): print(f"\nBatch {i}") assert isinstance(batch['audio'], torch.Tensor) assert isinstance(batch['text'], torch.Tensor) B, T_audio = batch['audio'].shape B2, T_text = batch['text'].shape assert B == B2, "Batch size mismatch between audio and text" assert T_audio == batch['audio_lengths'].max().item(), "Audio padding mismatch" assert T_text == batch['text_lengths'].max().item(), "Text padding mismatch" print(f" audio: {batch['audio'].shape}, text: {batch['text'].shape}") if i >= 10: break # 只测试前几个 batch if __name__ == "__main__": test_data_pipeline() # PYTHONPATH=./ python test/dataset_test.py