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