如何将ndarray转换部分赋值给自身?Rust实现PyTorch逻辑避拷贝
解决ndarray中对称赋值避免内存拷贝的方案
因为你提到的赋值区域(i<j和j<i的对应位置)完全无重叠,完全可以避免内存拷贝,直接在原数组上做原地操作即可,下面分两种场景给出具体实现:
1. 从无到有创建符合要求的数组
用Array4::uninit初始化未填充的数组,直接遍历赋值,全程不需要拷贝:
use ndarray::{Array4, Axis}; fn create_symmetric_array<F>(l: usize, a: usize, init_fn: F) -> Array4<f64> where F: Fn(usize, usize, usize, usize) -> f64, { // 初始化未填充的4维数组,无内存拷贝 let mut arr = Array4::uninit((l, l, a, a)); for i in 0..l { for j in 0..l { if i == j { // 对角线位置直接置0 for a_idx in 0..a { for b_idx in 0..a { arr[[i, j, a_idx, b_idx]].write(0.0); } } } else if i < j { // 只计算一次值,同时给(i,j,a,b)和(j,i,b,a)赋值 for a_idx in 0..a { for b_idx in 0..a { let val = init_fn(i, j, a_idx, b_idx); arr[[i, j, a_idx, b_idx]].write(val); arr[[j, i, b_idx, a_idx]].write(val); } } } // i>j的情况已经被i<j时处理,直接跳过 } } // 标记数组为已初始化,无拷贝 unsafe { arr.assume_init() } }
2. 对已存在的数组应用对称规则和对角线置0
用ndarray的切片视图直接操作原数组,不需要转成owned数组:
use ndarray::Array4; fn apply_symmetry_rules(arr: &mut Array4<f64>) { let (l, _, a, _) = arr.dim(); // 对角线位置批量置0 for i in 0..l { arr.slice_mut(s![i, i, .., ..]).fill(0.0); } // 处理对称关系:遍历i<j,将(j,i,b,a)设为(i,j,a,b)的值 for i in 0..l { for j in (i+1)..l { // 获取源区域的视图和目标区域的可变视图 let src_view = arr.slice(s![i, j, .., ..]); let mut dst_view = arr.slice_mut(s![j, i, .., ..]); // 直接将源视图的转置赋值给目标视图,原地修改无拷贝 dst_view.assign(&src_view.t()); } } }
关键说明
你之前用into_owned()产生拷贝,是因为把视图转换成了拥有独立内存的数组——但完全没必要这么做。因为操作区域无重叠,直接通过切片视图访问原数组的内存区域,原地修改即可,全程不会产生额外的内存拷贝。
内容的提问来源于stack exchange,提问作者benjamin-lieser
相关产品推荐
相关产品推荐

