如何将图像字符串数据集重塑为(7049,96,96)格式数组?
解决方法:将Series中的二维数组合并为三维数组
你遇到的问题很常见——当你用apply处理pandas Series时,返回的每个元素是二维数组,但整个Series的形状还是一维的。要把这些二维数组堆叠成一个三维数组,只需要一步额外的操作:用numpy.stack(或者直接转成numpy数组)来组合它们。
修正后的代码
首先,确保你的函数把字符串转换成数值类型(因为split得到的是字符串,直接转数组是字符串类型,后续处理可能有问题),然后用np.stack来合并:
import numpy as np import pandas as pd def split_reshape(row): # 注意:如果你的字符串开头的数字是图像索引(比如示例里的0、1、2),需要去掉第一个元素 # 比如改成:pixel_strings = row.split(' ')[1:] pixel_strings = row.split(' ') # 转换为整数类型数组,再reshape return np.array(pixel_strings, dtype=np.int32).reshape(96, 96) # 处理每个图像字符串 processed_series = train_x.apply(split_reshape) # 将Series中的所有二维数组合并为三维数组 result_array = np.stack(processed_series.values) # 验证形状 print(result_array.shape) # 输出 (7049, 96, 96)
为什么之前的方法不行?
当你调用train_x.apply(split_reshape)时,返回的是一个pandas Series,其中每个元素是一个(96,96)的numpy数组。这个Series本身的形状是(7049,),因为它是一维的容器,里面装着二维数组。np.stack会把这些二维数组沿着新的轴(默认是第0轴)堆叠起来,最终形成一个三维数组。
额外注意点:检查像素数量
一定要确认每个字符串split后的元素数量是96*96=9216个。从你的示例输出看,每个字符串开头有一个数字(比如第一个是0,第二个是1),这可能是图像的索引,而不是像素值。如果是这样的话,你需要在split后去掉第一个元素,否则reshape(96,96)会因为元素数量不对而报错。你可以用下面的代码验证:
# 检查第一个字符串的元素数量 print(len(train_x.iloc[0].split(' '))) # 如果输出是9217,就需要修改split_reshape函数,取[1:]
内容的提问来源于stack exchange,提问作者Allen Yeh
相关产品推荐
相关产品推荐

