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

加载预训练ESRGAN模型时遇KeyError: 'model_state_dict'问题求助

解决方案

1. 修正权重加载逻辑

公开预训练的ESRGAN模型大多直接保存权重字典,而非封装在model_state_dict键下。修改加载代码:

# 替换原加载部分
model_data = torch.load(model_path, map_location=device)
# 直接使用model_data作为状态字典
model.load_state_dict(model_data)

2. 修正模型结构不匹配问题

你的模型定义和官方ESRGAN结构存在多处差异,导致权重无法匹配:

  • RRDBBlock结构错误:官方RRDB由多个残差密集块(RDB)组成,而非直接堆叠卷积层,你的实现会导致通道数持续增长,与预训练模型权重形状不符。
  • RRDB融合逻辑错误:官方ESRGAN中RRDB的输出是累积残差,而非用1x1卷积融合所有RRDB输出。
  • 输入输出通道不匹配:你下载的是黑白漫画模型,输入应为单通道灰度图,输出也应为单通道,而非当前定义的3通道。

修正后的model_definition.py:

import torch
import torch.nn as nn

# 残差密集块(RDB)
class RDBBlock(nn.Module):
    def __init__(self, channels, growth_rate=32):
        super().__init__()
        self.layers = nn.Sequential(
            nn.Conv2d(channels, growth_rate, kernel_size=3, stride=1, padding=1),
            nn.ReLU(inplace=True),
            nn.Conv2d(channels + growth_rate, growth_rate, kernel_size=3, stride=1, padding=1),
            nn.ReLU(inplace=True),
            nn.Conv2d(channels + 2*growth_rate, growth_rate, kernel_size=3, stride=1, padding=1),
            nn.ReLU(inplace=True),
            nn.Conv2d(channels + 3*growth_rate, growth_rate, kernel_size=3, stride=1, padding=1),
            nn.ReLU(inplace=True),
            nn.Conv2d(channels + 4*growth_rate, channels, kernel_size=1, stride=1, padding=0),
        )
    def forward(self, x):
        return x + self.layers(x)

# 残差中的残差密集块(RRDB)
class RRDBBlock(nn.Module):
    def __init__(self, channels, growth_rate=32):
        super().__init__()
        self.rdb1 = RDBBlock(channels, growth_rate)
        self.rdb2 = RDBBlock(channels, growth_rate)
        self.rdb3 = RDBBlock(channels, growth_rate)
    def forward(self, x):
        out = self.rdb1(x)
        out = self.rdb2(out)
        out = self.rdb3(out)
        return out * 0.2 + x  # 残差缩放

# 修正后的ESRGAN模型
class ESRGANModel(nn.Module):
    def __init__(self, num_rrdb_blocks=16, channels=64, growth_rate=32):
        super().__init__()
        # 单通道输入适配黑白模型
        self.conv_input = nn.Conv2d(1, channels, kernel_size=3, stride=1, padding=1)
        self.rrdb_blocks = nn.Sequential(*[RRDBBlock(channels, growth_rate) for _ in range(num_rrdb_blocks)])
        self.conv_mid = nn.Conv2d(channels, channels, kernel_size=3, stride=1, padding=1)
        self.upsample = nn.Sequential(
            nn.Conv2d(channels, channels * 4, kernel_size=3, stride=1, padding=1),
            nn.PixelShuffle(2),
            nn.ReLU(inplace=True),
            nn.Conv2d(channels, channels * 4, kernel_size=3, stride=1, padding=1),
            nn.PixelShuffle(2),
            nn.ReLU(inplace=True),
            nn.Conv2d(channels, 1, kernel_size=3, stride=1, padding=1)  # 单通道输出
        )
    def forward(self, x):
        out = self.conv_input(x)
        residual = out
        out = self.rrdb_blocks(out)
        out = self.conv_mid(out)
        out += residual
        out = self.upsample(out)
        return out

3. 调整图像加载逻辑适配单通道

修改upscale_image函数,加载灰度图并处理单通道张量:

def upscale_image(input_image_path, output_image_path):
    # 加载为灰度图
    image = Image.open(input_image_path).convert("L")
    transform = transforms.ToTensor()
    image_tensor = transform(image).unsqueeze(0).to(device)
    with torch.no_grad():
        upscaled_image = model(image_tensor).clamp(0.0, 1.0)
    # 转换回灰度PIL图像
    upscaled_image = transforms.ToPILImage()(upscaled_image.squeeze(0).cpu())
    upscaled_image.save(output_image_path)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 03:39:52