如何在Python的TensorFlow或Keras中导入GeoTIFF作为输入数据?
在TensorFlow/Keras中导入GeoTIFF的可行方法
我之前在做遥感图像深度学习项目时,也碰到过TensorFlow不直接支持GeoTIFF的问题,用GDAL踩过坑之后,总结了几个Python环境下靠谱的解决方案,分享给你:
方法1:使用rasterio库读取(最推荐)
rasterio是专门处理栅格地理数据的库,对GeoTIFF的支持非常友好,而且能直接输出numpy数组,无缝对接TensorFlow。
步骤:
- 先安装rasterio:
pip install rasterio
- 读取GeoTIFF并转换为TensorFlow张量:
import rasterio import tensorflow as tf # 读取GeoTIFF文件 with rasterio.open('your_geotiff_file.tif') as src: # 读取所有波段,得到形状为 (bands, height, width) 的numpy数组 geotiff_data = src.read() # 转换为TensorFlow需要的 (height, width, bands) 格式(多波段场景必备) geotiff_data = tf.transpose(geotiff_data, perm=[1, 2, 0]) # 转换为适合模型输入的数据类型,比如float32 geotiff_tensor = tf.cast(geotiff_data, tf.float32) # 现在geotiff_tensor就可以直接作为模型的输入了
如果需要保留地理坐标信息,可以用src.transform和src.crs获取后续处理。
方法2:正确使用GDAL读取(解决你之前的失败问题)
你之前用GDAL转换失败,大概率是没处理好波段顺序或者数据格式,试试下面的正确流程:
- 安装GDAL(如果没装的话):
pip install gdal
- 读取并转换:
from osgeo import gdal import numpy as np import tensorflow as tf # 打开GeoTIFF文件 ds = gdal.Open('your_geotiff_file.tif') # 获取波段数量 num_bands = ds.RasterCount # 读取每个波段并拼接成数组 bands = [] for i in range(1, num_bands+1): band = ds.GetRasterBand(i) bands.append(band.ReadAsArray()) # 得到 (bands, height, width) 的数组 geotiff_data = np.array(bands) # 转换为HWC格式 geotiff_data = tf.transpose(geotiff_data, perm=[1, 2, 0]) # 转成张量 geotiff_tensor = tf.cast(geotiff_data, tf.float32) # 记得关闭数据集释放资源 ds = None
注意:GDAL读取波段是从1开始计数的,这点容易踩坑;另外如果你的GeoTIFF有NoData值,记得用band.GetNoDataValue()获取后替换为合理数值,避免影响模型训练。
方法3:使用tifffile + PIL(适合简单场景)
如果你的GeoTIFF是单波段或者少数波段,也可以用tifffile读取后转成PIL图像,再转换为张量:
- 安装依赖:
pip install tifffile pillow
- 代码示例:
import tifffile from PIL import Image import tensorflow as tf # 读取GeoTIFF img = tifffile.imread('your_geotiff_file.tif') # 转成PIL图像(多波段场景需要调整通道顺序) pil_img = Image.fromarray(img.transpose(1,2,0)) if len(img.shape)==3 else Image.fromarray(img) # 转换为TensorFlow张量 geotiff_tensor = tf.convert_to_tensor(pil_img, dtype=tf.float32)
这个方法比较轻量,但对于复杂的多波段GeoTIFF或者带有地理信息的场景,还是推荐前两种方法。
额外注意事项
- 通道顺序:GeoTIFF通常是
(bands, height, width)的格式,而TensorFlow/Keras的输入一般要求(height, width, bands),所以一定要记得用tf.transpose调整顺序。 - 数据类型:GeoTIFF常见的是uint8、uint16类型,模型输入一般需要float32,记得用
tf.cast转换。 - NoData值处理:很多GeoTIFF会有无效值(比如-9999),读取后要替换成0或者模型能处理的数值,避免干扰训练。
内容的提问来源于stack exchange,提问作者송준석
相关产品推荐
相关产品推荐

