如何正确定义带有长度属性的PyTorch Dataset类型?
如何正确定义带有长度属性的PyTorch Dataset类型?
这个问题我之前也踩过坑,PyTorch的Dataset基类确实没在类型定义里包含__len__,导致静态类型检查器(比如mypy)没法识别哪些子类实现了这个方法——你自己定义DatasetWithLength后,传入现有子类会报类型不匹配警告,就是因为这个原因。下面给你几个实用的解决办法:
最优方案:用Protocol实现结构类型检查
从Python 3.8开始,标准库typing模块的Protocol支持鸭子类型——只要类的结构符合协议要求(实现了指定方法),类型检查器就会认为它兼容,不需要显式继承。这刚好适配你的场景,毕竟大部分自定义PyTorch Dataset都会实现__len__和__getitem__。
代码示例如下:
from typing import Protocol, TypeVar from torch.utils.data import Dataset, IterableDataset # 定义协变类型变量,适配Dataset的输出类型 T_co = TypeVar('T_co', covariant=True) class SizedDataset(Protocol[T_co]): # 声明协议需要的方法,用...代替具体实现 def __len__(self) -> int: ... def __getitem__(self, index: int) -> T_co: ... class MetaDataset(Dataset): def __init__(self, regular_dataset: SizedDataset, iterable_dataset: IterableDataset): self.regular_dataset = regular_dataset self.iterable_dataset = iterable_dataset # 你的其他初始化逻辑...
只要你的FirstDataset实现了__len__和__getitem__(这几乎是所有非Iterable Dataset的标准操作),类型检查器就会自动识别它符合SizedDataset类型,不会再报之前的警告。
备选方案:抽象基类+注册现有子类
如果你用的是Python 3.8之前的版本,或者更习惯抽象基类的方式,可以定义一个带__len__抽象方法的ABC,然后把你要使用的现有Dataset子类注册到这个ABC上:
from abc import ABC, abstractmethod from torch.utils.data import Dataset, IterableDataset # 导入你的自定义Dataset类 from your_module import FirstDataset class DatasetWithLength(Dataset, ABC): @abstractmethod def __len__(self) -> int: ... # 把已有的Dataset子类注册到我们的ABC上 DatasetWithLength.register(FirstDataset) class MetaDataset(Dataset): def __init__(self, regular_dataset: DatasetWithLength, iterable_dataset: IterableDataset): self.regular_dataset = regular_dataset self.iterable_dataset = iterable_dataset # 你的其他初始化逻辑...
不过这个方法需要手动注册每个你要用到的Dataset子类,灵活性不如Protocol,适合子类数量不多的场景。
临时应急:类型断言或忽略警告
如果你只是想快速消除警告,不想改类型定义,可以用类型断言或者直接忽略警告,但这种方法会失去类型检查的优势,不推荐长期使用:
方式1:类型断言
from typing import cast from your_module import FirstDataset, FirstIterableDataset foo = MetaDataset( cast(DatasetWithLength, FirstDataset()), FirstIterableDataset(), )
方式2:忽略警告
foo = MetaDataset( FirstDataset(), # type: ignore FirstIterableDataset(), )
总的来说,Protocol是最优雅、最灵活的解决方案,推荐优先使用。
备注:内容来源于stack exchange,提问作者Richie Bendall
相关产品推荐
相关产品推荐

