如何将给定的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
相关产品推荐
相关产品推荐

