加载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+instancenorm+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
相关产品推荐
相关产品推荐

