如何用Python从JSON结构生成多级决策树及核心功能实现
从JSON规则构建Python决策树的实用方案
示例结构与TreeNode类
示例JSON规则
{ "rule": "A > 5", "true": { "rule": "B < 3", "true": { "rule": "C == 1", "true": "Leaf 1", "false": "Leaf 2" }, "false": { "rule": "D != 4", "true": "Leaf 3", "false": "Leaf 4" } }, "false": { "rule": "E >= 2", "true": { "rule": "F < 6", "true": "Leaf 5", "false": "Leaf 6" }, "false": { "rule": "G == 0", "true": "Leaf 7", "false": "Leaf 8" } } }
TreeNode类定义
class TreeNode: def __init__(self, rule, truebranch, falsebranch): self.rule = rule self.truebranch = truebranch self.falsebranch = falsebranch def evaluate(self, conditions): # 待实现逻辑 pass
1. 规则解析:高效识别运算符与值
核心思路
优先匹配长运算符(如>=、<=、==、!=),避免被拆分为单个符号;用正则表达式一次性提取变量名、运算符、目标值,同时处理值的类型转换(支持整数/浮点数)。
实现代码
import re # 定义运算符匹配顺序(长运算符在前,避免拆分) OPERATORS = ['>=', '<=', '==', '!=', '>', '<'] OP_PATTERN = '|'.join(re.escape(op) for op in OPERATORS) RULE_REGEX = re.compile(rf'^(\w+)\s*({OP_PATTERN})\s*([\d.]+)$') def parse_rule(rule_str): """解析规则字符串,返回(变量名, 运算符, 转换后的值)""" match = RULE_REGEX.match(rule_str.strip()) if not match: raise ValueError(f"无效规则格式: {rule_str}") var_name, op, value_str = match.groups() # 尝试转换为数值类型,优先整数再浮点数 try: value = int(value_str) except ValueError: try: value = float(value_str) except ValueError: raise ValueError(f"规则值无法转换为数值: {value_str}") return var_name, op, value
关键说明
- 正则表达式严格匹配
变量 运算符 数值格式,自动忽略空格 - 类型转换逻辑兼容整数、浮点数两种规则值类型
- 不符合格式的规则直接抛出明确的错误信息,便于排查
2. evaluate方法实现:递归遍历决策树
核心思路
递归处理每个节点:解析当前规则→根据输入条件判断分支→若分支是字符串则返回叶子结果,若为TreeNode实例则继续递归调用evaluate。
实现代码
class TreeNode: def __init__(self, rule, truebranch, falsebranch): self.rule = rule self.truebranch = truebranch self.falsebranch = falsebranch def evaluate(self, conditions): # 解析当前规则 var_name, op, target_val = parse_rule(self.rule) # 检查输入条件是否包含所需变量 if var_name not in conditions: raise KeyError(f"条件中缺少变量: {var_name}") current_val = conditions[var_name] # 统一转换为数值类型,确保比较合法 if not isinstance(current_val, (int, float)): try: current_val = int(current_val) except ValueError: current_val = float(current_val) # 执行运算符比较 result = False if op == '>': result = current_val > target_val elif op == '<': result = current_val < target_val elif op == '>=': result = current_val >= target_val elif op == '<=': result = current_val <= target_val elif op == '==': result = current_val == target_val elif op == '!=': result = current_val != target_val # 递归处理分支 next_branch = self.truebranch if result else self.falsebranch if isinstance(next_branch, str): return next_branch elif isinstance(next_branch, TreeNode): return next_branch.evaluate(conditions) else: raise TypeError(f"无效分支类型: {type(next_branch)}")
关键说明
- 递归终止条件为遇到字符串类型的叶子节点
- 自动转换输入条件的变量类型,避免类型不兼容的比较错误
- 分支类型判断确保只处理合法的节点或叶子,提前拦截非法格式
3. 边缘情况处理:规避常见错误
常见错误场景与解决方案
- 无效规则格式:比如规则写成
A > five或A >= B(非数值比较),通过parse_rule的正则匹配和类型转换抛出异常,可在构建树前批量校验所有规则 - 输入条件缺少变量:evaluate时直接抛出
KeyError,或提前在调用evaluate前检查条件的键是否覆盖所有规则中的变量 - 叶节点类型非法:比如叶节点是数字/字典而非字符串,在evaluate的分支判断中抛出
TypeError,或在构建TreeNode时提前校验分支类型 - 类型不兼容的比较:比如输入条件中的变量是字符串但无法转成数值,在evaluate的类型转换步骤抛出异常,或添加默认值 fallback 逻辑
- 循环引用的JSON规则:构建TreeNode时添加递归深度检测,避免无限递归
额外优化:构建树时的前置校验
在从JSON生成TreeNode实例的阶段,提前校验所有规则和分支的合法性:
def build_tree(json_data): """从JSON数据递归构建TreeNode实例,同时前置校验""" if isinstance(json_data, str): # 叶子节点直接返回 return json_data # 检查节点必填字段 required_fields = ['rule', 'true', 'false'] if not all(field in json_data for field in required_fields): raise ValueError(f"节点缺少必填字段,需包含: {required_fields}") # 提前解析规则,校验有效性 parse_rule(json_data['rule']) # 递归构建子节点 true_branch = build_tree(json_data['true']) false_branch = build_tree(json_data['false']) return TreeNode(json_data['rule'], true_branch, false_branch)
内容的提问来源于stack exchange,提问作者kiruthikpurpose
相关产品推荐
相关产品推荐

