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

C++反向模式自动微分中父节点内存提前释放问题:如何正确使用共享指针解决?

C++反向模式自动微分中父节点内存提前释放问题:如何正确使用共享指针解决?

嘿,我仔细看了你的代码和问题,确实是临时变量的生命周期问题在搞鬼!当你把所有运算写在一行(比如Variable z = (x / y) - (x.tangent() * y.cosine()) + x.exponential();)的时候,中间生成的临时Variable对象(比如x/y的结果、x.tangent()的结果)会在这条语句执行完毕后立刻被销毁。但你的子节点还持有这些临时对象的裸指针,等到反向传播的时候访问这些已经被释放的内存,必然会导致崩溃或者错误结果。

用shared_ptr来管理Variable的生命周期是完全正确的思路,它能自动维护对象的引用计数,只要还有子节点持有父节点的shared_ptr,父节点就不会被销毁。下面我一步步给你修改代码:

核心修改点

  1. 让Variable继承std::enable_shared_from_this<Variable>,这样每个实例都能获取自身的shared_ptr,避免悬空指针。
  2. 把_parents成员从vector<Variable*>改成vector<std::shared_ptr<Variable>>,用共享指针管理父节点引用。
  3. 运算符和成员函数不再返回Variable对象,而是返回std::shared_ptr<Variable>,确保中间结果被共享指针持有,不会提前销毁。
  4. 修改lambda捕获方式,捕获父节点的shared_ptr而不是裸指针或引用,保证反向传播时父节点依然存在。
  5. 去掉不必要的const_cast,通过共享指针的const接口处理只读操作。

修改后的完整代码

#include <iostream>
#include <vector>
#include <cmath>
#include <functional>
#include <memory>

using namespace std;

class Variable : public enable_shared_from_this<Variable> {
private:
    double value;
    double grad;
    function<void()> _backward;
    vector<shared_ptr<Variable>> _parents;
    bool visited;

public:
    // 私有构造函数,通过静态create方法创建实例
    Variable(double value) : value(value), grad(0.0), visited(false) {
        this->_backward = [](){};
    }

    // 静态方法创建shared_ptr实例,强制使用智能指针管理对象
    static shared_ptr<Variable> create(double value) {
        return shared_ptr<Variable>(new Variable(value));
    }

    shared_ptr<Variable> operator + (const shared_ptr<Variable>& other) {
        cout << "add" << endl;
        auto out = create(this->value + other->value);
        out->_parents.push_back(shared_from_this());
        out->_parents.push_back(other);

        out->_backward = [self = shared_from_this(), other, out]() {
            self->grad += out->grad;
            other->grad += out->grad;
            cout << "addition " << self->grad << "\t" << other->grad << endl;
        };
        return out;
    }

    shared_ptr<Variable> operator * (const shared_ptr<Variable>& other) {
        cout << "mul" << endl;
        auto out = create(this->value * other->value);
        out->_parents.push_back(shared_from_this());
        out->_parents.push_back(other);

        out->_backward = [self = shared_from_this(), other, out]() {
            self->grad += other->value * out->grad;
            other->grad += self->value * out->grad;
            cout << self->grad << "\t" << other->grad << endl;
        };
        return out;
    }

    shared_ptr<Variable> operator - (const shared_ptr<Variable>& other) {
        cout << "sub" << endl;
        auto out = create(this->value - other->value);
        out->_parents.push_back(shared_from_this());
        out->_parents.push_back(other);

        out->_backward = [self = shared_from_this(), other, out]() {
            self->grad += out->grad;
            other->grad -= out->grad;
            cout << "Subtraction: " << self->grad << "\t" << other->grad << endl;
        };
        return out;
    }

    shared_ptr<Variable> operator / (const shared_ptr<Variable>& other) {
        cout << "div" << endl;
        auto out = create(this->value / other->value);
        out->_parents.push_back(shared_from_this());
        out->_parents.push_back(other);

        out->_backward = [self = shared_from_this(), other, out]() {
            self->grad += out->grad / other->value;
            other->grad -= (self->value / pow(other->value, 2)) * out->grad;
            cout << "Division: " << self->grad << "\t" << other->grad << endl;
        };
        return out;
    }

