如何优化含NaN值的大型4D数组以降低内存占用?
问题描述
我有一个形状为shape=(1, 3, 1000000, 112)的大型数组data_neighbors,其中包含大量NaN值。示例数组如下:
array([[[[ 88.769226, 80.62714 , 75.95856 ]],...[[ nan, nan, nan]]]], dtype=float32)
如何移除该数组中的所有NaN值以提升内存使用效率?需注意最后一个维度的NaN数量不固定,例如data_neighbors[0,0,0].shape=3,data_neighbors[0,0,1].shape=112,无法得到规整数组,是否可用数组嵌套列表形式?
补充背景
脚本目标是实现多点重网格化,为网格A每个点分配网格B中其周围x公里内的对应值,该值由ind_regrid(形状1000000*112)确定,该变量包含网格A各索引对应的待整合网格B点,其112个潜在重网格化索引可能含NaN值。相关代码:
nc_conf = Dataset(fic_regril, 'r') print('-> Read regrid file '+str(fic_regril)) ind_regrid = nc_conf.variables['inds_regrid'][:] nc_conf.close() masked_indices = np.ma.getmaskarray(ind_regrid) data_neighbors = data[:,:,:,np.where(~masked_indices,ind_regrid,0)] data_neighbors[masked_indices] = np.nan data_neighbors_list.append(data_neighbors) #pt, regrid, param, time
解决方案
完全可以用嵌套列表存储去NaN后的数据,既能节省内存,又适配每个位置长度不固定的情况,以下是两种实现方式:
1. 从已生成的data_neighbors数组转换
如果已经得到了含NaN的数组,直接逐元素过滤生成嵌套列表:
# 压缩掉第一个长度为1的维度,简化后续遍历 squeezed_data = data_neighbors.squeeze(axis=0) # shape变为(3, 1000000, 112) filtered_list = [] # 遍历每个参数维度 for param_group in squeezed_data: point_list = [] # 遍历每个网格点 for point_data in param_group: # 过滤当前点的所有NaN值 valid_values = point_data[~np.isnan(point_data)] point_list.append(valid_values.tolist()) filtered_list.append(point_list)
最终的filtered_list结构为[参数1的网格点数据列表, 参数2的网格点数据列表, 参数3的网格点数据列表],每个网格点对应不含NaN的数值列表。
2. 优化原代码,跳过NaN索引直接生成有效数据(更高效)
不需要先生成全量含NaN的数组,直接在索引阶段过滤无效值,能大幅节省内存和计算资源:
nc_conf = Dataset(fic_regril, 'r') print('-> Read regrid file '+str(fic_regril)) ind_regrid = nc_conf.variables['inds_regrid'][:] # shape=(1000000, 112) nc_conf.close() filtered_list = [] # 遍历data的参数维度(对应原代码中data的第二个维度) for param_idx in range(data.shape[1]): current_param_data = data[0, param_idx] # 取出当前参数的全量数据 param_point_list = [] # 遍历每个网格点 for point_idx in range(ind_regrid.shape[0]): # 获取当前网格点的有效索引(排除NaN) valid_indices = ind_regrid[point_idx][~np.isnan(ind_regrid[point_idx])] # 直接提取有效数据,无需填充NaN valid_values = current_param_data[:, valid_indices].tolist() param_point_list.append(valid_values) filtered_list.append(param_point_list)
这种方式避免了生成中间大数组,直接构建有效数据的嵌套列表,内存效率提升明显。
内存优化效果对比
- 原
data_neighbors数组(float32类型)全量内存占用:1*3*1000000*112*4 ≈ 1.3GB,若大部分是NaN,内存浪费严重。 - 嵌套列表仅存储有效数值,假设每个网格点平均有效数据为10个,内存占用约为
3*1000000*10*4 ≈ 120MB,内存占用降低90%以上。
内容的提问来源于stack exchange,提问作者Alexis Vdv
相关产品推荐
相关产品推荐

