如何为同时使用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
相关产品推荐
相关产品推荐

