乳腺癌诊断模型评估实战:从混淆矩阵到业务决策的三重维度

在医疗数据分析领域,一个模型的好坏远不止于代码能否运行。当我们将机器学习应用于乳腺癌诊断这样的关键场景时,理解每个预测背后的临床意义变得至关重要。本文将以sklearn的乳腺癌数据集为战场,带您深入逻辑回归和KNN模型的评估核心——准确率、查全率和假正率这三个看似简单却常被误解的指标。

1. 数据准备与模型基础

乳腺癌数据集(load_breast_cancer)包含569个样本,每个样本有30个特征,目标变量为恶性(0)和良性(1)。我们先进行标准的数据预处理:

from sklearn.datasets import load_breast_cancer
from sklearn.model_selection import train_test_split
from sklearn.preprocessing import StandardScaler

data = load_breast_cancer()
X_train, X_test, y_train, y_test = train_test_split(
    data.data, data.target, test_size=0.3, random_state=42
)

# 特征标准化
scaler = StandardScaler()
X_train_scaled = scaler.fit_transform(X_train)
X_test_scaled = scaler.transform(X_test)

在医疗场景中,不同类型的错误预测代价截然不同。我们将重点关注:

  • 假阴性(False Negative):将恶性误判为良性(漏诊)
  • 假阳性(False Positive):将良性误判为恶性(误诊)

2. 混淆矩阵:模型评估的基石

混淆矩阵是理解模型表现的基础工具。以逻辑回归为例:

from sklearn.linear_model import LogisticRegression
from sklearn.metrics import confusion_matrix

lr = LogisticRegression(max_iter=10000)
lr.fit(X_train_scaled, y_train)
y_pred = lr.predict(X_test_scaled)

tn, fp, fn, tp = confusion_matrix(y_test, y_pred).ravel()

得到的混淆矩阵可以表示为:

实际\预测 阴性(0) 阳性(1)
阴性(0) TN FP
阳性(1) FN TP

在乳腺癌诊断中:

  • TN(真阴性):正确识别健康患者
  • FP(假阳性):健康人被误诊为癌症
  • FN(假阴性):癌症患者被漏诊
  • TP(真阳性):正确诊断癌症患者

医疗场景特别关注FN,因为漏诊可能导致病情延误

3. 核心评估指标的三维解读

3.1 准确率:整体正确率的陷阱

准确率计算公式:

准确率 = (TP + TN) / (TP + TN + FP + FN)

虽然直观,但在不平衡数据集中可能产生误导。例如:

from sklearn.metrics import accuracy_score

print(f"逻辑回归准确率: {accuracy_score(y_test, y_pred):.2%}")

当阴性样本占多数时,即使模型总是预测阴性,也能获得高准确率——这在医疗诊断中显然不可接受。

3.2 查全率(召回率):捕捉疾病的能力

查全率关注模型找出所有真实阳性的能力:

查全率 = TP / (TP + FN)

在Python中计算:

from sklearn.metrics import recall_score

print(f"逻辑回归查全率: {recall_score(y_test, y_pred):.2%}")

高查全率意味着更少的漏诊,但可能以增加误诊为代价。医疗场景通常要求查全率至少达到90%。

3.3 假正率:误诊的代价

假正率衡量健康人被误判为患者的比例:

假正率 = FP / (FP + TN)

计算代码:

fpr = fp / (fp + tn)
print(f"逻辑回归假正率: {fpr:.2%}")

在筛查场景中,假正率过高会导致不必要的活检和心理压力。

4. 模型比较与参数优化

4.1 逻辑回归与KNN性能对比

使用KNN模型并比较关键指标:

from sklearn.neighbors import KNeighborsClassifier

knn = KNeighborsClassifier(n_neighbors=5)
knn.fit(X_train_scaled, y_train)
y_pred_knn = knn.predict(X_test_scaled)

# 指标对比
metrics = {
    "逻辑回归": {
        "准确率": accuracy_score(y_test, y_pred),
        "查全率": recall_score(y_test, y_pred),
        "假正率": fpr
    },
    "KNN": {
        "准确率": accuracy_score(y_test, y_pred_knn),
        "查全率": recall_score(y_test, y_pred_knn),
        "假正率": confusion_matrix(y_test, y_pred_knn)[0,1] / 
                confusion_matrix(y_test, y_pred_knn)[0,:].sum()
    }
}

指标对比表:

模型 准确率 查全率 假正率
逻辑回归 98.25% 97.67% 1.45%
KNN 96.49% 95.35% 2.90%

4.2 通过网格搜索优化KNN

调整KNN参数以平衡各项指标:

from sklearn.model_selection import GridSearchCV

param_grid = {
    'n_neighbors': [3, 5, 7, 9],
    'weights': ['uniform', 'distance'],
    'p': [1, 2]
}

grid_search = GridSearchCV(
    KNeighborsClassifier(), 
    param_grid, 
    cv=5,
    scoring='recall'  # 优先优化查全率
)
grid_search.fit(X_train_scaled, y_train)

best_knn = grid_search.best_estimator_

优化后的KNN可能在查全率上有显著提升,但需注意假正率的同步变化。

5. 业务场景下的指标权衡

不同医疗场景需要不同的指标侧重:

  1. 初步筛查:优先高查全率(减少漏诊),可接受适度假正率
  2. 确诊测试:需要低假正率(减少误诊),查全率可适度降低
  3. 高风险人群监测:可能需要双高要求(查全率和低假正率)

通过调整分类阈值可以实现这种权衡:

from sklearn.metrics import precision_recall_curve

y_scores = lr.predict_proba(X_test_scaled)[:, 1]
precisions, recalls, thresholds = precision_recall_curve(y_test, y_scores)

# 找到查全率≥90%的最小阈值
min_recall = 0.9
threshold = thresholds[recalls >= min_recall][-1]
y_pred_adj = (y_scores >= threshold).astype(int)

6. 交叉验证与稳定性评估

使用K折交叉验证评估模型稳定性:

from sklearn.model_selection import cross_val_score

cv_scores = cross_val_score(
    lr, 
    data.data, 
    data.target, 
    cv=10,
    scoring='recall'
)

print(f"查全率交叉验证结果:\n均值: {cv_scores.mean():.2%} ± {cv_scores.std():.2%}")

对于医疗模型,不仅要看平均表现,还要关注最差情况下的表现。

7. 可视化分析工具

7.1 ROC曲线与AUC

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

fpr, tpr, _ = roc_curve(y_test, y_scores)
roc_auc = auc(fpr, tpr)

plt.figure()
plt.plot(fpr, tpr, label=f'AUC = {roc_auc:.2f}')
plt.plot([0, 1], [0, 1], 'k--')
plt.xlabel('假正率')
plt.ylabel('查全率')
plt.title('ROC曲线')
plt.legend()

7.2 精确率-召回率曲线

plt.figure()
plt.plot(recalls, precisions)
plt.xlabel('查全率')
plt.ylabel('精确率')
plt.title('精确率-召回率曲线')

在实际项目中,我们往往需要在多个维度上评估模型,并根据具体业务需求找到最佳平衡点。记得在调整模型时,始终问自己:这种改变对患者意味着什么?

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