Numpy as_strided处理float32图像分块触发segmentation fault问题求助
问题原因与修复方案
错误根源
代码核心问题出在new_strides的计算逻辑上:
NumPy数组的strides属性本身就代表对应维度每前进一个元素需要跳过的字节数,你在计算时额外乘以arr.itemsize相当于把步长放大了N倍(N为单元素占用的字节数):
- uint8类型的itemsize为1,额外相乘后结果和原生strides一致,所以运行正常
- float32类型的itemsize为4,步长直接放大4倍,访问数组时会超出合法内存范围,触发segmentation fault。
修复代码
只需要删除new_strides计算时多余的arr.itemsize *即可,修正后的函数如下:
import numpy as np def patch_sample(arr, patch_size=(128, 128), stride=64): arr = np.ascontiguousarray(arr) H, W, C = arr.shape h, w = patch_size patches_h = int((H - h) / stride + 1) patches_w = int((W - w)/ stride + 1) new_shape = (patches_h, patches_w, h, w, C) # 移除多余的arr.itemsize乘法 new_strides = np.array( [ W * arr.strides[1] * stride, arr.strides[1] * stride, W * arr.strides[1], arr.strides[1], 1 ] ) patches_out = np.lib.stride_tricks.as_strided( arr, shape=new_shape, strides=new_strides, writeable=False ) patches_out_cp = np.ascontiguousarray(patches_out) patches_out_cp = patches_out_cp.reshape((-1, h, w, C)) return patches_out_cp
测试两种类型的输入都可以正常运行:
# 正常运行 test_array = np.random.randint(0, 100, (1000, 1500, 3), dtype=np.uint8) print(patch_sample(test_array).shape) # 正常运行,不再段错误 test_array = np.random.rand(1000, 1500, 3).astype(np.float32) print(patch_sample(test_array).shape)
更安全的替代实现
如果你的NumPy版本 >= 1.20.0,推荐直接使用内置的sliding_window_view函数实现重叠分块,不需要手动计算strides,避免出错:
def patch_sample_safe(arr, patch_size=(128,128), stride=64): H, W = arr.shape[:2] h, w = patch_size # 生成滑动窗口视图 windows = np.lib.stride_tricks.sliding_window_view(arr, (h,w), axis=(0,1)) # 按步长采样 patches = windows[::stride, ::stride] # 调整维度顺序为 (n_patches_h, n_patches_w, h, w, C) 后展平 patches = patches.transpose(0,1,3,4,2).reshape(-1, h, w, arr.shape[-1]) return patches
内容的提问来源于stack exchange,提问作者michaal94
相关产品推荐
相关产品推荐

