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

