PyTorch二分类模型训练和评估:ROC曲线和混淆矩阵
import pandas as pd
import torch
import torch.nn as nn
from sklearn.preprocessing import StandardScaler
import matplotlib.pyplot as plt
from sklearn.metrics import roc_curve, roc_auc_score, confusion_matrix
import numpy as np
# 读取数据
data = pd.read_excel('C:\Users\lenovo\Desktop\数据测试\output_data1.xlsx')
# 标准化处理
scaler = StandardScaler()
X = scaler.fit_transform(data.iloc[:, 1:].values) # 特征矩阵
y = data.iloc[:, 0].values # 标签向量
# 将numpy数组转为PyTorch张量
X = torch.tensor(X, dtype=torch.float32)
y = torch.tensor(y, dtype=torch.float32)
# 获取特征数量
num_features = X.shape[1]
# 定义模型
class Net(nn.Module):
def __init__(self):
super(Net, self).__init__()
self.fc1 = nn.Linear(num_features , 64)
self.fc2 = nn.Linear(64, 32)
self.fc3 = nn.Linear(32, 16)
self.fc4 = nn.Linear(16, 1)
self.relu = nn.ReLU()
self.sigmoid = nn.Sigmoid()
def forward(self, x):
x = self.relu(self.fc1(x))
x = self.relu(self.fc2(x))
x = self.relu(self.fc3(x))
x = self.sigmoid(self.fc4(x))
return x
# 训练模型
net = Net()
criterion = nn.BCELoss()
optimizer = torch.optim.Adam(net.parameters(), lr=0.001, weight_decay=0.001)
# 存储loss和accuracy
losses = []
accuracies = []
epochs=1000
for epoch in range(epochs):
optimizer.zero_grad()
outputs = net(X)
loss = criterion(outputs, y.view(-1, 1))
loss.backward()
optimizer.step()
# 计算训练准确率
with torch.no_grad():
predicted = net(X)
predicted = predicted.round()
accuracy = (predicted == y.view(-1, 1)).sum().item() / len(y)
# 保存loss和accuracy
losses.append(loss.item())
accuracies.append(accuracy)
if epoch % 1 == 0:
print('Epoch {}, Loss: {:.4f}, Accuracy: {:.2f}%'.format(epoch, loss.item(), accuracy * 100))
# 绘制loss和accuracy曲线
plt.plot(losses, label='Loss')
plt.plot(accuracies, label='Accuracy')
plt.legend()
plt.xlabel('Epoch')
plt.ylabel('Value')
plt.show()
# 计算预测概率和真实标签
with torch.no_grad():
predicted = net(X)
predicted_prob = predicted.numpy()
true_labels = y.numpy()
# 计算ROC曲线的假正率和真正率
fpr, tpr, thresholds = roc_curve(true_labels, predicted_prob)
# 计算AUC值
auc = roc_auc_score(true_labels, predicted_prob)
# 绘制ROC曲线
plt.plot(fpr, tpr, label='ROC curve (area = %0.2f)' % auc)
plt.plot([0, 1], [0, 1], 'k--')
plt.xlim([0.0, 1.0])
plt.ylim([0.0, 1.05])
plt.xlabel('False Positive Rate')
plt.ylabel('True Positive Rate')
plt.title('Receiver operating characteristic')
plt.legend(loc='lower right')
plt.show()
# 计算混淆矩阵
threshold = 0.5
predicted_binary = (predicted >= threshold).numpy().astype(int)
confusion = confusion_matrix(true_labels, predicted_binary)
print('Confusion Matrix:
', confusion)
原文地址: https://www.cveoy.top/t/topic/nKIR 著作权归作者所有。请勿转载和采集!