PyTorch 训练代码:带验证功能的模型训练与早期停止
for epoch in range(args.epochs):
t = time.time()
# for train
model.train()
optimizer.zero_grad()
output = model(features, adjtensor)
# 平均输出
areout = output[1]
loss_xy = 0
loss_ncl = 0
for k in range(len(output[0])):
# print('k = ' + str(k))
# print(F.nll_loss(output[0][k][idx_train], labels[idx_train]))
# print(F.mse_loss(output[0][k][idx_unlabel], areout[idx_unlabel]))
loss_xy += F.nll_loss(output[0][k][idx_train], labels[idx_train])
loss_ncl += F.mse_loss(output[0][k][idx_unlabel], areout[idx_unlabel])
loss_train = (1-args.lamd)* loss_xy - args.lamd * loss_ncl
# loss_train = (1 - args.lamd) * loss_xy + args.lamd * 1 / loss_ncl
# loss_train = (1 - args.lamd) * loss_xy + args.lamd * (torch.exp(-loss_ncl))
print(loss_xy)
print(loss_ncl)
print(torch.exp(-loss_ncl))
print((1 - args.lamd) * loss_xy)
print(args.lamd * (torch.exp(-loss_ncl)))
print(epoch)
print(loss_train)
print('.............')
acc_train = accuracy(areout[idx_train], labels[idx_train])
loss_train.backward()
optimizer.step()
# for val
if validate:
# print('no')
model.eval()
output = model(features, adjtensor)
areout = output[1]
vl_step = len(idx_val)
loss_val = F.nll_loss(areout[idx_val], labels[idx_val])
acc_val = accuracy(areout[idx_val], labels[idx_val])
# vl_step = len(idx_train)
# loss_val = F.nll_loss(areout[idx_train], labels[idx_train])
# acc_val = accuracy(areout[idx_train], labels[idx_train])
cost_val.append(loss_val)
# 原始GCN的验证
# if epoch > args.early_stopping and cost_val[-1] > torch.mean(torch.stack(cost_val[-(args.early_stopping + 1):-1])):
# # print('Early stopping...')
# print(epoch)
# break
# print(epoch)
# GAT的验证
if acc_val/vl_step >= vacc_mx or loss_val/vl_step <= vlss_mn:
if acc_val/vl_step >= vacc_mx and loss_val/vl_step <= vlss_mn:
vacc_early_model = acc_val/vl_step
vlss_early_model = loss_val/vl_step
torch.save(model, checkpt_file)
vacc_mx = np.max((vacc_early_model, vacc_mx))
vlss_mn = np.min((vlss_early_model, vlss_mn))
curr_step = 0
else:
curr_step += 1
# print(curr_step)
if curr_step == args.early_stopping:
# print('Early stop! Min loss: ', vlss_mn, ', Max accuracy: ', vacc_mx)
# print('Early stop model validation loss: ', vlss_early_model, ', accuracy: ', vacc_early_model)
break
代码解释:
- 循环遍历
epochs:代码使用for epoch in range(args.epochs)循环来遍历每个 epoch。 - 计算训练损失和准确率:代码计算了训练集上的损失和准确率,并打印输出以监控训练过程。
- 模型验证:代码通过设置
validate参数为True来开启模型验证功能。在每个 epoch 结束时,模型会在验证集上进行推断,计算验证集上的损失和准确率。 - 模型保存和早期停止:代码会根据验证集的性能进行模型保存或提前结束训练。
- 如果验证集上的准确率达到了当前最高准确率,或者验证集上的损失达到了当前最低损失,就会保存当前模型,并更新最高准确率和最低损失的值。
- 如果连续
args.early_stopping个 epoch 验证集上的准确率都没有超过当前最高准确率,或者验证集上的损失连续args.early_stopping个 epoch 都没有下降,就会提前结束训练过程。
早期停止的作用:
- 防止模型过拟合:在训练过程中,模型可能会在训练集上表现良好,但在测试集上表现不佳,这就是过拟合。早期停止可以帮助防止模型过拟合,从而提高模型的泛化能力。
代码示例:
import torch
import numpy as np
# 假设定义了模型、优化器、数据加载器、验证集等
# ...
vacc_mx = 0
vlss_mn = float('inf')
curr_step = 0
cost_val = []
checkpt_file = 'best_model.pth'
for epoch in range(args.epochs):
# ... 训练过程
# 验证过程
if validate:
# ... 验证代码
总结:
这段代码演示了使用 PyTorch 训练模型并包含验证功能以实现早期停止。通过验证集上的性能评估,可以有效地防止模型过拟合,并提高模型的泛化能力。
原文地址: https://www.cveoy.top/t/topic/ihBe 著作权归作者所有。请勿转载和采集!