    shared_ptr<Variable> power(const shared_ptr<Variable>& other) {
        cout << "pow" << endl;
        auto out = create(pow(this->value, other->value));
        out->_parents.push_back(shared_from_this());
        out->_parents.push_back(other);

        out->_backward = [self = shared_from_this(), other, out]() {
            self->grad += other->value * pow(self->value, (other->value-1)) * out->grad;
            other->grad += pow(self->value, other->value) * log(self->value) * out->grad;
        };
        return out;
    }

    shared_ptr<Variable> sine() {
        cout << "sin" << endl;
        auto out = create(sin(this->value));
        out->_parents.push_back(shared_from_this());

        out->_backward = [self = shared_from_this(), out]() {
            self->grad += cos(self->value) * out->grad;
            cout << "Sine: " << self->grad << endl;
        };
        return out;
    }

    shared_ptr<Variable> cosine() {
        cout << "cos" << endl;
        auto out = create(cos(this->value));
        out->_parents.push_back(shared_from_this());

        out->_backward = [self = shared_from_this(), out]() {
            self->grad -= sin(self->value) * out->grad;
            cout << "Cosine: " << self->grad << endl;
        };
        return out;
    }

    shared_ptr<Variable> tangent() {
        cout << "tan" << endl;
        auto out = create(tan(this->value));
        out->_parents.push_back(shared_from_this());

        out->_backward = [self = shared_from_this(), out]() {
            self->grad += (1 / pow(cos(self->value), 2)) * out->grad;
            cout << "Tangent: " << self->grad << endl;
        };
        return out;
    }

    shared_ptr<Variable> exponential() {
        cout << "exp" << endl;
        auto out = create(exp(this->value));
        out->_parents.push_back(shared_from_this());

        out->_backward = [self = shared_from_this(), out]() {
            self->grad += exp(self->value) * out->grad;
            cout << "Exponential: " << self->grad << endl;
        };
        return out;
    }

    void backward() {
        this->grad = 1.0;
        vector<shared_ptr<Variable>> funcs;
        funcs.push_back(shared_from_this());

        while(!funcs.empty()){
            auto f = funcs.back();
            funcs.pop_back();

            if (f->visited){
                continue;
            }
            f->visited = true;
            if (f->_backward){
                f->_backward();
            }
            for (auto& parent: f->_parents){
                funcs.push_back(parent);
            }
        }
    }

    double getGrad() const {
        return this->grad;
    }

    double getValue() const {
        return this->value;
    }
};

int main(){
    auto x = Variable::create(2);
    auto y = Variable::create(3);
    auto z = (*x / *y) - (*x->tangent() * *y->cosine()) + *x->exponential();

    cout << endl << "Backward Pass" << endl << endl;
    z->backward();

    cout << "Value of z: " << z->getValue() << endl;
    cout << "Gradient with respect to x: " << x->getGrad() << endl;
    cout << "Gradient with respect to y: " << y->getGrad() << endl;
    return 0;
}

关键说明

  • 用enable_shared_from_this让每个Variable实例能安全获取自身的shared_ptr,避免手动创建shared_ptr导致的引用计数错误。
  • 所有运算符和成员函数都返回shared_ptr<Variable>,中间临时结果会被自动管理,不会在语句结束后立刻销毁。
  • Lambda表达式中捕获shared_ptr(比如self = shared_from_this()),保证反向传播时父节点的内存依然有效。
  • 构造函数改成私有,通过静态create方法创建实例,强制使用shared_ptr管理对象,彻底避免栈上创建Variable导致的生命周期问题。

这样修改后,不管你是分步骤创建变量还是一行链式调用,父节点的内存都不会被提前释放,反向传播就能正常运行了!

备注:内容来源于stack exchange,提问作者Dipin Adhikari

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.15 09:39:52