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

基于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)

解决方案

问题定位

原预处理函数存在两个关键问题:

  1. 重复颜色空间转换:在处理3通道图像时,连续执行了两次cv2.COLOR_BGR2RGB转换,导致原本正确的RGB图像被再次反转,颜色空间错乱,模型接收到异常输入后生成模糊结果。
  2. 插值方式不明确: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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 00:22:24