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

在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 03:07:17