python装饰器

31

前言

最近在重写训练框架,正好借这个机会整理一下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函数,这是语法上一个需要注意的点。