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
相关产品推荐
相关产品推荐

