You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何让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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.14 18:08:07