You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

微调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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.29 03:57:26