Python调用GDAL做预测?滑窗裁切怎么写?

GIS基础理论
Dr.GIS
wowwwai GIS研习社 · 工具流程与项目排障

引言

很多做遥感分类、地物提取或栅格回归的同学都会遇到一个问题:Python调用GDAL做预测?滑窗裁切怎么写? 模型已经训练好,输入是一幅大影像,但不能一次性把整幅 GeoTIFF 读进内存,只能按窗口逐块读取、预测、再写回结果。

这篇文章用一个可复用的思路讲清楚:如何用 Python 调用 GDAL 做栅格滑窗裁切,如何处理边界窗口,如何把预测结果按原始空间参考写成新的 GeoTIFF。重点不是某一个深度学习框架,而是 GIS 预测流程里最容易出错的“读窗口、算窗口、写窗口”。

背景

在 GIS 和遥感场景中,GDAL 常用于读取 GeoTIFF、IMG、VRT 等栅格数据。实际项目里,影像通常很大,例如几 GB 到几十 GB。如果直接使用以下方式读取整幅影像:

arr = dataset.ReadAsArray()

小影像可以运行,大影像就可能出现内存不足、程序卡死或预测速度很慢的问题。因此,Python GDAL滑窗预测通常采用分块方式:

  • 按固定窗口大小读取影像块。
  • 对每个窗口做预处理和模型预测。
  • 把预测结果写入输出栅格的对应位置。
  • 保留原始影像的投影、仿射变换和像元对齐关系。

这类需求常见于土地利用分类、建筑物提取、植被指数预测、灾害范围识别和栅格分割结果生成。

Python调用GDAL做预测 Python GDAL滑窗预测流程图
Python 调用 GDAL 做滑窗预测的基本流程:读取窗口、模型预测、写回结果,并保留空间参考。

原理

GDAL 栅格数据读取的核心函数是 ReadAsArray(xoff, yoff, xsize, ysize)。它表示从左上角开始,读取一个指定位置和大小的窗口。

  • xoff:窗口左上角的列号偏移。
  • yoff:窗口左上角的行号偏移。
  • xsize:窗口宽度,也就是读取多少列。
  • ysize:窗口高度,也就是读取多少行。

对应写入结果时使用 WriteArray(array, xoff, yoff),它会把预测后的数组写到输出栅格的对应位置。

滑窗裁切的关键点有三个:

  • 窗口不能越界:靠近影像右边和下边时,窗口大小可能小于设定的 window_size
  • 输出尺寸要对齐:预测结果的行列数必须和写入窗口一致。
  • 空间参考要继承:输出 GeoTIFF 必须复制原始影像的投影和仿射变换,否则结果虽然能打开,但位置可能不对。

如果模型要求固定输入大小,例如 256×256,但边界窗口不足 256×256,就需要补边或者跳过。对于 GIS 预测结果,通常建议补边后预测,再裁回真实窗口大小写入。

步骤

1. 安装和导入 GDAL

如果你使用的是 Conda 环境,推荐通过 conda-forge 安装 GDAL,避免 DLL 或动态库问题。

conda install -c conda-forge gdal numpy

Python 中导入:

from osgeo import gdal
import numpy as np

2. 打开输入影像并读取基本信息

下面代码会打开输入 GeoTIFF,并读取宽度、高度、波段数、投影和仿射变换。

input_path = "input.tif"
output_path = "prediction.tif"

src_ds = gdal.Open(input_path, gdal.GA_ReadOnly)
if src_ds is None:
    raise RuntimeError("无法打开输入影像: {}".format(input_path))

width = src_ds.RasterXSize
height = src_ds.RasterYSize
band_count = src_ds.RasterCount
geo_transform = src_ds.GetGeoTransform()
projection = src_ds.GetProjection()

print("width:", width)
print("height:", height)
print("bands:", band_count)

如果这里打印的波段数和模型输入要求不一致,需要先确认模型到底需要哪些波段。例如 RGB 模型通常需要 3 个波段,多光谱模型可能需要 4 个或更多波段。

3. 创建输出预测栅格

假设预测结果是单波段分类图,像元类型使用 GDT_Byte。如果是概率图或连续值回归结果,可以改为 GDT_Float32

driver = gdal.GetDriverByName("GTiff")

out_ds = driver.Create(
    output_path,
    width,
    height,
    1,
    gdal.GDT_Byte,
    options=["COMPRESS=LZW", "TILED=YES"]
)

out_ds.SetGeoTransform(geo_transform)
out_ds.SetProjection(projection)

out_band = out_ds.GetRasterBand(1)
out_band.SetNoDataValue(0)

这里的重点是 SetGeoTransformSetProjection。它们决定输出预测图能否和原始影像、矢量边界、底图正确叠加。

4. 编写一个示例预测函数

为了让代码可以直接理解,先写一个占位预测函数。实际使用时,你可以把这里替换成 PyTorch、TensorFlow、scikit-learn 或自己的模型推理代码。

