Python中修改字典值类型时的类型提示最佳实践
问题原因
- 示例代码中的
to_tensor函数采用了原地修改输入字典的实现,输入参数my_dict的类型标注为Dict[str, npt.NDArray[np.uint8]],意味着该字典的值只能是numpy数组,原地将值替换为torch张量的操作本身就违反了输入参数的类型约定。 - 初始定义的
test_dict会被mypy自动推断为Dict[str, npt.NDArray[np.float64]],哪怕后续将to_tensor的返回值重新赋值给test_dict,mypy也会因为类型冲突,依然将其识别为存储numpy数组的字典,而numpy数组没有size()方法,因此抛出attr-defined错误。 - 额外的类型不匹配问题:代码中用
np.random.random生成的是浮点类型数组,和函数参数要求的np.uint8类型不匹配,本身也会触发类型错误。
最优解决方案
最优方案是避免原地修改输入字典,改为在函数内部构造新的字典返回,彻底规避原地修改带来的类型矛盾,同时符合无副作用的函数设计规范,修改后可完全通过mypy检查:
from typing import Dict import numpy as np import numpy.typing as npt import torch import torchvision.transforms as T def to_tensor(input_dict: Dict[str, npt.NDArray[np.float64]]) -> Dict[str, torch.FloatTensor]: # 新建字典存储转换后的张量,不修改原输入 output_dict = {} for key, val in input_dict.items(): output_dict[key] = T.functional.to_tensor(val) return output_dict test_dict = {"foo": np.random.random((3,10,10)), "bar": np.random.random((3, 10, 10))} test_dict = to_tensor(test_dict) print(test_dict['foo'].size())
如果你确实有特殊场景需要原地修改字典,可使用
typing.cast强制声明类型,但该做法跳过了类型校验,仅建议在确定逻辑无问题的场景使用:from typing import cast # 函数返回后强制转换类型 test_dict = cast(Dict[str, torch.FloatTensor], to_tensor(test_dict))
内容的提问来源于stack exchange,提问作者Andrew Stewart
相关产品推荐
相关产品推荐

