欢迎来到尧图网

客户服务 关于我们

您的位置:首页 > 科技 > 名人名企 > PyTorch计算机视觉入门:测试模型与评估,对单帧图片进行推理

PyTorch计算机视觉入门:测试模型与评估,对单帧图片进行推理

2024/10/25 1:26:46 来源:https://blog.csdn.net/laukal/article/details/139693762  浏览:    关键词:PyTorch计算机视觉入门:测试模型与评估,对单帧图片进行推理

在完成模型的训练之后,对模型进行测试与评估是至关重要的一步,它能帮助我们理解模型在未知数据上的泛化能力。本篇指南将带您了解如何使用PyTorch进行模型测试,并对测试结果进行分析。我们将基于之前训练好的模型,演示如何加载数据、进行预测、计算指标以及可视化结果。

准备工作

假设您已经有一个训练好的模型,保存在.pth文件中,以及一个用于测试的自定义数据集。我们将继续使用前文提到的自定义数据集CustomDataset类,并引入一些新的概念和代码。

加载测试数据集

与训练过程类似,首先需要加载测试数据集,并对其进行适当的预处理。确保您的测试集遵循与训练集相同的数据结构和预处理步骤。

test_dataset = CustomImageDataset(data_path="./data/", model= "test", transform = transform)
test_loader = DataLoader(test_dataset, batch_size=64, shuffle=True)

测试模型

训练完成后,使用测试数据集来评估模型的性能

def test(model, device, test_loader):model.eval()test_loss = 0correct = 0with torch.no_grad():for data, target in test_loader:data, target = data.to(device), target.to(device)output = model(data)test_loss += criterion(output, target)# max 函数返回两个值,一个是是数值,一个是indexpred = output.max(1, keepdim=True)[1] # 找到概率最大的下标 correct += pred.eq(target.view_as(pred)).sum().item()test_loss /= len(test_loader.dataset)print('\nTest set: Average loss: {:.4f}, Accuracy: {}/{} ({:.0f}%)\n'.format(test_loss, correct, len(test_loader.dataset),100. * correct / len(test_loader.dataset)))
Test set: Average loss: 0.0018, Accuracy: 9671/10000 (97%)

单帧图片进行测试

# test single image 
img = Image.open("./data/data_test/1.jpg")
img_t = transform(img)
img_t = img_t.unsqueeze(0)  # 变为[1, 1, 28, 28]
img_t = img_t.to(device)
model.eval()
output = model(img_t)
_, predicted_class = torch.max(output, 1)
print(predicted_class)
tensor([2], device='cuda:0')

通过以上步骤,我们可以全面地评估和分析PyTorch模型在计算机视觉任务中的表现,从而确保模型在实际应用中的有效性和可靠性。

关注我的公众号Ai fighting, 第一时间获取更新内容。

版权声明:

本网仅为发布的内容提供存储空间,不对发表、转载的内容提供任何形式的保证。凡本网注明“来源:XXX网络”的作品,均转载自其它媒体,著作权归作者所有,商业转载请联系作者获得授权,非商业转载请注明出处。

我们尊重并感谢每一位作者,均已注明文章来源和作者。如因作品内容、版权或其它问题,请及时与我们联系,联系邮箱:809451989@qq.com,投稿邮箱:809451989@qq.com