def predict_block(block):
    """
    block形状通常为:
    - 多波段: (bands, rows, cols)
    - 单波段: (rows, cols)

    这里仅做示例:使用第1个波段大于均值的位置作为类别1。
    实际项目中请替换为自己的模型预测逻辑。
    """
    if block.ndim == 3:
        first_band = block[0]
    else:
        first_band = block

    valid = first_band > 0
    threshold = first_band[valid].mean() if np.any(valid) else 0

    pred = np.zeros(first_band.shape, dtype=np.uint8)
    pred[first_band > threshold] = 1
    return pred

如果你的模型输入需要形状为 (1, channels, height, width),可以在这个函数里做维度转换。GDAL 读出的多波段数组一般是 (bands, rows, cols),而很多深度学习模型需要 (batch, channels, rows, cols)(batch, rows, cols, channels)

5. 编写 GDAL 滑窗裁切预测主循环

下面是一个完整的 Python调用GDAL做预测滑窗代码。它会自动处理右边和下边不足一个完整窗口的情况。

window_size = 512

for yoff in range(0, height, window_size):
    for xoff in range(0, width, window_size):
        xsize = min(window_size, width - xoff)
        ysize = min(window_size, height - yoff)

        block = src_ds.ReadAsArray(xoff, yoff, xsize, ysize)

        if block is None:
            continue

        pred = predict_block(block)

        if pred.shape != (ysize, xsize):
            raise RuntimeError(
                "预测结果尺寸不匹配: pred={}, expected={}".format(
                    pred.shape, (ysize, xsize)
                )
            )

        out_band.WriteArray(pred, xoff, yoff)

out_band.FlushCache()
out_ds.FlushCache()

src_ds = None
out_ds = None

print("预测完成:", output_path)

这段代码适合模型能够接受任意窗口大小的情况。很多传统机器学习或简单规则预测都可以这样写。

6. 如果模型必须输入固定大小窗口

深度学习模型经常要求固定输入尺寸,例如 256×256 或 512×512。此时边界窗口可能不够大,需要先补边,预测后再裁剪回真实大小。

def pad_to_window(block, target_size):
    """
    将GDAL读取的block补到target_size。
    返回补边后的block,以及原始窗口高度和宽度。
    """
    if block.ndim == 2:
        rows, cols = block.shape
        padded = np.zeros((target_size, target_size), dtype=block.dtype)
        padded[:rows, :cols] = block
    else:
        bands, rows, cols = block.shape
        padded = np.zeros((bands, target_size, target_size), dtype=block.dtype)
        padded[:, :rows, :cols] = block

    return padded, rows, cols

主循环可以改成这样:

window_size = 512

for yoff in range(0, height, window_size):
    for xoff in range(0, width, window_size):
        xsize = min(window_size, width - xoff)
        ysize = min(window_size, height - yoff)

        block = src_ds.ReadAsArray(xoff, yoff, xsize, ysize)
        if block is None:
            continue

        padded_block, real_rows, real_cols = pad_to_window(block, window_size)

        padded_pred = predict_block(padded_block)

        pred = padded_pred[:real_rows, :real_cols]

        out_band.WriteArray(pred, xoff, yoff)

这就是很多“滑窗裁切怎么写”问题的核心:不是只会循环,而是要保证边界窗口不会越界,预测结果不会错位。

7. 多类别分类结果写出建议

如果模型输出的是多类别分类图,通常可以把类别编号写成单波段整数栅格:

  • 0:NoData 或背景。
  • 1:水体。
  • 2:建筑。
  • 3:植被。
  • 4:裸地。

这种结果适合后续在 QGIS、ArcGIS Pro 或 Python 中做面积统计、矢量化和精度评价。

如果模型输出的是每个类别的概率,可以写成多波段 Float32 GeoTIFF,每个波段对应一个类别概率。但文件体积会更大,后处理也更复杂。

常见坑

1. ReadAsArray 的 xoff 和 yoff 写反

GDAL 的窗口参数顺序是 xoff, yoff, xsize, ysize。其中 xoff 是列方向,yoff 是行方向。很多错误来自把行列顺序写反,导致读取窗口错位。

2. 多波段数组维度和模型输入维度不一致

GDAL 读取多波段栅格时,返回数组通常是 (bands, rows, cols)。但深度学习模型可能要求:

  • PyTorch 常见输入:(batch, channels, height, width)
  • TensorFlow 常见输入:(batch, height, width, channels)

如果不转换维度,模型可能报错,也可能输出完全错误的预测结果。

3. 输出结果没有投影或位置偏移

如果创建输出影像后没有设置 GeoTransformProjection,预测结果会丢失空间位置。打开后可能显示在错误位置,或者无法和原图叠加。

4. 边界窗口尺寸不匹配

当影像宽高不能被 window_size 整除时,右边和下边的窗口会变小。如果预测函数仍然返回固定的 512×512,而直接写入一个 300×200 的边界区域,就会报错或产生错位。

5. NoData 没有处理

遥感影像边缘、裁剪区域外或云掩膜区域可能存在 NoData。做 Python GDAL滑窗预测时,如果不处理 NoData,模型可能把无效区域也预测成有效类别。

