PyTorch用torch.as_strided实现Toeplitz矩阵负步长报错及相关问题
问题1:对as_strided的理解是否正确
你的理解存在偏差。as_strided本质是直接复用输入张量的底层内存生成新视图,而非主动填充值到新数组。它的工作逻辑是:以输入张量的第一个元素地址为起始位置,按照指定的size(输出张量形状)和stride(输出张量每个维度每前进1个元素,底层内存地址偏移的元素个数)直接映射取值,不会自动复制或填充数据。你理解的“按步长放值到新下标”是后续copy()操作触发的行为,as_strided本身仅生成视图。PyTorch目前不支持负步长的原因是负步长意味着要访问起始地址之前的内存空间,容易触发越界问题,也不符合PyTorch的张量内存布局设计。
问题2:规避负步长的改写方案
可以完全不用as_strided实现等价功能,以下是和scipy.linalg.toeplitz逻辑完全对齐的可运行版本:
import torch def toeplitz_torch(c, r=None): # 用as_tensor代替tensor,保留输入梯度 c = torch.as_tensor(c).ravel() if r is None: r = torch.conj(c) else: r = torch.as_tensor(r).ravel() n_col = len(r) n_row = len(c) # 构造包含所有托普利茨元素的一维数组 vals = torch.cat([torch.flip(c, dims=(0,)), r[1:]]) # 生成行列索引,自动适配输入设备 row_idx = torch.arange(n_row, device=c.device).unsqueeze(1) col_idx = torch.arange(n_col, device=c.device).unsqueeze(0) # 计算对应vals的正索引 vals_idx = col_idx - row_idx + (n_row - 1) return vals[vals_idx]
该实现和原版本逻辑本质一致,只是把as_strided的内存映射逻辑换成了显式索引取值,完全避开了负步长限制,同时支持CPU/GPU多设备运行。
问题3:梯度支持情况
上述改写后的函数完全支持相对于c和r的梯度传递。函数中所有操作(as_tensor、ravel、conj、flip、cat、索引取值)都是PyTorch官方支持的可微操作,只要输入c和r本身是需要梯度的张量,反向传播时梯度可以正常回传。
内容的提问来源于stack exchange,提问作者MRicci
相关产品推荐
相关产品推荐

