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

如何并行训练8个Keras集成模型并实现闭环训练流程?

多Keras模型闭环训练优化方案

方案一:多进程并行训练+主进程管控闭环流程

用Python multiprocessing 模块创建8个独立进程,每个进程绑定一块GPU、加载专属模型并执行训练,主进程负责监控训练完成状态、更新数据集、重启训练流程,全程避免繁琐的多脚本调用,进程间通过队列通信,数据集用内存映射文件优化读写效率。

实现步骤

  1. GPU绑定逻辑:每个子进程启动时,通过tf.config.set_visible_devices指定专属GPU,避免设备冲突。
  2. 进程通信:主进程创建一个队列,每个子进程训练完成后向队列发送完成信号,主进程等待所有信号接收完毕后触发数据集更新。
  3. 闭环管控:主进程循环执行「启动子进程→等待全部完成→更新数据集→重启子进程」流程。

代码示例

主进程代码

import multiprocessing as mp
import tensorflow as tf
from mmap import mmap

def worker_process(gpu_id, model_config, dataset_path, queue):
    # 绑定专属GPU
    gpus = tf.config.list_physical_devices('GPU')
    tf.config.set_visible_devices(gpus[gpu_id], 'GPU')
    tf.config.experimental.set_memory_growth(gpus[gpu_id], True)

    # 加载对应模型(根据配置生成不同结构的Keras模型)
    model = create_model(model_config)

    # 从内存映射文件加载数据集
    with open(dataset_path, 'r+b') as f:
        mm = mmap(f.fileno(), 0)
        dataset = load_dataset_from_mmap(mm)

    # 执行训练
    model.fit(dataset, epochs=10)

    # 向主进程发送完成信号
    queue.put(gpu_id)

def update_dataset(dataset_path):
    # 增删训练数据,直接操作内存映射文件
    with open(dataset_path, 'r+b') as f:
        mm = mmap(f.fileno(), 0)
        modify_dataset_in_mmap(mm)

if __name__ == '__main__':
    # 初始化内存映射形式的数据集
    init_dataset_to_mmap('shared_dataset.dat')

    # 8个模型的差异化配置
    model_configs = [get_model_config(i) for i in range(8)]

    while True:
        queue = mp.Queue()
        processes = []
        # 启动8个训练进程
        for gpu_id in range(8):
            p = mp.Process(
                target=worker_process,
                args=(gpu_id, model_configs[gpu_id], 'shared_dataset.dat', queue)
            )
            processes.append(p)
            p.start()

        # 等待所有进程完成
        completed_count = 0
        while completed_count < 8:
            queue.get()
            completed_count += 1

        # 回收子进程资源
        for p in processes:
            p.join()

        # 更新数据集并启动下一轮训练
        update_dataset('shared_dataset.dat')
        print("数据集已更新,启动下一轮训练...")

核心辅助函数说明

  • create_model(config):根据传入的配置生成不同结构的Keras模型(比如调整层数、激活函数、 dropout率等)。
  • load_dataset_from_mmap(mm):从内存映射对象中解析数据,构建tf.data.Dataset,避免磁盘IO开销。
  • modify_dataset_in_mmap(mm):直接在内存映射中增删样本,修改后自动同步到磁盘,无需额外读写操作。

方案二:subprocess精细化管控独立训练脚本

如果更倾向于保留独立训练脚本,用subprocess.Popen替代os.system,主进程可实时监控子进程状态,数据集同样用内存映射优化。

实现示例

主进程管控代码

import subprocess

def run_training_script(gpu_id):
    # 通过环境变量指定GPU
    env = {'CUDA_VISIBLE_DEVICES': str(gpu_id)}
    return subprocess.Popen(
        ['python', 'training_script.py', '--gpu-id', str(gpu_id)],
        env=env,
        stdout=subprocess.PIPE,
        stderr=subprocess.PIPE
    )

if __name__ == '__main__':
    while True:
        # 启动8个训练进程
        processes = [run_training_script(i) for i in range(8)]
        
        # 等待所有进程完成
        for idx, p in enumerate(processes):
            p.wait()
            stdout, stderr = p.communicate()
            print(f"GPU {idx} 训练完成,输出:{stdout.decode()}")
        
        # 更新数据集
        update_dataset('shared_dataset.dat')
        print("数据集已更新,启动下一轮训练...")

独立训练脚本(training_script.py)核心逻辑

import tensorflow as tf
import sys

gpu_id = int(sys.argv[1])
gpus = tf.config.list_physical_devices('GPU')
tf.config.set_visible_devices(gpus[gpu_id], 'GPU')

# 加载对应模型、从内存映射文件读取数据集、执行训练
model = create_model_by_id(gpu_id)
with open('shared_dataset.dat', 'r+b') as f:
    mm = mmap(f.fileno(), 0)
    dataset = load_dataset_from_mmap(mm)
model.fit(dataset, epochs=10)

关键优化点

  • GPU隔离:每个进程通过tf.config.set_visible_devices或CUDA_VISIBLE_DEVICES绑定专属GPU,避免资源竞争。
  • 数据集提速:用内存映射文件替代普通磁盘文件,数据增删操作直接在内存完成,同步磁盘开销极低。
  • 闭环自动化:主进程全程管控训练启停,无需人工干预,训练完成后自动触发数据集更新与下一轮训练。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 19:45:01