Numpy是否存在原地拼接方法?原数组修改需同步至拼接后数组
Numpy中实现"修改原数组同步到拼接数组"的方案
Numpy本身并没有提供像你示例中in_place_concatenate这样的内置函数,原因是:普通的独立Numpy数组在内存中是分散的块,无法直接生成一个连续的内存视图来同时覆盖多个独立数组。不过可以通过两种方式实现你想要的同步效果:
1. 基于共享内存的原生方案(推荐,若场景允许)
如果你的原数组本来就是从同一个大数组中切分出来的,直接使用这个大数组作为"拼接结果"即可,所有切片都会共享大数组的内存,修改原切片会同步反映到大数组上:
import numpy as np # 创建一个大数组 big_arr = np.zeros(6) # 从大数组中切分出两个子数组 arr1 = big_arr[:3] arr2 = big_arr[3:] # big_arr就是等效的"拼接数组" arr3 = big_arr # 修改原数组,同步反映到arr3 arr1[0] = 10 assert arr3[0] == 10 # 不会触发AssertionError
2. 自定义类模拟"可同步的拼接数组"
如果原数组是完全独立的(内存不连续),可以自定义一个类,通过索引映射来间接访问原数组,实现修改同步的效果:
import numpy as np class InPlaceConcatenator: def __init__(self, arrays): self.arrays = arrays self.total_len = sum(arr.size for arr in arrays) def __getitem__(self, idx): if isinstance(idx, int): current_pos = 0 for arr in self.arrays: if idx < current_pos + arr.size: return arr[idx - current_pos] current_pos += arr.size raise IndexError("索引超出范围") elif isinstance(idx, slice): start, stop, step = idx.indices(self.total_len) segments = [] current_pos = 0 for arr in self.arrays: seg_start = max(0, start - current_pos) seg_stop = min(arr.size, stop - current_pos) if seg_start < seg_stop: segments.append(arr[seg_start:seg_stop:step]) current_pos += arr.size return np.concatenate(segments) else: raise NotImplementedError("暂不支持该索引类型") def __setitem__(self, idx, value): if isinstance(idx, int): current_pos = 0 for arr in self.arrays: if idx < current_pos + arr.size: arr[idx - current_pos] = value return current_pos += arr.size raise IndexError("索引超出范围") elif isinstance(idx, slice): start, stop, step = idx.indices(self.total_len) current_pos = 0 val_idx = 0 for arr in self.arrays: seg_start = max(0, start - current_pos) seg_stop = min(arr.size, stop - current_pos) if seg_start < seg_stop: seg_len = seg_stop - seg_start arr[seg_start:seg_stop:step] = value[val_idx:val_idx+seg_len] val_idx += seg_len current_pos += arr.size else: raise NotImplementedError("暂不支持该索引类型") # 测试代码 arr1 = np.zeros(3) arr2 = np.zeros(3) arr3 = InPlaceConcatenator([arr1, arr2]) arr1[0] = 10 assert arr3[0] == 10 # 正常通过 arr3[4] = 20 assert arr2[1] == 20 # 反向修改也生效
这个类会把对arr3的索引操作映射到对应的原数组上,从而实现修改同步的效果。
内容的提问来源于stack exchange,提问作者Alexander Soare
相关产品推荐
相关产品推荐

