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

双波段大TIFF文件分块优化:过滤无效块并提升处理速度

问题描述

我有一个包含两个波段的大型TIFF文件,文件中大部分区域为NaN值或0值,仅部分区域存在有效信息。我尝试将该图像切割为64×64的图像块,仅保存非全NaN/全0的块,同时跳过已生成的块。但目前仍会生成全NaN的块,且处理速度极慢,还有多个同类文件需要处理,求优化方案。

原代码如下:

import os
from osgeo import gdal
import numpy as np

# 存放TIFF文件的输入文件夹
input_folder = 'path/to/input_folder'

# 图像块的输出文件夹
output_folder = 'path/to/output_folder'

# 图像块尺寸
tile_size = 64

# 若输出文件夹不存在则创建
if not os.path.exists(output_folder):
    os.makedirs(output_folder)

# 遍历输入文件夹中的所有文件
for file in os.listdir(input_folder):
    # 检查是否为TIFF文件
    if file.endswith(".tif"):
        # 打开输入TIFF文件
        ds = gdal.Open(os.path.join(input_folder, file))

        # 获取图像的宽度和高度
        width = ds.RasterXSize
        height = ds.RasterYSize

        # 计算图像块的行列数量
        num_cols = width // tile_size
        num_rows = height // tile_size

        # 遍历图像块的行和列
        for i in range(num_rows):
            for j in range(num_cols):
                # 计算图像块的x、y偏移量
                xoff = j * tile_size
                yoff = i * tile_size
                
                # 生成输出文件名
                output_file = os.path.join(output_folder, f"{file}_row{i}_col{j}_tile.tif")
                
                # 若输出文件已存在则跳过
                if os.path.exists(output_file):
                    continue
                
                # 使用gdal_translate裁剪图像块
                tile_ds = gdal.Translate(output_file, ds, srcWin=[xoff, yoff, tile_size, tile_size])
                
                # 读取所有波段的图像块数据
                tile_data = [tile_ds.GetRasterBand(band).ReadAsArray() for band in range(1, ds.RasterCount + 1)]
                
                # 检查所有波段是否均为全NaN或全0
                if all(np.all(np.isnan(band_data)) or np.all(band_data == 0) for band_data in tile_data):
                    # 若所有波段均无效则删除输出文件
                    os.remove(output_file)
                    continue  # 跳过保存该无效图像块
                
                tile_ds = None  # 关闭数据集
                
        # 获取所有行的最后一块图像
        for i in range(num_rows):
            xoff = num_cols * tile_size
            yoff = i * tile_size
            output_file = os.path.join(output_folder, f"{file}_row{i}_last_col.tif")
            remaining_width = width - (num_cols * tile_size)
            gdal.Translate(output_file, ds, srcWin=[xoff, yoff, remaining_width, tile_size])
            
        ds = None  # 关闭数据集

print('Finished!')
优化方案及代码

核心优化点

  • 先判断后写入:避免先生成文件再删除的无效IO操作,先读取区域数据判断有效性,再决定是否保存
  • 减少GDAL文件操作:直接用ReadAsArray读取指定区域,跳过中间临时文件生成步骤
  • 向量化判断逻辑:用numpy批量操作替代循环判断,提升速度
  • 优化存在性检查:提前缓存已存在的文件名,减少磁盘IO
  • 覆盖剩余区域的有效性判断:补全原代码中未处理的最后一列块判断逻辑,避免无效块生成

优化后代码

import os
from osgeo import gdal, gdal_array
import numpy as np

input_folder = 'path/to/input_folder'
output_folder = 'path/to/output_folder'
tile_size = 64

if not os.path.exists(output_folder):
    os.makedirs(output_folder)

# 提前缓存已存在的输出文件名,减少循环内磁盘查询开销
existing_files = set(os.listdir(output_folder))

