求助:基于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
相关产品推荐
相关产品推荐

