使用参数自定义激活

use*_*665 3 python machine-learning keras activation-function keras-layer

我正在尝试在Keras中创建一个激活函数,该函数可以采用如下所示的参数beta

from keras import backend as K
from keras.utils.generic_utils import get_custom_objects
from keras.layers import Activation

class Swish(Activation):

    def __init__(self, activation, beta, **kwargs):
        super(Swish, self).__init__(activation, **kwargs)
        self.__name__ = 'swish'
        self.beta = beta


def swish(x):
    return (K.sigmoid(beta*x) * x)

get_custom_objects().update({'swish': Swish(swish, beta=1.)})
Run Code Online (Sandbox Code Playgroud)

它在没有beta参数的情况下可以正常运行,但是如何在激活定义中包含参数?我也model.to_json()喜欢在激活ELU 时保存该值。


更新:我根据@today的答案编写了以下代码:

from keras.layers import Layer
from keras import backend as K

class Swish(Layer):
    def __init__(self, beta, **kwargs):
        super(Swish, self).__init__(**kwargs)
        self.beta = K.cast_to_floatx(beta)
        self.__name__ = 'swish'

    def call(self, inputs):
        return K.sigmoid(self.beta * inputs) * inputs

    def get_config(self):
        config = {'beta': float(self.beta)}
        base_config = super(Swish, self).get_config()
        return dict(list(base_config.items()) + list(config.items()))

    def compute_output_shape(self, input_shape):
        return input_shape

from keras.utils.generic_utils import get_custom_objects
get_custom_objects().update({'swish': Swish(beta=1.)})
gnn = keras.models.load_model("Model.h5")
arch = gnn.to_json()
with open(directory + 'architecture.json', 'w') as arch_file:
    arch_file.write(arch)
Run Code Online (Sandbox Code Playgroud)

但是,它当前不将beta值保存在.json文件中。如何保存价值?

tod*_*day 5

由于您希望在序列化模型时保存激活函数的参数,因此我认为最好将激活函数定义为类似于Keras中定义高级激活的层。您可以这样做:

from keras.layers import Layer
from keras import backend as K

class Swish(Layer):
    def __init__(self, beta, **kwargs):
        super(Swish, self).__init__(**kwargs)
        self.beta = K.cast_to_floatx(beta)

    def call(self, inputs):
        return K.sigmoid(self.beta * inputs) * inputs

    def get_config(self):
        config = {'beta': float(self.beta)}
        base_config = super(Swish, self).get_config()
        return dict(list(base_config.items()) + list(config.items()))

    def compute_output_shape(self, input_shape):
        return input_shape
Run Code Online (Sandbox Code Playgroud)

然后,可以像使用Keras层一样使用它:

# ...
model.add(Swish(beta=0.3))
Run Code Online (Sandbox Code Playgroud)

由于get_config()方法已在其定义中实现,因此beta使用诸如to_json()或方法时将保存参数save()