加载预训练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
相关产品推荐
相关产品推荐

