File size: 1,373 Bytes
d8bfe4a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 | 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
|