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

如何在Keras中实现带下采样的共享ResNet模块?

修正Keras共享ResNet模块的权重共享与下采样问题

原代码无法实现权重共享的核心原因是:每次调用SharedResBlock都会重新创建全新的卷积层,而非复用已实例化的层对象。同时下采样逻辑也存在维度不匹配、步长设置不合理的问题,下面是修正后的完整方案:

1. 重写可复用的残差块

首先要让残差块成为接收输入张量、复用同一层实例的结构,而不是每次调用都生成新层。同时修正下采样逻辑:

from tensorflow.keras.layers import Input, Conv2D, Add, Activation, BatchNormalization
from tensorflow.keras.models import Model

def build_shared_res_block(filters, kernel_size, downsample=False):
    # 定义可复用的层实例(只创建一次,后续复用)
    conv1 = Conv2D(filters, kernel_size, padding="same", strides=(2 if downsample else 1))
    bn1 = BatchNormalization()
    conv2 = Conv2D(filters, kernel_size, padding="same")
    bn2 = BatchNormalization()
    
    # 定义shortcut路径的层(仅当下采样时需要)
    shortcut_conv = Conv2D(filters, 1, strides=(2 if downsample else 1), padding="same") if downsample else None
    shortcut_bn = BatchNormalization() if downsample else None
    
    def res_block(input_tensor):
        # 主路径
        x = conv1(input_tensor)
        x = bn1(x)
        x = Activation("relu")(x)
        
        x = conv2(x)
        x = bn2(x)
        
        # Shortcut路径:下采样时需要调整维度,否则直接用输入
        if downsample:
            shortcut = shortcut_conv(input_tensor)
            shortcut = shortcut_bn(shortcut)
        else:
            shortcut = input_tensor
        
        # 残差连接
        x = Add()([x, shortcut])
        x = Activation("relu")(x)
        return x
    
    return res_block

2. 构建共享模型

先实例化共享的残差块,再分别用于两个输入分支,确保权重复用:

# 初始化输入(补充Conv2D所需的通道维度)
input_A = Input(shape=(64, 64, 1), name="iA")
input_B = Input(shape=(64, 64, 1), name="iB")

# 共享的初始卷积层
shared_initial_conv = Conv2D(16, 5, strides=2, activation='relu', padding="same")

# 处理两个输入分支的初始层
convPA = shared_initial_conv(input_A)
convPB = shared_initial_conv(input_B)

filterSize = [32, 64, 128]  # 可根据需求调整滤波器尺寸列表
cnt = 0

# 循环创建共享残差块
for i in range(1, 19):
    if i % 2 == 0:
        cnt += 1
    current_filters = filterSize[min(cnt, len(filterSize)-1)]
    # 实例化共享的残差块(只创建一次,两个分支复用)
    shared_res_block = build_shared_res_block(current_filters, 3, downsample=(i%2==0))
    # 复用同一残差块处理两个分支
    convPA = shared_res_block(convPA)
    convPB = shared_res_block(convPB)

# 示例:添加输出层(可根据任务需求调整)
output_A = Conv2D(1, 1, activation='sigmoid')(convPA)
output_B = Conv2D(1, 1, activation='sigmoid')(convPB)

model = Model(inputs=[input_A, input_B], outputs=[output_A, output_B])
model.summary()

关键修正点说明

  • 权重共享实现:通过先实例化共享的层对象(如conv1、conv2),再将不同输入传入同一层对象,确保所有分支复用同一组权重。
  • 下采样逻辑修正:
    • 主路径第一个卷积用strides=2实现下采样,符合ResNet标准设计(原代码strides=4会过度压缩特征)。
    • Shortcut路径在需要下采样时,用1x1卷积调整通道数和尺寸,保证与主路径输出维度一致,避免Add层报错。
  • 输入维度修正:Conv2D层要求输入是3D张量(高度、宽度、通道),原代码输入shape=(64,64)缺少通道维度,需补充(如(64,64,1)代表单通道灰度图)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 21:30:59