您可以找到确切的decorator location 来了解这个想法。
def weak_script_method(fn):
weak_script_methods[fn] = {
"rcb": createResolutionCallback(frames_up=2),
"original_method": fn
}
return fn
但是,您不必担心那个装饰器。这个装饰器是 JIT 内部的。
技术上用@weak_script_method修饰的方法会被添加到前面创建的weak_script_methods字典中,像这样:
weak_script_methods = weakref.WeakKeyDictionary()
dict 跟踪方法以避免循环依赖问题;创建 PyTorch 图时调用其他方法的方法。
这真的没有多大意义,除非你对 TorchScript 的概念有大体的了解。
TorchScript 的想法是在 PyTorch 中训练模型并将模型导出到另一个支持静态类型的非 Python 生产环境(阅读:C++/C/Cuda)。
PyTorch 团队在有限的 Python 基础上制作了 TorchScript,以支持静态类型。
默认情况下,Python 是 动态 类型的语言,但通过一些技巧(read:checks)它可以成为 静态 类型的语言。
所以 TorchScript 函数是 Python 的静态类型子集,包含 PyTorch 的所有内置张量操作。这种差异允许 TorchScript 模块代码在不需要 Python 解释器的情况下运行。
您可以使用跟踪(torch.jit.trace() 方法)将现有的 PyTorch 方法转换为 TorchScript,或者使用 @torch.jit.script 装饰器手动创建您的 TorchScript。
如果您使用跟踪,最后您将获得一个类模块。示例如下:
import inspect
import torch
def foo(x, y):
return x + y
traced_foo = torch.jit.trace(foo, (torch.rand(3), torch.rand(3)))
print(type(traced_foo)) #<class 'torch.jit.TopLevelTracedModule'>
print(traced_foo) #foo()
print(traced_foo.forward) #<bound method TopLevelTracedModule.forward of foo()>
lines = inspect.getsource(traced_foo.forward)
print(lines)
输出:
<class 'torch.jit.TopLevelTracedModule'>
foo()
<bound method TopLevelTracedModule.forward of foo()>
def forward(self, *args, **kwargs):
return self._get_method('forward')(*args, **kwargs)
您可以使用检查模块进行进一步调查。这只是展示如何使用跟踪转换一个函数。