如何导入导出为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
相关产品推荐
相关产品推荐

