前言
最近在重写训练框架,正好借这个机会整理一下python装饰器的语法
什么是装饰器
在某些代码中会看到类似这样的语法
@decorator
def example_function: ...
其中@decorator就是所谓的装饰器,是python提供的一个语法糖,其作用就是给原本的函数添加一些额外的功能,类似一个外挂模块。python是可以把函数作为参数进行传递的,装饰器可以理解为一个参数和返回值都是函数的函数。
这样说是不是太抽象了?我们来举一个具体的例子,假设我想要给多个函数增加猫叫的功能,也就是:
def mewo(func):
print("Mewo")
return func
这样就实现了一个猫叫装饰器,我们把它装饰到一个现有的函数上:
@meow
def hello_world():
print("hello world")
实际上,上面这个带有装饰器的函数相当于:
def hello_world():
print("hello world")
hello_world = mewo(hello_world)
先执行了mewo函数,然后把hello_world函数返回
具体的执行顺序是,在定义装饰器@mewo时,执行mewo函数,在具体调用hello_world函数时才会执行hello_world。注意,装饰器加的功能不是在被装饰函数执行时触发,而是定义装饰器时触发的。也就是说后面无论调用多少次hello_world,都不会再输出Mewo了。
如果想要实现每次调用hello_world都会输出一次Mewo,需要这样写:
def mewo(func):
def wrapper():
print("Mewo")
return func()
return wrapper
@meow
def hello_world()
print("hello world")
这样写时发生了什么呢?
带有装饰器的函数仍然相当于
def hello_world():
print("hello world")
hello_world = mewo(hello_world) #返回的实际上是wapper
不同之处在于,此时mewo返回的函数不再是原始的hello_world,而是wrapper,相当于hello_world函数被wrapper函数替换了,而wrapper函数是先输出Mewo再执行原来的函数,就实现了每次调用hello_world都会猫叫,可喜可贺可喜可贺。
一个使用场景
接下来是一个略微复杂一些的例子,展示装饰器的一个具体使用场景。
在深度学习训练时,经常会尝试多个网络模型或者在多个数据集上进行测试,为了方便实验,需要实现一个从配置文件中读取模型名字的功能,这样只要修改配置文件就能使用不同的模型进行测试。为了实现这样的功能,需要将所有的模型先注册到一个字典中,需要模型的时候就根据配置文件中的模型名去字典中寻找对应的类。
为了实现这样的功能,就可以用装饰器:
MODELS = {} # 储存模型类的字典
def register_module(name=None):
def _register(cls):
key = name or cls.__name__ # 没有指定名字就用用类名作为key
MODELS[key] = cls
return cls # 将类储存在字典中并返回
return _register
# 注册网络模型
@register_module()
class ResNet:
def __init__(self, depth):
self.depth = depth
print(f"构建了一个 ResNet,depth={depth}")
@register_module(name='vgg16')
class VGG:
def __init__(self, depth):
self.depth = depth
print(f"构建了一个 VGG,depth={depth}")
此时
print(MODELS)
# {'ResNet': <class 'ResNet'>, 'vgg16': <class 'VGG'>}
此时假设配置文件中写了:
model: vgg16
就可以通过一个函数去MODELS中查出vgg16对应的类并实例化
这里的写法和上面的喵喵函数有一点不同
@mewo
@register_module()
register_module后面多了一个括号,有没有这个括号装饰器的展开是不同的,我们以register_module来解释一下
@register_module()
class ResNet: ...
# 展开后相当于
temp = register_module()
ResNet = temp()(ResNet)
# 再展开就相当于
temp = register_module()
result = temp(ResNet)
ResNet = result
不加括号的写法:
@register_module
class ResNet: ...
# 展开后相当于
ResNet = register_module(ResNet)
不加括号时,@语法会把下面的类(或者函数)作为唯一参数传给register_module,而加了括号后传给的是register_module函数的返回值,也就是上文中的_register函数,这是语法上一个需要注意的点。