TensorFlow MultiWorkerMirroredStrategy在HPC多节点OOM后挂起不终止问题
问题结论
TensorFlow 2.3版本的MultiWorkerMirroredStrategy多节点分布式训练场景下,部分GPU触发OOM后确实会出现进程持续挂起不终止的问题,属于该版本分布式策略的已知缺陷。
根因说明
- 多节点分布式训练依赖所有worker节点的集体通信同步(如AllReduce、AllGather算子),所有节点需要在同一个训练步保持步调一致,完成数据交换后才能进入下一个训练步。
- 当单个节点的某张GPU触发OOM抛出
ResourceExhaustedError时,该节点的训练进程会中断当前计算步,但TensorFlow 2.3的MultiWorkerMirroredStrategy没有完善的异常传播和故障终止机制,其余正常运行的节点会持续等待故障节点的通信数据,不会主动退出,最终导致全局任务无限挂起。 - 单节点场景下使用的是同进程内的
MirroredStrategy,单GPU抛出OOM异常后会被同一进程的异常逻辑捕获,直接终止整个进程,因此不会出现挂起问题。
解决方案
- 优先升级TensorFlow版本到2.6及以上,该版本优化了多worker的故障容错逻辑,单节点抛出异常后会主动向集群所有节点发送终止信号,避免无限等待。
- 若受环境限制无法升级框架,可在训练代码的
model.fit外层添加异常捕获逻辑,检测到ResourceExhaustedError时直接调用os._exit(1)强制终止当前进程,Slurm集群检测到任务的某一个进程异常退出后,会自动终止整个任务的所有节点进程,避免资源浪费。 - 训练前开启显存动态分配,减少不必要的显存占用:
import tensorflow as tf gpus = tf.config.list_physical_devices('GPU') for gpu in gpus: tf.config.experimental.set_memory_growth(gpu, True)
- 提前在单节点环境测试不同batch size的显存占用,匹配多节点的全局batch设置,从源头避免OOM问题。
- 可在Slurm提交脚本中配置任务超时阈值,超出预期运行时长后自动终止任务。
内容的提问来源于stack exchange,提问作者ysl
相关产品推荐
相关产品推荐

