C++反向模式自动微分中父节点内存提前释放问题:如何正确使用共享指针解决?
C++反向模式自动微分中父节点内存提前释放问题:如何正确使用共享指针解决?
嘿,我仔细看了你的代码和问题,确实是临时变量的生命周期问题在搞鬼!当你把所有运算写在一行(比如Variable z = (x / y) - (x.tangent() * y.cosine()) + x.exponential();)的时候,中间生成的临时Variable对象(比如x/y的结果、x.tangent()的结果)会在这条语句执行完毕后立刻被销毁。但你的子节点还持有这些临时对象的裸指针,等到反向传播的时候访问这些已经被释放的内存,必然会导致崩溃或者错误结果。
用shared_ptr来管理Variable的生命周期是完全正确的思路,它能自动维护对象的引用计数,只要还有子节点持有父节点的shared_ptr,父节点就不会被销毁。下面我一步步给你修改代码:
核心修改点
- 让
Variable继承std::enable_shared_from_this<Variable>,这样每个实例都能获取自身的shared_ptr,避免悬空指针。 - 把
_parents成员从vector<Variable*>改成vector<std::shared_ptr<Variable>>,用共享指针管理父节点引用。 - 运算符和成员函数不再返回
Variable对象,而是返回std::shared_ptr<Variable>,确保中间结果被共享指针持有,不会提前销毁。 - 修改lambda捕获方式,捕获父节点的
shared_ptr而不是裸指针或引用,保证反向传播时父节点依然存在。 - 去掉不必要的
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
相关产品推荐
相关产品推荐

