CNN猫狗大战

好买网 www.goodmai.com IT技术交易平台 使用VGG模型进行猫狗大战 大赛简介 Kaggle 中的猫狗大战竞赛题目。在这个比赛中,有25000张标记好的猫和狗的图片用做训练,有12500张图片用做测试。这个竞赛是2013年开展的,如果你能够达到80%的准确率,在当年是一个 state-of-the-art 的成绩。

数据准备 在这里其实出了问题,由于研习社的题目给的是rar格式的压缩包,所以没办法和zip一样解压,我开始直接改成

1!wget https://static.leiphone.com/cat_dog.rar 2!unzip cat_dog.rar

显然是不行的,报错结果如下: image 可以看到需要加入< Comands >中的x,然后需要加入目录地址/content/cat_dog.rar,为什么是这个地址,请看下面一张图: image 然后就能愉快的下载解压数据了:

1!wget https://static.leiphone.com/cat_dog.rar 2!unrar x /content/cat_dog.rar 3 4!wget http://fenggao-image.stor.sinaapp.com/dogscats.zip 5!unzip dogscats.zip

解压完成后cat_dog数据集目录: 2. test:最后通过训练好的模型来识别的测试图片。 train:用来训练模型的图片。 val:文件夹下的图片有确定的标签,用来测试模型训练效果。 3. 代码实战

1import numpy as np 2import matplotlib.pyplot as plt 3import os 4import torch 5import torch.nn as nn 6import torchvision 7from torchvision import models,transforms,datasets 8import time 9import json 10 11# 判断是否存在GPU设备 12device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") 13print('Using gpu: %s ' % torch.cuda.is_available())

数据预处理

1normalize = transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) 2 3vgg_format = transforms.Compose([ 4 transforms.CenterCrop(224), 5 transforms.ToTensor(), 6 normalize, 7 ]) 8 9data_dir = '/content/dogscats' 10 11dsets = {x: datasets.ImageFolder(os.path.join(data_dir, x), vgg_format) 12 for x in ['train', 'valid']} 13 14dset_sizes = {x: len(dsets[x]) for x in ['train', 'valid']} 15dset_classes = dsets['train'].classes
1# 通过下面代码可以查看 dsets 的一些属性 2print(dsets['train'].classes) 3print(dsets['train'].class_to_idx) 4print(dsets['train'].imgs[:5]) 5print('dset_sizes: ', dset_sizes) 6 7loader_train = torch.utils.data.DataLoader( 8dsets['train'], batch_size=64, shuffle=True, num_workers=6) 9 10loader_valid = torch.utils.data.DataLoader( 11dsets['valid'], batch_size=5, shuffle=False, num_workers=6) 12 13''' 14valid 数据一共有2000张图,每个batch是5张,因此,下面进行遍历一共会输出到 400 15同时,把第一个 batch 保存到 inputs_try, labels_try,分别查看 16''' 17count = 1 18for data in loader_valid: 19 print(count, end='\n') 20 if count == 1: 21 inputs_try,labels_try = data 22 count +=1 23 24print(labels_try) 25print(inputs_try.shape)

打印预览数据

1# 显示图片的小程序 2def imshow(inp, title=None): 3# Imshow for Tensor. 4 inp = inp.numpy().transpose((1, 2, 0)) 5 mean = np.array([0.485, 0.456, 0.406]) 6 std = np.array([0.229, 0.224, 0.225]) 7 inp = np.clip(std * inp + mean, 0,1) 8 plt.imshow(inp) 9 if title is not None: 10 plt.title(title) 11 plt.pause(0.001) # pause a bit so that plots are updated 12 13# 显示 labels_try 的5张图片,即valid里第一个batch的5张图片 14out = torchvision.utils.make_grid(inputs_try) 15imshow(out, title=[dset_classes[x] for x in labels_try])

image 构建VGG 我们直接使用预训练好的 VGG 模型。

1model_vgg = models.vgg16(pretrained=True) 2print(model_vgg)

我们的目标是使用预训练好的模型,并且只对全连接层的最后一层进行改写nn.Linear(4096, 2)使得最后输出的结果只有两个,即分辨猫与狗。

为了在训练中冻结前面层的参数,需要设置 required_grad=False。这样,前面层的权重就不会自动更新了。训练的时候只会更新最后一层的参数。

1model_vgg_new = model_vgg; 2 3for param in model_vgg_new.parameters(): 4 param.requires_grad = False #冻结参数 5''' 6更改最后一层输出层 7''' 8model_vgg_new.classifier._modules['6'] = nn.Linear(4096, 2) 9model_vgg_new.classifier._modules['7'] = torch.nn.LogSoftmax(dim = 1) 10 11model_vgg_new = model_vgg_new.to(device) 12 13''' 14输出新的vgg模型 15''' 16print(model_vgg_new.classifier)

image 使用Adam优化器对模型进行优化

