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

如何将给定的Keras图像加载代码转换为PyTorch实现

PyTorch 版本对齐实现方案

原Keras代码的核心逻辑为目录文件遍历、指定尺寸缩放图像、沿宽度轴拆分图像、归集数据转为数组,PyTorch生态下可直接复现全部逻辑,且能做到和原代码输出格式完全对齐。


核心API对应关系

不需要额外记忆复杂的新API,两者的图像加载底层都依赖Pillow库,对应关系非常直接:

  • Keras的load_img:本质是封装Pillow的图像读取+缩放操作,直接用Pillow原生接口或者torchvision的transforms接口即可实现同等效果
  • Keras的img_to_array:作用是把PIL格式图像转为*高度-宽度-通道(HWC)*格式的numpy数组,像素值保持0-255范围,Pillow图像直接转numpy的结果和它完全一致
  • 数组拆分逻辑:只要保持维度顺序一致,numpy切片、PyTorch张量切片的写法和原代码几乎没有区别

1:1 对齐原Keras输出的实现版本

这个版本返回结果和你提供的Keras代码完全一致,都是两个形状为[样本数, 256, 256, 通道数]的numpy数组,不需要修改后续原有数据处理逻辑即可直接替换。
首先安装依赖:
pip install pillow numpy

对应代码:

import os
import numpy as np
from PIL import Image
from numpy import asarray

def load_images(path, size=(256, 512)):
    src_list, tar_list = list(), list()
    # 遍历指定目录下所有文件,原逻辑默认目录内全部为有效图像文件
    for filename in os.listdir(path):
        # 跨系统兼容的路径拼接,避免Windows/Linux路径分隔符不统一报错
        file_path = os.path.join(path, filename)
        # 加载图像并按指定尺寸缩放,用双线性插值和Keras默认插值逻辑对齐
        with Image.open(file_path) as img:
            # 注意Pillow的resize入参顺序为(宽度, 高度),和size定义的(高度,宽度)顺序相反
            resized_img = img.resize((size[1], size[0]), resample=Image.BILINEAR)
            # 转为HWC格式numpy数组,像素值范围0-255,和原img_to_array输出完全一致
            pixels = asarray(resized_img)
        # 沿宽度轴对半拆分卫星图、地图,切片逻辑和原代码完全一致
        sat_img, map_img = pixels[:, :256], pixels[:, 256:]
        src_list.append(sat_img)
        tar_list.append(map_img)
    # 归集为numpy数组返回,和原Keras返回格式无差异
    return [asarray(src_list), asarray(tar_list)]

适配PyTorch模型训练的优化版本

如果你后续需要直接把数据输入PyTorch模型,可以用torchvision的工具函数做适配,直接返回PyTorch原生张量,省掉手动转换格式的步骤。
额外安装依赖:
pip install torch torchvision

对应代码:

import os
import torch
from PIL import Image
from torchvision.transforms import functional as F

def load_images_pytorch(path, size=(256, 512), normalize=True):
    src_list, tar_list = list(), list()
    for filename in os.listdir(path):
        file_path = os.path.join(path, filename)
        with Image.open(file_path) as img:
            # 按指定尺寸缩放,默认双线性插值和原逻辑对齐
            resized_img = F.resize(img, size)
            # 直接转为PyTorch默认的*通道-高度-宽度(CHW)*格式张量,默认自动把像素值归一化到0-1范围
            pixels = F.to_tensor(resized_img)
        # torch张量的宽度维度为最后一维,对应调整切片位置
        sat_img, map_img = pixels[:, :, :256], pixels[:, :, 256:]
        # 如果不需要归一化、要保持0-255的像素值范围,取消下面两行注释即可
        # sat_img = sat_img * 255
        # map_img = map_img * 255
        src_list.append(sat_img)
        tar_list.append(map_img)
    # 沿batch维度拼接,最终输出形状为[样本数, 通道数, 256, 256]的张量,可直接输入模型
    src_tensor = torch.stack(src_list, dim=0)
    tar_tensor = torch.stack(tar_list, dim=0)
    return [src_tensor, tar_tensor]

额外注意事项

  • 如果遍历的目录中混有非图像文件,可以在遍历逻辑里加一层后缀判断,比如只处理.jpg、.png、.jpeg后缀的文件,避免读取无效文件报错
  • 如果需要和原Keras做像素级结果对齐,注意统一插值方式,不要随便换成最近邻、双三次插值,避免缩放后的像素值出现差异
  • 小批量数据集可以直接用上面的函数一次性加载到内存,如果是超大数据集建议改用PyTorch的Dataset+DataLoader做按需加载,避免内存溢出

内容的提问来源于stack exchange,提问作者Francesco Conti

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 08:33:11