基于CycleGAN-and-pix2pix的单图推理结果模糊问题求助
问题:CycleGAN单图推理结果模糊,与批量推理效果不符
我在Google Colab上基于CycleGAN-and-pix2pix开源API训练了CycleGAN模型,训练命令为:
!python train.py --dataroot /content/drive/MyDrive/project/dataset --name F2F --model cycle_gan --display_id -1
通过数据加载器从文件夹批量推理的代码运行正常,生成结果良好,但自行编写单图推理的预处理函数后,生成的图像非常模糊,达不到批量推理的效果,希望解决这个问题。
批量推理代码
opt = TestOptions() # defined options occurs here dataset = create_dataset(opt) # Initialize the model model = create_model(opt) model.setup(opt) model.eval() data_iter = iter(dataset.dataloader) data_dict = next(data_iter) input_image_tensor = data_dict['A'] data = {'A': input_image_tensor,'A_paths': ''} model.set_input(data) model.test() visuals = model.get_current_visuals() output_image = visuals['fake'] output_image_np = output_image.squeeze().cpu().numpy().transpose(1, 2, 0) output_image_np = ((output_image_np - output_image_np.min()) / (output_image_np.max() - output_image_np.min()) * 255).astype(np.uint8) output_image_np = cv2.cvtColor(output_image_np, cv2.COLOR_BGR2RGB) cv2_imshow(output_image_np)
自定义预处理及单图推理代码
预处理函数
def preprocess(image): if image.ndim == 2 or image.shape[2] == 1: image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB) elif image.shape[2] == 4: image = cv2.cvtColor(image, cv2.COLOR_BGRA2BGR) elif image.shape[2] == 3: image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) pil_image = transforms.ToPILImage()(image) transform_pipeline = transforms.Compose([ transforms.Resize(286), transforms.CenterCrop(256), transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)) ]) image_tensor = transform_pipeline(pil_image) image_tensor = image_tensor.unsqueeze(0) return image_tensor
单图推理代码
input_image = cv2.imread('/content/drive/MyDrive/dataset/testA/image_1.jpg') input_image_tensor = preprocess(input_image) data = {'A': input_image_tensor,'A_paths': ''} model.set_input(data) model.test() visuals = model.get_current_visuals() output_image = visuals['fake'] output_image_np = output_image.squeeze().cpu().numpy().transpose(1, 2, 0) output_image_np = ((output_image_np - output_image_np.min()) / (output_image_np.max() - output_image_np.min()) * 255).astype(np.uint8) output_image_np = cv2.cvtColor(output_image_np, cv2.COLOR_BGR2RGB) cv2_imshow(output_image_np)
解决方案
问题定位
原预处理函数存在两个关键问题:
- 重复颜色空间转换:在处理3通道图像时,连续执行了两次
cv2.COLOR_BGR2RGB转换,导致原本正确的RGB图像被再次反转,颜色空间错乱,模型接收到异常输入后生成模糊结果。 - 插值方式不明确:
transforms.Resize未显式指定插值方式,虽然默认是双线性,但为了和批量推理时数据集的预处理逻辑完全对齐,需要明确指定。
修正后的预处理函数
def preprocess(image): # 处理不同通道数的图像 if image.ndim == 2 or image.shape[2] == 1: image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB) elif image.shape[2] == 4: image = cv2.cvtColor(image, cv2.COLOR_BGRA2BGR) elif image.shape[2] == 3: # 仅执行一次BGR到RGB的转换 image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) pil_image = transforms.ToPILImage()(image) # 显式指定双线性插值,与CycleGAN默认预处理逻辑一致 transform_pipeline = transforms.Compose([ transforms.Resize(286, interpolation=transforms.InterpolationMode.BILINEAR), transforms.CenterCrop(256), transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)) ]) image_tensor = transform_pipeline(pil_image) image_tensor = image_tensor.unsqueeze(0) return image_tensor
额外验证建议
- 检查输入图像的宽高比,确保经过
Resize(286)和CenterCrop(256)后没有出现严重变形,若原图比例特殊,可调整预处理逻辑适配。 - 对比批量推理时输入张量和单图预处理后的张量,查看数值范围、维度是否完全一致,进一步排查输入差异。
内容的提问来源于stack exchange,提问作者yair yatzkan
相关产品推荐
相关产品推荐

