列表转NumPy数组因形状不均报错,求数据加载代码修复方案
问题
我尝试用load_data_new函数从topomaps/文件夹读取拓扑图数据、从labels/文件夹读取标签数据,二者都是.npy格式文件。其中topomaps/下的文件形状不一致,比如s01_trial03.npy包含128个拓扑图,s01_trial12.npy包含2944个拓扑图。需求是训练集仅保留标签为0的拓扑图,测试集可包含标签为0、1、2的数据。运行代码时,在执行x = np.array(topomaps)步骤触发报错:
Traceback (most recent call last):
File "/Users/alex/PycharmProjects/VAE-EEG-XAI/vae.py", line 574, in
(x_train, y_train), (x_test, y_test) = load_data_new("topomaps", "labels")
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/Users/alex/PycharmProjects/VAE-EEG-XAI/vae.py", line 60, in load_data_new
x = np.array(topomaps)
^^^^^^^^^^^^^^^^^^
ValueError: setting an array element with a sequence. The requested array has an inhomogeneous shape after 1 dimensions. The detected shape was (851,) + inhomogeneous part.
报错原因是topomaps列表内的元素形状不一致,而NumPy数组要求元素形状统一,请问该如何修复?
解决方案
方法1:扁平化数据,统一维度(推荐,适用于单张拓扑图特征维度一致的场景)
如果每个拓扑图的特征维度(比如都是(32,32))是统一的,只是每个文件内的拓扑图数量不同,直接把所有拓扑图逐个展开到列表,再转成NumPy数组,同时对应处理标签:
import numpy as np import os def load_data_new(topomaps_dir, labels_dir): all_topomaps = [] all_labels = [] # 遍历所有npy文件,确保topomaps和labels文件名一一对应 for filename in os.listdir(topomaps_dir): if not filename.endswith('.npy'): continue # 读取拓扑图和对应标签 topo_path = os.path.join(topomaps_dir, filename) label_path = os.path.join(labels_dir, filename) topo_data = np.load(topo_path) label_data = np.load(label_path) # 校验数据匹配性 assert len(topo_data) == len(label_data), f"文件{filename}的拓扑图与标签数量不匹配" # 逐个添加元素到总列表 all_topomaps.extend(topo_data) all_labels.extend(label_data) # 转成统一形状的NumPy数组 x = np.array(all_topomaps) y = np.array(all_labels) # 划分训练集(仅保留标签0)和测试集(保留0/1/2) train_mask = y == 0 x_train, y_train = x[train_mask], y[train_mask] # 测试集可选择从全量数据中划分,这里示例用训练集之外的样本 test_mask = np.logical_not(train_mask) x_test, y_test = x[test_mask], y[test_mask] return (x_train, y_train), (x_test, y_test)
方法2:填充/截断到统一长度(适用于需要保留文件级结构的场景)
如果必须保留每个文件作为一个样本维度(比如按trial划分),可以将所有文件的拓扑图填充到最大数量,或截断到最小数量:
import numpy as np import os def load_data_new(topomaps_dir, labels_dir): all_topomaps = [] all_labels = [] max_length = 0 # 第一步:获取所有文件中拓扑图的最大数量 for filename in os.listdir(topomaps_dir): if filename.endswith('.npy'): topo_data = np.load(os.path.join(topomaps_dir, filename)) max_length = max(max_length, len(topo_data)) # 第二步:处理每个文件,统一长度 for filename in os.listdir(topomaps_dir): if not filename.endswith('.npy'): continue topo_path = os.path.join(topomaps_dir, filename) label_path = os.path.join(labels_dir, filename) topo_data = np.load(topo_path) label_data = np.load(label_path) # 填充或截断拓扑图 if len(topo_data) < max_length: # 用0填充,可根据需求替换为均值等其他值 pad_topo = np.pad(topo_data, ((0, max_length - len(topo_data)), (0,0), (0,0)), mode='constant') # 用-1标记填充的无效标签 pad_label = np.pad(label_data, (0, max_length - len(topo_data)), mode='constant', constant_values=-1) else: pad_topo = topo_data[:max_length] pad_label = label_data[:max_length] all_topomaps.append(pad_topo) all_labels.append(pad_label) x = np.array(all_topomaps) y = np.array(all_labels) # 提取训练集有效样本(标签0且非填充值) train_mask = (y == 0) & (y != -1) x_train = x[train_mask].reshape(-1, *x.shape[2:]) y_train = y[train_mask].reshape(-1) # 提取测试集有效样本(标签0/1/2且非填充值) test_mask = (y >= 0) & (y <= 2) x_test = x[test_mask].reshape(-1, *x.shape[2:]) y_test = y[test_mask].reshape(-1) return (x_train, y_train), (x_test, y_test)
方法3:直接按标签分类收集(适用于灵活划分数据集的场景)
遍历每个拓扑图和标签,直接将符合条件的样本分到训练集或测试集,最后统一转成数组:
import numpy as np import os def load_data_new(topomaps_dir, labels_dir): train_topomaps = [] train_labels = [] test_topomaps = [] test_labels = [] for filename in os.listdir(topomaps_dir): if not filename.endswith('.npy'): continue topo_data = np.load(os.path.join(topomaps_dir, filename)) label_data = np.load(os.path.join(labels_dir, filename)) # 逐个判断标签,分配到对应数据集 for topo, label in zip(topo_data, label_data): if label == 0: train_topomaps.append(topo) train_labels.append(label) elif label in [1,2]: test_topomaps.append(topo) test_labels.append(label) # 若需要将部分标签0的样本放入测试集,可添加随机划分逻辑 # 转成统一形状的数组 x_train = np.array(train_topomaps) y_train = np.array(train_labels) x_test = np.array(test_topomaps) y_test = np.array(test_labels) return (x_train, y_train), (x_test, y_test)
内容的提问来源于stack exchange,提问作者tail