1''' 2第一步:创建损失函数和优化器 3 4损失函数 NLLLoss() 的 输入 是一个对数概率向量和一个目标标签. 5它不会为我们计算对数概率,适合最后一层是log_softmax()的网络. 6''' 7criterion = nn.NLLLoss() 8 9# 学习率 10lr = 0.001
1# 这里使用Adam优化器 2optimizer_vgg = torch.optim.Adam(model_vgg_new.classifier[6].parameters(),lr = lr) 3 4''' 5第二步:训练模型并保存 6model: 训练的模型 7dataloader: 训练集 8size: 训练集大小 9epochs: 训练次数 10optimizer: 优化器 11''' 12def train_model(model,dataloader,size,epochs=1,optimizer=None): 13 model.train() #用于模型训练 14 15 for epoch in range(epochs): 16 epoch_acc_max = 0 17 running_loss = 0.0 18 running_corrects = 0 19 count = 0 20 21 for inputs,classes in dataloader: 22 inputs = inputs.to(device) 23 classes = classes.to(device) 24 25 outputs = model(inputs) #参数前向传播 26 27 loss = criterion(outputs,classes) 28 optimizer = optimizer 29 optimizer.zero_grad() #优化器梯度初始化 30 loss.backward() #梯度反向传播 31 optimizer.step() 32 _,preds = torch.max(outputs.data,1) #得到预测结果 33 # statistics 34 running_loss += loss.data.item() 35 running_corrects += torch.sum(preds == classes.data) 36 37 count += len(inputs) 38 print('Training: No. ', count, ' process ... total: ', size) 39 40 epoch_loss = running_loss / size 41 epoch_acc = running_corrects.data.item() / size 42 43 if epoch_acc > epoch_acc_max: 44 epoch_acc_max = epoch_acc 45 torch.save(model, 'model_best.pth') #保存最好模型 46 47 print('Loss: {:.4f} Acc: {:.4f}'.format( 48 epoch_loss, epoch_acc)) 49 50 51# 模型训练 52train_model(model_vgg_new, loader_train,size = dset_sizes['train'], 53 epochs = 5, optimizer=optimizer_vgg)

测试test

1dsets = datasets.ImageFolder('/content/cat_dog', vgg_format) 2 3final = {} #结果数组 4 5loader_test = torch.utils.data.DataLoader(dsets, batch_size=1, shuffle=False, num_workers=0) 6 7model_vgg_new = torch.load("/content/model_best.pth") 8 9def test(model,dataloader,size): 10 model.eval() #参数固定 11 12 cnt = 0 #count 13 for inputs,_ in dataloader: 14 if cnt < size: 15 inputs = inputs.to(device) 16 outputs = model(inputs) 17 _,preds = torch.max(outputs.data,1) #预测值最大化 18 key = dsets.imgs[cnt][0].split("/")[-1].split('.')[0] #对目录项进行分割 19 final[key] = preds[0] 20 cnt += 1 21 else: 22 break; 23test(model_vgg_new,loader_test,size=2000) 24 25'''' 26写表格 27'''' 28with open("/content/test.csv",'a+') as f: 29 for key in range(2000): 30 f.write("{},{}\n".format(key,final[str(key)]))

测试结果 image

点赞
收藏

评论区

加载中...

相关推荐

MySQL:[Err] 1292 - Incorrect datetime value: ‘0000-00-00 00:00:00‘ for column ‘CREATE_TIME‘ at row 1

文章目录问题用navicat导入数据时,报错:原因这是因为当前的MySQL不支持datetime为0的情况。解决修改sql\mode:sql\mode:SQLMode定义了MySQL应支持的SQL语法、数据校验等,这样可以更容易地在不同的环境中使用MySQL。全局s

Oracle 分组与拼接字符串同时使用

SELECTT.,ROWNUMIDFROM(SELECTT.EMPLID,T.NAME,T.BU,T.REALDEPART,T.FORMATDATE,SUM(T.S0)S0,MAX(UPDATETIME)CREATETIME,LISTAGG(TOCHAR(

皕杰报表之UUID

​在我们用皕杰报表工具设计填报报表时,如何在新增行里自动增加id呢?能新增整数排序id吗?目前可以在新增行里自动增加id,但只能用uuid函数增加UUID编码,不能新增整数排序id。uuid函数说明:获取一个UUID,可以在填报表中用来创建数据ID语法:uuid()或uuid(sep)参数说明:sep布尔值,生成的uuid中是否包含分隔符'',缺省为

手写Java HashMap源码

HashMap的使用教程HashMap的使用教程HashMap的使用教程HashMap的使用教程HashMap的使用教程22

双十一预售活动分析

2022年双十一促销活动已经开始,大家应该都提前开始关注今年双十一活动的时间表了吧?2022年10月24日晚8:00天猫双11预售时间,第一波销售时间10月31日晚8:0,第二波销售时间11月10日晚8:00;天猫双11的优惠力度是跨店每满30050

2020年前端实用代码段,为你的工作保驾护航

有空的时候,自己总结了几个代码段,在开发中也经常使用,谢谢。1、使用解构获取json数据let jsonData  id: 1,status: "OK",data: 'a', 'b';let  id, status, data: number   jsonData;console.log(id, status, number )