如何从Google VertexAI的aiplatform库获取混淆矩阵或正负样本标签?
获取Vertex AI文本分类模型的混淆矩阵及标签
一、提取混淆矩阵(或TP/FP/FN/TN)
Vertex AI的模型评估结果中,classificationEvaluationMetrics字段包含完整的混淆矩阵数据,可通过高阶或低阶Python SDK直接提取:
1. 使用aiplatform高阶SDK
from google.cloud import aiplatform # 初始化客户端 aiplatform.init(project="你的项目ID", location="你的区域") # 获取模型及评估对象 model = aiplatform.Model("你的模型资源名称") model_evaluation = model.get_model_evaluation() # 解析分类评估指标 class_metrics = model_evaluation.metrics["classificationEvaluationMetrics"] confusion_matrix = class_metrics["confusionMatrix"] # 提取标签与矩阵值 actual_labels = confusion_matrix["rowLabels"] predicted_labels = confusion_matrix["columnLabels"] matrix_rows = confusion_matrix["rows"] # 二分类场景提取TP/FP/FN/TN if len(actual_labels) == 2: tn = matrix_rows[0][0] fp = matrix_rows[0][1] fn = matrix_rows[1][0] tp = matrix_rows[1][1] print(f"TP: {tp}, FP: {fp}, FN: {fn}, TN: {tn}") # 多分类场景遍历所有组合 for actual_idx, row in enumerate(matrix_rows): for pred_idx, count in enumerate(row): print(f"实际标签: {actual_labels[actual_idx]}, 预测标签: {predicted_labels[pred_idx]}, 样本数: {count}")
2. 使用aiplatform_v1低阶SDK
如果需要更精细的proto结构控制,可使用低阶API:
from google.cloud import aiplatform_v1 # 初始化客户端 client = aiplatform_v1.ModelServiceClient() # 构建请求并获取评估结果 request = aiplatform_v1.GetModelEvaluationRequest( name="你的模型评估资源名称" ) evaluation = client.get_model_evaluation(request) # 提取混淆矩阵数据 confusion_matrix = evaluation.metrics.classification_evaluation_metrics.confusion_matrix actual_labels = confusion_matrix.row_labels predicted_labels = confusion_matrix.column_labels # 遍历输出所有分类组合 for i, row in enumerate(confusion_matrix.rows): for j, count in enumerate(row.counts): print(f"实际标签: {actual_labels[i]}, 预测标签: {predicted_labels[j]}, 样本数: {count}")
二、获取训练/评估的标签集合
可通过模型或训练作业的元数据提取标签信息:
from google.cloud import aiplatform # 从模型对象获取标签相关元数据 model = aiplatform.Model("你的模型资源名称") # AutoML模型可从评估指标的labelClasses字段获取所有标签 class_metrics = model.get_model_evaluation().metrics["classificationEvaluationMetrics"] all_labels = class_metrics.get("labelClasses") # 从训练作业获取标签列配置 training_job = aiplatform.CustomJob("你的训练作业资源名称") label_column = training_job.job_spec.parameters.get("label_column")
注:确保模型评估已完成,且任务类型为文本分类,否则classificationEvaluationMetrics字段不会存在。
内容的提问来源于stack exchange,提问作者mitkrieg
相关产品推荐
相关产品推荐