常见处理方式包括:

  • 读取输入波段的 NoData 值。
  • 构建有效像元掩膜。
  • 预测后将无效区域重新赋值为 0 或指定 NoData。

6. 没有 FlushCache 导致文件不完整

写入完成后建议执行 FlushCache(),并把数据集对象设为 None。否则在某些环境下,文件可能还没有完全写入磁盘。

方法比较

方法 适用场景 优点 注意事项
整幅影像一次读取 小影像、测试代码 代码简单,调试方便 大影像容易内存不足
GDAL 固定窗口滑窗 大多数栅格预测任务 内存稳定,容易控制写出位置 需要处理边界窗口和维度转换
GDAL 重叠滑窗 深度学习分割、边缘误差明显的模型 可以减少拼接缝和边缘效应 需要融合重叠区域,代码更复杂
Rasterio window 读取 偏 Pythonic 的栅格处理流程 窗口 API 清晰,和 NumPy 结合方便 已有 GDAL 项目迁移时需要改写接口
先切片后批量预测 需要保存中间瓦片、分布式处理 便于断点续跑和人工检查 会产生大量中间文件,占用磁盘

如果你的目标是快速把模型应用到一幅大 GeoTIFF 上,GDAL 固定窗口滑窗是最直接的方案。如果模型在切片边缘预测不稳定,可以进一步改成带重叠区域的滑窗预测。

检查清单

在运行正式预测前,建议按下面清单检查一遍。

  • 输入影像能否用 gdal.Open 正常打开。
  • 输入影像的波段数是否和模型训练时一致。
  • 输入影像的数据类型、取值范围是否和训练数据一致。
  • ReadAsArray 的参数顺序是否为 xoff, yoff, xsize, ysize
  • 多波段数组是否完成模型所需的维度转换。
  • 边界窗口是否使用 min(window_size, width - xoff)min(window_size, height - yoff) 处理。
  • 固定输入模型是否对边界窗口做了补边。
  • 预测结果尺寸是否和写入窗口尺寸一致。
  • 输出栅格是否复制了原图的投影和仿射变换。
  • 输出数据类型是否适合预测结果,例如分类用 Byte 或 UInt16,概率用 Float32。
  • NoData 区域是否被正确保留或重新赋值。
  • 输出结果是否能在 QGIS 或 ArcGIS Pro 中与原图正确叠加。

FAQ

Python调用GDAL做预测时,窗口大小设置多少合适?

常见窗口大小有 256、512、1024。选择时主要看模型输入尺寸、显存或内存大小、影像波段数。深度学习模型如果训练时使用 512×512,预测时通常也建议使用 512×512。传统机器学习或规则判断可以根据内存适当调大。

滑窗裁切怎么写才不会漏掉右边和下边?

循环时不要假设每个窗口都是完整尺寸,而要使用 min 计算真实窗口大小:

xsize = min(window_size, width - xoff)
ysize = min(window_size, height - yoff)

这样即使影像宽高不能被窗口大小整除,也不会漏掉边缘区域。

为什么预测图在 QGIS 里和原图对不上?

最常见原因是输出 GeoTIFF 没有继承原始影像的 GeoTransformProjection。创建输出文件后必须执行:

out_ds.SetGeoTransform(src_ds.GetGeoTransform())
out_ds.SetProjection(src_ds.GetProjection())

GDAL读取栅格后为什么数组维度和我想的不一样?

单波段栅格通常返回二维数组 (rows, cols),多波段栅格通常返回三维数组 (bands, rows, cols)。如果模型需要通道在最后,需要使用 np.transpose 转换维度。

能不能一边滑窗一边批量预测,提高速度?

可以。更高效的做法是先收集多个窗口组成 batch,再送入模型预测,然后逐个写回。但批量预测需要额外记录每个窗口的 xoffyoffxsizeysize,否则写回时容易错位。

分类结果应该保存成 Byte 还是 Float32?

如果输出是类别编号,例如 0、1、2、3,通常保存为 GDT_ByteGDT_UInt16。如果输出是概率、指数或连续回归值,建议保存为 GDT_Float32

滑窗预测一定要有重叠区域吗?

不一定。传统分类或像元级模型通常不需要重叠。深度学习语义分割模型如果在切片边缘效果差,可以使用重叠滑窗,再对重叠区域做平均、投票或中心区域裁剪。

结论

Python调用GDAL做预测的核心并不复杂:用 ReadAsArray 按窗口读取,用模型对窗口预测,再用 WriteArray 写回对应位置。真正需要注意的是边界窗口、数组维度、NoData、输出数据类型和空间参考。

如果只是普通大影像预测,本文的固定窗口代码已经可以作为模板使用。如果你的模型要求固定输入尺寸,就增加补边和裁回逻辑;如果模型边缘误差明显,再进一步改成重叠滑窗。只要这些细节处理正确,Python GDAL滑窗预测就可以稳定用于实际 GIS 和遥感生产流程。