You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.16 18:20:12