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

如何导入导出为TXT的sklearn回归树并在测试集上预测?

从export_text导出的TXT文件还原决策树并实现预测

1. 正确读取TXT文件

pd.read_csv不适合解析这种树结构文本,直接用Python内置文件读取方法读取所有有效行:

# 读取树的文本内容,过滤空行
with open(path + "myTree.txt", "r") as f:
    tree_lines = [line.strip() for line in f if line.strip()]

2. 手动实现基于文本树的预测逻辑

export_text的输出有固定格式:

  • 非叶子节点以|---开头,格式为|--- 特征名 <=/> 阈值
  • 叶子节点包含value: [预测值]
  • 缩进级别(行首的| 数量)对应树的层级

下面是直接基于文本行实现的预测函数,遍历规则找到样本对应的叶子节点:

def predict_from_tree_text(sample, tree_lines):
    current_idx = 0
    current_depth = 0

    while True:
        line = tree_lines[current_idx]
        # 计算当前节点的层级:统计行首的`|   `数量
        depth = line.count('|   ')
        
        # 回退到父节点时跳过当前行
        if depth < current_depth:
            current_idx += 1
            continue
        
        # 叶子节点:提取预测值
        if 'value:' in line:
            return float(line.split('[')[1].split(']')[0])
        
        # 非叶子节点:解析判断条件
        condition = line.split('|--- ')[1]
        feature, op, threshold = condition.split()
        threshold = float(threshold)
        sample_val = sample[feature]

        # 根据条件选择子节点路径
        if op == '<=':
            if sample_val <= threshold:
                current_depth = depth + 1
                current_idx += 1
            else:
                # 跳转到同层级的右分支
                current_idx += 1
                while current_idx < len(tree_lines) and tree_lines[current_idx].count('|   ') != depth:
                    current_idx += 1
        elif op == '>':
            if sample_val > threshold:
                current_depth = depth + 1
                current_idx += 1
            else:
                current_idx += 1
                while current_idx < len(tree_lines) and tree_lines[current_idx].count('|   ') != depth:
                    current_idx += 1
        else:
            raise ValueError(f"不支持的操作符: {op}")

3. 测试预测功能

用测试数据验证,和原sklearn模型的预测结果对比:

# 文本树预测结果
text_preds = [predict_from_tree_text(row, tree_lines) for _, row in df_test.iterrows()]
# 原模型预测结果
original_preds = tree.predict(df_test)

print("文本树预测结果:", text_preds)
print("原模型预测结果:", original_preds)

补充说明

  • 该方法无需额外依赖,直接解析export_text的原生输出
  • 如果树结构复杂,也可以考虑将文本转换成嵌套字典的结构后再做预测,逻辑更清晰,但遍历效率和上述方法相当

内容的提问来源于stack exchange,提问作者user6703592

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 15:00:57