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

RuntimeError: CUDA内存不足(append操作引发)及解决方法咨询

问题分析与修复

核心问题排查

你遇到的copy.deepcopy报错,大概率是以下原因之一:

  • 模型权重w中包含无法被深拷贝的对象(比如某些自定义非序列化张量、资源句柄等)
  • 代码逻辑存在重复覆盖w_locals的问题,导致深拷贝时出现引用混乱
  • w_glob本身是不可深拷贝的对象,初始化w_locals时埋下隐患

代码修复方案

方案1:替换深拷贝方式(针对权重无法深拷贝的情况)

如果是PyTorch模型权重,直接利用state_dict()的特性做拷贝,无需依赖copy.deepcopy:

w_totals = []
for iter in range(epochs):
    # 初始化空列表,避免复用w_glob引用引发的问题
    w_locals = []
    
    for idxs_task in range(task_num):
        dataset = load_data() 
        
        # 存储当前任务的5个客户端权重
        current_task_w = []
        for client_idx in range(5):  # 直接写5更清晰,避免变量歧义
            dataloader_train = load_data_train()
            # 复制全局模型并训练
            local_net = copy.deepcopy(net_glob).to(device)
            train(local_net, dataloader_train)
            # 直接获取状态字典,无需额外深拷贝
            current_task_w.append(local_net.state_dict())
        
        # 将当前任务的客户端权重加入总列表
        w_totals.append(current_task_w)

方案2:修复原代码的深拷贝逻辑

如果必须保留深拷贝操作,确保只对可序列化的内容执行拷贝:

import copy

w_totals = []
for iter in range(epochs):
    # 初始化空列表,避免引用共享问题
    w_locals = []
    
    for idxs_task in range(task_num):
        dataset = load_data() 
        
        w_locals.clear()  # 清空当前任务的本地权重列表
        for client_idx in range(5):
            dataloader_train = load_data_train()
            local_net = copy.deepcopy(net_glob).to(device)
            w = train(local_net, dataloader_train)
            # 只深拷贝权重的状态字典,而非整个模型对象
            w_locals.append(copy.deepcopy(w.state_dict()))
        
        # 深拷贝整个本地权重列表加入总列表
        w_totals.append(copy.deepcopy(w_locals))

额外优化建议

  • 避免用[w_glob for i in range(client_5)]初始化列表,这种写法会导致列表中所有元素引用同一个w_glob对象,后续赋值虽会覆盖,但初始化阶段容易引发引用混乱
  • 如果你的需求是遍历50个客户端,每次取5个训练并保存,可调整循环逻辑实现分组训练:
w_totals = []
# 将50个客户端分成10组,每组5个
client_groups = [range(i, i+5) for i in range(0, 50, 5)]  

for iter in range(epochs):
    for group in client_groups:
        w_locals = []
        for client_idx in group:
            # 传入客户端索引加载对应数据
            dataloader_train = load_data_train(client_idx)  
            local_net = copy.deepcopy(net_glob).to(device)
            train(local_net, dataloader_train)
            w_locals.append(local_net.state_dict())
        w_totals.append(w_locals)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 12:02:50