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

求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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 09:17:12