在PyTorch中如何创建带方法和属性的TypedDict等效类?
实现兼具类型提示、自定义方法与属性的TypedDict等效类
针对你在PyTorch中使用TypedDict的痛点,以下是三种可行的解决方案,兼顾类型提示、自定义方法/属性访问,同时适配PyTorch默认的collate函数:
方案一:继承dict + 属性封装(最贴合TypedDict的dict本质)
这种方案本质仍是dict,完美兼容PyTorch默认collate,同时通过@property实现属性访问,避免字符串键拼写错误,还能直接添加自定义方法。
from typing import List import torch from torch import Tensor class Item(dict): # 类型注解保留完整类型提示 source: str encoding: Tensor # CHW result: float def __init__(self, source: str, encoding: Tensor, result: float): super().__init__(source=source, encoding=encoding, result=result) # 属性封装:替代字符串键访问 @property def source(self) -> str: return self["source"] @source.setter def source(self, value: str): self["source"] = value @property def encoding(self) -> Tensor: return self["encoding"] @encoding.setter def encoding(self, value: Tensor): self["encoding"] = value @property def result(self) -> float: return self["result"] @result.setter def result(self, value: float): self["result"] = value # 自定义方法示例 def normalize_encoding(self): self.encoding = self.encoding / self.encoding.max() class Batch(dict): source: List[str] encoding: Tensor # NCHV prediction: Tensor # Shape N def __init__(self, source: List[str], encoding: Tensor, prediction: Tensor): super().__init__(source=source, encoding=encoding, prediction=prediction) @property def source(self) -> List[str]: return self["source"] @source.setter def source(self, value: List[str]): self["source"] = value @property def encoding(self) -> Tensor: return self["encoding"] @encoding.setter def encoding(self, value: Tensor): self["encoding"] = value @property def prediction(self) -> Tensor: return self["prediction"] @prediction.setter def prediction(self, value: Tensor): self["prediction"] = value # 自定义方法:从批次提取单个Item def item_at(self, index: int) -> Item: return Item( source=self.source[index], encoding=self.encoding[index], result=float(self.prediction[index].item()) ) # 批量处理方法示例 def normalize_encodings(self): self.encoding = self.encoding / self.encoding.max(dim=(1,2,3), keepdim=True)[0]
优势:
- 完全兼容PyTorch默认collate函数:
List[Item]传入DataLoader后,会自动被整理成Batch结构,和TypedDict行为一致。 - 属性访问替代字符串键,避免拼写错误。
- 类型提示完整,IDE可自动补全和类型检查。
- 无额外依赖,原生Python实现。
不足:
- 需要为每个字段编写
@property和setter,代码略显繁琐(可通过装饰器或代码生成工具简化)。
方案二:dataclass + 自定义collate(代码最简洁)
利用Python标准库的dataclasses,无需手写属性封装,自动生成初始化、打印等方法,只需少量修改即可适配collate。
from dataclasses import dataclass, asdict from typing import List import torch from torch import Tensor from torch.utils.data import DataLoader @dataclass class Item: source: str encoding: Tensor # CHW result: float # 自定义方法 def normalize_encoding(self): self.encoding = self.encoding / self.encoding.max() @dataclass class Batch: source: List[str] encoding: Tensor # NCHV prediction: Tensor # Shape N # 自定义方法:从批次提取单个Item def item_at(self, index: int) -> Item: return Item( source=self.source[index], encoding=self.encoding[index], result=float(self.prediction[index].item()) ) # 适配collate:先转dict再处理,最后转回Batch对象 def collate_items(items: List[Item]) -> Batch: collated_dict = torch.utils.data.default_collate([asdict(item) for item in items]) return Batch(**collated_dict) # 使用示例 if __name__ == "__main__": items = [ Item(source="text1", encoding=torch.rand(3, 224, 224), result=0.5), Item(source="text2", encoding=torch.rand(3, 224, 224), result=0.8) ] loader = DataLoader(items, batch_size=2, collate_fn=collate_items) for batch in loader: print(batch.source) # ['text1', 'text2'] print(batch.encoding.shape) # torch.Size([2, 3, 224, 224])
优势:
- 代码极简,无需手写属性和初始化方法。
- 类型提示完整,IDE支持良好。
- 原生标准库,无额外依赖。
不足:
- 需要自定义collate函数,比dict继承方案多一步适配。
- Batch/Item不是dict类型,若需兼容原有dict代码,需用
asdict()转换。
方案三:Pydantic BaseModel(带类型验证的增强方案)
使用Pydantic库实现,自带类型验证、属性访问,支持复杂类型转换,适合需要严格类型检查的场景。
from pydantic import BaseModel from typing import List import torch from torch import Tensor from torch.utils.data import DataLoader class Item(BaseModel): source: str encoding: Tensor # CHW result: float # 自定义方法 def normalize_encoding(self): self.encoding = self.encoding / self.encoding.max() # 允许PyTorch Tensor等自定义类型 class Config: arbitrary_types_allowed = True class Batch(BaseModel): source: List[str] encoding: Tensor # NCHV prediction: Tensor # Shape N # 自定义方法:从批次提取单个Item def item_at(self, index: int) -> Item: return Item( source=self.source[index], encoding=self.encoding[index], result=float(self.prediction[index].item()) ) class Config: arbitrary_types_allowed = True # 适配collate函数 def collate_items(items: List[Item]) -> Batch: collated_dict = torch.utils.data.default_collate([item.dict() for item in items]) return Batch(**collated_dict) # 使用示例 if __name__ == "__main__": items = [ Item(source="text1", encoding=torch.rand(3, 224, 224), result=0.5), Item(source="text2", encoding=torch.rand(3, 224, 224), result=0.8) ] loader = DataLoader(items, batch_size=2, collate_fn=collate_items) for batch in loader: print(batch.source) print(batch.encoding.shape)
优势:
- 自带类型验证,初始化时类型不匹配会直接抛出错误,提前规避问题。
- 属性访问简洁,支持自动类型转换。
- 可轻松转成dict,兼容原有代码。
不足:
- 需要额外安装Pydantic库(
pip install pydantic)。 - 同样需要自定义collate函数。
方案选择建议
- 若需完全兼容原有TypedDict的dict行为,选方案一。
- 若追求代码简洁、无额外依赖,选方案二。
- 若需要严格类型验证和高级特性,选方案三。
内容的提问来源于stack exchange,提问作者Xiiryo
相关产品推荐
相关产品推荐

