如何将已删除的numpy数组元素按原正确位置重新插入?
需求场景
需要将从数组中按指定索引删除的元素插回原位置,典型应用为参数优化场景:部分参数固定不参与优化,优化器仅接收非固定参数子集运算,最终需要拼接完整参数集传入损失函数计算。
原有方案的问题
直接使用np.delete删除元素后用同一份索引调用np.insert无法还原数组,原因是np.delete处理多索引时会自动对索引升序排序后执行删除,而np.insert会严格按照传入的索引顺序插入,当索引不是升序时就会出现位置错位,示例如下:
import numpy as np pars = np.array([0,10,20,30,40,50]) exclude = [3,1] # 乱序的待排除索引 parsSubset = np.delete(pars,exclude) excludedPars = pars[exclude] parsRecreated = np.insert(parsSubset,exclude,excludedPars) print(parsRecreated) # 输出:[ 0 30 20 10 40 50],和原数组不一致
优雅实现方案
方案1:掩码标记法(最推荐)
预先构造和原数组长度一致的布尔掩码,标记待优化参数的位置,恢复时直接按掩码赋值即可,完全规避索引顺序问题,逻辑简单且性能更高。
import numpy as np pars = np.array([0,10,20,30,40,50]) exclude = [3,1] # 构造掩码:True为参与优化的参数位置 mask = np.ones(len(pars), dtype=bool) mask[exclude] = False parsSubset = pars[mask] excludedPars = pars[exclude] # 恢复完整数组 parsRecreated = np.zeros_like(pars) parsRecreated[mask] = parsSubset # 赋值优化后的参数 parsRecreated[~mask] = excludedPars # 赋值固定参数 print(parsRecreated) # 输出:[ 0 10 20 30 40 50],完全还原
方案2:索引排序适配法
如果一定要沿用np.delete+np.insert的实现逻辑,只需要在插入前对排除索引做升序排序即可和np.delete的逻辑对齐。
import numpy as np pars = np.array([0,10,20,30,40,50]) exclude = [3,1] parsSubset = np.delete(pars,exclude) excludedPars = pars[exclude] # 对索引升序排序,同时对应调整被排除元素的顺序 sorted_idx = np.argsort(exclude) sorted_exclude = np.array(exclude)[sorted_idx] sorted_excludedPars = np.array(excludedPars)[sorted_idx] parsRecreated = np.insert(parsSubset, sorted_exclude, sorted_excludedPars) print(parsRecreated) # 输出:[ 0 10 20 30 40 50]
内容的提问来源于stack exchange,提问作者andrea m.
相关产品推荐
相关产品推荐

