From 130c1b5796d9876a5f853fa0bea88e808cfda4ad Mon Sep 17 00:00:00 2001 From: yolain Date: Thu, 30 Apr 2026 03:24:51 +0800 Subject: [PATCH] fix(math): enhance evaluate_formula to handle list inputs and return results accordingly --- py/libs/math.py | 43 +++++++++++++++++++++++++++++-------------- py/nodes/logic.py | 7 +++++-- 2 files changed, 34 insertions(+), 16 deletions(-) diff --git a/py/libs/math.py b/py/libs/math.py index 07b9ec0..a651327 100644 --- a/py/libs/math.py +++ b/py/libs/math.py @@ -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: diff --git a/py/nodes/logic.py b/py/nodes/logic.py index c4ec3e1..929ed74 100755 --- a/py/nodes/logic.py +++ b/py/nodes/logic.py @@ -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