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

求助:基于Keras实现Wav文件批量分类的DataGenerator(含自定义频谱图函数)

实现自定义音频数据生成器(类似flow_from_directory)处理大规模WAV文件

完全理解你的痛点——上万条音频文件直接加载到内存里肯定撑不住,用Keras的Sequence类来实现自定义数据生成器是最优解,它能像ImageDataGenerator.flow_from_directory那样按目录结构自动读取数据,还能在批量加载时调用你的自定义频谱图函数,完全不用一次性把所有数据塞进内存。

第一步:先写你的自定义频谱图生成函数

这里给个示例用librosa生成梅尔频谱图,你可以直接替换成自己的spectrogram逻辑:

import librosa
import numpy as np

def custom_spectrogram(file_path, sr=16000, n_mels=128, n_fft=2048, hop_length=512):
    # 加载音频文件
    y, sr = librosa.load(file_path, sr=sr)
    # 生成梅尔频谱图
    mel_spec = librosa.feature.melspectrogram(y=y, sr=sr, n_mels=n_mels, n_fft=n_fft, hop_length=hop_length)
    # 转换成对数刻度(更符合人耳听觉)
    log_mel_spec = librosa.power_to_db(mel_spec, ref=np.max)
    # 增加通道维度(适配CNN输入,比如(128, T, 1))
    log_mel_spec = log_mel_spec[..., np.newaxis]
    return log_mel_spec

这个函数输入WAV文件路径,输出预处理好的频谱图张量,你可以根据自己的需求调整参数或者替换成其他频谱生成方式(比如STFT)。

第二步:实现自定义DataGenerator类

继承Keras的Sequence类,它会自动处理批量加载、多线程,还能保证每个epoch的 shuffle 是安全的:

from tensorflow.keras.utils import Sequence
import os
from sklearn.preprocessing import LabelEncoder
import numpy as np

class AudioDataGenerator(Sequence):
    def __init__(self, directory, batch_size=32, sr=16000, n_mels=128, shuffle=True, spectrogram_func=custom_spectrogram):
        # 初始化参数
        self.directory = directory
        self.batch_size = batch_size
        self.sr = sr
        self.n_mels = n_mels
        self.shuffle = shuffle
        self.spectrogram_func = spectrogram_func
        
        # 获取所有文件路径和对应的标签
        self.file_paths, self.labels = self._get_file_paths_and_labels()
        
        # 标签编码(把字符串标签转成数字)
        self.label_encoder = LabelEncoder()
        self.encoded_labels = self.label_encoder.fit_transform(self.labels)
        
        # 初始化时打乱数据(如果需要)
        self.on_epoch_end()
    
    def _get_file_paths_and_labels(self):
        # 遍历目录结构,每个子目录对应一个类别
        file_paths = []
        labels = []
        for label_dir in os.listdir(self.directory):
            label_dir_path = os.path.join(self.directory, label_dir)
            if os.path.isdir(label_dir_path):
                for file_name in os.listdir(label_dir_path):
                    if file_name.endswith('.wav'):
                        file_paths.append(os.path.join(label_dir_path, file_name))
                        labels.append(label_dir)
        return file_paths, labels
    
    def __len__(self):
        # 计算每个epoch有多少个batch
        return int(np.ceil(len(self.file_paths) / self.batch_size))
    
    def __getitem__(self, index):
        # 获取当前batch的文件路径和标签
        batch_paths = self.file_paths[index*self.batch_size : (index+1)*self.batch_size]
        batch_labels = self.encoded_labels[index*self.batch_size : (index+1)*self.batch_size]
        
        # 加载并预处理当前batch的音频
        batch_X = []
        for path in batch_paths:
            spec = self.spectrogram_func(path, sr=self.sr, n_mels=self.n_mels)
            batch_X.append(spec)
        
        # 转换成numpy数组,保证形状一致(如果你的频谱图长度不一,这里需要做padding或者截断)
        batch_X = np.array(batch_X)
        # 标签转成one-hot(如果用categorical_crossentropy的话)
        batch_y = np.eye(len(self.label_encoder.classes_))[batch_labels]
        
        return batch_X, batch_y
    
    def on_epoch_end(self):
        # 每个epoch结束后打乱数据顺序
        if self.shuffle:
            indices = np.arange(len(self.file_paths))
            np.random.shuffle(indices)
            self.file_paths = [self.file_paths[i] for i in indices]
            self.encoded_labels = self.encoded_labels[indices]

关键细节说明:

  • 目录结构要求:和flow_from_directory一样,你的音频文件需要按类别放在子目录里,比如:
    audio_data/
        class_0/
            file1.wav
            file2.wav
            ...
        class_1/
            fileA.wav
            fileB.wav
            ...
        ...
    
  • 频谱图形状统一:如果你的音频时长不一样,生成的频谱图宽度(时间轴)会有差异,这时候需要在__getitem__里做padding或者截断,保证每个batch的输入形状一致。比如可以固定截取前N帧,或者补零到最长的长度。
  • 多线程支持:在训练时,model.fit里可以设置workers参数(比如workers=4),生成器会自动开启多线程加载数据,提升效率。

第三步:使用生成器进行训练

from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Conv2D, MaxPooling2D, Flatten, Dense

# 初始化生成器
train_generator = AudioDataGenerator(
    directory='path/to/train/audio_data',
    batch_size=32,
    sr=16000,
    n_mels=128,
    shuffle=True
)

# 构建一个简单的CNN模型(示例)
model = Sequential([
    Conv2D(32, (3,3), activation='relu', input_shape=(128, None, 1)),
    MaxPooling2D((2,2)),
    Conv2D(64, (3,3), activation='relu'),
    MaxPooling2D((2,2)),
    Flatten(),
    Dense(128, activation='relu'),
    Dense(len(train_generator.label_encoder.classes_), activation='softmax')
])

model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])

# 训练模型
model.fit(
    train_generator,
    epochs=10,
    workers=4,  # 开启多线程加载
    use_multiprocessing=True  # 如果workers>1,建议开启这个
)

额外优化建议

  • 缓存预处理结果:如果你的频谱图生成比较耗时,可以考虑把预处理好的频谱图保存成numpy文件,下次训练直接加载numpy数组,速度会快很多。
  • 数据增强:可以在生成器里加入音频增强逻辑(比如加噪声、变速、变调),在生成频谱图之前对音频进行处理,提升模型泛化能力。
  • 动态调整batch大小:如果某些batch的音频处理后内存占用过高,可以适当减小batch_size。

内容的提问来源于stack exchange,提问作者xerxes01

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.12 04:34:38