NumPy保存训练数据报错:非均匀维度形状问题排查
解决Numpy保存异构训练数据的ValueError错误
错误原因
你遇到的问题是因为training_data里的每个元素是**(图像数组, 按键标签列表)的组合,numpy的np.save()要求数组必须是同构形状**的(即所有元素的维度和大小完全一致),但你的数据里,图像是(160,120)的二维数组,标签是长度为3的一维列表,直接转numpy数组会出现形状异构的问题,导致无法保存。
解决方案
下面提供两种可行的修改方案:
方案一:分开保存图像和标签(推荐,适合后续CNN训练)
把图像数据和标签数据拆分成两个独立的numpy数组分别保存,既符合numpy的要求,也方便后续模型加载训练。
修改后的代码:
import numpy as np from grabscreen import grab_screen import cv2 import time from getkeys import key_check import os def keys_to_output(keys): output = [0,0,0] if 'A' in keys: output[0] = 1 elif 'D' in keys: output[2] = 1 else: output[1] = 1 return output # 定义两个保存文件 img_file = 'training_images.npy' label_file = 'training_labels.npy' training_data = [] if os.path.isfile(img_file) and os.path.isfile(label_file): print('文件已存在,加载历史数据!') # 加载图像和标签,重新组合成training_data列表 images = np.load(img_file) labels = np.load(label_file) training_data = list(zip(images, labels)) else: print('文件不存在,开始新采集!') training_data = [] def main(): for i in list(range(4))[::-1]: print(i+1) time.sleep(1) paused = False while(True): if not paused: screen = grab_screen(region=(0,40,800,640)) screen = cv2.cvtColor(screen, cv2.COLOR_BGR2GRAY) screen = cv2.resize(screen, (160,120)) keys = key_check() output = keys_to_output(keys) training_data.append([screen, output]) if len(training_data) % 1000 == 0: print(f'当前数据量:{len(training_data)}') # 拆分图像和标签为独立数组 images = np.array([item[0] for item in training_data]) labels = np.array([item[1] for item in training_data]) # 分别保存 np.save(img_file, images) np.save(label_file, labels) print('数据保存完成!') keys = key_check() if 'T' in keys: if paused: paused = False print('已取消暂停!') time.sleep(1) else: print('正在暂停!') paused = True time.sleep(1) main()
方案二:使用Pickle保存原生Python列表
如果不想拆分数据,可以用Python的pickle模块直接保存整个training_data列表,它支持保存任意Python结构,包括异构的元素组合。
修改后的代码:
import numpy as np import pickle from grabscreen import grab_screen import cv2 import time from getkeys import key_check import os def keys_to_output(keys): output = [0,0,0] if 'A' in keys: output[0] = 1 elif 'D' in keys: output[2] = 1 else: output[1] = 1 return output file_name = 'training_data.pkl' training_data = [] if os.path.isfile(file_name): print('文件已存在,加载历史数据!') with open(file_name, 'rb') as f: training_data = pickle.load(f) else: print('文件不存在,开始新采集!') training_data = [] def main(): for i in list(range(4))[::-1]: print(i+1) time.sleep(1) paused = False while(True): if not paused: screen = grab_screen(region=(0,40,800,640)) screen = cv2.cvtColor(screen, cv2.COLOR_BGR2GRAY) screen = cv2.resize(screen, (160,120)) keys = key_check() output = keys_to_output(keys) training_data.append([screen, output]) if len(training_data) % 1000 == 0: print(f'当前数据量:{len(training_data)}') with open(file_name, 'wb') as f: pickle.dump(training_data, f) print('数据保存完成!') keys = key_check() if 'T' in keys: if paused: paused = False print('已取消暂停!') time.sleep(1) else: print('正在暂停!') paused = True time.sleep(1) main()
注意事项
- 方案一的优势是后续训练时可以直接加载numpy数组,无需额外转换,更适合深度学习框架(如TensorFlow/PyTorch)的输入要求。
- 方案二适合需要保持原始数据结构的场景,但加载速度可能略慢于numpy数组。
- 如果你之前已经有损坏的
training_data.npy文件,建议先删除它,避免加载时出现异常。
内容的提问来源于stack exchange,提问作者DIKTOR
相关产品推荐
相关产品推荐

