微调Whisper模型替换自定义数据集遇内存过高问题求助
自定义韩语数据集微调Whisper时内存占用过高、预处理耗时久的问题解决
你的核心问题是提前将所有音频解码为数组并加载到内存中,完全违背了Hugging Face Dataset的懒加载设计。而Common Voice的Arrow文件之所以高效,是因为它仅存储音频路径、元数据,音频解码是在预处理阶段按需执行并缓存的,不会一次性把所有音频数据塞进内存。
你的DataLoader实现存在两个关键问题:
- 在
getData方法中,直接把每个音频文件解码成float32数组并全部存入data_list,转成Dataset后,大数量级的音频数组会直接占满内存。 - 返回的Dataset包含完整音频数组,后续执行
map时,每个进程还要对这些数组做进一步处理,内存占用叠加导致运行缓慢甚至崩溃。
优化步骤1:创建仅包含路径和元数据的Dataset
修改DataLoader,只收集音频路径、文本标签和其他元数据,不提前解码音频:
class DataLoader_AIHub: def __init__(self, rootPath): self.rootPath = rootPath def getData(self, max_files_to_load, startPoint=0): rootPath_audio = os.path.join(self.rootPath, 'audio') audioDirPaths = getDirList(rootPath_audio) total_files_loaded = 0 data_list = [] for audioDir in audioDirPaths: audioFileNames = getFileList(audioDir) audioFilePaths = [os.path.join(audioDir, item) for item in audioFileNames] labelFilePaths = [item.replace('/audio/', '/label/').replace('.wav', '.json') for item in audioFilePaths] for audioPath, labelPath in zip(audioFilePaths, labelFilePaths): jsonInfo = getJson(labelPath) stt_text = jsonInfo['발화정보']['stt'] # 过滤带括号的文本 if '(' in stt_text: continue if startPoint > total_files_loaded: total_files_loaded += 1 continue # 仅存储路径和元数据,不解码音频 data_dict = { 'audio_path': audioPath, 'sentence': re.sub('\r\n', '', stt_text), 'age': jsonInfo['녹음자정보']['age'], 'gender': jsonInfo['녹음자정보']['gender'] } data_list.append(data_dict) total_files_loaded += 1 if total_files_loaded >= max_files_to_load + startPoint: return Dataset.from_list(data_list) return Dataset.from_list(data_list)
优化步骤2:在prepare_dataset中按需解码音频
修改prepare_dataset函数,在预处理阶段才解码音频,利用Dataset的懒加载和自动缓存机制:
import soundfile as sf import numpy as np import librosa def prepare_dataset(batch): # 按需解码音频 audio, sr = sf.read(batch['audio_path']) # 统一转换为Whisper要求的16kHz采样率 if sr != 16000: audio = librosa.resample(audio, orig_sr=sr, target_sr=16000) batch['audio'] = { 'array': audio.astype(np.float32), 'sampling_rate': 16000 } # 清理文本(根据韩语需求调整) batch['sentence'] = batch['sentence'].strip() return batch
优化步骤3:保存为Arrow格式复用
处理完数据集后,保存为Arrow格式,后续直接加载即可获得和Common Voice一样的高效体验:
# 加载自定义数据集 dataset = DataLoader_AIHub("your_root_path").getData(max_files_to_load=10000) # 执行预处理(按需解码音频,自动缓存) dataset = dataset.map(prepare_dataset, num_proc=4) # 保存到磁盘 dataset.save_to_disk("korean_custom_dataset") # 后续加载直接调用 from datasets import load_from_disk dataset = load_from_disk("korean_custom_dataset")
额外优化建议
- 根据CPU核心数调整
num_proc数值,加快预处理速度,同时避免内存过载。 - 若数据集过大,可分批次处理,避免一次性处理过多数据。
- 提前统一音频格式为16kHz单声道,减少预处理阶段的格式转换开销。
内容的提问来源于stack exchange,提问作者C yp
相关产品推荐
相关产品推荐

