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

如何创建规范PyTorch模型类型的Protocol类?解决继承nn.Module的类型错误

解决PyTorch模型Protocol类型定义的错误

你遇到的类型错误是因为Protocol是用于定义结构类型的抽象接口,不能直接继承非Protocol类(比如torch.nn.Module),同时继承两者会触发类型检查器的报错。

正确的做法是让Protocol只继承Protocol本身,专注定义你需要的方法签名(比如特定的forward方法),实际的模型类依然继承nn.Module并实现这个方法即可。

正确代码示例

from typing import Protocol, Optional, Tensor, runtime_checkable
import torch
import torch.nn as nn

# 定义协议,仅继承Protocol
@runtime_checkable  # 可选,添加后支持运行时检查实例是否符合协议
class DiffusionModelProtocol(Protocol):
    def forward(
        self,
        x: torch.Tensor,
        h: dict[str, torch.Tensor],
        node_mask: torch.Tensor,
        edge_mask: Optional[torch.Tensor] = None,
        context: Optional[torch.Tensor] = None,
    ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
        ...

# 实际的扩散模型类,继承nn.Module并实现符合协议的forward
class MyDiffusionModel(nn.Module):
    def forward(
        self,
        x: torch.Tensor,
        h: dict[str, torch.Tensor],
        node_mask: torch.Tensor,
        edge_mask: Optional[torch.Tensor] = None,
        context: Optional[torch.Tensor] = None,
    ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
        # 这里写你的模型逻辑
        return x, h

为什么这样可行?

PyTorch的所有模型都必须继承nn.Module,所以只要你的模型类继承了nn.Module并实现了协议中定义的forward方法签名,类型检查器会自动识别它符合DiffusionModelProtocol的类型要求。

如果需要明确约束协议的实例必须是nn.Module的子类,也可以在协议中添加一个类型注解的属性来暗示:

class DiffusionModelProtocol(Protocol):
    # 暗示该实例是nn.Module的子类
    __class__: type[nn.Module]
    
    def forward(...):
        ...

不过通常第一种方式已经足够满足类型检查的需求。

内容的提问来源于stack exchange,提问作者Dan Jackson

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 20:40:15