在Sympy中为自定义符号实现正确的高阶导数(模拟r=√(x²+y²+z²)的导数行为)
在Sympy中为自定义符号实现正确的高阶导数(模拟r=√(x²+y²+z²)的导数行为)
你的思路方向是对的,但只子类化Symbol并重写diff方法还不够——问题出在当r和其他符号结合成Mul这类表达式后,Sympy的默认求导逻辑会把r当成独立变量处理,忽略了它和x/y/z的依赖关系,所以二阶求导时漏掉了链式法则的部分。
更合适的解决方案:子类化Function而非Symbol
在Sympy里,Function类专门用来表示依赖于其他变量的函数,而非独立符号。用它来实现你的自定义r,Sympy会自动帮你处理链式法则,高阶导数也能正确计算。
方案1:带表达式展开的版本(需要时可以展开成√(x²+y²+z²))
from sympy import symbols, Function, diff, sqrt x, y, z = symbols("x y z") class r(Function): @classmethod def eval(cls): # 定义r的具体表达式,调用doit()时会自动展开 return sqrt(x**2 + y**2 + z**2) def diff(self, var): if var == x: return x / self elif var == y: return y / self elif var == z: return z / self # 对其他变量求导时用默认逻辑 return super().diff(var) # 实例化自定义r r_obj = r()
测试一下:
>>> diff(r_obj, x) x/r() >>> diff(diff(r_obj, x), x) 1/r() - x**2/r()**3
完全符合预期的二阶导数结果。如果需要展开r()成具体表达式,只需要调用r_obj.doit()即可。
方案2:纯抽象符号版本(不需要展开表达式)
如果只需要r保持抽象符号,同时拥有正确的导数行为,去掉eval方法就行:
from sympy import symbols, Function, diff x, y, z = symbols("x y z") class r(Function): def diff(self, var): if var in (x, y, z): return var / self return super().diff(var)
测试结果一样:
>>> diff(r(), x) x/r() >>> diff(diff(r(), x), x) 1/r() - x**2/r()**3
为什么原来的Symbol子类化方案不行?
Sympy的Symbol默认被设计为独立变量,当你得到x/r这个Mul对象后,求导时Sympy会默认认为r和x无关,所以计算d/dx(1/r)时直接返回0,只保留了x的导数乘以1/r的部分,也就是1/r,漏掉了链式法则的项。如果非要用Symbol子类化,需要手动注册自定义的求导规则,这会复杂很多,不如用Function来得直接。
备注:内容来源于stack exchange,提问作者Antimon
相关产品推荐
相关产品推荐

