AzureML中Keras+Horovod多节点训练时如何导出正确的损失/精度曲线?
解决AzureML + Horovod多GPU Keras训练的损失/精度曲线异常问题
我之前在AzureML上用Horovod跑双GPU的Keras CNN训练时,也碰到过曲线跳变、重复点的异常问题——本质是每个GPU节点都会独立记录训练数据,最后汇总时就会出现混乱。下面是我亲测有效的解决方案,核心思路是只让主节点(rank 0)负责记录全局平均的metrics,避免多节点数据冲突:
1. 核心原理:主节点独占日志记录权
Horovod训练时,每个GPU会处理一部分batch数据,直接让所有节点记录metrics的话,日志里会出现同一epoch对应多条局部数据的情况,绘制曲线自然会异常。我们需要:
- 用Horovod的
MetricAverageCallback收集所有节点的metrics并取平均 - 只让rank为0的主节点保存日志、生成记录
2. 自定义回调函数过滤非主节点操作
我们可以封装Keras原生的CSVLogger,让它只在主节点生效:
import horovod.tensorflow.keras as hvd from tensorflow.keras.callbacks import CSVLogger class RankedCSVLogger(CSVLogger): def __init__(self, filename, separator=',', append=False): super().__init__(filename, separator, append) def on_epoch_end(self, epoch, logs=None): # 仅主节点执行日志写入 if hvd.rank() == 0: super().on_epoch_end(epoch, logs)
3. 完整训练脚本配置
在你的AzureML训练脚本中,按以下步骤配置:
import tensorflow as tf import horovod.tensorflow.keras as hvd from azureml.core import Run # 初始化Horovod hvd.init() # 配置当前节点的GPU(每个节点只使用分配的GPU) gpus = tf.config.experimental.list_physical_devices('GPU') for gpu in gpus: tf.config.experimental.set_memory_growth(gpu, True) if gpus: tf.config.experimental.set_visible_devices(gpus[hvd.local_rank()], 'GPU') # 构建你的CNN模型... model = tf.keras.Sequential([...]) # 用Horovod包裹优化器,自动缩放学习率 base_lr = 0.001 opt = tf.keras.optimizers.Adam(learning_rate=base_lr * hvd.size()) opt = hvd.DistributedOptimizer(opt) # 编译模型 model.compile(loss='categorical_crossentropy', optimizer=opt, metrics=['accuracy']) # 基础回调:Horovod必需的同步回调 callbacks = [ # 广播初始模型变量到所有节点 hvd.callbacks.BroadcastGlobalVariablesCallback(0), # 收集所有节点的metrics并取平均,传给主节点 hvd.callbacks.MetricAverageCallback(), # 可选:学习率预热回调 hvd.callbacks.LearningRateWarmupCallback(warmup_epochs=5, verbose=1), ] # 仅主节点添加日志和可视化回调 if hvd.rank() == 0: callbacks.append(RankedCSVLogger('training_metrics.csv')) callbacks.append(tf.keras.callbacks.TensorBoard('./tb_logs')) # 启动训练(注意batch_size是单节点的batch大小,全局batch是这个值*hvd.size()) model.fit( train_dataset, epochs=50, batch_size=32, callbacks=callbacks, validation_data=val_dataset ) # 仅主节点上传日志到AzureML if hvd.rank() == 0: run = Run.get_context() run.upload_file('training_metrics.csv', 'training_metrics.csv') run.upload_folder('tb_logs', './tb_logs')
4. 绘制正确的曲线
训练结束后,从AzureML工作室下载主节点上传的training_metrics.csv,用pandas绘制曲线:
import pandas as pd import matplotlib.pyplot as plt df = pd.read_csv('training_metrics.csv') # 绘制损失曲线 plt.figure(figsize=(10,5)) plt.plot(df['epoch'], df['loss'], label='Training Loss') plt.plot(df['epoch'], df['val_loss'], label='Validation Loss') plt.xlabel('Epoch') plt.ylabel('Loss') plt.title('Training vs Validation Loss') plt.legend() plt.show() # 绘制精度曲线 plt.figure(figsize=(10,5)) plt.plot(df['epoch'], df['accuracy'], label='Training Accuracy') plt.plot(df['epoch'], df['val_accuracy'], label='Validation Accuracy') plt.xlabel('Epoch') plt.ylabel('Accuracy') plt.title('Training vs Validation Accuracy') plt.legend() plt.show()
常见坑点提醒
- 不要忘记添加
MetricAverageCallback:否则主节点记录的只是自身节点的局部metrics,不是全局平均值 - 禁止多节点同时写日志:会导致日志文件出现重复epoch条目,曲线直接报废
- AzureML中要确保从主节点获取日志:只有rank 0的节点上传了正确的记录文件
内容的提问来源于stack exchange,提问作者Garrett Edmondson
相关产品推荐
相关产品推荐

