能否在PyTorch JIT函数中追加列表元素?问题咨询
PyTorch JIT列表使用报错原因及解决方案
问题原因
PyTorch JIT对空列表的类型推断有严格限制:当初始化空列表my_list = []时,JIT会默认将其推断为List[Tensor]类型(因为PyTorch核心对象是Tensor)。但你后续追加的是int类型的元素,类型不匹配,因此触发类型错误。
这并非PyTorch JIT完全禁止使用列表,而是要求列表的元素类型必须明确且一致。
解决方案
1. 明确标注列表元素类型
使用torch.jit.annotate结合typing.List,提前告知JIT列表的元素类型,避免自动推断错误:
import torch from typing import List @torch.jit.script def my_function(x): # 明确标注这是一个int类型的列表 my_list = torch.jit.annotate(List[int], []) for i in range(int(x)): my_list.append(i) return my_list a = my_function(10) print(a) # 输出:[0, 1, 2, 3, 4, 5, 6, 7, 8, 9]
2. 针对Tensor元素的场景(适配参数优化需求)
如果后续要填充的是需优化的Tensor参数,直接标注列表为List[torch.Tensor]即可:
import torch from typing import List @torch.jit.script def my_function(x): my_list = torch.jit.annotate(List[torch.Tensor], []) for i in range(int(x)): # 追加带梯度的Tensor元素 my_list.append(torch.tensor(i, dtype=torch.float32, requires_grad=True)) return my_list a = my_function(10) print(a)
3. 替代方案:用Tensor动态构建(批量场景更高效)
如果不需要动态追加的灵活性,可直接预分配Tensor并填充,避免使用列表:
import torch @torch.jit.script def my_function(x): length = int(x) # 预分配Tensor存储空间 my_tensor = torch.zeros(length, dtype=torch.int32) for i in range(length): my_tensor[i] = i return my_tensor a = my_function(10) print(a) # 输出:tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9], dtype=torch.int32)
总结
PyTorch JIT允许使用列表,但必须保证列表元素类型明确且统一。空列表默认推断为List[Tensor],若存储其他类型(如int、float),需通过torch.jit.annotate提前标注;若存储Tensor,明确标注也能避免类型歧义。
内容的提问来源于stack exchange,提问作者Shep Bryan
相关产品推荐
相关产品推荐

