Keras自定义数据生成器返回嵌套数组引发模型输入形状错误
问题根因
你遇到的报错本质是TensorFlow要求输入的批次数据必须是维度对齐的规则张量,你当前按视频分组返回数据的逻辑下,同批次不同视频的滑动窗口数量不同,导致生成了嵌套的不规则数组,无法转换为合法的张量类型。
解决方案
这里提供两种可直接落地的调整思路,根据你的模型训练逻辑选择即可:
方案1:扁平化批次(适配逐窗口训练场景,推荐)
如果你的模型是对单个滑动窗口做分类/预测,不需要保留同一个视频的窗口分组关系,直接将所有窗口展开为批次的单样本即可,彻底避免不规则数组问题。
修改后的__getitem__核心逻辑如下:
def __getitem__(self, idx): classes = self.classes shape = self.target_shape nbframe = self.nbframe batchImages = [] batchLabels = [] indexes = self.vid_info[idx*self.batch_size:(idx+1)*self.batch_size] for i in indexes: # 直接遍历每个窗口,追加到全局批次列表,不做视频层级的分组 for x in i: vid = x folderPath = vid.get('name') classname = self._get_classname(folderPath) label = np.zeros(len(classes)) col = classes.index(classname) label[col] = 1. window_images = vid['images'] batchLabels.append(label) batchImages.append(window_images) # 此时所有样本维度统一,可直接转换为float32数组 batchImages = np.asarray(batchImages).astype(np.float32) batchLabels = np.asarray(batchLabels).astype(np.float32) return batchImages, batchLabels
调整后批次的维度为(总窗口数, 10, 128, 128, 2)和(总窗口数, 2),完全符合TensorFlow的张量要求。
方案2:填充对齐(适配需保留同视频窗口分组的场景)
如果你的模型需要输入完整视频的所有窗口作为单个样本,就需要对短视频的窗口做填充,让同批次所有样本的窗口数对齐到当前批次的最大窗口数:
def __getitem__(self, idx): classes = self.classes shape = self.target_shape nbframe = self.nbframe batchImages = [] batchLabels = [] max_window = 0 # 记录当前批次最大窗口数 indexes = self.vid_info[idx*self.batch_size:(idx+1)*self.batch_size] # 第一轮遍历计算最大窗口数 for i in indexes: max_window = max(max_window, len(i)) # 第二轮遍历填充对齐 for i in indexes: fileLabels = [] fileImages = [] for x in i: vid = x folderPath = vid.get('name') classname = self._get_classname(folderPath) label = np.zeros(len(classes)) col = classes.index(classname) label[col] = 1. window_images = vid['images'] fileLabels.append(label) fileImages.append(window_images) # 填充到最大窗口数 pad_len = max_window - len(fileImages) # 图像填充全0数组,维度和单窗口一致 fileImages += [np.zeros((10,128,128,2), dtype=np.float32) for _ in range(pad_len)] # 标签填充全0数组,可根据模型逻辑调整为占位值 fileLabels += [np.zeros((len(classes),), dtype=np.float32) for _ in range(pad_len)] batchLabels.append(fileLabels) batchImages.append(fileImages) batchImages = np.asarray(batchImages).astype(np.float32) batchLabels = np.asarray(batchLabels).astype(np.float32) return batchImages, batchLabels
注意:如果用填充方案,建议在模型输入层后添加Masking(mask_value=0.)层,忽略填充位的计算,避免填充数据干扰训练效果。
内容的提问来源于stack exchange,提问作者Japi Sandhu
相关产品推荐
相关产品推荐

