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

如何对std::shared_ptr进行深拷贝?(节点树求导场景)

问题描述

我有一个节点树的简单类MultNode,它是抽象类Node的子类。_pRight和_pLeft是std::shared_ptr<Node>对象。MultNode::DerFun()用于实现乘法节点的求导功能,遵循乘积求导法则:$\frac{d}{dx}[f(x)g(x)]=f(x)g'(x)+f'(x)g(x)$。我需要创建临时节点pTmpR和pTmpL,且只能使用智能指针,但不知道如何对std::shared_ptr进行深拷贝,当前代码如下:

void MultNode::DerFun() {
    auto pTmpR =  _pRight;
    auto pTmpL = _pLeft;
    _pRight->DerFun(); // sin(pi/4) * cos(pi/4);
    _pLeft->DerFun();  // der(sin(pi/4) * cos(pi/4));
    _pRight = std::make_shared<AddNode>(
        std::make_shared<MultNode>(_pRight, pTmpL),
        std::make_shared<MultNode>(pTmpR, _pLeft));
    _pLeft = std::make_shared<NumNode>(1);
}
解决方案

当前代码的问题在于pTmpR = _pRight是浅拷贝,两个智能指针指向同一对象,后续调用_pRight->DerFun()修改原节点时,pTmpR指向的内容也会被改变,无法保留求导前的原节点状态。要实现std::shared_ptr的深拷贝,需要给抽象基类添加克隆方法,让每个子类实现自身的深拷贝逻辑。

  1. 给抽象类Node添加纯虚克隆方法
class Node {
public:
    // 其他纯虚接口
    virtual std::shared_ptr<Node> clone() const = 0;
    virtual ~Node() = default;
};
  1. 各子类实现clone方法(递归深拷贝节点树)
  • 数值节点NumNode示例:
class NumNode : public Node {
private:
    double _value;
public:
    NumNode(double val) : _value(val) {}
    std::shared_ptr<Node> clone() const override {
        return std::make_shared<NumNode>(_value);
    }
    // 其他成员方法
};
  • 乘法节点MultNode示例:
class MultNode : public Node {
private:
    std::shared_ptr<Node> _pLeft;
    std::shared_ptr<Node> _pRight;
public:
    MultNode(std::shared_ptr<Node> left, std::shared_ptr<Node> right) 
        : _pLeft(std::move(left)), _pRight(std::move(right)) {}
    std::shared_ptr<Node> clone() const override {
        // 递归克隆左右子节点,实现整个子树的深拷贝
        return std::make_shared<MultNode>(_pLeft->clone(), _pRight->clone());
    }
    // 其他成员方法
};
  1. 修改MultNode::DerFun()使用深拷贝的临时节点
void MultNode::DerFun() {
    // 深拷贝原左右节点,保存求导前的状态
    auto pTmpR = _pRight->clone();
    auto pTmpL = _pLeft->clone();
    
    // 对原节点执行求导
    _pRight->DerFun();
    _pLeft->DerFun();
    
    // 按照乘积求导法则构建新节点树
    _pRight = std::make_shared<AddNode>(
        std::make_shared<MultNode>(_pRight, pTmpL),
        std::make_shared<MultNode>(pTmpR, _pLeft));
    _pLeft = std::make_shared<NumNode>(1);
}
说明

克隆方法的核心是递归复制整个节点树:每个子类负责复制自身的成员变量,对于子节点则调用其clone()方法完成深拷贝。这样pTmpR和pTmpL指向的是完全独立于原节点的拷贝,后续修改原节点的求导结果不会影响这两个临时节点,完全符合乘积求导法则中保留原函数值的需求。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 09:05:35