使用Torchaudio加载CommonVoice数据集时出现张量尺寸不匹配错误
解决Torchaudio加载CommonVoice数据集时张量尺寸不一致的问题
问题描述
使用Torchaudio加载CommonVoice德语数据集,通过DataLoader批量遍历数据时,触发错误:
RuntimeError: stack expects each tensor to be equal size, but got [150822] at entry 0 and [127008] at entry 1
原因是CommonVoice数据集中的音频时长各不相同,原始数据集返回的音频张量长度不一致,而DataLoader默认会尝试将同批次张量堆叠成大张量,要求所有张量维度完全一致,因此报错。
解决方法
方法1:自定义collate_fn处理批量数据
编写自定义批量处理函数,对每个批次的音频进行填充(补0)到当前批次最长音频的长度,确保同批次张量维度一致。
示例代码:
import torch from torchaudio.datasets import COMMONVOICE from torch.utils.data import DataLoader def collate_fn(batch): # 拆分批次中的音频和标签 audios, labels = zip(*batch) # 找到当前批次最长音频的长度 max_len = max(audio.shape[0] for audio in audios) # 对每个音频补0至max_len长度 padded_audios = [] for audio in audios: pad_len = max_len - audio.shape[0] padded = torch.nn.functional.pad(audio, (0, pad_len)) padded_audios.append(padded) # 堆叠成批量张量,转换标签为张量 return torch.stack(padded_audios), torch.tensor(labels) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") train_dataset = COMMONVOICE(root='/home/mr/Downloads/cv-corpus-7.0-2021-07-21/de/', tsv='train.tsv') # 传入自定义collate_fn train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=0, collate_fn=collate_fn) print(device) for inputs, targets in train_loader: print(inputs.shape, targets)
方法2:自定义数据集类统一音频长度
继承COMMONVOICE类,在加载每个样本时直接将音频裁剪或填充到固定长度,从源头上保证张量维度一致。
示例代码:
import torch from torchaudio.datasets import COMMONVOICE from torch.utils.data import DataLoader class FixedLengthCommonVoice(COMMONVOICE): def __init__(self, root, tsv, fixed_len=48000): super().__init__(root, tsv=tsv) self.fixed_len = fixed_len # 假设采样率16kHz,对应3秒音频 def __getitem__(self, idx): audio, _, label, _, _ = super().__getitem__(idx) # 裁剪过长音频,补0过短音频 if audio.shape[0] > self.fixed_len: audio = audio[:self.fixed_len] else: pad_len = self.fixed_len - audio.shape[0] audio = torch.nn.functional.pad(audio, (0, pad_len)) return audio, label device = torch.device("cuda" if torch.cuda.is_available() else "cpu") train_dataset = FixedLengthCommonVoice(root='/home/mr/Downloads/cv-corpus-7.0-2021-07-21/de/', tsv='train.tsv') train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=0) print(device) for inputs, targets in train_loader: print(inputs.shape, targets)
注意事项
- 填充操作建议在音频末尾补0,避免破坏音频开头的有效内容;裁剪过长音频时,可根据任务需求选择随机裁剪或固定从开头裁剪。
- 若用于语音识别等任务,需结合标签处理逻辑(如CTC算法无需标签与音频严格对齐)调整代码。
内容的提问来源于stack exchange,提问作者BR BR
相关产品推荐
相关产品推荐

