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

在TensorFlow中计算召回率、精确率、F1指标的方法与可视化方案

预训练句子编码器嵌入的评估与可视化方案

首先明确:召回率、精确率、F1均为有监督指标,必须依赖标注数据(类别标签或文本配对的相似性标注),以下分两种最常见的应用场景给出具体步骤。


场景一:文本分类任务(每条文本对应类别标签)

步骤1:准备标注数据

给2000条文本配上对应的类别标签,比如二分类场景下的labels = [0, 1, 0, 1, ...],多分类标签同理。

步骤2:基于嵌入向量训练分类器

嵌入向量是文本的特征表示,需要搭配分类模型完成预测。这里以sklearn的逻辑回归为例:

import numpy as np
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import train_test_split

# 将Tensor格式的嵌入转为numpy数组(关键步骤)
embeddings_np = np.array([emb.numpy().flatten() for emb in my_embeddings])

# 划分训练集与测试集
X_train, X_test, y_train, y_test = train_test_split(embeddings_np, labels, test_size=0.2, random_state=42)

# 训练分类器
clf = LogisticRegression(max_iter=1000)
clf.fit(X_train, y_train)

# 获取测试集预测结果
y_pred = clf.predict(X_test)
y_pred_proba = clf.predict_proba(X_test)[:, 1]  # 二分类场景下的正类概率,用于可视化

步骤3:计算精确率、召回率、F1

使用sklearn的metrics模块直接计算:

from sklearn.metrics import precision_score, recall_score, f1_score, classification_report

precision = precision_score(y_test, y_pred)
recall = recall_score(y_test, y_pred)
f1 = f1_score(y_test, y_pred)

print(f"精确率: {precision:.4f}")
print(f"召回率: {recall:.4f}")
print(f"F1值: {f1:.4f}")
print("\n分类报告:\n", classification_report(y_test, y_pred))

多分类场景下,需给precision_score等函数指定average参数,比如average='macro'(宏平均)或average='weighted'(加权平均)。

步骤4:可视化

混淆矩阵(直观展示分类对错分布)

from sklearn.metrics import confusion_matrix
import seaborn as sns
import matplotlib.pyplot as plt

cm = confusion_matrix(y_test, y_pred)
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',
            xticklabels=['类别0', '类别1'], yticklabels=['类别0', '类别1'])
plt.xlabel('预测标签')
plt.ylabel('真实标签')
plt.title('混淆矩阵')
plt.show()

PR曲线(精确率-召回率曲线,适合不平衡数据集)

from sklearn.metrics import precision_recall_curve

precision_curve, recall_curve, _ = precision_recall_curve(y_test, y_pred_proba)
plt.plot(recall_curve, precision_curve, linewidth=2)
plt.xlabel('召回率')
plt.ylabel('精确率')
plt.title('PR曲线')
plt.grid(True)
plt.show()

场景二:语义检索任务(判断文本对是否相似)

步骤1:准备配对标注数据

构建文本对数据集,格式为pairs = [(文本A, 文本B, 标签), ...],其中标签1表示文本A与B相似,0表示不相似。可从2000条文本中生成正负样本对(比如随机配对生成负样本,人工标注或基于业务规则生成正样本)。

步骤2:计算文本对的相似度分数

用嵌入向量的余弦相似度作为匹配分数:

from sklearn.metrics.pairwise import cosine_similarity
import numpy as np

# 先将所有嵌入转为numpy数组
embeddings_np = np.array([emb.numpy().flatten() for emb in my_embeddings])
# 建立文本到嵌入的映射字典
text_to_emb = dict(zip(my_texts, embeddings_np))

# 计算每个文本对的相似度
scores = []
true_labels = []
for text_a, text_b, label in pairs:
    emb_a = text_to_emb[text_a].reshape(1, -1)
    emb_b = text_to_emb[text_b].reshape(1, -1)
    sim_score = cosine_similarity(emb_a, emb_b)[0][0]
    scores.append(sim_score)
    true_labels.append(label)

步骤3:确定阈值并计算指标

相似度是连续值,需设定阈值将其转为预测标签(比如阈值设为0.5,分数≥0.5则预测为相似):

from sklearn.metrics import precision_score, recall_score, f1_score

threshold = 0.5
y_pred = [1 if s >= threshold else 0 for s in scores]

precision = precision_score(true_labels, y_pred)
recall = recall_score(true_labels, y_pred)
f1 = f1_score(true_labels, y_pred)

print(f"精确率: {precision:.4f}")
print(f"召回率: {recall:.4f}")
print(f"F1值: {f1:.4f}")

可根据业务需求调整阈值:想要高召回率(尽量不漏掉相似样本)就降低阈值;想要高精确率(尽量减少误判)就提高阈值。

步骤4:可视化

ROC曲线(展示不同阈值下的召回率与假阳性率)

from sklearn.metrics import roc_curve, auc
import matplotlib.pyplot as plt

fpr, tpr, thresholds = roc_curve(true_labels, scores)
roc_auc = auc(fpr, tpr)

plt.plot(fpr, tpr, linewidth=2, label=f'AUC = {roc_auc:.4f}')
plt.plot([0, 1], [0, 1], 'k--')  # 随机猜测基准线
plt.xlabel('假阳性率')
plt.ylabel('召回率(真阳性率)')
plt.title('ROC曲线')
plt.legend()
plt.grid(True)
plt.show()

相似度分布直方图(对比正负样本的相似度差异)

import matplotlib.pyplot as plt

# 拆分正负样本的相似度分数
pos_scores = [s for s, l in zip(scores, true_labels) if l == 1]
neg_scores = [s for s, l in zip(scores, true_labels) if l == 0]

plt.hist(pos_scores, bins=20, alpha=0.5, label='相似样本')
plt.hist(neg_scores, bins=20, alpha=0.5, label='不相似样本')
plt.xlabel('余弦相似度')
plt.ylabel('样本数量')
plt.title('正负样本相似度分布')
plt.legend()
plt.show()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 01:37:22