如何让PyTorch中nn.Module的__call__()自动继承forward()的类型提示与文档字符串?
如何让PyTorch中nn.Module的__call__()自动继承forward()的类型提示与文档字符串?
我完全懂你的痛点——每次写PyTorch模块都要重复复制forward的类型提示和文档到__call__里,不仅麻烦还容易出错。下面给你几个实用的解决方案,不用再重复写冗余代码:
方法一:用类装饰器自动同步(最省心)
写一个通用的装饰器,给你的nn.Module子类加上之后,就能自动把forward的文档字符串和类型注解同步到__call__上,一劳永逸:
from functools import wraps import torch.nn as nn def sync_call_forward(cls): # 同步forward的文档字符串到__call__ if cls.forward.__doc__: cls.__call__.__doc__ = cls.forward.__doc__ # 同步forward的类型注解到__call__ if hasattr(cls.forward, '__annotations__'): cls.__call__.__annotations__ = cls.forward.__annotations__.copy() return cls # 使用示例 @sync_call_forward class MyModule(nn.Module): def __init__(self) -> None: super().__init__() # 这里放你的层初始化逻辑 def forward(self, x: torch.FloatTensor, y: torch.FloatTensor) -> tuple[torch.FloatTensor, torch.IntTensor]: """ Args: x (FloatTensor): in shape of BxTxE y (FloatTensor): in shape of BxE Returns: tuple[FloatTensor, IntTensor]: (sth. in shape of BxT, sth. in shape of B) """ # 你的forward逻辑 return x.sum(dim=-1), y.argmax(dim=-1)
这个装饰器会在类定义完成后,自动把forward的文档和类型提示复制给__call__,VSCode之类的IDE就能识别到__call__的提示了,而且你只需要写一次forward的内容。
方法二:利用functools.wraps简化__call__的定义
如果不想用装饰器,也可以在子类里用@wraps来快速同步,比手动复制文档和类型要方便很多:
from functools import wraps import torch.nn as nn from torch import FloatTensor, IntTensor class MyModule(nn.Module): def __init__(self) -> None: super().__init__() def forward(self, x: FloatTensor, y: FloatTensor) -> tuple[FloatTensor, IntTensor]: """ Args: x (FloatTensor): in shape of BxTxE y (FloatTensor): in shape of BxE Returns: tuple[FloatTensor, IntTensor]: (sth. in shape of BxT, sth. in shape of B) """ # 你的forward逻辑 return x.sum(dim=-1), y.argmax(dim=-1) @wraps(forward) def __call__(self, *args, **kwargs): return super().__call__(*args, **kwargs)
@wraps(forward)会自动把forward的文档字符串、类型注解甚至函数元信息都复制到__call__上,你只需要写一行返回父类__call__的代码就行,比手动复制省不少事。
方法三:泛型基类实现强类型绑定
如果你追求更严谨的类型提示,可以用Python的泛型来定义一个基类,让子类的__call__类型和forward严格绑定:
from typing import TypeVar, Generic import torch.nn as nn from torch import FloatTensor, IntTensor # 定义输入输出的类型变量 InputType = TypeVar('InputType') OutputType = TypeVar('OutputType') class TypedModule(nn.Module, Generic[InputType, OutputType]): def __call__(self, *args: InputType.args, **kwargs: InputType.kwargs) -> OutputType: return super().__call__(*args, **kwargs) # 使用示例 class MyModule(TypedModule[tuple[FloatTensor, FloatTensor], tuple[FloatTensor, IntTensor]]): def __init__(self) -> None: super().__init__() def forward(self, x: FloatTensor, y: FloatTensor) -> tuple[FloatTensor, IntTensor]: """ Args: x (FloatTensor): in shape of BxTxE y (FloatTensor): in shape of BxE Returns: tuple[FloatTensor, IntTensor]: (sth. in shape of BxT, sth. in shape of B) """ # 你的forward逻辑 return x.sum(dim=-1), y.argmax(dim=-1)
这个方案需要你显式指定输入输出的类型,优点是类型提示非常精准,IDE能完美识别__call__的参数和返回值类型,适合对类型严谨性要求高的场景。
另外提一句,如果你用VSCode,确保安装了Pylance插件,它对PyTorch的类型提示支持会更好,配合上面的方法就能完美解决你的问题。
备注:内容来源于stack exchange,提问作者LibrarristShalinward
相关产品推荐
相关产品推荐

