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

分布式训练模型保存异常:MultiWorkerMirroredStrategy模型合并疑问

分布式训练中MLflow多模型保存异常的原因与解决方法

问题背景

使用spark-tensorflow-distributor结合TensorFlow的MultiWorkerMirroredStrategy进行多服务器分布式训练,代码如下:

import sys
from spark_tensorflow_distributor import MirroredStrategyRunner
import mlflow.keras
mlflow.keras.autolog()
mlflow.log_param("learning_rate", 0.001)
import tensorflow as tf
import time
from sklearn.model_selection import train_test_split
from sklearn.datasets import load_breast_cancer 

def train():
  strategy = tf.distribute.experimental.MultiWorkerMirroredStrategy()
  #tf.distribute.experimental.CollectiveCommunication.NCCL
  model = None
  with strategy.scope():

    data = load_breast_cancer()
    X_train, X_test, y_train, y_test = train_test_split(data.data, data.target, test_size=0.3)
    N, D = X_train.shape # number of observation and variables
    from sklearn.preprocessing import StandardScaler
    scaler = StandardScaler()
    X_train = scaler.fit_transform(X_train)
    X_test = scaler.transform(X_test)
    model = tf.keras.models.Sequential([
      tf.keras.layers.Input(shape=(D,)),
      tf.keras.layers.Dense(1, activation='sigmoid') # use sigmoid function for every epochs
    ])

    model.compile(optimizer='adam', # use adaptive momentum
        loss='binary_crossentropy',
        metrics=['accuracy']) 

    # Train the Model
    r = model.fit(X_train, y_train, validation_data=(X_test, y_test))

    mlflow.keras.log_model(model, "mymodel")

MirroredStrategyRunner(num_slots=4, use_custom_strategy=True).run(train)

问题现象

  • 设置num_slots=4时,Databricks实验中生成4个预测效果较差的模型;
  • 设置num_slots=1时,仅保存1个预测效果良好的模型;
  • 预期仅主节点保存模型,疑惑是否需要合并模型或操作有误。

解答

你不需要合并模型,问题出在两个核心操作错误:

1. 所有Worker都执行了模型保存逻辑

MultiWorkerMirroredStrategy下,所有worker会同步运行train函数内的全部代码,包括最后的mlflow.keras.log_model,因此4个slot会触发4次模型保存。分布式训练中,只有主节点的模型保存是有效的,其他worker的保存属于冗余操作。

2. 每个Worker的训练数据不一致

你在每个worker里单独执行load_breast_cancer和train_test_split,且未固定随机种子,导致每个worker拆分的训练/测试数据完全不同。虽然分布式训练会同步参数,但数据不一致会让模型训练过程偏离预期,最终保存的模型效果自然变差。


修正方案

步骤1:固定全局随机种子

确保所有worker的数据拆分、模型初始化完全一致,在train函数开头添加随机种子固定逻辑:

def train():
  # 固定全局随机种子
  tf.random.set_seed(42)
  import numpy as np
  np.random.seed(42)
  import random
  random.seed(42)
  from sklearn.utils import shuffle
  shuffle.random_state = 42
  
  strategy = tf.distribute.experimental.MultiWorkerMirroredStrategy()
  # 后续代码...

步骤2:仅主节点执行模型保存

通过TensorFlow的分布式API判断当前是否为主worker,仅主节点执行模型保存:

# Train the Model
    r = model.fit(X_train, y_train, validation_data=(X_test, y_test))

    # 仅主节点保存模型
    if strategy.cluster_resolver.task_type == "worker" and strategy.cluster_resolver.task_id == 0:
      mlflow.keras.log_model(model, "mymodel")

步骤3:优化数据处理(可选但推荐)

避免每个worker重复加载预处理数据,建议在Spark层面先完成数据加载、拆分和预处理,再分发给TensorFlow worker,确保所有worker使用完全一致的训练数据。


修改后的完整代码

import sys
from spark_tensorflow_distributor import MirroredStrategyRunner
import mlflow.keras
mlflow.keras.autolog()
mlflow.log_param("learning_rate", 0.001)
import tensorflow as tf
import time
import numpy as np
import random
from sklearn.model_selection import train_test_split
from sklearn.datasets import load_breast_cancer 
from sklearn.utils import shuffle
from sklearn.preprocessing import StandardScaler

def train():
  # 固定全局随机种子
  tf.random.set_seed(42)
  np.random.seed(42)
  random.seed(42)
  shuffle.random_state = 42

  strategy = tf.distribute.experimental.MultiWorkerMirroredStrategy()
  model = None
  with strategy.scope():
    data = load_breast_cancer()
    X_train, X_test, y_train, y_test = train_test_split(data.data, data.target, test_size=0.3, random_state=42)
    N, D = X_train.shape # number of observation and variables
    
    scaler = StandardScaler()
    X_train = scaler.fit_transform(X_train)
    X_test = scaler.transform(X_test)
    
    model = tf.keras.models.Sequential([
      tf.keras.layers.Input(shape=(D,)),
      tf.keras.layers.Dense(1, activation='sigmoid')
    ])

    model.compile(optimizer='adam',
        loss='binary_crossentropy',
        metrics=['accuracy']) 

    # Train the Model
    r = model.fit(X_train, y_train, validation_data=(X_test, y_test))

    # 仅主节点保存模型
    if strategy.cluster_resolver.task_type == "worker" and strategy.cluster_resolver.task_id == 0:
      mlflow.keras.log_model(model, "mymodel")

MirroredStrategyRunner(num_slots=4, use_custom_strategy=True).run(train)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 03:15:38