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

复制Keras PReLU层到自定义脚本时与官方版本行为差异原因

问题描述

在Mac M1设备上使用TensorFlow 2.9.2版本开发时,将框架内置的PReLU层源码复制到自有脚本中重命名使用,出现两处异常表现:

  • 复制代码编写的自定义PReLU层,在测试模型中未统计到任何可训练参数,模型summary将层内运算全部拆分为独立的TFOpLambda节点展示,结构冗余
  • 同模型中直接调用的官方打包PReLU层可正常统计到5个可训练参数,summary展示简洁
问题根因

核心错误是自定义层上误用了TensorFlow内部的@keras_export装饰器:

  • 这个装饰器是TensorFlow官方开发时用来标记内置API、导出公开接口用的,自定义层不能使用
  • 加了这个装饰器后,Keras用Functional API构建模型时,会把这个类错判成内置无状态算子,不会把层实例当成独立的模型节点管理,而是直接穿透call方法,把里面的TensorFlow算子直接拍平到计算图里
  • 因为层实例没进入模型的层管理列表,层里用add_weight创建的alpha参数就不会被模型的参数统计逻辑收集,最终出现参数为0、summary把层内算子全拆出来的问题。
修复方法

做两处修改即可恢复正常:

  • 删掉自定义层上方的@keras_export('keras.layers.PReLUcopy')装饰器
  • (可选,保障模型保存、加载的兼容性)给自定义层加上注册装饰器:
@tf.keras.utils.register_keras_serializable(package="Custom")
class PReLUcopy(Layer):
    # 原有层代码无需改动

改完之后重新运行,summary里会正常显示PReLUcopy层条目,可训练参数统计为5个,和官方内置PReLU表现完全一致。

复现代码
import tensorflow as tf

from tensorflow.keras.layers import *
from tensorflow.keras.models import Model

from tensorflow.python.framework import dtypes
from tensorflow.python.keras import backend
from tensorflow.python.keras import constraints
from tensorflow.python.keras import initializers
from tensorflow.python.keras import regularizers
from tensorflow.python.keras.engine.base_layer import Layer
from tensorflow.python.keras.engine.input_spec import InputSpec
from tensorflow.python.keras.utils import tf_utils
from tensorflow.python.ops import math_ops
from tensorflow.python.util.tf_export import keras_export

# 存在问题的自定义层:错误添加了@keras_export装饰器
@keras_export('keras.layers.PReLUcopy')
class PReLUcopy(Layer):
  """Parametric Rectified Linear Unit.

  It follows:
f(x) = alpha * x for x < 0
f(x) = x for x >= 0
where `alpha` is a learned array with the same shape as x.

Input shape:
  Arbitrary. Use the keyword argument `input_shape`
  (tuple of integers, does not include the samples axis)
  when using this layer as the first layer in a model.

Output shape:
  Same shape as the input.

Args:
  alpha_initializer: Initializer function for the weights.
  alpha_regularizer: Regularizer for the weights.
  alpha_constraint: Constraint for the weights.
  shared_axes: The axes along which to share learnable
    parameters for the activation function.
    For example, if the incoming feature maps
    are from a 2D convolution
    with output shape `(batch, height, width, channels)`,
    and you wish to share parameters across space
    so that each filter only has one set of parameters,
    set `shared_axes=[1, 2]`.
"""

def __init__(self,
             alpha_initializer='zeros',
             alpha_regularizer=None,
             alpha_constraint=None,
             shared_axes=None,
             **kwargs):
  super(PReLUcopy, self).__init__(**kwargs)
  self.supports_masking = True
  self.alpha_initializer = initializers.get(alpha_initializer)
  self.alpha_regularizer = regularizers.get(alpha_regularizer)
  self.alpha_constraint = constraints.get(alpha_constraint)
  if shared_axes is None:
    self.shared_axes = None
  elif not isinstance(shared_axes, (list, tuple)):
    self.shared_axes = [shared_axes]
  else:
    self.shared_axes = list(shared_axes)

@tf_utils.shape_type_conversion
def build(self, input_shape):
  param_shape = list(input_shape[1:])
  if self.shared_axes is not None:
    for i in self.shared_axes:
      param_shape[i - 1] = 1
  self.alpha = self.add_weight(
      shape=param_shape,
      name='alpha',
      initializer=self.alpha_initializer,
      regularizer=self.alpha_regularizer,
      constraint=self.alpha_constraint)
  # Set input spec
  axes = {}
  if self.shared_axes:
    for i in range(1, len(input_shape)):
      if i not in self.shared_axes:
        axes[i] = input_shape[i]
  self.input_spec = InputSpec(ndim=len(input_shape), axes=axes)
  self.built = True

def call(self, inputs):
  pos = backend.relu(inputs)
  neg = -self.alpha * backend.relu(-inputs)
  return pos + neg

def get_config(self):
  config = {
      'alpha_initializer': initializers.serialize(self.alpha_initializer),
      'alpha_regularizer': regularizers.serialize(self.alpha_regularizer),
      'alpha_constraint': constraints.serialize(self.alpha_constraint),
      'shared_axes': self.shared_axes
  }
  base_config = super(PReLUcopy, self).get_config()
  return dict(list(base_config.items()) + list(config.items()))

@tf_utils.shape_type_conversion
def compute_output_shape(self, input_shape):
  return input_shape


def test1():
  A_in = Input(shape=(5,), name='A_in')
  out = PReLUcopy()(A_in)
  out2 = PReLU()(A_in)
  model = Model(inputs=[A_in], outputs=[out,out2])
  model.compile(optimizer='adam', loss='mean_squared_error')
  print( model.summary() )



if __name__ == '__main__':
  test1()
异常运行输出
__________________________________________________________________________________________________
 Layer (type)                   Output Shape         Param #     Connected to                     
==================================================================================================
 A_in (InputLayer)              [(None, 5)]          0           []                               
                                                                                                  
 tf.math.negative (TFOpLambda)  (None, 5)            0           ['A_in[0][0]']                   
                                                                                                  
 tf.nn.relu_1 (TFOpLambda)      (None, 5)            0           ['tf.math.negative[0][0]']       
                                                                                                  
 tf.nn.relu (TFOpLambda)        (None, 5)            0           ['A_in[0][0]']                   
                                                                                                  
 tf.math.multiply (TFOpLambda)  (None, 5)            0           ['tf.nn.relu_1[0][0]']           
                                                                                                  
 tf.__operators__.add (TFOpLamb  (None, 5)           0           ['tf.nn.relu[0][0]',             
 da)                                                              'tf.math.multiply[0][0]']       
                                                                                                  
 p_re_lu (PReLU)                (None, 5)            5           ['A_in[0][0]']                   
                                                                                                  
==================================================================================================
Total params: 5
Trainable params: 5
Non-trainable params: 0
__________________________________________________________________________________________________

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 12:01:15