You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.21 10:34:57