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

联邦学习环境中为Python Client类配置GPU显存占比上限的问题

联邦学习客户端GPU显存占比配置问题

我正在开发一个联邦学习环境,包含Server类和10个Client类实例,训练时会依次启动各个Client,需要为每个Client实例指定不同的GPU显存使用上限。

我的Client类定义如下:

class Client:
    def __init__(self, id:int, neighbors:list, trainloader, valloader, testloader, gpu_fraction, device):
        self.id=id
        self.neighbors = neighbors
        self.model=None
        self.trainloader = trainloader
        self.valloader = valloader
        self.testloader = testloader
        self.gpu_fraction = gpu_fraction
        self.GPU_usage_table=None
        self.device = device   

我尝试在每个Client执行训练函数前,用以下代码限制显存占比,但没有效果:

import tensorflow as tf
gpu_options = tf.compat.v1.GPUOptions(per_process_gpu_memory_fraction=self.gpu_fraction)
config = tf.compat.v1.ConfigProto(gpu_options=gpu_options)
session = tf.compat.v1.Session(config=config)

我用以下代码检测GPU使用情况:

nvidia_smi.nvmlInit()
deviceCount = nvidia_smi.nvmlDeviceGetCount()
for i in range(deviceCount):
    handle = nvidia_smi.nvmlDeviceGetHandleByIndex(i)
    res = nvidia_smi.nvmlDeviceGetUtilizationRates(handle)
    mem = nvidia_smi.nvmlDeviceGetMemoryInfo(handle)
    table.add_row([i, f"{mem.free/1024**2:5.2f}MB/{mem.total/1024**2:5.2f}MB", f"{res.gpu:3.1%}", f"{res.memory:3.1%}"])

寻求有效的GPU显存占比配置方法。


有效配置方案

1. TensorFlow 2.x 适配(单进程多Client场景)

如果用TF2.x且所有Client在同一进程内依次训练,必须在每个Client训练前重置GPU配置(TF2显存配置为全局生效):

import tensorflow as tf

# 重置之前的GPU配置
tf.config.experimental.reset_memory_stats('GPU:0')
tf.config.experimental.set_memory_growth('GPU:0', False)

# 计算并设置显存上限
gpus = tf.config.list_physical_devices('GPU')
if gpus:
    total_mem = tf.config.experimental.get_memory_info(gpus[0])['total']
    mem_limit = int(total_mem * self.gpu_fraction)
    tf.config.set_logical_device_configuration(
        gpus[0],
        [tf.config.LogicalDeviceConfiguration(memory_limit=mem_limit)]
    )

注意:每次切换Client训练时都要重新执行这段代码,确保显存限制生效。

2. 进程级显存限制(跨框架通用)

如果每个Client以独立进程启动,直接通过环境变量限制最可靠:

import os
import nvidia_smi

# 获取GPU总显存(MB)
nvidia_smi.nvmlInit()
handle = nvidia_smi.nvmlDeviceGetHandleByIndex(0)
total_mem_mb = nvidia_smi.nvmlDeviceGetMemoryInfo(handle).total // (1024**2)
mem_limit_mb = int(total_mem_mb * self.gpu_fraction)

# 设置环境变量
os.environ['CUDA_VISIBLE_DEVICES'] = '0'  # 指定GPU索引
# TF专属配置
os.environ['TF_PER_GPU_MEMORY_LIMIT'] = str(mem_limit_mb)
os.environ['TF_FORCE_GPU_ALLOW_GROWTH'] = 'false'
# PyTorch专属配置
os.environ['PYTORCH_CUDA_ALLOC_CONF'] = f'max_split_size_mb:{mem_limit_mb},garbage_collection_threshold:0.6'

这种方式直接限制整个进程的显存使用,无论用什么框架都能生效。

3. PyTorch 专属细粒度控制

如果用PyTorch 1.8+版本,支持直接按比例限制进程显存:

import torch

# 确保设备正确
torch.cuda.set_device(self.device)
# 按比例设置显存上限
torch.cuda.memory.set_per_process_memory_fraction(self.gpu_fraction, self.device)

这段代码要在模型初始化和数据加载前执行,每次切换Client时重新设置。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 01:02:21