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

如何在联邦学习Transformer代码中获取分类指标并绘制性能曲线

解决方案

1. 添加F1、Recall、Precision评估指标

Keras原生支持多分类场景下的精确率、召回率、F1分数计算,直接导入对应指标并配置到模型编译流程即可。混淆矩阵则需要通过模型预测结果与真实标签对比生成。

修改模型编译代码

将原有的metrics='acc'替换为包含多分类指标的列表,注意根据你的分类任务调整num_classes和average参数:

from tensorflow.keras.metrics import CategoricalAccuracy, Precision, Recall, F1Score

# 初始化多分类指标,macro表示对每个类别计算后取平均
metrics = [
    CategoricalAccuracy(name='acc'),
    Precision(name='precision', average='macro'),
    Recall(name='recall', average='macro'),
    F1Score(name='f1_score', average='macro', num_classes=你的类别数量)
]

local_model.compile(
    loss=tf.keras.losses.CategoricalCrossentropy(),
    optimizer=tf.keras.optimizers.Adam(learning_rate=0.001),
    metrics=metrics
)

计算混淆矩阵

在全局模型测试阶段,新增混淆矩阵计算逻辑:

from sklearn.metrics import confusion_matrix
import numpy as np

def test_model(test_x, test_y, global_model, round_num):
    # 生成模型预测结果并转换为类别索引
    y_pred = global_model.predict(test_x, verbose=0)
    y_pred_classes = np.argmax(y_pred, axis=1)
    y_true_classes = np.argmax(test_y, axis=1)
    
    # 计算并打印混淆矩阵
    cm = confusion_matrix(y_true_classes, y_pred_classes)
    print(f"Comm Round {round_num} - Confusion Matrix:\n{cm}")
    
    # 获取compile时配置的所有指标结果
    results = global_model.evaluate(test_x, test_y, verbose=0)
    global_loss = results[0]
    global_acc = results[1]
    global_precision = results[2]
    global_recall = results[3]
    global_f1 = results[4]
    
    print(f"Comm Round {round_num} - Loss: {global_loss:.4f}, Acc: {global_acc:.4f}, Precision: {global_precision:.4f}, Recall: {global_recall:.4f}, F1: {global_f1:.4f}")
    
    return global_acc, global_loss, global_precision, global_recall, global_f1

2. 记录训练与测试的历史指标

在联邦学习循环外部初始化全局列表,用于存储每一轮的训练/测试数据:

# 初始化历史记录容器
train_loss_history = []
train_acc_history = []
test_loss_history = []
test_acc_history = []
test_precision_history = []
test_recall_history = []
test_f1_history = []

修改客户端训练循环,记录每个客户端的训练指标(若需要全局训练均值,可收集每轮所有客户端数据后取平均):

# 在客户端训练循环内,训练完成后记录指标
history = local_model.fit(clients_batched[client], epochs=1, verbose=0, callbacks=[checkpoint_callback])
train_loss_history.append(history.history['loss'][0])
train_acc_history.append(history.history['acc'][0])

在全局测试阶段更新测试指标记录:

# 替换原有的测试循环逻辑
global_acc, global_loss, global_precision, global_recall, global_f1 = test_model(test_x, test_y, global_model, comm_round + 1)
test_loss_history.append(global_loss)
test_acc_history.append(global_acc)
test_precision_history.append(global_precision)
test_recall_history.append(global_recall)
test_f1_history.append(global_f1)

3. 绘制训练/测试指标折线图

使用matplotlib绘制损失、准确率等指标的变化曲线:

import matplotlib.pyplot as plt

# 绘制损失与准确率对比图
plt.figure(figsize=(12, 5))
plt.subplot(1, 2, 1)
plt.plot(train_loss_history, label='Train Loss')
plt.plot(test_loss_history, label='Test Loss')
plt.title('Training vs Test Loss')
plt.xlabel('Communication Rounds')
plt.ylabel('Loss')
plt.legend()

plt.subplot(1, 2, 2)
plt.plot(train_acc_history, label='Train Accuracy')
plt.plot(test_acc_history, label='Test Accuracy')
plt.title('Training vs Test Accuracy')
plt.xlabel('Communication Rounds')
plt.ylabel('Accuracy')
plt.legend()

plt.tight_layout()
plt.show()

# 绘制Precision、Recall、F1分数变化图
plt.figure(figsize=(12, 8))
plt.subplot(2, 2, 1)
plt.plot(test_precision_history, label='Test Precision')
plt.title('Test Precision Over Rounds')
plt.xlabel('Communication Rounds')
plt.ylabel('Precision')
plt.legend()

plt.subplot(2, 2, 2)
plt.plot(test_recall_history, label='Test Recall')
plt.title('Test Recall Over Rounds')
plt.xlabel('Communication Rounds')
plt.ylabel('Recall')
plt.legend()

plt.subplot(2, 2, 3)
plt.plot(test_f1_history, label='Test F1 Score')
plt.title('Test F1 Score Over Rounds')
plt.xlabel('Communication Rounds')
plt.ylabel('F1 Score')
plt.legend()

plt.tight_layout()
plt.show()

完整代码整合示例

将上述修改整合到你的原有代码中,最终代码如下:

import tensorflow as tf
from tensorflow.keras.metrics import CategoricalAccuracy, Precision, Recall, F1Score
from sklearn.metrics import confusion_matrix
import numpy as np
import matplotlib.pyplot as plt
from tensorflow.keras import backend as K
import random

