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

如何为包装NdArray的Tensor类型实现iter()迭代器?

问题:Tensor行迭代器的生命周期错误解决

问题背景

开发Rust张量库时,二维Tensor内部用Rc<RefCell<TensorData>>包装ndarray::Array2<f64>。需要实现iter()方法迭代张量的行用于随机梯度下降,但当前实现触发E0515生命周期错误,且写法繁琐。

Tensor定义代码

use ndarray::prelude::*;

pub struct Tensor(Rc<RefCell<TensorData>>);

pub struct TensorData {
    pub data: Array2<f64>,
    pub grad: Array2<f64>,
    // other fields...
}

impl TensorData {
    fn new(data: Array2<f64>) -> TensorData {
        let shape = data.raw_dim();
        TensorData {
            data,
            grad: Array2::zeros(shape),
            // other fields...
        }
    }
}

impl Tensor {
    pub fn new(array: Array2<f64>) -> Tensor {
        Tensor(Rc::new(RefCell::new(TensorData::new(array))))
    }

    pub fn data(&self) -> impl Deref<Target = Array2<f64>> + '_ {
        Ref::map((*self.0).borrow(), |mi| &mi.data)
    }
}

错误的iter()实现

impl Tensor {
    // other methods...
    
    pub fn iter(&self) -> impl Iterator<Item = Tensor> + '_ {
        self.data().outer_iter().map(|el| {
            let reshaped_and_cloned_el = el.into_shape((el.shape()[0], 1)).unwrap().mapv(|el| el.clone());
            reshaped_and_cloned_el
        }).map(|el| Tensor::new(el))
    }
}

编译错误

error[E0515]: cannot return value referencing temporary value
   --> src/tensor/mod.rs:348:9
    |
348 |           self.data().outer_iter().map(|el| {
    |           ^----------
    |           |
    |  _________temporary value created here
    | |
349 | |             let reshaped_and_cloned_el = el.into_shape((el.shape()[0], 1)).unwrap().mapv(...
350 | |             reshaped_and_cloned_el
351 | |         }).map(|el| Tensor::new(el))
    | |____________________________________^ returns a value referencing data owned by the current function
    |
    = help: use `.collect()` to allocate the iterator

解决方案

问题根源

self.data()返回的是临时的Ref<'_, Array2<f64>>对象,outer_iter()生成的迭代器依赖该临时对象的生命周期,但返回的迭代器会超出临时对象的存活范围,导致生命周期不匹配。同时原代码的克隆逻辑冗余。

优化后的实现

方案1:绑定Ref生命周期(推荐)

先持有RefCell的借用,确保迭代器的生命周期与&self一致,同时优化行的克隆与重塑:

impl Tensor {
    // other methods...
    
    pub fn iter(&self) -> impl Iterator<Item = Tensor> + '_ {
        // 持有Ref,确保其生命周期覆盖整个迭代器
        let data_ref = self.0.borrow();
        // 基于持有的Ref生成行迭代器
        data_ref.data.outer_iter().map(|row| {
            // 直接获取行的所有权,再重塑为二维张量格式
            let row_owned = row.to_owned();
            let row_2d = row_owned.into_shape((row_owned.len(), 1)).unwrap();
            Tensor::new(row_2d)
        })
    }
}
说明
  • self.0.borrow()获取的Ref<TensorData>绑定到data_ref,其生命周期与&self相同,避免临时对象失效问题。
  • row.to_owned()直接复制行数据生成拥有所有权的Array1,比mapv(|x| x.clone())更高效简洁。
  • 拥有所有权的数组进行into_shape操作无生命周期限制,生成的Array2可直接传入Tensor::new。

方案2:提前收集为Vec(适合需要并发修改场景)

如果需要在迭代原张量的同时允许其他地方修改它,可以先将所有行克隆为Vec<Tensor>,再返回其迭代器:

impl Tensor {
    pub fn iter(&self) -> std::vec::IntoIter<Tensor> {
        let data_ref = self.0.borrow();
        let tensors = data_ref.data.outer_iter()
            .map(|row| {
                let row_2d = row.to_owned().into_shape((row.len(), 1)).unwrap();
                Tensor::new(row_2d)
            })
            .collect::<Vec<_>>();
        // 释放Ref,允许后续修改
        drop(data_ref);
        tensors.into_iter()
    }
}
说明
  • 提前将所有行转换为Tensor并收集到Vec中,之后可以立即释放Ref,不阻塞其他对TensorData的可变借用。
  • 缺点是会提前分配内存存储所有行的副本,适合数据量较小的场景。

内容的提问来源于stack exchange,提问作者JS4137

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 15:43:19