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

如何为同时使用Protocol与类自身属性的类方法添加类型提示?

解决mypy识别self同时属于Protocol与父类的问题

问题描述

我正在实现一个基于PyTorch Lightning的LightningDataModule通用类,用于搭建训练/验证/测试DataLoader。该类提供通用功能,将train_ds、val_ds、test_ds属性的初始化留给子类实现。我尝试通过HasTrainValTestDatasets Protocol约束子类必须实现这些属性,但为类方法的self添加类型提示时,mypy报错提示HasTrainValTestDatasets没有_batch_size_train等类自身的属性。需要让mypy识别self同时属于HasTrainValTestDatasets Protocol和GenericTrainValTestDataModule类。

代码示例

from typing import Protocol
import pytorch_lightning as pl
from torch.utils.data import Dataset, DataLoader


class HasTrainValTestDatasets(Protocol):
    @property
    def train_ds(self) -> Dataset: ...

    @property
    def val_ds(self) -> Dataset: ...

    @property
    def test_ds(self) -> Dataset: ...


class GenericTrainValTestDataModule(pl.LightningDataModule):
    def __init__(
        self,
        batch_size_train: int,
        batch_size_eval: int,
        num_workers: int = 0,
    ):
        self._batch_size_train = batch_size_train
        self._batch_size_eval = batch_size_eval
        self._num_workers = num_workers

    def train_dataloader(self: HasTrainValTestDatasets) -> DataLoader:
        return DataLoader(self.train_ds, batch_size=self._batch_size_train, num_workers=self._num_workers)

    def val_dataloader(self: HasTrainValTestDatasets) -> DataLoader:
        return DataLoader(self.val_ds, batch_size=self._batch_size_eval, num_workers=self._num_workers)

    def test_dataloader(self: HasTrainValTestDatasets) -> DataLoader:
        return DataLoader(self.test_ds, batch_size=self._batch_size_eval, num_workers=self._num_workers)

mypy报错信息

d.py:30: error: "HasTrainValTestDatasets" has no attribute "_batch_size_train"  [attr-defined]
d.py:30: error: "HasTrainValTestDatasets" has no attribute "_num_workers"  [attr-defined]
d.py:35: error: "HasTrainValTestDatasets" has no attribute "_batch_size_eval"  [attr-defined]
d.py:35: error: "HasTrainValTestDatasets" has no attribute "_num_workers"  [attr-defined]
d.py:40: error: "HasTrainValTestDatasets" has no attribute "_batch_size_eval"  [attr-defined]
d.py:40: error: "HasTrainValTestDatasets" has no attribute "_num_workers"  [attr-defined]

复现步骤

virtualenv venv
pip install torch pytorch-lightning
mypy --install-types d.py

解决方案

方法1:使用Python 3.11+的Self类型 + 类型交集

Python 3.11引入了typing.Self,结合&运算符可以直接表示self同时属于当前类和目标Protocol。修改方法的self类型提示即可:

from typing import Protocol, Self
import pytorch_lightning as pl
from torch.utils.data import Dataset, DataLoader


class HasTrainValTestDatasets(Protocol):
    @property
    def train_ds(self) -> Dataset: ...

    @property
    def val_ds(self) -> Dataset: ...

    @property
    def test_ds(self) -> Dataset: ...


class GenericTrainValTestDataModule(pl.LightningDataModule):
    def __init__(
        self,
        batch_size_train: int,
        batch_size_eval: int,
        num_workers: int = 0,
    ):
        self._batch_size_train = batch_size_train
        self._batch_size_eval = batch_size_eval
        self._num_workers = num_workers

    def train_dataloader(self: Self & HasTrainValTestDatasets) -> DataLoader:
        return DataLoader(self.train_ds, batch_size=self._batch_size_train, num_workers=self._num_workers)

    def val_dataloader(self: Self & HasTrainValTestDatasets) -> DataLoader:
        return DataLoader(self.val_ds, batch_size=self._batch_size_eval, num_workers=self._num_workers)

    def test_dataloader(self: Self & HasTrainValTestDatasets) -> DataLoader:
        return DataLoader(self.test_ds, batch_size=self._batch_size_eval, num_workers=self._num_workers)

方法2:定义包含父类属性的新Protocol(兼容Python 3.10及以下)

针对旧Python版本,可定义一个同时继承HasTrainValTestDatasets并声明父类属性的新Protocol,以此作为self的类型:

from typing import Protocol
import pytorch_lightning as pl
from torch.utils.data import Dataset, DataLoader


class HasTrainValTestDatasets(Protocol):
    @property
    def train_ds(self) -> Dataset: ...

    @property
    def val_ds(self) -> Dataset: ...

    @property
    def test_ds(self) -> Dataset: ...

# 定义同时包含数据集属性和父类内部属性的Protocol
class GenericDataModuleProtocol(HasTrainValTestDatasets, Protocol):
    _batch_size_train: int
    _batch_size_eval: int
    _num_workers: int


class GenericTrainValTestDataModule(pl.LightningDataModule):
    def __init__(
        self,
        batch_size_train: int,
        batch_size_eval: int,
        num_workers: int = 0,
    ):
        self._batch_size_train = batch_size_train
        self._batch_size_eval = batch_size_eval
        self._num_workers = num_workers

    def train_dataloader(self: GenericDataModuleProtocol) -> DataLoader:
        return DataLoader(self.train_ds, batch_size=self._batch_size_train, num_workers=self._num_workers)

    def val_dataloader(self: GenericDataModuleProtocol) -> DataLoader:
        return DataLoader(self.val_ds, batch_size=self._batch_size_eval, num_workers=self._num_workers)

    def test_dataloader(self: GenericDataModuleProtocol) -> DataLoader:
        return DataLoader(self.test_ds, batch_size=self._batch_size_eval, num_workers=self._num_workers)

方法3:使用typing_extensions.Self(兼容旧Python版本)

如果需要在Python 3.10及以下版本使用类似Python 3.11的写法,可安装typing_extensions库,导入其中的Self类型:

首先安装依赖:

pip install typing_extensions

修改后的代码:

from typing import Protocol
from typing_extensions import Self
import pytorch_lightning as pl
from torch.utils.data import Dataset, DataLoader


class HasTrainValTestDatasets(Protocol):
    @property
    def train_ds(self) -> Dataset: ...

    @property
    def val_ds(self) -> Dataset: ...

    @property
    def test_ds(self) -> Dataset: ...


class GenericTrainValTestDataModule(pl.LightningDataModule):
    def __init__(
        self,
        batch_size_train: int,
        batch_size_eval: int,
        num_workers: int = 0,
    ):
        self._batch_size_train = batch_size_train
        self._batch_size_eval = batch_size_eval
        self._num_workers = num_workers

    def train_dataloader(self: Self & HasTrainValTestDatasets) -> DataLoader:
        return DataLoader(self.train_ds, batch_size=self._batch_size_train, num_workers=self._num_workers)

    def val_dataloader(self: Self & HasTrainValTestDatasets) -> DataLoader:
        return DataLoader(self.val_ds, batch_size=self._batch_size_eval, num_workers=self._num_workers)

    def test_dataloader(self: Self & HasTrainValTestDatasets) -> DataLoader:
        return DataLoader(self.test_ds, batch_size=self._batch_size_eval, num_workers=self._num_workers)

适用场景

  • Python 3.11+版本优先选择方法1,写法最简洁直观。
  • Python 3.10及以下版本可选择方法2或方法3,方法3更贴近新版本的语法风格。

内容的提问来源于stack exchange,提问作者Maciej Ziaja

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 09:14:54