# 初始化历史记录列表
train_loss_history = []
train_acc_history = []
test_loss_history = []
test_acc_history = []
test_precision_history = []
test_recall_history = []
test_f1_history = []

def weight_scalling_factor(clients_batched, client):
    # 你的原有实现
    pass

def scale_model_weights(weights, scalar):
    # 你的原有实现
    pass

def sum_scaled_weights(scaled_weights):
    # 你的原有实现
    pass

def test_model(test_x, test_y, global_model, round_num):
    y_pred = global_model.predict(test_x, verbose=0)
    y_pred_classes = np.argmax(y_pred, axis=1)
    y_true_classes = np.argmax(test_y, axis=1)
    
    cm = confusion_matrix(y_true_classes, y_pred_classes)
    print(f"Comm Round {round_num} - Confusion Matrix:\n{cm}")
    
    results = global_model.evaluate(test_x, test_y, verbose=0)
    global_loss = results[0]
    global_acc = results[1]
    global_precision = results[2]
    global_recall = results[3]
    global_f1 = results[4]
    
    print(f"Comm Round {round_num} - Loss: {global_loss:.4f}, Acc: {global_acc:.4f}, Precision: {global_precision:.4f}, Recall: {global_recall:.4f}, F1: {global_f1:.4f}")
    
    return global_acc, global_loss, global_precision, global_recall, global_f1

# 假设你的Transformer模型、客户端数据、测试数据已定义
# global_model = Transformer(...)
# clients_batched = {...}
# test_x, test_y = ...
# total_comm_rounds = 10  # 替换为你的通信轮数

for comm_round in range(total_comm_rounds):
    global_weights = global_model.get_weights()
    scaled_local_weight_list = list()
    client_names = list(clients_batched.keys())
    random.shuffle(client_names)

    # 初始化多分类指标,替换为你的实际类别数量
    metrics = [
        CategoricalAccuracy(name='acc'),
        Precision(name='precision', average='macro'),
        Recall(name='recall', average='macro'),
        F1Score(name='f1_score', average='macro', num_classes=10)
    ]

    for client in client_names:
        local_model = Transformer
        local_model.compile(
            loss=tf.keras.losses.CategoricalCrossentropy(),
            optimizer=tf.keras.optimizers.Adam(learning_rate=0.001),
            metrics=metrics
        )

        global_model.set_weights(global_weights)
        local_model.set_weights(global_weights)

        history = local_model.fit(clients_batched[client], epochs=1, verbose=0, callbacks=[checkpoint_callback])

        # 记录客户端训练指标
        train_loss_history.append(history.history['loss'][0])
        train_acc_history.append(history.history['acc'][0])

        scaling_factor = weight_scalling_factor(clients_batched, client)
        scaled_weights = scale_model_weights(local_model.get_weights(), scaling_factor)
        scaled_local_weight_list.append(scaled_weights)

        K.clear_session()

    average_weights = sum_scaled_weights(scaled_local_weight_list)
    global_model.set_weights(average_weights)

    # 全局测试并记录指标
    global_acc, global_loss, global_precision, global_recall, global_f1 = test_model(test_x, test_y, global_model, comm_round + 1)
    test_loss_history.append(global_loss)
    test_acc_history.append(global_acc)
    test_precision_history.append(global_precision)
    test_recall_history.append(global_recall)
    test_f1_history.append(global_f1)

# 绘制曲线
plt.figure(figsize=(12, 5))
plt.subplot(1, 2, 1)
plt.plot(train_loss_history, label='Train Loss')
plt.plot(test_loss_history, label='Test Loss')
plt.title('Training vs Test Loss')
plt.xlabel('Communication Rounds')
plt.ylabel('Loss')
plt.legend()

plt.subplot(1, 2, 2)
plt.plot(train_acc_history, label='Train Accuracy')
plt.plot(test_acc_history, label='Test Accuracy')
plt.title('Training vs Test Accuracy')
plt.xlabel('Communication Rounds')
plt.ylabel('Accuracy')
plt.legend()

plt.tight_layout()
plt.show()

# 绘制Precision/Recall/F1曲线
plt.figure(figsize=(12, 8))
plt.subplot(2, 2, 1)
plt.plot(test_precision_history, label='Test Precision')
plt.title('Test Precision Over Rounds')
plt.xlabel('Communication Rounds')
plt.ylabel('Precision')
plt.legend()

plt.subplot(2, 2, 2)
plt.plot(test_recall_history, label='Test Recall')
plt.title('Test Recall Over Rounds')
plt.xlabel('Communication Rounds')
plt.ylabel('Recall')
plt.legend()

plt.subplot(2, 2, 3)
plt.plot(test_f1_history, label='Test F1 Score')
plt.title('Test F1 Score Over Rounds')
plt.xlabel('Communication Rounds')
plt.ylabel('F1 Score')
plt.legend()

plt.tight_layout()
plt.show()

注意事项

  • 替换num_classes=10为你的实际分类类别数量;若为二分类任务,将average='macro'改为average='binary',并设置num_classes=2。
  • 若需要更准确的全局训练指标,可收集每轮所有客户端的训练数据后取均值,而非直接追加单个客户端数据:
# 在客户端循环外初始化临时存储
round_train_loss = []
round_train_acc = []

for client in client_names:
    # ...训练代码...
    round_train_loss.append(history.history['loss'][0])
    round_train_acc.append(history.history['acc'][0])

# 记录每轮的平均训练指标
train_loss_history.append(np.mean(round_train_loss))
train_acc_history.append(np.mean(round_train_acc))

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 17:57:01