TensorFlow2.9运行含自定义平方激活函数的模型时触发AttributeError
环境信息
- TensorFlow版本:2.9.1
- Keras版本:2.9.0
问题描述
运行搭载自定义平方激活函数的模型代码时,前序所有命令均可正常执行,执行到添加Conv2D层的代码行时触发报错。
报错详情
- 触发报错的代码行:
model.add(Conv2D(32, kernel_size=(3, 3), activation=custom_activation, input_shape=((input_shape)))) - 完整报错内容(已翻译):
AttributeError:调用层"conv2d_4"(类型Conv2D)时遇到异常
模块keras.api._v2.keras.backend不存在属性x
层"conv2d_4"(类型Conv2D)接收到的调用参数:
- inputs=tf.Tensor(shape=(None, 28, 28, 1), dtype=float32)
完整复现代码
import tensorflow from tensorflow.keras.datasets import mnist from tensorflow.keras import backend as K from keras.utils.generic_utils import get_custom_objects from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Dense, Dropout, Flatten, Activation from tensorflow.keras.layers import Conv2D, MaxPooling2D import numpy as np import matplotlib.pyplot as plt # 自定义激活函数 def custom_activation(x): return K.cast(K.x**2) # 曾尝试调用Square(x)实现 # 创建模型前注册自定义对象 get_custom_objects().update({'custom_activation': Activation(custom_activation)}) # 模型配置参数 img_width, img_height = 28, 28 batch_size = 32 no_epochs = 5 no_classes = 10 verbosity = 1 # 加载MNIST数据集 (input_train, target_train), (input_test, target_test) = mnist.load_data() # 数据维度调整 input_train = input_train.reshape(input_train.shape[0], img_width, img_height, 1) input_test = input_test.reshape(input_test.shape[0], img_width, img_height, 1) input_shape = (img_width, img_height, 1) # 转换数据类型为float input_train = input_train.astype('float32') input_test = input_test.astype('float32') # 数据归一化到[0,1]区间 input_train = input_train / 255 input_test = input_test / 255 # 标签转换为独热编码格式 target_train = tensorflow.keras.utils.to_categorical(target_train, no_classes) target_test = tensorflow.keras.utils.to_categorical(target_test, no_classes) # 创建模型 model = Sequential() model.add(Conv2D(32, kernel_size=(3, 3), activation=custom_activation, input_shape=((input_shape))))
故障原因与修复方法
核心故障原因
自定义激活函数的写法存在语法和逻辑错误,同时存在导入路径不统一的兼容隐患:
- 代码中写的
K.x属于错误调用,K是Keras后端模块,该模块下不存在名为x的属性,x是自定义函数接收的输入张量参数,不需要加K.前缀 K.cast函数调用时缺少必填的目标数据类型参数,即使修正了K.x的问题,这行代码依然会触发参数缺失报错- 代码中混用了独立Keras和TensorFlow内置Keras的导入路径,在TF2.9版本环境下容易触发兼容问题
修复步骤
- 修正自定义激活函数逻辑,直接对传入的输入张量做平方计算即可,修正后的代码:
def custom_activation(x): return x ** 2
- 统一导入路径,将原代码中
from keras.utils.generic_utils import get_custom_objects替换为以下代码,避免跨版本兼容问题:from tensorflow.keras.utils import get_custom_objects
完成以上修改后,代码即可正常运行,Conv2D层可以正确识别并调用自定义的平方激活函数。
内容的提问来源于stack exchange,提问作者Anmar Ali
相关产品推荐
相关产品推荐

