PyTorch中哪些函数或模块要求输入为contiguous连续张量?
PyTorch 连续张量(contiguous)常见问题解答
当某些函数或模块需要连续张量时,若未显式调用
tensor.contiguous()转换,会抛出如下异常:RuntimeError: invalid argument 1: input is not contiguous at .../src/torch/lib/TH/generic/THTensor.c:231
哪些函数/模块要求输入为连续张量?官方是否有相关记录?
- 要求输入为连续张量的算子多为底层C/C++/CUDA实现、依赖连续内存访问优化性能的算子,常见场景包括:
- 执行
view()等对内存布局对齐有严格要求的张量操作 - 部分早期版本的卷积、RNN算子,以及
grid_sample()、Embedding参数校验、部分排序操作等
- 执行
- PyTorch官方目前没有统一整理全量要求连续输入的算子清单,相关要求要么散落在对应算子文档的注意事项中,要么在触发报错时通过错误提示告知用户。
哪些场景下需要调用contiguous方法?
- 运行代码抛出
input is not contiguous类错误时,需要显式在报错算子前对输入张量调用contiguous() - 你已经通过
transpose()、permute()、步长不为1的切片、narrow()、expand()等操作得到非连续张量,且后续需要调用内存敏感的底层算子时,可以提前调用避免报错 - 性能优化场景下,若后续会频繁访问张量内存,提前转换为连续张量可以降低访存开销,效率高于算子内部隐式转换。
Conv1d是否要求输入为连续张量?
当前新版本PyTorch的Conv1d(含Conv2d、Conv3d等)均已兼容非连续输入,算子内部会自动判断是否需要做连续转换,无需用户显式处理。如果官方文档没有明确标注要求连续输入,默认算子已做兼容,仅当实际运行抛出非连续相关报错时再处理即可。
为什么PyTorch不像Theano一样自动转换非连续输入?
PyTorch设计上优先给用户提供内存开销的控制权,隐式自动转换会在用户无感知的情况下产生额外的内存拷贝开销,大张量场景下对性能影响非常明显,因此PyTorch仅在高频常用算子内部做了自动兼容,其余场景会通过报错提醒用户显式处理。
内容的提问来源于stack exchange,提问作者Albert
相关产品推荐
相关产品推荐

