如何在TensorFlow中读取32位TIFF图像并构建TF Dataset
解决TensorFlow Dataset读取32位浮点TIFF栅格的方案
针对你遇到的问题,核心原因分为两类:
tfio.experimental.image.decode_tiff报错多是版本不兼容问题;- 使用rasterio/skimage/cv2时,Python函数在
map()中无法被TensorFlow自动追踪形状信息,导致输出形状未知。
以下是两种可行的解决方案:
方案一:修复TFIO的使用(推荐,性能更优)
TensorFlow 2.10需要匹配特定版本的TensorFlow IO,先确保安装兼容版本:
pip uninstall -y tensorflow_io pip install tensorflow_io==0.28.0
然后使用以下代码读取32位浮点TIFF:
import tensorflow as tf import tensorflow_io as tfio def decode_32bit_tiff(tiff_path): # 读取文件字节 file_bytes = tf.io.read_file(tiff_path) # 解码为单波段32位浮点图像 img = tfio.experimental.image.decode_tiff(file_bytes, dtype=tf.float32) # 显式指定形状(若所有图像尺寸固定,替换None为具体数值,比如[512, 512, 1]) img = tf.ensure_shape(img, [None, None, 1]) return img # 构建Dataset file_paths = tf.data.Dataset.from_tensor_slices(["path/to/dem1.tif", "path/to/dem2.tif"]) dataset = file_paths.map(decode_32bit_tiff, num_parallel_calls=tf.data.AUTOTUNE) # 验证输出形状 for elem in dataset.take(1): print(elem.shape)
方案二:用tf.py_function包装非TF函数(兼容现有代码)
如果坚持使用rasterio/skimage/cv2,可通过tf.py_function包装Python函数,并显式指定输出形状和类型:
以rasterio为例:
import tensorflow as tf import rasterio import numpy as np def load_dem_rasterio(tiff_path): # 将TF张量转换为Python字符串路径 path = tiff_path.numpy().decode('utf-8') with rasterio.open(path) as src: dem = src.read(1) # 读取单波段数据 # 增加通道维度,转为HxWx1格式 dem = np.expand_dims(dem, axis=-1) return dem.astype(np.float32) def tf_wrap_load_dem(tiff_path): # 包装Python函数,指定输出类型 dem = tf.py_function( func=load_dem_rasterio, inp=[tiff_path], Tout=tf.float32 ) # 强制设置形状(固定尺寸则替换None为具体数值) dem.set_shape([None, None, 1]) return dem # 构建Dataset file_paths = tf.data.Dataset.from_tensor_slices(["path/to/dem1.tif", "path/to/dem2.tif"]) dataset = file_paths.map(tf_wrap_load_dem, num_parallel_calls=tf.data.AUTOTUNE) # 验证输出形状 for elem in dataset.take(1): print(elem.shape)
关键提示
- 若所有DEM图像尺寸固定,务必将
set_shape或ensure_shape中的None替换为实际尺寸(如[512, 512, 1]),这能让TensorFlow更好地优化训练 pipeline,避免形状相关的报错。 - 若图像尺寸不固定,后续可通过
tf.image.resize在map()中统一尺寸,确保模型输入形状一致。
内容的提问来源于stack exchange,提问作者Phil Wernette
相关产品推荐
相关产品推荐

