基于xarray优化TensorFlow CNN训练的卫星图像生成效率
优化建议
1. 用纯TensorFlow逻辑重构核心流程,替换tf.py_function
tf.py_function会脱离TensorFlow计算图,无法进行图优化和预取加速,必须把采样、索引计算等逻辑转成TensorFlow原生操作:
预处理多边形与卫星数据
提前提取所有多边形的边界、标签,以及卫星影像的坐标,转成TensorFlow张量存储,避免在数据生成阶段反复调用geopandas或xarray:
# 提前提取多边形边界和标签 poly_bounds = tf.convert_to_tensor(gdf.bounds.values, dtype=tf.float32) labels = tf.convert_to_tensor(gdf["label"].values, dtype=tf.float32) # 提前转换卫星影像坐标为TensorFlow张量 xr_x = tf.convert_to_tensor(dataset.x.values, dtype=tf.float32) xr_y = tf.convert_to_tensor(dataset.y.values, dtype=tf.float32) # 提前加载整个卫星影像为TensorFlow张量(关键优化,避免动态读取xarray数据) sat_img_tensor = tf.convert_to_tensor(dataset.band_data.values, dtype=tf.uint8)
TensorFlow版本的索引查找与点采样
替换原有的numpy/shapely实现为TensorFlow原生函数:
@tf.function def tf_find_nearest_idx(array, value): abs_diff = tf.abs(array - value) return tf.cast(tf.argmin(abs_diff, axis=0), tf.int32) @tf.function def tf_random_point_in_polygon(bounds): minx, miny, maxx, maxy = bounds # 若需精确判断点是否在多边形内,可使用tensorflow_io的tfio.experimental.geometry.contains_point # 这里先实现边界内采样,后续可结合掩码优化 x = tf.random.uniform((), minx, maxx) y = tf.random.uniform((), miny, maxy) return x, y
重构数据生成函数为tf.function
@tf.function def tf_get_data(i): # 获取当前多边形的边界和标签 bounds = poly_bounds[i] label = labels[i] # 采样随机点 x, y = tf_random_point_in_polygon(bounds) # 查找最近栅格索引 idx_x = tf_find_nearest_idx(xr_x, x) idx_y = tf_find_nearest_idx(xr_y, y) # 计算图像切片范围 half_size = tf.cast(img_size / 2, tf.int32) idx_x_min = idx_x - half_size idx_x_max = idx_x + half_size idx_y_min = idx_y - half_size idx_y_max = idx_y + half_size # 从预加载的张量中切片获取图像 image = tf.slice(sat_img_tensor, [idx_y_min, idx_x_min, 0], [img_size, img_size, -1]) return image, label
2. 优化tf.data管道,提升并行与预取能力
调整数据集管道,加入打乱、批处理、预取等操作,最大化GPU利用率:
# 创建索引数据集 dataset = tf.data.Dataset.range(gdf.shape[0]) # 使用tf.function的map替换tf.py_function,开启并行处理 dataset = dataset.map(tf_get_data, num_parallel_calls=tf.data.AUTOTUNE) # 打乱数据(缓冲区大小设为数据集规模,保证充分打乱) dataset = dataset.shuffle(buffer_size=gdf.shape[0]) # 设置批次大小(根据GPU显存调整) dataset = dataset.batch(32) # 预取数据,让GPU始终有数据可处理 dataset = dataset.prefetch(tf.data.AUTOTUNE)
3. 优化复杂多边形的点采样效率
如果多边形是复杂非凸形状,循环采样点的效率极低,可提前为每个多边形生成采样点池,训练时直接从池中随机选取:
# 预处理阶段为每个多边形生成1000个有效随机点 sampled_points = [] for poly in gdf.geometry: points = [] for _ in range(1000): x, y = random_point_from_geometry(poly) points.append([x[0], y[0]]) sampled_points.append(points) sampled_points_tensor = tf.convert_to_tensor(sampled_points, dtype=tf.float32) # 替换采样函数为从预生成池中取点 @tf.function def tf_random_point_from_pool(i): points = sampled_points_tensor[i] idx = tf.random.uniform((), 0, tf.shape(points)[0], dtype=tf.int32) return points[idx][0], points[idx][1]
内容的提问来源于stack exchange,提问作者Queeno11
相关产品推荐
相关产品推荐

