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

多进程在多训练集上训练Keras模型是否可行?附环境及报错

解决TensorFlow 1.12 + Keras多进程训练的序列化与Eager模式问题

我来帮你拆解下遇到的两个核心问题,再给出适配TensorFlow 1.12版本的可行解决方案:

第一个报错:Eager模式下的Optimizer类型错误

当你添加tf.enable_eager_execution()后出现ValueError: optimizer must be an instance of tf.train.Optimizer, not a <class 'str'>,原因很明确:TensorFlow 1.x的Eager模式下,Keras的compile()方法不支持用字符串指定优化器。Graph模式下Keras可以自动解析'adam'这类字符串,但Eager模式要求你传入tf.train模块下的优化器实例。

比如把optimizer='adam'改成optimizer=tf.train.AdamOptimizer()就能解决这个类型错误,但这只是第一步,Eager模式和多进程配合还会有其他潜在问题,所以这个方案并非最优选择。

第二个报错:模型序列化失败(NotImplementedError: numpy() is only available when eager execution is enabled)

这个报错的根源是:TensorFlow 1.x的Graph模式下,Keras模型绑定了TensorFlow计算图的引用,而multiprocessing的Pool在传递进程间结果时需要用pickle序列化对象,但TF1的计算图和模型对象无法被正确序列化。当你尝试返回完整的Sequential模型时,序列化过程中会尝试访问Tensor的numpy()方法——而Graph模式下的Tensor并没有这个方法(仅Eager模式支持),因此抛出了这个错误。

针对TF1.12的最优解决方案

考虑到你使用的是TF1.12这类较老版本,最稳妥的方式是避免跨进程传递完整模型,转而返回模型权重数组,或者将模型保存到本地文件后在主进程中加载。这样就能绕开序列化模型对象的坑。

下面是修改后的可运行代码示例:

from multiprocessing import Pool
import numpy as np
import os

def fit_model(dataset, save_path=None):
    # 每个进程独立导入TF和Keras模块,避免进程间计算图冲突
    from tensorflow.python.keras.models import Sequential
    from tensorflow.python.keras.layers import Dense
    
    # 修正数据集形状:单行数据转为(样本数, 特征数)格式
    x = dataset[:, 0:3].reshape(-1, 3)
    y = dataset[:, 3].reshape(-1, 1)
    
    model = Sequential()
    model.add(Dense(3, input_dim=3, activation='relu'))
    model.add(Dense(3, activation='relu'))
    model.add(Dense(1, activation='linear'))
    model.compile(loss='mean_squared_error', optimizer='adam', metrics=['accuracy'])
    # 关闭训练日志,避免多进程输出混乱
    model.fit(x, y, epochs=3, verbose=0)
    
    if save_path:
        # 方案1:保存模型到文件,返回路径
        model.save(save_path)
        return save_path
    else:
        # 方案2:返回模型权重,主进程重建模型
        return model.get_weights()

if __name__ == "__main__":
    # 加载数据集并修正形状
    dataset1 = np.loadtxt('ts1.txt').reshape(1, -1)
    dataset2 = np.loadtxt('ts2.txt').reshape(1, -1)
    dataset3 = np.loadtxt('ts3.txt').reshape(1, -1)
    datasets = [dataset1, dataset2, dataset3]
    
    # ------------------- 方案1:返回权重,主进程重建模型 -------------------
    with Pool() as p:
        weights_list = p.map(fit_model, datasets)
    
    # 主进程中重建模型并加载权重
    from tensorflow.python.keras.models import Sequential
    from tensorflow.python.keras.layers import Dense
    
    def build_base_model():
        model = Sequential()
        model.add(Dense(3, input_dim=3, activation='relu'))
        model.add(Dense(3, activation='relu'))
        model.add(Dense(1, activation='linear'))
        model.compile(loss='mean_squared_error', optimizer='adam', metrics=['accuracy'])
        return model
    
    trained_models = []
    for weights in weights_list:
        model = build_base_model()
        model.set_weights(weights)
        trained_models.append(model)
    
    # ------------------- 方案2:保存模型到文件,主进程加载 -------------------
    # 创建保存目录
    os.makedirs('trained_models', exist_ok=True)
    save_paths = [f'trained_models/model_{i}.h5' for i in range(3)]
    # 用starmap传递多个参数给fit_model
    with Pool() as p:
        saved_paths = p.starmap(fit_model, zip(datasets, save_paths))
    
    # 主进程加载模型
    from tensorflow.python.keras.models import load_model
    models_from_files = [load_model(path) for path in saved_paths]

额外注意事项

  1. 数据集形状修正:你的训练集是单行数据,必须用reshape(1, -1)转换成(样本数, 特征数)的格式,否则Keras会报输入形状不匹配的错误。
  2. 进程间计算图隔离:每个进程独立导入TF和Keras模块是正确的做法,能避免多个进程共享同一个计算图导致的冲突。
  3. 训练日志控制:多进程训练时建议关闭verbose(设为0),否则多个进程的日志会混杂在一起,难以阅读。
  4. TF版本建议:TF1.x的多进程支持本身就有诸多局限性,如果后续有机会升级到TF2.x,体验会好很多——TF2默认Eager模式,且模型序列化更友好。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 09:14:00