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

Keras中共享层实现咨询:多任务神经网络构建方法

嘿,我来帮你理清Keras里共享层的实现思路,还有keras.layers.concatenate的正确用法~

Keras中共享层的核心逻辑

首先得明确:共享层的本质是复用同一个层实例,而不是每次都新建一个同名层。不管你用不用concatenate,这个逻辑都是实现共享的基础——只要多个输入分支调用同一个层对象,它们就会共享该层的权重,训练时这些权重会被所有分支共同更新。

针对你需求的代码实现(3个带共享/非共享层的网络)

假设你的场景是:3个网络输入输出形状一致,每个网络包含「独有层(彩色)+ 共享层 + 独有输出层」,下面是具体代码示例:

from keras.layers import Input, Dense, concatenate
from keras.models import Model

# ----------------------
# 第一步:定义共享层(所有分支共用)
# ----------------------
# 比如定义两个共享的全连接层,你也可以换成Conv2D等其他层
shared_dense_1 = Dense(64, activation='relu')
shared_dense_2 = Dense(32, activation='relu')

# ----------------------
# 第二步:构建3个分支(每个分支有自己的独有层)
# ----------------------
# 分支1:独有输入 + 独有层 + 共享层 + 独有输出
input_branch1 = Input(shape=(100,))  # 假设输入维度是100
# 独有层(对应你说的彩色层,每个分支单独实例化)
branch1_private = Dense(128, activation='relu')(input_branch1)
# 接入共享层(直接调用之前定义好的共享层实例)
x1 = shared_dense_1(branch1_private)
x1 = shared_dense_2(x1)
# 分支1的独有输出层
output_branch1 = Dense(10, activation='softmax')(x1)

# 分支2:逻辑和分支1完全一致,只是独有层是新实例
input_branch2 = Input(shape=(100,))
branch2_private = Dense(128, activation='relu')(input_branch2)
x2 = shared_dense_1(branch2_private)
x2 = shared_dense_2(x2)
output_branch2 = Dense(10, activation='softmax')(x2)

# 分支3:同理
input_branch3 = Input(shape=(100,))
branch3_private = Dense(128, activation='relu')(input_branch3)
x3 = shared_dense_1(branch3_private)
x3 = shared_dense_2(x3)
output_branch3 = Dense(10, activation='softmax')(x3)

# ----------------------
# 第三步:用concatenate合并分支结果(可选)
# ----------------------
# 如果需要把三个分支的中间结果或输出合并,就用concatenate
# 它的作用是把多个张量在指定轴上拼接(默认是最后一个轴)
merged_output = concatenate([x1, x2, x3])
# 可以在合并后加一个全连接层做后续处理
final_dense = Dense(10, activation='softmax')(merged_output)

# ----------------------
# 第四步:构建完整模型
# ----------------------
# 模型可以同时接受多个输入,输出多个结果
# 比如同时输出三个分支的单独结果+合并后的结果
model = Model(
    inputs=[input_branch1, input_branch2, input_branch3],
    outputs=[output_branch1, output_branch2, output_branch3, final_dense]
)
# 如果不需要合并结果,只输出三个分支的单独结果也可以
# model = Model(inputs=[input_branch1, input_branch2, input_branch3], outputs=[output_branch1, output_branch2, output_branch3])
关键细节说明
  • 共享层的验证:你可以打印shared_dense_1.get_weights(),会发现三个分支调用它后,权重是完全相同的,训练时更新也会同步。
  • concatenate的定位:它只是一个张量拼接工具,和共享层本身没有直接关联——你可以用它合并共享层输出的张量,也可以合并其他任何张量。
  • 独有层的注意点:每个分支的独有层必须是新的层实例(比如branch1_private、branch2_private是不同的Dense对象),这样它们的权重才会各自独立。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 04:05:04