使用自定义生成器进行模型预测时返回异常数组尺寸
问题原因与解决方案
1. 自定义JoinedGen序列生成器的错误
你的__getitem__方法没有按批次返回样本,而是每次返回单个样本并手动增加batch维度,同时__len__的计算逻辑也和实际取数逻辑不匹配,导致预测时样本数量错误。
修正后的JoinedGen代码:
class JoinedGen(tf.keras.utils.Sequence): def __init__(self, input_gen1, input_gen2, input_gen3, target_gen, batch_size, preds=False, shuffle=False): self.gen1 = input_gen1 self.gen2 = input_gen2 self.gen3 = input_gen3 self.target_gen = target_gen self.batch_size = batch_size self.preds = preds self.shuffle = shuffle self.on_epoch_end() assert len(input_gen1) == len(target_gen) def __len__(self): # 计算总批次数量,向上取整避免遗漏样本 return int(np.ceil(len(self.gen1) / self.batch_size)) def __getitem__(self, i): # 计算当前批次的起始和结束索引 start = i * self.batch_size end = min(start + self.batch_size, len(self.gen1)) # 取当前批次的所有样本 x1 = self.gen1[start:end] x2 = self.gen2[start:end] x3 = self.gen3[start:end] if self.preds: return [x1, x2, x3] y = self.target_gen[start:end] return [x1, x2, x3], y def on_epoch_end(self): self.indices = np.arange(len(self.gen1)) if self.shuffle: np.random.shuffle(self.indices) # 打乱样本顺序 self.gen1 = self.gen1[self.indices] self.gen2 = self.gen2[self.indices] self.gen3 = self.gen3[self.indices] self.target_gen = self.target_gen[self.indices]
关键修正点:
__getitem__按批次索引截取连续的batch_size个样本,不再处理单个样本__len__改用np.ceil确保所有样本都被覆盖on_epoch_end中增加了对样本数组的打乱逻辑(原代码只打乱了索引但没实际重排样本)
2. 模型结构中的维度错误
你的模型里使用了Concatenate(axis=0),这会在batch维度拼接输入,导致后续张量的batch大小翻倍,最终输出形状不符合预期。你应该在通道维度(axis=-1)拼接,或者根据实际需求调整拼接轴。
修正后的模型build_model方法:
def build_model(self): inputs1 = Input((self.height, self.width, self.channels)) inputs2 = Input((self.height, self.width, self.channels)) inputs3 = Input((self.height, self.width, self.channels)) M1M2 = Subtract()([inputs2, inputs1]) M1M2E1 = Add()([M1M2, inputs3]) # 改为在通道维度拼接(axis=-1),避免改变batch大小 concat = Concatenate(axis=-1)([M1M2, inputs3]) x = Conv2D(32, 1, activation='relu')(inputs1) sr1 = Conv2D(32, 1, activation='relu')(inputs2) sr2 = Conv2D(32, 1, activation='relu')(x) sr2 = Conv2D(32, 1, activation='relu')(sr2) # 拼接后的通道数是3+3=6,这里Conv2D要处理6通道的输入,原代码的32核没问题 sr2 = Conv2D(32, 1, activation='relu')(concat) addition = Add()([sr1, sr2]) # 删除未使用的冗余输入层 f3 = Conv2D(1, 1, activation='relu')(addition) outputs = Conv2D(3, 1)(f3) outputs = Conv2D(3, 1)(outputs) model = Model([inputs1, inputs2, inputs3], outputs) return model
关键修正点:
- 将
Concatenate(axis=0)改为Concatenate(axis=-1),保持batch维度不变,仅合并通道 - 删除了未使用的
f1_input相关冗余代码,避免混淆
3. 验证预测结果
完成上述修正后,运行以下代码验证:
mygen_preds = JoinedGen(x1, x2, x3, y, 16, preds=True) preds = model.predict(mygen_preds) print(preds.shape) # 输出应为(64, 16, 16, 3)
或者直接传入数组预测(无需生成器):
preds = model.predict([x1, x2, x3], batch_size=16) print(preds.shape) # 输出应为(64, 16, 16, 3)
内容的提问来源于stack exchange,提问作者George
相关产品推荐
相关产品推荐

