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

如何让自定义Keras层支持接收多个输入?

问题描述

我创建了一个自定义Keras层MaskedBatchNormalization,该层需要接收两个输入(一个数值张量和一个掩码),实现代码如下:

class MaskedBatchNormalization(tf.keras.layers.BatchNormalization):
   def __init__(self, **kwargs):
      super().__init__(fused=False,**kwargs)

   def build(self, input_shapes):
      super().build(input_shapes[0])
   
   # Overriding this method is the whole point of this class
   def _calculate_mean_and_var(self, inputs, reduction_axes, keep_dims):
      return tf.nn.weighted_moments(inputs,reduction_axes,self.mask, keepdims=keep_dims)

   def call(self, inputs, training=None):
      input, mask = inputs
      self.mask = mask
      # do something with input
      output = input*2
      return tf.where(mask, output, 0)

单独调用该层时(x = MaskedBatchNormalization()((x, mask)))运行正常,但将其整合进模型后:

def create_embedding(x:tf.Tensor, edge_mask:tf.Tensor, params:dict,name:str="Embedding") -> tf.Tensor:
    with tf.name_scope(name):
        x = tf.keras.layers.Dense(params["dim"],name=name)(x) #(N,P,P,dim)
        if(params["batchnorm"]):
            x = MaskedBatchNormalization(name=f"{name:s}/MaskedBatchNorm")((x, edge_mask)) #(N,P,P,dim)
        return x
inputs = tf.keras.Inputs(name='inputs', shape=(30,30,4))
mask = tf.keras.Inputs(name='mask', shape=(30,30,1))
x = create_embedding(inputs,mask, {"dim": 30, "batchnorm":True})
# do some more stuff with x, ending in a softmax layer
model = tf.keras.Model(inputs=[inputs,mask],output=x,name="MyModel")
model.compile(optimizer='adam', loss='categorical_crossentropy')
# fine up until here
model.fit(training_data_generator, validation_data=validation_data_generator,epochs=30) # Crash

模型编译正常,但调用model.fit时Keras抛出错误:

Layer "MyLayer" expects 1 input(s), but it received 2 input tensors. Inputs received: [<tf.Tensor 'truediv:0' shape=(None, 30, 30, 20) dtype=float32>, <tf.Tensor 'NotEqual:0' shape=(None, 30, 30, 1) dtype=bool>]

已验证编译时层确实收到两个张量,请问如何让该层正确“期望”两个输入,或有什么可行的解决方法?


解决方案

问题根源在于继承BatchNormalization时,父类默认期望单个输入,自定义层的多输入逻辑未被Keras的输入规范系统正确识别,以下是两种可行解决方法:

方法一:修正继承自BatchNormalization的输入规范

重写compute_output_shape和确保build方法正确处理多输入结构,让Keras明确该层接收双输入:

class MaskedBatchNormalization(tf.keras.layers.BatchNormalization):
    def __init__(self, **kwargs):
        super().__init__(fused=False, **kwargs)

    def build(self, input_shapes):
        # input_shapes是包含输入张量、掩码形状的元组
        super().build(input_shapes[0])

    def _calculate_mean_and_var(self, inputs, reduction_axes, keep_dims):
        return tf.nn.weighted_moments(inputs, reduction_axes, self.mask, keepdims=keep_dims)

    def call(self, inputs, training=None):
        input_tensor, mask = inputs
        self.mask = mask
        # 调用父类批量归一化逻辑,再应用掩码
        normalized_input = super().call(input_tensor, training=training)
        return tf.where(mask, normalized_input, 0.0)

    def compute_output_shape(self, input_shapes):
        # 输出形状与输入张量保持一致
        return input_shapes[0]

    def get_config(self):
        config = super().get_config()
        return config

模型构建时保持现有传入双输入的逻辑即可,Keras会通过重写的方法识别输入数量。

方法二:改用基础Layer类完全自定义

如果继承父类的冲突难以彻底解决,直接继承tf.keras.layers.Layer,手动实现带掩码的批量归一化逻辑:

class MaskedBatchNormalization(tf.keras.layers.Layer):
    def __init__(self, axis=-1, epsilon=1e-3, momentum=0.99, **kwargs):
        super().__init__(**kwargs)
        self.axis = axis
        self.epsilon = epsilon
        self.momentum = momentum

    def build(self, input_shapes):
        input_shape = input_shapes[0]
        # 创建批量归一化的可训练参数
        self.gamma = self.add_weight(
            name='gamma',
            shape=input_shape[self.axis:],
            initializer='ones',
            trainable=True
        )
        self.beta = self.add_weight(
            name='beta',
            shape=input_shape[self.axis:],
            initializer='zeros',
            trainable=True
        )
        # 创建移动平均的非训练参数
        self.moving_mean = self.add_weight(
            name='moving_mean',
            shape=input_shape[self.axis:],
            initializer='zeros',
            trainable=False
        )
        self.moving_var = self.add_weight(
            name='moving_var',
            shape=input_shape[self.axis:],
            initializer='ones',
            trainable=False
        )
        super().build(input_shapes)

    def call(self, inputs, training=None):
        input_tensor, mask = inputs
        if training:
            # 训练阶段计算带掩码的均值和方差
            mean, var = tf.nn.weighted_moments(
                input_tensor,
                axes=[i for i in range(len(input_tensor.shape)) if i != self.axis],
                frequency_weights=mask,
                keepdims=True
            )
            # 更新移动平均
            self.moving_mean.assign(self.momentum * self.moving_mean + (1 - self.momentum) * mean)
            self.moving_var.assign(self.momentum * self.moving_var + (1 - self.momentum) * var)
        else:
            # 推理阶段使用移动平均的值
            mean = self.moving_mean
            var = self.moving_var
        # 应用归一化和缩放偏移
        normalized = tf.nn.batch_normalization(
            input_tensor, mean, var, self.beta, self.gamma, self.epsilon
        )
        # 应用掩码
        return tf.where(mask, normalized, 0.0)

    def compute_output_shape(self, input_shapes):
        return input_shapes[0]

    def get_config(self):
        config = super().get_config()
        config.update({
            'axis': self.axis,
            'epsilon': self.epsilon,
            'momentum': self.momentum
        })
        return config

这种方式完全自主控制输入输出逻辑,从根源避免了父类输入规范的冲突。

额外注意事项

  • 确保掩码形状与输入张量形状广播兼容,否则tf.where会抛出维度不匹配错误
  • 训练时数据生成器返回的输入需与模型输入列表对应,即格式为([输入张量, 掩码张量], 标签张量)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 11:02:56