You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用自定义生成器进行模型预测时返回异常数组尺寸

问题原因与解决方案

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.02 09:30:56