如何将元组列表转换为3D Numpy数组?解决形状不匹配报错
如何将分组元组列表转换为3D NumPy数组
问题背景
给定元组列表数据集:
data = [(1, 65, -18, -1, -1 ), (1, -18,-1, -1,-1), (2, 65, -19, -1, -1), (2, 65, -18, -1, -1), (3, 62, -18, -1, -1)]
需要转换为结构如下的3D NumPy数组:
array([[[[65], [-18], [-1], [-1]], [[-18], [-1], [-1], [-1]]], [[[65], [-19], [-1], [-1]], [[65], [-18], [-1], [-1]]], [[[62], [-18], [-1], [-1]]]])
尝试以下代码后,得到的是数组列表而非单个3D数组;使用np.stack时因分组长度不一致报错ValueError: all input arrays must have the same shape:
from collections import defaultdict import numpy as np d = defaultdict(list) for item in data: d[item[0]].append((item[1:5])) values_list = [np.array(v) for v in d.values()] my_array = np.array(values_list) print(my_array)
解决方案
方案1:生成不规则结构的3D数组(object类型)
如果允许数组内的子分组长度不同,可先将每个元组的后4个元素转换为(4,1)维度的数组,再按分组组合:
from collections import defaultdict import numpy as np data = [(1, 65, -18, -1, -1 ), (1, -18,-1, -1,-1), (2, 65, -19, -1, -1), (2, 65, -18, -1, -1), (3, 62, -18, -1, -1)] # 按第一个元素分组,同时将每个子项转为(4,1)数组 grouped = defaultdict(list) for item in data: sub_arr = np.array(item[1:5]).reshape(-1, 1) grouped[item[0]].append(sub_arr) # 组合成分组数组,再转为最终的object类型3D数组 result = np.array([np.array(g) for g in grouped.values()], dtype=object) print(result)
输出结构与期望完全一致,整体为形状(3,)的object数组,其中每个元素的形状分别为(2,4,1)、(2,4,1)、(1,4,1)。
方案2:生成形状统一的3D数组(填充补全)
如果需要严格统一的数组形状,可找到最大分组长度,用填充值(如np.nan)补全短分组:
from collections import defaultdict import numpy as np data = [(1, 65, -18, -1, -1 ), (1, -18,-1, -1,-1), (2, 65, -19, -1, -1), (2, 65, -18, -1, -1), (3, 62, -18, -1, -1)] grouped = defaultdict(list) for item in data: sub_arr = np.array(item[1:5]).reshape(-1, 1) grouped[item[0]].append(sub_arr) # 获取最长分组的长度 max_group_len = max(len(g) for g in grouped.values()) # 对每个分组进行填充补全 padded_groups = [] for group in grouped.values(): # 计算需要填充的数量 pad_count = max_group_len - len(group) # 生成填充数组,形状与分组内元素一致 padding = np.full((pad_count, 4, 1), np.nan) # 拼接原分组与填充数组 padded_group = np.concatenate([np.array(group), padding], axis=0) padded_groups.append(padded_group) # 转为统一形状的3D数组 result = np.array(padded_groups) print(result) print(result.shape) # 输出 (3, 2, 4, 1)
此方案得到的数组形状完全统一,缺失项用np.nan填充,适合后续批量数值运算。
失败原因说明
之前的代码中,分组后的数组形状为(2,4)、(2,4)、(1,4),缺少每个元素的最后一维(即未将单个值转为[val]形式);同时np.stack要求所有输入数组的形状完全一致,因此分组长度不同会触发报错。
内容的提问来源于stack exchange,提问作者Hadis
相关产品推荐
相关产品推荐

