如何让自定义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

