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

基于Ray+TensorFlow的联邦学习GPU OOM问题求助

联邦学习GPU OOM问题排查方案(Ray+TensorFlow+Flower)

问题背景

我正在使用Ray结合TensorFlow与Flower框架开展联邦学习任务,通过Ray实现GPU在多客户端间共享,将GPU划分为N份(N=客户端数量+1,包含服务端)。当前使用2张32GB V100 GPU进行MNIST分类任务,理论上显存资源足够,但前3个客户端完成首轮训练后,后续3个客户端启动训练时触发GPU Out Of Memory(OOM)错误,疑似前一轮训练的内存未被释放。日志显示GPU仅被分配452MB显存,即使将batch size调整为1也无法解决该问题。

相关代码与日志

Ray初始化代码

import ray
# 假设N=6,划分6份GPU资源
ray.init(num_gpus=2, resources={"custom_gpu": 6})

客户端训练函数代码

import tensorflow as tf
from flwr.client import Client

class FlowerClient(Client):
    def __init__(self):
        self.model = self.build_model()
    
    def build_model(self):
        model = tf.keras.Sequential([
            tf.keras.layers.Flatten(input_shape=(28,28)),
            tf.keras.layers.Dense(128, activation='relu'),
            tf.keras.layers.Dense(10, activation='softmax')
        ])
        model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy'])
        return model
    
    def fit(self, parameters, config):
        self.model.set_weights(parameters)
        # 加载本地MNIST数据
        (x_train, y_train), _ = tf.keras.datasets.mnist.load_data()
        x_train = x_train / 255.0
        self.model.fit(x_train, y_train, batch_size=1, epochs=1)
        return self.model.get_weights(), len(x_train), {}

运行日志片段

2024-XX-XX XX:XX:XX INFO ray[worker.py] GPU memory allocated: 452MB
2024-XX-XX XX:XX:XX ERROR ray[task.py] Task failed with error: OutOfMemoryError: GPU memory exhausted

排查与解决方向

1. TensorFlow显存资源强制清理

  • 在客户端训练完成后,显式清除TensorFlow计算图与会话资源,在fit方法末尾添加:
    tf.keras.backend.clear_session()
    import gc
    gc.collect()
    
  • 为每个客户端训练任务创建独立的TensorFlow计算图,避免全局图累积占用显存:
    def fit(self, parameters, config):
        with tf.Graph().as_default():
            self.model = self.build_model()
            self.model.set_weights(parameters)
            # 训练逻辑...
            weights = self.model.get_weights()
            tf.keras.backend.clear_session()
            return weights, len(x_train), {}
    

2. Ray资源分配与回收优化

  • 确认每个客户端任务请求的GPU份额与划分的N份匹配,比如N=6时,每个客户端请求1/6的GPU资源:
    # 为客户端任务分配对应份额的GPU资源
    flower_client = ray.remote(num_gpus=1/6)(FlowerClient).remote()
    
  • 客户端训练任务结束后,强制销毁Ray客户端实例,避免残留资源占用:
    # 训练完成后执行资源回收
    ray.kill(flower_client)
    

3. 显存碎片化与动态显存配置

  • 启用TensorFlow动态显存增长,避免一次性占用固定显存块,在代码开头添加:
    gpus = tf.config.list_physical_devices('GPU')
    if gpus:
        for gpu in gpus:
            tf.config.experimental.set_memory_growth(gpu, True)
    
  • 训练过程中避免创建不必要的大张量,确保临时张量在训练结束后被显式删除。

4. Flower客户端生命周期管理

  • 避免复用客户端实例,每个训练任务创建全新的FlowerClient对象,防止模型变量累积占用显存:
    # 每次训练任务生成新的客户端实例
    def create_client():
        return FlowerClient()
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 20:15:59