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

Keras如何水平合并含预训练分支的同形状模型(分支相加)

Keras并行双分支逐元素加和合并模型实现

核心实现逻辑

  • 两个独立分支共用同一个模型输入层,保证输入数据同步传入两个子模型
  • 预训练分支提前加载训练好的权重完成初始化,另一分支保持随机初始化状态
  • 两个分支输出直接传入Add()层做逐元素相加,不使用拼接操作
  • 构建整体模型时以共享输入为入口,加和结果为输出,训练时可按需选择是否冻结预训练分支参数

完整实现代码

import tensorflow as tf
from tensorflow.keras.layers import Input, Conv1D, Lambda, Add, ReLU, Dense
from tensorflow.keras import Model

def define_model_a(input_shape, initializer, outputs = 1):
   input_layer = Input(shape=(input_shape))
   path10 = input_layer
   path11 = Conv1D(filters=1, kernel_size=3, strides=1, padding="same", use_bias = True, kernel_initializer=initializer)(path10)
   path12 = Lambda(lambda x: abs(x))(path11)
   output = Add()([path10, path12])
   define_model_a = Model(inputs=input_layer, outputs=output)
   define_model_a._name = 'model_a'
   return define_model_a

def define_model_b(input_shape, initializer, outputs = 1):
    input_layer = Input(shape=(input_shape))
    path10 = input_layer
    path11 = Conv1D(filters=1, kernel_size=3, strides=1, padding="same", use_bias = True, kernel_initializer=initializer)(path10)
    path12 = ReLU()(path11)
    path13 = Dense(1, use_bias = True)(path12)
    output = path13
    define_model_b = Model(inputs=input_layer, outputs=output)
    define_model_b._name = 'model_b'
    return define_model_b    

def define_merge_interpretation(input_shape, model_a, model_b):
    # 构建共享输入层
    input_layer = Input(shape=input_shape)
    # 输入分别传入两个并行分支
    branch_a_output = model_a(input_layer)
    branch_b_output = model_b(input_layer)
    # 分支输出逐元素相加
    merged_output = Add()([branch_a_output, branch_b_output])
    # 组装最终模型
    merged_model = Model(inputs=input_layer, outputs=merged_output)
    return merged_model


# 初始化配置
input_shape = (100, 1) # 替换为实际输入形状
load_path = "your_model_a_weights_path.h5" # 替换为预训练权重实际存储路径
initializer = tf.keras.initializers.HeNormal()    
model_a = define_model_a(input_shape, initializer, outputs = 1)
model_b = define_model_b(input_shape, initializer, outputs = 1)   
# 预训练分支加载权重
model_a.load_weights(load_path)

# 生成合并模型
merge_interpretation = define_merge_interpretation(input_shape, model_a, model_b)
# 打印结构校验
merge_interpretation.summary()

# 常规调用fit开展训练即可
# history = merge_interpretation.fit(...)

补充说明

如果需要冻结预训练分支权重、训练时只更新随机初始化分支的参数,在model_a.load_weights(load_path)后添加代码model_a.trainable = False即可。

结构参考

双并行分支逐元素加和模型结构

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 18:31:04