| 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 |
|
|
| if __name__ == "__main__": |
| test_data_pipeline() |
| |
|
|