fix(math): enhance evaluate_formula to handle list inputs and return results accordingly

This commit is contained in:
yolain
2026-04-30 03:24:51 +08:00
parent 3cf9ab4e63
commit 130c1b5796
2 changed files with 34 additions and 16 deletions
+29 -14
View File
@@ -4,7 +4,7 @@ Math utility functions for formula evaluation
import math
import re
def evaluate_formula(formula: str, a=0, b=0, c=0, d=0) -> float:
def evaluate_formula(formula: str, a=0, b=0, c=0, d=0):
"""
计算字符串数学公式
@@ -23,7 +23,7 @@ def evaluate_formula(formula: str, a=0, b=0, c=0, d=0) -> float:
d: 变量d的值
Returns:
计算结果
如果任意输入为list则返回list[float],否则返回float
Examples:
>>> evaluate_formula("a + b", 1, 2)
@@ -60,19 +60,34 @@ def evaluate_formula(formula: str, a=0, b=0, c=0, d=0) -> float:
# 常量
'pi': math.pi,
'e': math.e,
# 变量
'a': float(a),
'b': float(b),
'c': float(c),
'd': float(d),
}
try:
# 使用eval计算公式,限制可用的函数和变量
result = eval(formula, {"__builtins__": {}}, safe_dict)
return float(result)
except Exception as e:
raise ValueError(f"公式计算错误: {str(e)}")
# 判断是否有 list 输入
list_inputs = {k: v for k, v in {'a': a, 'b': b, 'c': c, 'd': d}.items() if isinstance(v, (list, tuple))}
scalar_inputs = {k: v for k, v in {'a': a, 'b': b, 'c': c, 'd': d}.items() if not isinstance(v, (list, tuple))}
def _eval_single(vals: dict) -> float:
env = dict(safe_dict)
env.update({k: float(v) for k, v in vals.items()})
try:
result = eval(formula, {"__builtins__": {}}, env)
return float(result)
except Exception as e:
raise ValueError(f"公式计算错误: {str(e)}")
if not list_inputs:
# 全是标量
return _eval_single({k: v for k, v in {'a': a, 'b': b, 'c': c, 'd': d}.items()})
# 有 list 输入,逐元素计算
max_len = max(len(v) for v in list_inputs.values())
results = []
for i in range(max_len):
vals = {k: float(v) for k, v in scalar_inputs.items()}
for k, v in list_inputs.items():
vals[k] = float(v[i] if i < len(v) else v[-1])
results.append(_eval_single(vals))
return results
def ceil_value(value: float) -> int:
+5 -2
View File
@@ -530,6 +530,9 @@ class simpleMath(io.ComfyNode):
def execute(cls, value, a=0, b=0, c=0):
try:
result = evaluate_formula(value, a, b, c)
if isinstance(result, list):
result_int = [int(r) for r in result]
return io.NodeOutput(result_int, result, [r != 0 for r in result])
result_int = int(result)
return io.NodeOutput(result_int, result, result_int != 0)
except Exception as e:
@@ -565,14 +568,14 @@ class simpleMathDual(io.ComfyNode):
def execute(cls, value1, value2, a=0, b=0, c=0, d=0):
try:
result1 = evaluate_formula(value1, a, b, c, d)
result1_int = int(result1)
result1_int = [int(r) for r in result1] if isinstance(result1, list) else int(result1)
except Exception as e:
log_node_warn(f"公式1计算错误: {str(e)}")
result1 = 0.0
result1_int = 0
try:
result2 = evaluate_formula(value2, a, b, c, d)
result2_int = int(result2)
result2_int = [int(r) for r in result2] if isinstance(result2, list) else int(result2)
except Exception as e:
log_node_warn(f"公式2计算错误: {str(e)}")
result2 = 0.0