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

加载CycleGAN预训练权重后图像去噪结果不准确求助

问题描述

我训练了一个用于图像去噪的CycleGAN,得到了.pth模型。自行编写测试代码加载模型时,无法输出准确的去噪图像,但在Jupyter Notebook中运行官方测试脚本却能正常工作:

%run pytorch-CycleGAN-and-pix2pix/test.py 
--dataroot testB/ 
--name cgan_date_snippet --model test --no_dropout --num_test 10

以下是我的测试代码:

import torch
from options.base_options import BaseOptions
from models.networks import define_G

# opt = BaseOptions()

# generator = define_G(input_nc=3, output_nc=3, ngf=64,
#                      netG='resnet_9blocks', norm='instance', use_dropout='store_true')
# generator = define_G(input_nc=3, output_nc=3, ngf=64, netG='unet_256', norm='batch', use_dropout='store_true',init_type='normal',init_gain=0.02)
generator = define_G(input_nc=3, output_nc=3, ngf=64, netG='resnet_9blocks', norm='instance', use_dropout=False)

# Load the pre-trained weights from a saved checkpoint
generator_checkpoint_path = 'latest_net_G_B.pth'

checkpoint = torch.load(generator_checkpoint_path, map_location=torch.device('cuda'))

# Print the keys in the pre-trained model's state_dict to understand its structure
print(checkpoint.keys())

# Load the generator state_dict
generator.load_state_dict(checkpoint, strict=False)

# Set the model to evaluation mode (important if using dropout during training)
generator.eval()

from PIL import Image
from torchvision import transforms
from IPython.display import display

# Load your input image
input_image_path = 'nt6.jpg'
input_image = Image.open(input_image_path).convert('RGB')

# Resize the input image to the expected size
input_image = input_image.resize((512, 512))

# Convert the input image to a PyTorch tensor
input_tensor = transforms.ToTensor()(input_image).unsqueeze(0)  # Add batch dimension

# Move the input tensor to the GPU if available
if torch.cuda.is_available():
    input_tensor = input_tensor.to('cuda')

# Set the generator to evaluation mode (if not already)
generator.eval()

# Move the generator to the same device as the input tensor
generator = generator.to(input_tensor.device)

# Generate the output image
with torch.no_grad():
    output_tensor = generator(input_tensor)

# Move the output tensor to the CPU if necessary
output_tensor = output_tensor.cuda()

# Convert the output tensor to a PIL image
output_image = transforms.ToPILImage()(output_tensor.squeeze(0))

# Display the generated image
display(output_image)

# Save the generated image
output_image.save('image2.jpg')

排查与解决步骤

1. 对齐模型初始化参数与训练配置

官方test.py会自动加载训练时保存的配置文件(路径为checkpoints/cgan_date_snippet/opt.txt),手动初始化define_G时必须和训练参数完全一致,否则模型结构不匹配,加载权重后无法正常工作:

  • 打开opt.txt,确认训练时的netG、norm、use_dropout、ngf等参数,比如训练时用resnet_9blocks+instance norm+use_dropout=True,手动初始化时必须完全对应。
  • 建议用官方Options加载配置,避免手动写死参数出错:
    from options.test_options import TestOptions
    opt = TestOptions().parse()
    generator = define_G(opt.input_nc, opt.output_nc, opt.ngf, opt.netG, opt.norm, not opt.no_dropout, opt.init_type, opt.init_gain)
    

2. 修正权重加载逻辑

加载latest_net_G_B.pth时需注意:

  • 如果训练时用了多GPU,权重键名会带有module.前缀,手动加载时需要去除,否则会出现结构不匹配:
    checkpoint = {k.replace('module.', ''): v for k, v in checkpoint.items()}
    generator.load_state_dict(checkpoint, strict=True)
    
  • 可以打印模型和权重的结构对比,确认每层名称、形状是否匹配:
    print("手动初始化模型结构:")
    for name, param in generator.named_parameters():
        print(name, param.shape)
    print("\n权重文件结构:")
    for name, param in checkpoint.items():
        print(name, param.shape)
    

3. 补全图像预处理/后处理步骤

官方测试代码会对图像做归一化,手动代码缺少这一步会导致输入分布与训练时不一致:

  • 预处理时添加归一化(CycleGAN默认将图像归一化到[-1,1]):
    transform = transforms.Compose([
        transforms.Resize((512, 512)),
        transforms.ToTensor(),
        transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
    ])
    input_tensor = transform(input_image).unsqueeze(0)
    
  • 后处理时反归一化,将输出从[-1,1]转回[0,1],再转到CPU处理:
    output_tensor = (output_tensor + 1) / 2.0  # 反归一化
    output_tensor = output_tensor.clamp(0, 1)  # 限制像素值范围
    output_tensor = output_tensor.cpu()  # 转到CPU
    output_image = transforms.ToPILImage()(output_tensor.squeeze(0))
    
  • 去掉代码中多余的output_tensor = output_tensor.cuda()语句。

4. 统一设备配置

提前将模型和权重放到同一设备,避免设备不匹配问题:

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
generator = generator.to(device)
checkpoint = torch.load(generator_checkpoint_path, map_location=device)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 11:30:13