Python SymPy中如何将StrictlyLessThan转换为整数或浮点数
问题根源
两个报错本质是同一个原因:使用sympy.numpy做运算时,比较表达式square_sum < 10不会返回原生布尔值或数值类型,而是返回SymPy的符号关系对象StrictlyLessThan。这类符号表达式既不能直接和浮点数做算术乘法,也不能直接被if语句判断真值,因此触发错误。
可行实现方案
方案1:转换flag为数值类型直接参与乘法
对符号关系对象做显式数值求值,转换为0/1对应的浮点数/整数后即可正常做乘法运算,适配你第一版代码的写法:
import sympy.numpy as snp def function(x): square_sum = x[0]**2 + x[1]**2 flag = square_sum < 10 # 对符号关系求值后转为浮点数,满足条件为1.0,不满足为0.0 flag_val = float(flag.evalf()) cubic_sum = x[0]**3 + x[1]**3 return flag_val * cubic_sum # 测试 print(function(snp.array([2, 3]))) # 输出0.0,因为2²+3²=13不满足小于10的条件 print(function(snp.array([1, 2]))) # 输出9.0,因为1²+2²=5满足条件,1³+2³=9
如果需要支持批量数组的矢量化运算,直接用snp.where做条件选择更稳妥,不需要手动做乘法:
import sympy.numpy as snp def function(x): square_sum = x[0]**2 + x[1]**2 cubic_sum = x[0]**3 + x[1]**3 # 逐元素判断,满足条件返回cubic_sum,否则返回0.0 return snp.where(square_sum < 10, cubic_sum, 0.0)
方案2:修复条件分支写法
如果要使用if-else分支逻辑,不能直接把符号关系对象放入判断条件,需要先求解出明确的布尔值:
import sympy.numpy as snp def function(x): square_sum = x[0]**2 + x[1]**2 flag = square_sum < 10 # 先求解关系表达式的布尔结果再做分支判断 if bool(flag.doit()): return x[0]**3 + x[1]**3 else: return 0.0
注意:if-else写法仅支持单组输入的标量判断,传入批量数组做矢量化运算时必须使用
snp.where写法,否则仍会触发真值判断错误。
内容的提问来源于stack exchange,提问作者Euler_Salter
相关产品推荐
相关产品推荐

