求fit_generator()具体示例:语音识别模型内存优化与生成器问题
优化语音识别模型的数据生成器,避免全量加载内存
问题拆解
你当前的代码核心问题是提前把所有音频数据加载到了X_train和Y_train列表中,完全浪费了生成器“按需加载、缓解内存压力”的设计初衷,不仅导致内存占用爆炸,预处理阶段的耗时也会拉得很长。而且生成器里的逻辑还有bug——比如create_X返回的是生成器对象,你直接把它append到列表里,后续根本没法正常处理成模型需要的输入数据,这也是程序长时间没输出的原因之一。
解决方案:真正实现按需加载的生成器
我们要彻底改造数据加载逻辑,让生成器每次只读取并处理一个batch的音频数据,而不是先把5GB的全量数据塞进内存。下面分步骤给出优化方案:
1. 先收集文件路径和标签(只存字符串,内存占用极小)
先把所有有效音频文件的路径和对应标签整理好,这一步只存文本信息,几乎不占内存:
import librosa import glob import tensorflow as tf import os import numpy as np from sklearn.preprocessing import OneHotEncoder # 配置路径和类别 audio_dir = "D:\\SpeechRecognitionData\\train\\audio\\" # 过滤掉背景噪声类别 class_names = [name for name in os.listdir(audio_dir) if name != '_background_noise_'] # 初始化OneHot编码器 encoder = OneHotEncoder(sparse_output=False) encoder.fit(np.array(class_names).reshape(-1, 1)) # 收集所有音频文件路径和对应标签 file_paths = [] labels = [] for class_name in class_names: class_dir = os.path.join(audio_dir, class_name) # 遍历当前类别下的所有wav文件 for wav_file in glob.glob(os.path.join(class_dir, "*.wav")): file_paths.append(wav_file) labels.append(class_name) # 打乱数据(训练前打乱很重要,避免模型学到顺序规律) indices = np.random.permutation(len(file_paths)) file_paths = np.array(file_paths)[indices] labels = np.array(labels)[indices]
2. 自定义Batch级别的生成器函数
这个生成器会循环读取每个batch的文件,实时加载音频、调整形状并生成OneHot标签,完全不需要提前把所有数据读入内存:
def audio_generator(file_paths, labels, encoder, batch_size=36, target_sr=22050): num_samples = len(file_paths) # 生成器需要无限循环,直到训练结束 while True: for start_idx in range(0, num_samples, batch_size): end_idx = min(start_idx + batch_size, num_samples) # 取出当前batch的文件路径和标签 batch_files = file_paths[start_idx:end_idx] batch_label_names = labels[start_idx:end_idx] # 加载并处理当前batch的音频数据 batch_x = [] for file in batch_files: # 加载音频,强制采样率为22050 wave, sr = librosa.load(file, sr=target_sr) # 统一音频长度:不够补零,过长截断 if len(wave) < target_sr: wave = np.pad(wave, (0, target_sr - len(wave)), mode='constant') elif len(wave) > target_sr: wave = wave[:target_sr] # 调整为模型需要的(22050, 1)形状 batch_x.append(wave.reshape(-1, 1)) # 转换标签为OneHot编码 batch_y = encoder.transform(np.array(batch_label_names).reshape(-1, 1)) yield np.array(batch_x), batch_y
3. 模型训练逻辑调整
现在用新生成器训练,内存占用会大幅降低,而且能实时看到训练日志:
input_shape = (22050, 1) model = tf.keras.models.Sequential([ tf.keras.layers.Conv1D(16, activation='relu', input_shape=input_shape, kernel_size=10), tf.keras.layers.MaxPool1D(), tf.keras.layers.Dropout(0.2), tf.keras.layers.Conv1D(32, activation='relu', kernel_size=10), tf.keras.layers.MaxPool1D(), tf.keras.layers.Dropout(0.2), tf.keras.layers.Conv1D(16, activation='relu', kernel_size=10), tf.keras.layers.MaxPool1D(), tf.keras.layers.Dropout(0.2), tf.keras.layers.Flatten(), tf.keras.layers.Dense(128, activation='relu'), tf.keras.layers.Dense(64, activation='relu'), # 这里用类别总数替代固定30,更灵活 tf.keras.layers.Dense(len(class_names), activation='softmax') ]) model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy']) # 初始化生成器 batch_size = 36 steps_per_epoch = len(file_paths) // batch_size train_generator = audio_generator(file_paths, labels, encoder, batch_size=batch_size) # 开始训练,verbose=1会实时输出日志,epochs指定训练轮数(避免无限训练) model.fit(train_generator, steps_per_epoch=steps_per_epoch, epochs=10, shuffle=True, verbose=1) model.save("model.h5")
额外优化建议
- 预处理缓存:如果需要重复训练,可以提前把处理好的
(22050,1)音频数据保存成小体积的npy文件,生成器直接读取这些文件,比每次加载wav格式更快。 - 切换到tf.data.Dataset:现在Keras更推荐用
tf.data.Dataset处理数据,它支持多线程并行加载,性能比自定义生成器更好,示例代码如下:
def load_audio(file_path, label): # 用TensorFlow原生API读取音频,更适配TF生态 audio = tf.io.read_file(file_path) audio, sr = tf.audio.decode_wav(audio, desired_channels=1) audio = tf.squeeze(audio, axis=-1) # 统一音频长度 audio = tf.cond(tf.shape(audio)[0] < 22050, lambda: tf.pad(audio, [[0, 22050 - tf.shape(audio)[0]]]), lambda: audio[:22050]) audio = tf.reshape(audio, (22050, 1)) # 转换标签为OneHot label_str = label.numpy().decode('utf-8') label_onehot = encoder.transform([[label_str]])[0] return audio, label_onehot # 创建Dataset并配置并行加载 dataset = tf.data.Dataset.from_tensor_slices((file_paths, labels)) dataset = dataset.map( lambda x, y: tf.py_function(load_audio, [x, y], [tf.float32, tf.float32]), num_parallel_calls=tf.data.AUTOTUNE ) dataset = dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE) # 训练 model.fit(dataset, epochs=10, verbose=1)
内容的提问来源于stack exchange,提问作者Darshan Deshpande
相关产品推荐
相关产品推荐

