双波段大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!')
关键优化说明
- 集合缓存已存在文件:将输出文件夹内的文件名存入集合,循环内直接做内存查询,比多次调用
os.path.exists减少磁盘IO开销 - 直接读取区域数据:用
ds.ReadAsArray一次性读取目标区域的多波段数据,避免gdal_translate生成临时文件的冗余操作 - 向量化有效性判断:通过numpy批量运算检查所有元素,替代原代码中逐波段循环判断的逻辑,大幅提升判断速度
- 补全剩余块判断:对每行最后一个非标准尺寸的块也执行有效性检查,彻底避免无效块生成
- 保留空间信息:创建输出文件时同步复制原文件的投影和地理变换参数,确保输出图像块的空间属性完整
内容的提问来源于stack exchange,提问作者SNunes
相关产品推荐
相关产品推荐

