在NumPy或PyTorch中,如何编写将strides转换为shape的函数?
如何在NumPy/PyTorch中将strides转换为shape?
首先先纠正并翻译你给出的「shape转strides」函数——原代码存在参数名错误、未定义变量的问题,按行优先(C顺序)逻辑修正后的正确实现如下:
def shape_to_strides(shape): strides = [1] # 从右往左遍历维度,计算每个维度的元素步长 for s in reversed(shape): strides.append(s * strides[-1]) # 移除最后一个元素(对应数组总元素数),反转后得到从左到右的strides return tuple(reversed(strides[:-1]))
注:这个函数返回的是元素步长,如果要得到NumPy的字节步长,需要再乘以对应数据类型的字节大小。
回到核心问题:仅靠strides无法唯一还原shape,必须结合两个关键信息:数组的总元素数(或总字节数+元素字节大小),以及strides的排列规则(行/列优先或自定义)。下面分别给出NumPy和PyTorch的实现方案:
NumPy中的实现
NumPy的strides返回的是字节步长,所以需要先转换为元素步长,再结合总元素数计算shape:
import numpy as np def strides_to_shape(strides, dtype, total_elements): itemsize = dtype.itemsize # 将字节步长转为元素步长 elem_strides = [s // itemsize for s in strides] shape = [] remaining = total_elements for i in range(len(elem_strides)): if i == len(elem_strides) - 1: # 最后一个维度的大小就是剩余元素数 shape.append(remaining) else: # 当前维度大小 = 下一个元素步长 // 当前元素步长 dim_size = elem_strides[i+1] // elem_strides[i] shape.append(dim_size) remaining = remaining // dim_size return tuple(shape)
示例验证
# 连续数组 arr = np.zeros((2, 3, 4), dtype=np.int32) print("原shape:", arr.shape) # 输出 (2, 3, 4) print("原strides:", arr.strides) # 输出 (48, 16, 4) recovered_shape = strides_to_shape(arr.strides, arr.dtype, arr.size) print("恢复的shape:", recovered_shape) # 输出 (2, 3, 4) # 切片后的非连续数组 arr_slice = arr[:, :, ::2] print("切片后shape:", arr_slice.shape) # 输出 (2, 3, 2) print("切片后strides:", arr_slice.strides) # 输出 (48, 16, 8) recovered_shape = strides_to_shape(arr_slice.strides, arr_slice.dtype, arr_slice.size) print("恢复的shape:", recovered_shape) # 输出 (2, 3, 2)
PyTorch中的实现
PyTorch的stride()返回的是元素步长,无需转换,直接结合总元素数(numel()方法获取)即可计算:
import torch def strides_to_shape(strides, total_elements): shape = [] remaining = total_elements for i in range(len(strides)): if i == len(strides) - 1: shape.append(remaining) else: dim_size = strides[i+1] // strides[i] shape.append(dim_size) remaining = remaining // dim_size return tuple(shape)
示例验证
# 连续张量 tensor = torch.zeros((2, 3, 4), dtype=torch.int32) print("原shape:", tensor.shape) # 输出 torch.Size([2, 3, 4]) print("原strides:", tensor.stride()) # 输出 (12, 4, 1) recovered_shape = strides_to_shape(tensor.stride(), tensor.numel()) print("恢复的shape:", recovered_shape) # 输出 (2, 3, 4) # 切片后的非连续张量 tensor_slice = tensor[:, :, ::2] print("切片后shape:", tensor_slice.shape) # 输出 torch.Size([2, 3, 2]) print("切片后strides:", tensor_slice.stride()) # 输出 (12, 4, 2) recovered_shape = strides_to_shape(tensor_slice.stride(), tensor_slice.numel()) print("恢复的shape:", recovered_shape) # 输出 (2, 3, 2)
注意事项
- 上述方法仅适用于规则排列的数组/张量(即使整体不连续,每个维度内部是连续块),对于高级索引生成的非规则数组,strides无法准确还原shape。
- 总元素数是必不可少的输入,否则无法确定最后一个维度的大小,也无法验证前面维度的计算是否正确。
内容的提问来源于stack exchange,提问作者olives
相关产品推荐
相关产品推荐

