Kedro项目中通过DataCatalog规范保存预测结果与测试标签的方法咨询
问题解答
你当前在节点内硬编码实例化CSVDataSet的写法不符合Kedro的开发规范。Kedro设计的核心原则之一就是计算逻辑和IO逻辑完全解耦,节点只需要关注核心计算逻辑,不需要关心数据从哪来、存到哪去,这些IO配置全部统一交给catalog.yml管理,也能避免硬编码路径带来的后续维护问题,同时可以直接复用Kedro内置的数据集版本管理、多存储适配、权限控制等能力。
以下是符合规范的实现方案:
实现步骤
步骤1:在catalog.yml中配置目标数据集
编辑conf/base/catalog.yml,添加你要保存的预测结果数据集配置:
# 预测结果数据集配置 test_prediction_result: type: pandas.CSVDataSet filepath: data/test.csv # 可选:开启数据集版本管理 # versioned: true
步骤2:修改节点函数,返回待保存的数据
原节点函数没有返回值,现在只需要把构造好的预测结果DataFrame作为返回值返回即可,不需要在节点内做任何保存操作:
import logging import numpy as np import pandas as pd def report_accuracy(predictions: np.ndarray, test_y: pd.DataFrame) -> pd.DataFrame: """Node for reporting the accuracy of the predictions performed by the previous node, and return prediction results for persistence. """ # 获取真实类别索引 target = np.argmax(test_y.to_numpy(), axis=1) # 计算预测准确率 accuracy = np.sum(predictions == target) / target.shape[0] # 记录模型准确率 log = logging.getLogger(__name__) log.info("Model accuracy on test set: %0.2f%%", accuracy * 100) # 直接返回待保存的结果即可,不需要在节点内处理保存逻辑 return pd.DataFrame({"target": target , "prediction": predictions})
步骤3:修改管道配置,绑定节点输出与catalog数据集
在管道定义的位置,将节点的输出名称设置为和catalog.yml中配置的数据集键名完全一致,Kedro运行时会自动匹配,拿到节点返回值后自动完成保存:
from kedro.pipeline import Pipeline, node from .nodes import report_accuracy def create_pipeline(**kwargs) -> Pipeline: return Pipeline( [ node( func=report_accuracy, inputs=["predictions", "test_y"], # 输出名称和catalog中配置的数据集键名完全匹配 outputs="test_prediction_result", name="report_accuracy_node", ) ] )
后续如果需要修改保存路径、更换存储介质(比如改为存到对象存储)、修改存储格式,只需要修改catalog.yml的配置即可,不需要改动任何节点逻辑。
内容的提问来源于stack exchange,提问作者BlueMango
相关产品推荐
相关产品推荐

