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

如何优化含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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 08:03:28