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

如何在Tensorflow Federated中提取y_true与y_pred生成分类报告

问题描述

处理多分类任务时,需要通过真实标签与预测标签生成分类报告,完成联邦学习模型效果评估,当前卡点为不清楚如何在联邦计算流程中提取y_true和y_pred。

现有联邦模型训练代码

for round_num in range(0, NUM_ROUNDS):
    train_metrics = eval_process(state.model, test_data)['eval']
    state, _= iterative_process.next(state, train_data)

    print(f'Round {round_num:3d}: {train_metrics}')
    data_frame = data_frame.append({'Round': round_num,
                                      **train_metrics}, ignore_index=True)
  

test_metrics = eval_process(state.model, test_data)
print("The final evaluation is: ")
print(test_metrics)

return data_frame

目标分类报告生成逻辑

from sklearn.metrics import classification_report

y_pred = model.predict(x_test, batch_size=64, verbose=1)
y_pred_bool = np.argmax(y_pred, axis=1)

print(classification_report(y_test, y_pred_bool))
解决方案

默认的eval_process接口仅返回聚合后的标量评估指标(比如准确率、损失值),不会保留单样本维度的预测值和真实标签,需要单独实现样本级结果的收集逻辑,步骤如下:

  • 自定义预测结果收集函数,逐batch遍历测试集,拿到每个样本的真实类别和预测类别:
import numpy as np
from sklearn.metrics import classification_report

def collect_sample_results(model, test_data):
    y_true_list = []
    y_pred_list = []
    for batch in test_data:
        x, y = batch
        # 模型推理得到输出logits
        logits = model.predict(x, verbose=0)
        # 多分类场景取概率最大的索引作为预测类别
        pred_cls = np.argmax(logits, axis=1)
        # 如果标签是one-hot编码,同步转成类别索引;如果是索引格式直接使用
        true_cls = np.argmax(y, axis=1) if len(y.shape) > 1 and y.shape[-1] > 1 else y
        y_true_list.extend(true_cls.tolist())
        y_pred_list.extend(pred_cls.tolist())
    return np.array(y_true_list), np.array(y_pred_list)
  • 在原有训练流程结束、拿到最终训练好的state.model后,调用上述函数拿到全量测试集的标签和预测结果,直接生成分类报告即可,插入位置参考:
for round_num in range(0, NUM_ROUNDS):
    train_metrics = eval_process(state.model, test_data)['eval']
    state, _= iterative_process.next(state, train_data)

    print(f'Round {round_num:3d}: {train_metrics}')
    data_frame = data_frame.append({'Round': round_num,
                                      **train_metrics}, ignore_index=True)
  

test_metrics = eval_process(state.model, test_data)
print("The final evaluation is: ")
print(test_metrics)

# 插入以下代码生成分类报告
y_true, y_pred = collect_sample_results(state.model, test_data)
print(classification_report(y_true, y_pred))

return data_frame
  • 适配不同联邦框架的注意事项:
    • 如果使用TFF、FATE等封装型联邦框架,需要自定义客户端评估逻辑,要求客户端返回本地分片的样本级预测结果和标签,不要在客户端侧直接聚合成指标标量再回传,否则服务端拿不到单样本数据。
    • 如果测试数据分散在多个客户端存储,需要遍历所有客户端的本地测试分片分别收集结果,在服务端拼接为完整数组后再生成报告,不要跨客户端做指标聚合。
    • 如果数据集标签本身就是类别索引格式,不需要做np.argmax转换,直接收集原始标签即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.31 13:01:20