基于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
相关产品推荐
相关产品推荐