for file in os.listdir(input_folder):
    if not file.endswith(".tif"):
        continue
        
    file_path = os.path.join(input_folder, file)
    ds = gdal.Open(file_path)
    if ds is None:
        print(f"无法打开文件: {file_path}")
        continue
        
    width = ds.RasterXSize
    height = ds.RasterYSize
    num_cols = width // tile_size
    num_rows = height // tile_size
    driver = gdal.GetDriverByName('GTiff')
    
    # 处理完整的64x64图像块
    for i in range(num_rows):
        for j in range(num_cols):
            xoff = j * tile_size
            yoff = i * tile_size
            output_filename = f"{file}_row{i}_col{j}_tile.tif"
            
            # 检查文件是否已存在,直接用集合查询加速
            if output_filename in existing_files:
                continue
                
            # 一次性读取该区域所有波段数据,避免生成临时文件
            tile_data = ds.ReadAsArray(xoff, yoff, tile_size, tile_size)
            # 批量判断是否存在有效数据:只要有一个元素非NaN且非0,就保留该块
            has_valid = np.any((~np.isnan(tile_data)) & (tile_data != 0))
            
            if not has_valid:
                continue
                
            # 创建输出数据集,保留原文件的空间信息
            output_path = os.path.join(output_folder, output_filename)
            out_ds = driver.Create(
                output_path, 
                tile_size, tile_size, 
                ds.RasterCount, 
                gdal_array.NumericTypeCodeToGDALTypeCode(tile_data.dtype)
            )
            out_ds.SetProjection(ds.GetProjection())
            out_ds.SetGeoTransform((
                ds.GetGeoTransform()[0] + xoff * ds.GetGeoTransform()[1],
                ds.GetGeoTransform()[1],
                ds.GetGeoTransform()[2],
                ds.GetGeoTransform()[3] + yoff * ds.GetGeoTransform()[5],
                ds.GetGeoTransform()[4],
                ds.GetGeoTransform()[5]
            ))
            
            # 写入各波段数据并复制NoData值
            for band_idx in range(ds.RasterCount):
                out_band = out_ds.GetRasterBand(band_idx + 1)
                out_band.WriteArray(tile_data[band_idx])
                out_band.SetNoDataValue(ds.GetRasterBand(band_idx + 1).GetNoDataValue())
            
            out_ds.FlushCache()
            out_ds = None
    
    # 处理每行剩余宽度的非标准尺寸块
    remaining_width = width - num_cols * tile_size
    if remaining_width > 0:
        for i in range(num_rows):
            xoff = num_cols * tile_size
            yoff = i * tile_size
            output_filename = f"{file}_row{i}_last_col.tif"
            
            if output_filename in existing_files:
                continue
                
            tile_data = ds.ReadAsArray(xoff, yoff, remaining_width, tile_size)
            has_valid = np.any((~np.isnan(tile_data)) & (tile_data != 0))
            
            if not has_valid:
                continue
                
            output_path = os.path.join(output_folder, output_filename)
            out_ds = driver.Create(
                output_path, 
                remaining_width, tile_size, 
                ds.RasterCount, 
                gdal_array.NumericTypeCodeToGDALTypeCode(tile_data.dtype)
            )
            out_ds.SetProjection(ds.GetProjection())
            out_ds.SetGeoTransform((
                ds.GetGeoTransform()[0] + xoff * ds.GetGeoTransform()[1],
                ds.GetGeoTransform()[1],
                ds.GetGeoTransform()[2],
                ds.GetGeoTransform()[3] + yoff * ds.GetGeoTransform()[5],
                ds.GetGeoTransform()[4],
                ds.GetGeoTransform()[5]
            ))
            
            for band_idx in range(ds.RasterCount):
                out_band = out_ds.GetRasterBand(band_idx + 1)
                out_band.WriteArray(tile_data[band_idx])
                out_band.SetNoDataValue(ds.GetRasterBand(band_idx + 1).GetNoDataValue())
            
            out_ds.FlushCache()
            out_ds = None
    
    ds = None

print('Finished!')

关键优化说明

  1. 集合缓存已存在文件:将输出文件夹内的文件名存入集合,循环内直接做内存查询,比多次调用os.path.exists减少磁盘IO开销
  2. 直接读取区域数据:用ds.ReadAsArray一次性读取目标区域的多波段数据,避免gdal_translate生成临时文件的冗余操作
  3. 向量化有效性判断:通过numpy批量运算检查所有元素,替代原代码中逐波段循环判断的逻辑,大幅提升判断速度
  4. 补全剩余块判断:对每行最后一个非标准尺寸的块也执行有效性检查,彻底避免无效块生成
  5. 保留空间信息:创建输出文件时同步复制原文件的投影和地理变换参数,确保输出图像块的空间属性完整

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 06:29:53