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

pickle加载Unet2Dmodel报错:找不到diffusers中的AttentionBlock

解决pickle加载diffusers模型时的AttributeError问题

问题原因

直接用pickle保存整个UNet2DModel对象时,pickle会记录类的完整模块路径(比如diffusers.models.attention.AttentionBlock)。即便切换回两个月前的diffusers版本,仍可能因以下原因加载失败:

  • 库内部模块结构微调(比如AttentionBlock被移动、重命名,或是模块路径发生变动)
  • 不同环境中库的安装路径、导入逻辑存在差异,导致pickle无法匹配到对应的类引用

解决方案

方案1:修复现有.pkl文件的类路径映射

如果必须使用已保存的model.pkl,可以通过自定义Unpickler重定向类的引用:

# 先确认当前环境中AttentionBlock对应的实际类(需对照旧版本diffusers代码)
# 示例:若旧版本的AttentionBlock现在改名为Attention,或位于diffusers.models.attention模块下
from diffusers.models.attention import Attention as AttentionBlock

import pickle

class CustomUnpickler(pickle.Unpickler):
    def find_class(self, module, name):
        # 将旧的类引用映射到当前环境中的实际类
        if name == 'AttentionBlock' and module == 'diffusers.models.attention':
            return AttentionBlock
        # 其他类保持默认查找逻辑
        return super().find_class(module, name)

def load_model(path):
    with open(path, 'rb') as f:
        return CustomUnpickler(f).load()

注意:需要先确认旧版本diffusers中AttentionBlock的定义位置,再对应到当前环境中的类路径,可能需要查看diffusers历史版本代码来确认。

方案2:改用diffusers官方推荐的保存加载方式(推荐)

pickle并非diffusers官方推荐的模型保存方式,改用save_pretrained和from_pretrained可彻底避免类路径依赖问题,兼容性更强:

保存模型

def save_model(model, save_dir):
    # 保存模型权重和配置文件到指定目录
    model.save_pretrained(save_dir)

加载模型

from diffusers import UNet2DModel

def load_model(save_dir):
    # 从保存的目录加载模型
    return UNet2DModel.from_pretrained(save_dir)

额外:恢复训练的状态保存

如果需要恢复训练,除模型本身外,需单独保存优化器等训练状态(避免pickle整个优化器对象):

import torch

# 保存训练状态
def save_training_state(model, optimizer, epoch, loss, save_dir):
    torch.save(
        {
            'epoch': epoch,
            'model_state_dict': model.state_dict(),
            'optimizer_state_dict': optimizer.state_dict(),
            'loss': loss,
        },
        f"{save_dir}/training_checkpoint.pt"
    )

# 加载训练状态
def load_training_state(model, optimizer, checkpoint_path):
    checkpoint = torch.load(checkpoint_path)
    model.load_state_dict(checkpoint['model_state_dict'])
    optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
    return checkpoint['epoch'], checkpoint['loss']

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 07:13:18