如何移除3D数组中含NaN的子数组并维持三维维度结构
问题:移除三维稀疏数组中的NaN并保留三维结构
我有一个形状为(863, 923, 2)的稀疏数组,其中包含大量NaN,数组示例如下:
[[[ 43.06010628 -11.01121568] [ 25.03068277 16.3949826 ] [-23.75853158 -10.95350074] ... [ 25.52110353 3.00428452] [ 32.66945663 9.76115107] [ 19.1341548 8.48547008]] [[ 19.08099208 11.27167832] [-29.4360534 -12.39131814] [ 11.24612069 14.38915742] ... [ 16.6897315 10.04601296] [ 30.09409518 17.09382562] [ -9.47312129 -9.57484782]] [[ 21.22006655 -5.01340343] [ 11.65512749 2.32398374] [-22.14668148 -11.05883399] ... [ nan nan] [ nan nan] [ nan nan]] ... [[ 32.32522443 -3.73563526] [ 30.88408144 -2.92184744] [ 37.44548043 -21.8209554 ] ... [ nan nan] [ nan nan] [ nan nan]] [[ 36.85471348 -7.86696711] [ 37.20204074 -6.32105844] [ 32.32522443 -3.73563526] ... [ nan nan] [ nan nan] [ nan nan]] [[ 34.21397091 -5.88930588] [ 35.88819735 -7.64992589] [ 35.48958094 -10.34708285] ... [ nan nan] [ nan nan] [ nan nan]]]
我需要移除所有包含NaN的子数组,同时保留三维维度结构,预期形状为类似(m, n, 2)的形式,但尝试以下代码时触发报错:
nonnanarr = arr[~np.isnan(arr).any(axis=-1)].reshape((863, -1, 2))
报错信息:
Traceback (most recent call last): File "c:\Users\username\Desktop\observables\my_script.py", line 167, in <module> main() File "c:\Users\username\Desktop\observables\my_script.py", line 104, in main time_stamp_num, agents_num, spatial_dimensions_num = dataframe_splitter() File "c:\Users\username\Desktop\observables\utilities.py", line 1351, in dataframe_splitter nonnan_arr = arr[~np.isnan(arr).any(axis=-1)].reshape( ValueError: cannot reshape array of size 226512 into shape (863,newaxis,2)
问题原因
你的代码尝试将过滤后的一维数组强行reshape为(863, -1, 2),但原数组每个第一维度(共863个)的子数组中,非NaN元素的数量不一致,导致总元素数226512无法被863*2整除,reshape失败。
解决方案
方案1:统一截断到最小非NaN元素数(规整三维数组)
如果可以接受截断部分有效数据,将所有第一维度的子数组统一保留到相同长度:
import numpy as np # 生成掩码:标记每个(2,)子数组是否不含NaN mask = ~np.isnan(arr).any(axis=-1) # 计算每个第一维度子数组的非NaN元素数量 non_nan_counts = mask.sum(axis=1) # 取所有子数组中的最小非NaN数量作为统一长度 min_valid_count = non_nan_counts.min() # 过滤每个子数组并截断到统一长度 nonnan_arr = np.array([sub_arr[mask[i]][:min_valid_count] for i, sub_arr in enumerate(arr)]) # 最终形状为(863, min_valid_count, 2)
方案2:保留所有有效元素(可变长度数组)
如果需要保留全部非NaN元素,无法形成规整三维数组,可以使用numpy的object类型数组:
import numpy as np mask = ~np.isnan(arr).any(axis=-1) # 每个元素是对应子数组的非NaN部分,形状为(n_i, 2) nonnan_arr = np.array([sub_arr[mask[i]] for i, sub_arr in enumerate(arr)], dtype=object) # 数组整体形状为(863,),每个元素是二维数组
方案3:填充NaN保留原形状
如果需要保持原数组的(863,923,2)形状,仅移除NaN并将有效元素前置:
import numpy as np mask = ~np.isnan(arr).any(axis=-1) # 初始化一个和原数组形状相同的NaN数组 nonnan_arr = np.full_like(arr, np.nan) for i in range(arr.shape[0]): # 获取当前子数组的所有有效元素 valid_elements = arr[i][mask[i]] # 将有效元素填充到新数组的对应位置 nonnan_arr[i][:len(valid_elements)] = valid_elements # 最终形状仍为(863,923,2)
内容的提问来源于stack exchange,提问作者Roosha
相关产品推荐
相关产品推荐

