EchoLoc / rendering /universal_tts /test /dataset_test.py
zsy814's picture
Initial EchoLoc code release
d8bfe4a verified
Raw
History Blame Contribute Delete
1.37 kB
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