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

Keras多分支非共享权重网络输入形状不匹配报错如何解决

问题根因

你遇到的报错核心有两个原因:

  1. 数据生成器输出的形状不符合多输入模型的要求,缺少batch维度,导致每个分支收到的输入少了一维,从预期的5维变成了4维
  2. 部分场景下分支输入层定义与融合模型的输入绑定错误也会触发该类报错

修复步骤

1. 修正数据生成器输出格式

多输入Keras模型要求数据生成器返回的输入格式为:

  • 长度与输入分支数相等的列表(此处为6个元素)
  • 列表的每个元素对应一个分支的整批数据,形状为(batch_size, 30, 64, 64, 3),而非单样本的(30,64,64,3)

生成器修正示例参考:

# 错误写法(你当前的逻辑)
def wrong_generator(batch_size):
    while True:
        batch_x = []
        batch_y = []
        for _ in range(batch_size):
            # 每个样本生成6个分支的输入
            sample_inputs = [np.random.rand(30,64,64,3) for _ in range(6)]
            sample_label = np.random.rand(1)
            batch_x.append(sample_inputs)
            batch_y.append(sample_label)
        # 此时batch_x维度为(batch_size,6,30,64,64,3),不符合模型要求
        yield batch_x, np.array(batch_y)

# 正确写法
def correct_generator(batch_size):
    while True:
        # 预先为6个分支各初始化存储整批数据的列表
        branch_inputs = [[] for _ in range(6)]
        batch_y = []
        for _ in range(batch_size):
            sample_inputs = [np.random.rand(30,64,64,3) for _ in range(6)]
            sample_label = np.random.rand(1)
            # 把同分支的样本放到对应列表中
            for i in range(6):
                branch_inputs[i].append(sample_inputs[i])
            batch_y.append(sample_label)
        # 每个分支的数据堆叠出batch维度
        processed_inputs = [np.array(branch) for branch in branch_inputs]
        yield processed_inputs, np.array(batch_y)

2. 确认多分支模型的输入定义正确

你需要保证融合模型的输入列表,就是6个独立vggLstmNet实例的输入层,不要错误复用输入层或者取错模型的输入属性:

# 正确多分支搭建示例
# 初始化6个独立分支(实现权重不共享)
branches = [vggLstmNet() for _ in range(6)]
# 取每个分支的输出做融合,示例为拼接操作
branch_outputs = [branch.output for branch in branches]
x = Concatenate()(branch_outputs)
# 后续添加自定义的分类/回归层
x = Dense(128, activation='relu')(x)
output = Dense(1, activation='sigmoid')(x)
# 融合模型的输入就是6个分支的输入层列表
class_models = Model(inputs=[branch.input for branch in branches], outputs=output)

3. 可选:显式指定输入层形状避免推断错误

如果修正生成器后仍有报错,可以把vggLstmNet里的输入层从shape参数改成batch_shape,强制指定形状,避免Keras动态推断形状时出现维度丢失:

inp = Input(batch_shape=(None, flameSize, size, size, 3))

内容的提问来源于stack exchange,提问作者PlanetRT

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 10:27:03