如何将SQLite导出元组生成的一维NumPy数组转为二维矩阵供Keras使用
你操作失败的核心原因是你生成的是带命名字段的NumPy结构化数组,而非普通的二维同质数值数组:
- 结构化数组本质是一维数组,每个元素对应你从数据库取出的一行记录,包含12个不同类型的命名字段,所以直接用二维索引语法
arr[:, 列号]会报错 - 结构化数组的
size属性统计的是外层元素的总数,也就是你的数据行数(7826行),并非所有字段的总数值个数,所以调用np.reshape(arr,(12,len(arr)))时,系统会判断12*7826=93912远大于数组总size 7826,直接抛出维度不匹配错误。
方法一:直接基于结构化数组的字段索引提取特征
不需要转换数组格式,直接按你定义的dtype字段名提取对应列即可:
# 提取第1-5列(对应100-1到100-5字段)作为第一个输入,形状为(样本数, 5) input_1 = arr[["100-1", "100-2", "100-3", "100-4", "100-5"]].view(int).reshape(len(arr), 5) # 提取第6-10列(对应200-1到200-5字段)作为第二个输入,形状为(样本数, 5) input_2 = arr[["200-1", "200-2", "200-3", "200-4", "200-5"]].view(int).reshape(len(arr), 5) # 如需标签可单独提取sideWon字段 label = arr["sideWon"]
这里的view(int)是将提取的结构化子数组转换为普通的int类型数组,再通过reshape调整为二维格式,可直接输入Keras模型。
方法二:转换为普通二维数组后按列号索引
如果你更习惯用数字索引列,可以先将数值字段统一转换为普通二维数组:
# 排除第一列字符串类型的matchId,提取所有11个数值字段转换为二维数组 numeric_arr = arr[list(dataType.names[1:])].view(int).reshape(len(arr), -1) # 现在numeric_arr的形状为(7826, 11),支持你熟悉的二维索引语法 input_1 = numeric_arr[:, :5] # 第0-4列对应原数据的第1-5列 input_2 = numeric_arr[:, 5:10] # 第5-9列对应原数据的第6-10列 label = numeric_arr[:, 10]
内容的提问来源于stack exchange,提问作者Michael Ngo
相关产品推荐
相关产品推荐

