- 在线时间
- 1630 小时
- 最后登录
- 2024-1-29
- 注册时间
- 2017-5-16
- 听众数
- 82
- 收听数
- 1
- 能力
- 120 分
- 体力
- 567236 点
- 威望
- 12 点
- 阅读权限
- 255
- 积分
- 175394
- 相册
- 1
- 日志
- 0
- 记录
- 0
- 帖子
- 5313
- 主题
- 5273
- 精华
- 3
- 分享
- 0
- 好友
- 163
TA的每日心情 | 开心 2021-8-11 17:59 |
|---|
签到天数: 17 天 [LV.4]偶尔看看III 网络挑战赛参赛者 网络挑战赛参赛者 - 自我介绍
- 本人女,毕业于内蒙古科技大学,担任文职专业,毕业专业英语。
 群组: 2018美赛大象算法课程 群组: 2018美赛护航培训课程 群组: 2019年 数学中国站长建 群组: 2019年数据分析师课程 群组: 2018年大象老师国赛优 |
【深度学习】 图像识别实战 102鲜花分类(flower 102)实战案例
8 a: {$ N, Y0 \. A6 r! n; ~- l8 d# y% C4 ~* F8 `; L
文章目录
4 [1 v* _7 T2 m7 K6 v% V- U w卷积网络实战 对花进行分类
0 Y1 T# j+ j0 P# ?% j6 Q' B数据预处理部分
' m. z( M# A% L/ r( w网络模块设置
) g7 ^3 `- }9 i4 v网络模型的保存与测试
- @. v' |7 S. G3 P" i$ q! p数据下载:7 ^( S# z. v* c3 H( h, S! v
1. 导入工具包
3 `( Z. P6 t ~, l- |/ K8 r3 X0 f2. 数据预处理与操作( ?+ n+ B8 F4 ~! t! D
3. 制作好数据源* P, r+ Y. z; _5 \3 M7 V r+ Q" o
读取标签对应的实际名字
7 `; N2 v% w4 c- l# |6 y4.展示一下数据. n9 j4 \7 V; N* m. i3 |
5. 加载models提供的模型,并直接用训练好的权重做初始化参数3 ?5 m2 U! J' I; J
6.初始化模型架构
1 @* y2 w1 q" _! }1 {$ s# h7. 设置需要训练的参数2 S- N% D2 ?; d, c% n8 t
7. 训练与预测
5 T8 H; k# b& g5 I7.1 优化器设置# w" h' t3 X# e5 n1 L
7.2 开始训练模型, \5 E0 G. `9 E- _* ]1 }
7.3 训练所有层. ?1 x: J% Z$ n: I0 `% e$ A
开始训练
8 l7 [% O/ D' b8. 加载已经训练的模型
. o2 H/ u+ \# d- k9. 推理 [1 L1 T& K4 |9 a) E4 s
9.1 计算得到最大概率; O) @* i' Z- [5 I* x; q) d |
9.2 展示预测结果
/ h, {) b/ R) d, F! g) O写在最后
2 X2 p5 t, g7 Q3 f! _卷积网络实战 对花进行分类
2 C1 a/ [8 ^5 i: p5 U本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能
9 ]5 y6 E5 d8 I) q. t6 b& f. m
/ u/ k: `4 q" H$ ~. w在文件夹中有102种花,我们主要要对这些花进行分类任务
x% g; m1 }4 A) V文件夹结构
7 G, ~' o! @4 ~1 S
2 A, s* k4 ?$ E& G/ |6 G; J3 Dflower_data
- k# c+ f0 Q! \" p5 j4 Y; _, D5 O3 U! W
train
2 k. \4 v. j# \3 L' O! X- X2 j C7 s
1(类别). e1 o& X+ r2 u% k0 L# Y
24 Z( |, S: R+ J4 ^
xxx.png / xxx.jpg' M! d) \( S% o, I
valid7 Q% z: O0 n4 q5 T$ _/ y
; r7 Q/ b8 Z. t) k4 s- M
主要分为以下几个大模块1 a6 f9 G n M8 K2 ], u# Q# C& n
8 t7 P; Q7 v& o& h$ w数据预处理部分5 F7 _7 u/ n- h. @7 V2 ^
数据增强
4 w x- d1 z, a# s数据预处理2 E8 r% w. B# m A6 m5 c
网络模块设置
% U4 C( |5 B% C1 |加载预训练模型,直接调用torchVision的经典网络架构, l, v+ _- ]1 f: m3 E# M
因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务
/ N A: J' |# C% H1 g% F$ H网络模型的保存与测试 h- \7 G, O8 w+ v: }& h9 h) D
模型保存可以带有选择性 r+ f$ Y5 B+ |$ X. A
数据下载:, f/ F. j3 @* o8 ~ c
https://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset
6 m% j; K& T2 p' p; |* _ d& W4 Y. f& O
改一下文件名,然后将它放到同一根目录就可以了
; q5 z% c) k; Q M9 P+ L6 }6 \2 e @( D$ B- h% Z0 A1 g
下面是我的数据根目录
5 f6 M! @- Q; U2 n( e8 t, y, W
+ Q* J+ y, Z, s: w
$ n g1 x) g* L1. 导入工具包
P* ?/ E4 L2 U+ {6 C7 ~# Uimport os
6 K) k7 j( \1 o' M1 f6 N5 Timport matplotlib.pyplot as plt; F5 l! F" f/ l1 z& d
# 内嵌入绘图简去show的句柄' }; S* [0 V* p2 g( W
%matplotlib inline , X+ p4 b$ C5 c
import numpy as np
* N! c7 E9 y3 D$ J9 v1 v$ W2 e9 k4 H8 Bimport torch
) h5 R" b4 ]/ sfrom torch import nn
" Q3 u! t) N1 g& i7 O& J' X4 Q9 X" O8 N* d
import torch.optim as optim8 s) ~, f1 T! Y5 ^
import torchvision3 z; D# N' S% f$ o4 p1 V, {
from torchvision import transforms, models, datasets
2 y3 S- r6 w$ W% P0 a6 `9 v1 Q# N- c( K7 V' k3 K
import imageio y6 `' X. m) B0 L9 q# @
import time
$ x7 X+ M3 B. j$ c0 U3 eimport warnings
. _- r) Z' b% K f3 kimport random
- S- W) h* v9 t) z* |) n$ i" b/ Rimport sys8 A+ _1 e* I' Z4 [4 v9 B3 j3 B* S' L
import copy
' P) Q6 {5 z; Q( H2 h, O1 Eimport json. l0 K& m$ Q' |# f; \8 b# M' x, r
from PIL import Image
6 K4 y2 b9 ~0 w! b9 N$ f: {% l) `( S
$ X0 H/ ^; e; B- X
1! W7 x, ~2 `" v6 L# L# G
2
% g% ~& E% U& U) x3
) `# D' b) U6 p8 i8 w/ @8 _) [( a4
) c. A+ x" [( `3 l* ^) X54 @ c$ P2 R6 N* z$ w9 o
6
7 n) B9 k& ~2 b9 ]/ d& Y7
6 M5 ~8 `% h& Q3 I8: i, ?6 c8 r5 q% Y+ o
9
& O7 e6 c: L! P+ g9 ~10
; v. Q* z5 g! e115 T T, Q* |) ^. o! { w1 c
12
/ _' M6 H% _/ X, v; m+ ]9 m13' A$ j2 _) c# ]1 e) ^( l. R
14
$ `- H' ?/ I# J2 ^* c9 g15" [% u1 Y! Z6 K
16
6 \' _9 T( e! h( x J& n# I4 ^17
4 w) W z+ H+ W$ v+ a" G9 Q18
: p7 e1 H8 _6 p7 I19
T; {: Q3 j, }$ b. p9 P& h20; ` N# K$ Y7 C- u4 J
21; q5 o7 B8 ^: k. {/ k. F
2. 数据预处理与操作% F2 o3 s: t1 Z7 n9 s9 ]9 K/ F
#路径设置
G3 g( Z* E5 T* a" D$ Y- a9 X9 Vdata_dir = './flower_data/' # 当前文件夹下的flowerdata目录$ h m Z1 N+ H! d ^: t2 C
train_dir = data_dir + '/train'
/ X* S2 y( K1 \7 Dvalid_dir = data_dir + '/valid'
8 G* K4 d% ~/ a' {: u. f1; G4 W. H# D$ b2 B; ?
2, U- R" h' G; b0 A* F, z) U' C' k
3; V/ v, X! c6 R( B6 s; s+ Y: w
4
3 j- Y$ `# j, h$ Q! B: ^$ upython目录点杠的组合与区别& j9 b. A2 C$ H! L' x
注: 里面注明了点杠和斜杠的操作
5 q6 @/ \* A/ H+ U+ [
) \: N) W3 x# v1 @7 S. A3. 制作好数据源" Y2 x' M- p5 m3 @ A2 L3 \
data_transforms中制定了所有图像预处理的操作
3 \# s! W( Y5 h& L) AImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片! r- |- L9 O# t
data_transforms = {
4 r+ g( J% h" {* q' d # 分成两部分,一部分是训练
% M& Q* {* A( E" Y/ f 'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间
( j; J7 a7 T. I ` transforms.CenterCrop(224), # 从中心处开始裁剪5 n* E- g0 O, ~) N( j
# 以某个随机的概率决定是否翻转 55开9 L9 D% l: b/ c3 L- g' Z
transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转
% x) {% T8 B8 s/ _: t- L transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转
! |( O8 |4 N. m4 w # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相- v* }3 R, C$ Z; i1 ?4 Y
transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),2 V+ s& k& e5 B) D3 f
transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB8 C( m2 s1 W8 Y4 [
# 灰度图转换以后也是三个通道,但是只是RGB是一样的7 f8 P. U k R
transforms.ToTensor(),
" }# U9 P, |3 k4 B transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差
% M6 @# l2 v8 j$ n W ]),
, Y9 N2 s8 ^: o* t5 g( t # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
+ z$ B0 t* v! M 'valid': transforms.Compose([transforms.Resize(256),
- R4 d. ~+ }8 T( e transforms.CenterCrop(224),2 r1 g8 ? g# i7 B1 [
transforms.ToTensor(),
* m) E7 }6 ^4 W7 }0 _. u3 V transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同
% Z$ I6 I, g' i L( w% @4 o ]),# v! L- A/ Y) B& s2 y9 l s8 D3 C+ {8 {
}0 m+ m2 r' Z7 s0 h
8 j' E5 V* w3 c- z" I
1
9 @! S% Y* X+ T5 N" l# |* ~1 \0 d2( M4 b) ^6 W0 N6 Q# n; D Y
3
$ t; ?, ^" f; i7 X" k0 ~4
6 c) ~7 B3 T: Z3 w5# d2 f& `1 K2 _ @# ?3 @
68 q9 f3 v1 ^. Q! b6 ^% {* i' i
7
0 J8 Q9 l$ b. }) F1 @8' n3 c/ G5 v2 ~7 R) P
9
5 L+ _: @( y% [1 k% H3 c6 e10
# g) O% F, @& w b0 @; ]) P) Z11
6 k! t1 _5 [! J/ j7 i12
( p6 U) T+ K% f- }13
* ~' H$ u' d2 z# q14
# U4 ~$ t' @# \, |" `2 B" z15
2 P r6 G. o6 `" q16
6 S6 P6 x" e* n173 Z' P4 Y4 R1 U* h
18
3 _* t0 h( G# ?8 R E19
2 Z. E/ T' ?- P9 R, l/ x' a5 P20
- P4 w3 s5 ]1 K1 e21( `) J3 t3 m& k. L* o! F
batch_size = 8
5 W `* }) W% h8 B8 D2 o# Y( J& Zimage_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}, {% m8 m2 v# m `# C+ [0 v1 N
dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}1 @8 O6 k' [$ U
dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']} ( A2 r- Y( l* u& {; c5 d; {6 Z$ X* u
class_names = image_datasets['train'].classes
1 S1 t7 @6 K' Y* i( \. o" K! z+ w
#查看数据集合) ]; J6 n4 f) a. C1 l
image_datasets8 c+ R! W* F+ l) ]" B) Q
/ \+ y. x$ _6 ~6 U2 |9 k3 y
1
" Y+ ?) ~% [' \+ }2( F( a p: j) Q% U. W7 V
3
$ k1 U/ T2 A0 ?' d0 Q; F40 J2 L. a! M; C
5! [; W/ d* A5 ]: `
6
& z5 x' `5 b$ E8 V/ Z7 _7! ?" b+ c' v( a$ m) D- ~
88 O4 t7 M6 O: H6 ^7 w$ z
9& A8 l3 k( Q9 ?+ c& O3 B
{'train': Dataset ImageFolder
. U. q5 q) h5 h( q# I: K2 E9 S Number of datapoints: 65525 B1 I! j3 @( r6 \7 q# g6 z% A9 ~
Root location: ./flower_data/train9 n( @' B k" K) D$ ]" q1 U1 y# P
StandardTransform
/ V8 j, I. J$ h: k* b: @- \3 t Transform: Compose(: t7 n, s* \" P; e- z1 R+ k+ H
RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)
) }' ^" l9 _+ ^1 F8 `( L CenterCrop(size=(224, 224))+ i' G; A d3 e0 T# E7 S
RandomHorizontalFlip(p=0.5)7 G. p* s8 r; {' m: [9 J
RandomVerticalFlip(p=0.5)
# c4 S: l. Z6 K3 v" L ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
$ ?/ U" f8 j& J0 ]0 a RandomGrayscale(p=0.025)" s3 q/ h' S! n! O5 |5 G7 e
ToTensor()
$ s' i+ T3 ~$ g Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])+ d+ _9 x# j; S u3 n P. w4 l
),+ D8 c7 L' l! T- ?# T
'valid': Dataset ImageFolder
# y# s* u8 h8 r: `& }4 l Number of datapoints: 818
- i+ [/ W" R0 F Root location: ./flower_data/valid9 c5 i' c: c: v- S* a, c! {
StandardTransform6 h8 \( t7 o. H9 h5 L; z
Transform: Compose(7 q& p Z$ \$ x: y
Resize(size=256, interpolation=bilinear, max_size=None, antialias=None). F1 \& B0 t+ X
CenterCrop(size=(224, 224))) h/ V; W+ X4 g. Y8 |1 E0 v% U) ]# ~
ToTensor()
- K9 T/ Y" d' U" e Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])4 z) o- h7 c+ D1 K8 j8 m
)}
: J) ?/ u8 G+ y# `2 r# Y
, e" |2 e6 g. r( g, p1
+ r5 N/ Q7 K- a( @2
# X( \ y3 M. D$ f p3
. s) W2 _+ Z$ n! D1 x1 e w! p2 r4
/ y- N$ m7 H0 o4 x2 b" F5" \8 q7 u8 U; y: O8 `, o& N
6
) I/ C! a7 W0 T0 f1 n( P7: |" D, o# L k$ z4 g. @
88 o9 V" K' S. G2 T: i4 Y
9
2 f" P7 u4 P- ]7 l10( R9 M% Q8 M6 j& F9 v
11
* b, f' l+ Q9 J+ u) F5 k# s12
$ S9 O) O! H! E" ?/ _$ i13* S; [2 m0 N, R _$ `
140 H! s3 n6 A: e! ~6 G4 |! J
15
4 I/ V% r) |0 F1 O- K165 U. k/ o9 t8 @4 E( J
17
% }6 g( v* a, d, f* c18
% }+ A7 @0 T: M2 T' d19
/ \/ `6 i/ f# ~( M! z2 S6 [7 V* s20
" P; X, @0 I7 \/ }$ C21
1 W1 L+ W, ]0 _ Z; R1 s" E22
2 Z8 y. f ?* C" r6 x' V23& [$ U/ c8 ^8 e5 Q6 r+ x+ Z& t' h
24( Z5 v' B7 S$ a' h
# 验证一下数据是否已经被处理完毕( z( A1 F* Q G, t; i. a. a" H2 w5 t
dataloaders
; O& i' ?6 m0 y( z3 y6 _& _) x1$ ]+ _ u$ ?0 a- G* b* R4 n
2
6 M D7 N; E. k7 J$ ~$ g, `$ M{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,
! Q) L4 y! |2 x4 F: n 'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}9 w" c2 n" }) y6 `5 t1 G
1
% @/ t' ^9 \) G1 J2
2 j) L9 t% K5 D' K- w' [9 Odataset_sizes. `" l! G; @' I/ k! a" v
1- t2 x% k, q2 E& g
{'train': 6552, 'valid': 818}, B& Z. T( I( P* V d) u
1
2 _% v+ ?. c2 N3 d% W9 n读取标签对应的实际名字
6 R- J) x" v+ _5 W, B3 R2 B; H使用同一目录下的json文件,反向映射出花对应的名字' A- `4 b0 t) f! S# z* d
! Y, D' m% O/ t4 j3 Z, w
with open('./flower_data/cat_to_name.json', 'r') as f:
% p& S+ E2 s0 M" R) M4 p6 T. h- Q; l! s cat_to_name = json.load(f)
, Y8 b6 u1 Q7 N' i1
* \/ W" E R( {7 J' d% F2% _. b, w' G& [5 U* k: s) D0 Z
cat_to_name
( {. A/ R# F7 l/ p1
7 P, D& y( }1 L6 S{'21': 'fire lily',
$ J J- j$ L' _2 y4 f: _/ r) ]/ ~ '3': 'canterbury bells',' w; ~" I2 e% O( X
'45': 'bolero deep blue',2 I( x$ E, Q5 R+ v1 z
'1': 'pink primrose',
5 b( T% s8 p( A" L8 x '34': 'mexican aster',
) a& C! b% f( x) S1 I+ z '27': 'prince of wales feathers',
2 m& i& l% Z. ?. z# M* } '7': 'moon orchid',' Y7 O/ `) p( ?( \ Q2 ?, K
'16': 'globe-flower',, q6 A* k; K$ z& H1 J
'25': 'grape hyacinth',: ^7 t# }' W! {( _
'26': 'corn poppy',
4 D# i! ]- C/ Y0 P/ z '79': 'toad lily',
7 U2 {* p: h( o0 W, [ '39': 'siam tulip',5 b/ [4 G1 T: Q3 h2 v0 n, }
'24': 'red ginger',
9 f2 C& r' L+ z; f d '67': 'spring crocus',1 }. I3 m$ X) ~1 x) N
'35': 'alpine sea holly',- ]/ G; g K7 x! w: O; o
'32': 'garden phlox',
/ r, i# F0 I; ]- Y q" N/ i '10': 'globe thistle',
6 n( W. {0 c2 T# s. H. x! p6 K$ q '6': 'tiger lily',$ ?9 s* d, Y Y8 r E
'93': 'ball moss',
Q6 S* H3 m* V$ t* V '33': 'love in the mist',
, R% W% U7 R& k1 Z0 I* D '9': 'monkshood',
/ I6 L/ p; {$ Q, L ` '102': 'blackberry lily',
% j; ?. o6 S1 y1 A# C9 k% W7 x6 q4 f/ E '14': 'spear thistle',1 C0 {' A( Y7 q$ |: n7 E$ o
'19': 'balloon flower',
. K% k# g4 B; ^) b0 z '100': 'blanket flower',& p: @0 a) l! ^: u4 [+ @2 R5 C* ^
'13': 'king protea',
6 Z7 }5 D+ {: e; ~5 n6 I '49': 'oxeye daisy',
* Y; J' `' H& f9 L5 L, ~" W: W '15': 'yellow iris',: d* t8 e! M: b) J8 e: r+ G
'61': 'cautleya spicata',
& [" S( J: ]/ t% W '31': 'carnation',- V) O; l# p2 ^8 s2 x2 _3 K E
'64': 'silverbush',3 L& c* b6 O! [8 d
'68': 'bearded iris',* M: Q8 }: \3 U9 X* I: _5 _0 n0 Q# q
'63': 'black-eyed susan',5 I8 K) E1 x8 g3 x) }
'69': 'windflower',
1 Z& H* E( o t& b$ T/ s7 b, n '62': 'japanese anemone',
T8 Y# b/ r' `3 ]: g '20': 'giant white arum lily',) M% ~8 [4 o; w0 Z
'38': 'great masterwort',
) ~+ c! p8 }4 }1 w) y# v5 p '4': 'sweet pea',' @! v7 H% h3 N6 F8 H
'86': 'tree mallow',
( p3 ~$ Y8 P. k '101': 'trumpet creeper',
, S: z: w, D9 b# S+ |# w2 ? '42': 'daffodil',. I' V4 r- j7 D8 o' K0 T6 N# g
'22': 'pincushion flower',
- T3 I7 _: }$ Y1 x' N, F' { '2': 'hard-leaved pocket orchid',
0 \2 v1 J0 l. ?2 ~7 D5 _ '54': 'sunflower',6 g$ ~6 Q$ [ k; y
'66': 'osteospermum',- [% I. C. a$ P' Q% m4 H% O2 ?
'70': 'tree poppy',* k1 H! G4 K6 Q0 |8 }% t
'85': 'desert-rose',% ?, g: x6 v* c9 J( \, q
'99': 'bromelia',& f/ \ J, v% v1 v1 ]# U1 u: B
'87': 'magnolia',
& t: ]3 Z- o. {8 p! B. D '5': 'english marigold',8 N. o( H3 _8 V! z& g
'92': 'bee balm',5 Q2 r7 @$ ?% _5 s5 H
'28': 'stemless gentian',
% ~6 ]+ q9 e# ?" `: [# A9 C. ~ '97': 'mallow',
7 ]5 Y5 v1 X' V2 t- n '57': 'gaura',
$ Q2 |' ~% Y! M h2 e '40': 'lenten rose',
; q, Q3 E7 g9 t '47': 'marigold',
* x4 z( M: B& H) O; h, a '59': 'orange dahlia',
% z' J0 ^6 \ E. d" F- e '48': 'buttercup',
- x5 I/ R/ n! \9 ~( @; |0 Q2 _3 x '55': 'pelargonium',
" l7 ?9 j8 J7 A1 m '36': 'ruby-lipped cattleya',5 P' u9 l4 G5 t! m$ W) j1 [
'91': 'hippeastrum',$ v" I: f# w! u: |5 E7 g: N* x
'29': 'artichoke',* o' e$ [& ]( q+ J
'71': 'gazania',1 y# n' H# D3 k
'90': 'canna lily',, r, o# y# V) W( e9 U, I6 `' J
'18': 'peruvian lily',
& C3 ]% T3 ^7 ?6 f* z2 n '98': 'mexican petunia',
( N2 S# ~! \; \) \- ~ '8': 'bird of paradise',
4 j. U) G! N' p '30': 'sweet william',
3 S$ v8 n9 y5 B '17': 'purple coneflower',: j) S8 H9 }" f; s; O, N
'52': 'wild pansy',
: l O4 k8 e4 v3 [0 x9 ] '84': 'columbine',9 E4 O" {0 D; I* [9 v
'12': "colt's foot",
: q D$ v) S! M3 N! h4 V O '11': 'snapdragon',
: ~' v) [, x( M9 q '96': 'camellia',
" j+ n/ \9 o$ Y0 @ '23': 'fritillary',
/ Z% t2 u+ k. t '50': 'common dandelion', z T F- H2 T" [
'44': 'poinsettia',4 f8 b+ Y2 w$ u) e8 U5 y
'53': 'primula',4 g# d( A2 M1 Z" h7 d K8 k Y
'72': 'azalea',/ ~, @: U9 X2 y, }8 f( c+ E3 R
'65': 'californian poppy',
$ g S. y, _2 ~( X0 e; L& j '80': 'anthurium',
+ O- z6 z1 ~# w '76': 'morning glory',0 P3 x& W& T9 t
'37': 'cape flower',) f* \5 G) S6 p; Q
'56': 'bishop of llandaff',- O+ y: u% k0 f$ }, N1 `% l; T2 W5 O
'60': 'pink-yellow dahlia',
! S6 s* f* g" R A" q '82': 'clematis',7 L9 J, u7 ?7 Q2 a& S( b' r& ?
'58': 'geranium',3 J2 ]2 ]9 E) c, ^ {
'75': 'thorn apple',
4 @$ k3 ^: I2 \" A- y '41': 'barbeton daisy',0 ]1 H% V" x" F' o- c1 m
'95': 'bougainvillea',/ U0 v0 M# X2 k6 |, ]. p
'43': 'sword lily',0 Y1 E8 N- w9 c- l- Q+ i( F" s. B
'83': 'hibiscus',: w2 l+ F* H- h3 Q1 j" W: j
'78': 'lotus lotus',: S r8 j8 d) W' [* m1 Y, h4 d
'88': 'cyclamen',0 _+ i, Q# F( O, f A9 H- w
'94': 'foxglove',
) F! b- c/ E0 e '81': 'frangipani',/ C$ `. P6 y- [" x0 d9 O6 w! T
'74': 'rose',
* X! x5 s( l1 W% ?( ?& ] '89': 'watercress',
2 q6 q: ]8 b) q! I+ N4 l4 V '73': 'water lily',, d( D+ n8 u q5 V p6 R
'46': 'wallflower',
" z n X1 q( f. V3 A) E& d '77': 'passion flower',. i7 s! v5 Q `) k( E
'51': 'petunia'}) J" C; L/ F9 w( B+ [8 F
$ E+ p7 ^% E+ p/ n1 a1
7 u6 R$ a4 e) X9 l7 K* Z* Q2
1 T" ]5 D/ i" G9 Z3
) T, D! Z+ d- l% q43 {* o$ K" U3 G- x* i! p+ o) Q
5
: q4 p j8 e. z) j1 E! B8 N# x% T6
: |; v/ X$ ^$ G& t/ v7
- E, [! Q3 |; u/ R y8
4 l: |& F; M/ O* H8 ~8 D6 { B9( j# G" e) _1 `6 Q$ D! Z
10' W0 q4 ~8 f) ?6 p* }
11
( ?3 g9 `$ u" M3 R, f120 q2 V5 J. X$ Y# h6 N
13) |3 k: _& h7 a7 e7 H
144 u& `( v' y, C" L$ s$ m
157 s. Z* `9 Q5 O/ I
16% z* P1 L6 W" R7 X
17/ z/ i+ {* u; t' r5 ]! K1 e! `7 X
18- L! m3 [! X, Y' F; m6 u: Q
19
( n3 y' H0 {3 j20
; ?# H* g% X2 N: ^21, R# l( ]. P( ?6 l& t6 e- h7 S
22
' B/ q! _' }) _# B$ H2 {5 B- |# F23% d5 C$ b) x( A1 X2 Q0 ]7 Y
248 C5 n2 q* l4 ]/ G
25, E R8 |4 h0 {. [/ g8 \
261 \ t! u$ i* A: \0 @0 [
27. h; L( i* S: s% `
286 E2 y4 j6 \# k% q. M- D
29
6 O; S, u# U7 l9 x, r$ Z' Y30
' @0 [. r4 _& J314 L; d l% ]2 {7 ?: m1 p& {- U
32& T* [( c0 ?& g5 C# s
33
/ y3 T& |3 S* U344 [ r8 a) R) r. j
35
7 @2 _/ p3 [9 a1 c6 i7 ]9 B36
; p0 i* i3 Z# s1 D `; I- F$ V37# f( b, R% L0 o- j3 W M. H- [ t
38
3 q* V3 ?# T I' m$ ^$ M39
' N, M% J& U2 r2 I" ?400 W2 q' `9 M2 F- {
41* K2 ?0 ] R6 _' X: Y" c5 ?! N( K+ \
42
. @! j& h |5 P9 A6 L43
7 d" [2 n* |+ Z; f44# i9 d8 k$ `' y( G0 J/ C: }9 _
45
, A- {' d* {4 Z% [. r4 J; e: \+ n, ]46
$ K$ G5 Q5 {4 O2 M47
8 I. i6 y+ g! x- P48
0 a0 w/ ~6 n; h/ N! q5 _49
! m7 ?* @9 ?/ G+ N50
% k) m! ^) U4 I( y. b* B0 ~3 S51" o" t: p+ r' a2 F$ U
527 D* W$ n# U- D8 y+ Q
53! X. q* t9 w7 {' z2 P
546 {" [' g4 W. T+ ~- [
55
, L0 X# x' Z) G1 ?/ B; b56
( H, x$ m+ Y }/ R2 y l57/ e7 X( S9 P1 G( f' J2 Z+ B) q
58! i" N# S% Z6 c. E
59, X( ?) Q: J) R' A$ ?
601 i4 s& |0 c! r# o( D
61
! ~. V) b! o1 A( S' x) x624 ?% [4 g0 V% x9 [( i
63; P7 {+ I$ w. m+ }3 G1 p
648 d8 I' m$ D* t
65# [3 Z/ [0 ]7 X, O5 Y; _
66, m/ I/ w2 ^2 `/ L! u
67
: m3 I2 ?( e" R682 v3 k3 ^+ D: K5 r9 J
69% ^$ A/ R& ]. d7 V2 }
70
+ \' w! a$ l4 F4 f3 A# e0 E71) w% S9 m# t% Z$ ]- }
728 K/ F: k$ `; Z* L- j4 E
73+ m" z$ y% R5 q3 }1 Y, r
74% F, E& V* y1 C4 e5 \. G" ?1 e
75
& G) o) a' W5 N/ N76
0 q4 w* a" u$ N; V* o776 S+ m+ ], v( v* j" b
78
[/ b9 D& g7 K) d79+ A- A' T* ?/ E' b+ v) [5 E6 K Q
802 I7 [4 Y2 G- _4 \- K
81
) X9 j& r: R" D( u- J0 [82
/ E+ P1 n% V% G83
/ e8 w& ~; W2 K8 @4 O84% c, H1 d- F# }- I0 _4 |; h1 j Y
85
3 |! \0 W" j2 V% {- b. ]1 v86( b* m; y9 D" K: x+ y' j
872 ^/ p7 i6 C% f C2 K, W2 R
88
3 w+ [/ K5 k3 G& A. I8 l) r! C2 U! S7 y898 Z- H. a' l& Y
907 l" E" A$ }) q8 b* M" [
91! H* }" Q; j g S
926 A2 ]2 P6 O# s7 s8 D1 K
93
/ R% Z3 m. r1 D1 h3 q94- B" }( d: l, c3 t
95
' N% b3 n* ?; u" H4 D) k96! ]' ?, C5 }9 M# i
97
3 s8 ~$ M2 U. h( U R) K98
" ~! b) V* I4 R! }% @99% z: [1 K4 O$ m3 W: ~$ Z- A
1003 s/ \9 Z& d% x
101
# t5 j) d& F8 {. N: L102
$ T' x+ ^: G6 `1 W4.展示一下数据- v% {5 S# T$ T
def im_convert(tensor):' z1 C- V2 g- V7 u; f
"""数据展示"""( [$ h& z, c1 M) ?# k5 R- }& ^
image = tensor.to("cpu").clone().detach()
4 a. U% `4 H; m7 C$ r image = image.numpy().squeeze()6 D. t% c1 V$ m0 G( W+ ?2 y
# 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图* t7 M, B, N& n1 B2 n8 I0 G5 ]4 i
# transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)
2 @0 b$ H, y% k& a image = image.transpose(1, 2, 0). @ q( i! a1 J+ h! a
# 反正则化(反标准化)
/ {1 V( L" |6 R8 y3 k! w; Q image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))5 F+ G5 V0 l/ E# X+ D' F2 q1 r6 J
' y/ p3 { c5 h+ D$ d% S9 `- Z5 Y # 将图像中小于0 的都换成0,大于的都变成14 O0 a: \2 c6 @4 k! ~3 T$ z* ?
image = image.clip(0, 1)
: _1 t; y2 V4 D* S; C
5 t C' d: U! i+ \ return image
/ o' a6 t% Q3 n6 [1
6 H# q/ h% ~. @7 p* |& \2) c. n' H( _! V' O" k
3
/ M( p. v% A7 T& N' p40 e0 W+ C0 [2 e% }- c
51 K) Y. h0 d3 ~/ C# {2 s
6
% O# Q4 g% h( [0 X! ~7% v1 h! g% v. h- j
8/ `- U, k: W1 Q" r7 O
9) ~0 p* `* h1 |1 N
10
4 l( @' I9 ?8 N11
x9 u1 e5 M% ~" @ m* F# P127 h! G+ X! L2 n& }$ h- ?
135 O; j; u8 }& l& V) S v
14
, @- F& Y W2 D8 \+ U# 使用上面定义好的类进行画图/ Q/ y( h# R9 ?" c5 X3 I
fig = plt.figure(figsize = (20, 12))+ Z$ k5 g% @" x# h: M
columns = 4$ [7 k% o% y7 G) d3 w
rows = 2
$ [+ H5 S/ ?+ }, R2 h3 K7 k8 y3 V$ q# }5 k
# iter迭代器' D1 f. q1 _% p( F C1 D. q p
# 随便找一个Batch数据进行展示
* t. a! h3 Y9 _6 L' adataiter = iter(dataloaders['valid'])
2 c# T9 m- R& {0 pinputs, classes = dataiter.next()
8 ^# b$ W8 H) i
: d& x. y; c& ^3 F- o* O" X6 E3 m5 @for idx in range(columns * rows):
) R5 a1 }. {) Y) l% {6 O6 l# b ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])& p3 s' b8 [& y9 Y" N8 U
# 利用json文件将其对应花的类型打印在图片中
; h; G4 d% _$ J6 K ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])
: }7 v2 r6 v8 }0 Y3 J! @ plt.imshow(im_convert(inputs[idx]))
: Z# m2 ~7 q' U0 gplt.show()
3 ^9 }; y# h$ B# o2 v. D# `
7 P8 U8 }' e: u, d$ ]1) g2 v, l( s+ Z2 h) O1 ^
2( K7 N- E5 `% v3 a5 \: q/ a
3" k( P. s0 a2 N+ l/ C5 ?: Z
4. {: x8 T- I: d* j
5. W' B# L- j: E7 p n* B
68 V3 @6 k) O% x. f# C
7* m8 d8 V! O. R- U, O# l" z
8/ q5 d( d8 a1 |1 M/ u' t0 I
9
8 J: S3 p/ d6 |1 d1 T' {' q5 K10
- e& ]( k8 |, @/ f3 o( ]116 Z- a u0 X! _) @ E
12
8 N2 o0 k K8 ^, ~3 x9 t13
( Q9 n- p1 s# i" t: H! Z14
+ c' a* B" c& Z, z6 @; I- S: f15. |$ P: `+ v0 K
16
6 w c1 X/ |2 x1 u2 f+ v4 \
* E% f0 ^: j' Q* ~1 C- P2 N8 N/ |& t; X1 F9 B4 q% z
5. 加载models提供的模型,并直接用训练好的权重做初始化参数
& L, O; z7 B: h/ [" jmodel_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']- ~& q) U& K6 k8 a& C k1 j5 K1 o
# 主要的图像识别用resnet来做 S4 J& { F& n9 o1 K4 x
# 是否用人家训练好的特征
# v2 ]( J) h, Dfeature_extract = True
8 U7 E" a0 K8 C' `# A% N1
8 x3 g4 Q9 h0 m: R, l9 H6 G2
) {; i8 @* {* G3
* _ g* m2 U& n x& Q1 ?4
- _" l5 c, M( |. C# u, Y" K# 是否用GPU进行训练& e7 Q! Q' k1 z( N' o
train_on_gpu = torch.cuda.is_available()2 x' h6 j" @, r$ O
o5 a) z* K, v# E. nif not train_on_gpu:
8 R1 [! \4 g: s7 } print('CUDA is not available. Training on CPU ...')! ]. o+ ?( X# j# ?
else:1 ^$ ^ n( c% ^* n) H0 O
print('CUDA is available! Training on GPU ...')2 m' } k$ T* H
% H2 G$ p( l- b: `: h* R5 E
device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')
( ^8 g2 N# U4 j' y0 d$ I7 u& s. @2 [1: j$ R& z% e; ~8 `: K/ l3 i; H
2# _) |5 \. f, D
3
0 ]) Q% o+ B& b1 {8 h# }, p4
! d- v( h, c* `& ^1 M52 f; D' q3 F) Z8 V
6
* p! ]! W4 R2 k& l* g. C) O0 ]7
8 q" P' @' L( h0 n! Q5 _8' [$ ~' i% D: ] ~
9, {! v8 y* _- V8 V$ \- @) v
CUDA is not available. Training on CPU ...
/ ?5 a, C! o# O2 }' j7 Z1 p4 p% \1" r! n5 O& O3 M v
# 将一些层定义为false,使其不自动更新
" q1 D3 C9 m* L% w) Udef set_parameter_requires_grad(model, feature_extracting):
" ?" K+ _- y8 g( e if feature_extracting:
# S; J5 E$ U+ @" M& z; Q; n for param in model.parameters():5 I3 I5 ?; s7 J- }4 B7 P
param.requires_grad = False
" K! X5 u2 G% F' |7 B( o$ x1+ l. z8 X# y6 s
2
0 Z: \& h. J! G0 D3
5 K( E; w/ P" Z1 r4
' r5 L0 @2 K3 g! p8 U57 C- [& x) W `' m2 D: i6 h, m; t X
# 打印模型架构告知是怎么一步一步去完成的# F ^+ c" z! k N( ?: f! ?
# 主要是为我们提取特征的
# s+ g: d& {8 o: l
* H7 }' `4 o, [3 `model_ft = models.resnet152()
' |5 D0 Z0 @" Z& F; Mmodel_ft. h2 u% X- J! j& ^, `
1
3 N$ n3 o1 G w+ _ g, d% @" r7 Y28 ~7 p$ K2 H/ _% `
3
. J! G# q; `3 i4 X G49 g, W; x3 N& w. B3 G3 I
5* y* f' j, N* M
ResNet(- q6 ^. e( T" e& U# m; i" s g( i
(conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
# x" L$ s4 o; [" h (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
1 m4 q$ p5 a; `; K0 r* ` (relu): ReLU(inplace=True)* Y3 D4 @6 @& i' N$ M) @
(maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
/ }! I7 c+ j# y. j (layer1): Sequential(
, E! Y) T. [! ~+ {4 s$ B (0): Bottleneck(
: V; C0 Q7 J# Q; Y5 i0 l (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)4 Q+ w+ A. `% H3 D2 N, y$ q* @
(bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
/ g! _7 T8 D7 }" m$ S# P (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
# n$ J0 a0 w) `% [9 H* M (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
$ ^8 z7 s9 u @/ W: l5 ]# b$ V2 v (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)( V: {% u& d7 k. U8 W
(bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
# a+ A3 I- r, H; J2 [- ~ (relu): ReLU(inplace=True)
* G" r, N w) H6 b2 b! K (downsample): Sequential(/ A' \3 ~7 j) d, s2 `
(0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
4 u2 y: a% _5 u* D3 b) U5 V5 Y (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
4 g5 M% U" \5 v/ G- b ); q7 K. W, s( p
)
9 n4 D0 v$ J9 W+ F3 } n8 L中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。0 Z7 c' R1 a9 S" Q3 C
(2): Bottleneck(4 v5 D2 T; }0 V. p) q4 \
(conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)
9 U& Q# H G0 m$ ]( N: f (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)% _4 j# ], X" _8 f) A) ~
(conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)/ R. O2 {* d1 O
(bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)* P1 K( k3 ?# M# U
(conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
7 W- f6 U/ A$ _; \" i9 K. P (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
- A) V, o( \' _- o2 W) s) R (relu): ReLU(inplace=True)3 h5 j: [" O* Y5 C4 W @; r, C+ T
)( J+ B1 T$ F9 ?9 o: B* V
)
7 f: J0 W6 K# L. {2 ?% g* G$ u (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))
2 ^2 o& T c5 E, b8 q4 v. a (fc): Linear(in_features=2048, out_features=1000, bias=True)& N1 j+ O% V- m
)- e2 u3 b& k) K4 T
: s3 V! t) ~/ z0 w* j
1
/ F' T7 c' r. b8 u8 g5 l; Q2- n3 \ v4 |/ A- W/ d. @
3
9 f# [7 W4 Q, K: y4
: N/ i" X; E; y; n5
9 Y' b. \7 c1 }( s8 S; w6% v5 { e% q- J- b# U5 W
7
' A2 A; Q3 M: ]9 Q; c* _/ ~8
- v+ I) o3 Z, b% \2 ?; G9 ?# g6 ?9) d- l$ d5 o0 N- _+ A! V d
108 }7 ?2 m; d' g5 s3 n8 j6 z6 O0 r
11
2 G7 e5 @! e o4 p12
6 E$ `, u: f- g! g/ Q( |9 T1 A13
$ \2 m! l" f' Q- F& q' p9 Z+ `14
4 |' q( A: C" Z! y/ c15
7 A1 ? W5 }' ~16/ |, M9 ?+ R% B& c1 P6 e: e7 [
17
. Q8 A- L7 P7 V, T% H18
, {' h3 v: a3 ` i9 m19
5 s; }9 _6 N6 _2 O5 k8 Y% v202 F$ K3 u) [+ J, g( k
21, V L9 H" ^/ w$ Z; M" Q3 a1 s
22
3 V, i9 y+ G1 b23
6 b+ F0 _: s% X" @24
s1 U6 a6 q5 T0 w25
! q* }4 b( [, k5 I* ?% ^+ r( [& U26
1 c) J B0 h) ]/ a+ ^$ ^27
9 n3 V. z/ J4 p28
4 v9 J& Z' O% n& |29
K+ {1 R4 v6 _ G7 g7 i2 e308 n' Q% V& ^2 y; K( x, G
31
4 R) X( G: z( D- n1 o, P328 u$ y1 ^8 P* z5 x+ B
33+ O' A! e( H+ A
最后是1000分类,2048输入,分为1000个分类
0 W) a5 {9 Q, `+ ]$ K( j而我们需要将我们的任务进行调整,将1000分类改为102输出, ^0 d7 W+ ?# N* o
2 B- [1 t* R4 y9 N" ]2 {
6.初始化模型架构
6 g7 a- m4 N( D) f; {步骤如下:) j: F& s. R* j# J( Y% `- e
- W# r" O9 l" R" }将训练好的模型拿过来,并pre_train = True 得到他人的权重参数
% E* `8 L0 G$ u3 P可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)9 h( i9 ]$ x; P* }& v
无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数& M$ O0 p$ X6 e4 x8 d
官方文档链接+ G' m# n; N- L' L
https://pytorch.org/vision/stable/models.html
( J4 i0 S# Y4 \
! H3 |# h4 }9 |% {# 将他人的模型加载进来
8 B0 I% w7 d: Q+ K/ fdef initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):
. W0 ^) g) A* ]8 w: {9 E # 选择适合的模型,不同的模型初始化参数不同
, T. `0 }. K! K8 q' T model_ft = None
9 y D! j, ]1 u' y& d0 \' U input_size = 0$ ?0 g0 F/ L+ k' T. X7 F
# B' G) p% E# i' t/ R# Y- Q( d' d* r
if model_name == "resnet": R+ ?. `8 U- a3 ?. y3 E$ C
"""
# G+ h& O0 A8 g& Z Resnet152/ I( p5 M' k* H# Y) K- [
""" H6 z! |) D2 V5 o: Y% d2 q% V( m
- t0 P, M/ R3 }8 s6 S2 f # 1. 加载与训练网络
5 t7 q: r; x' b6 v0 Z6 k" O model_ft = models.resnet152(pretrained = use_pretrained)
; z( I$ O" [# I) g # 2. 是否将提取特征的模块冻住,只训练FC层" n5 s) e% P* C
set_parameter_requires_grad(model_ft, feature_extract)
- p7 a! L% m8 z4 {/ f3 o4 o # 3. 获得全连接层输入特征
" x* J2 U- R+ R! w5 f num_frts = model_ft.fc.in_features/ o1 H' L# b' c3 R6 w4 E f
# 4. 重新加载全连接层,设置输出102
- q# F# ^5 I5 B ~ \ model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),7 S/ Z. `+ k' [' B6 x" Z" l
nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1
9 i( y% W& t! L. K& Q! y; f# ? input_size = 2244 N! N0 {5 n9 {
3 q. y8 C2 n, D& e
elif model_name == "alexnet": Z8 m2 \6 O g% M$ _
"""
* s. R3 I. P4 E Alexnet6 D$ N' @0 Y4 O; r( _4 v3 J
"""
% E/ X) U% C) s# l: m# q$ }) d model_ft = models.alexnet(pretrained = use_pretrained)! U) L& m# \" o4 |! L5 N
set_parameter_requires_grad(model_ft, feature_extract)
6 j) H) X" I) O9 X( j% z3 f0 l, _1 B7 G. y% n+ ? R# M8 g* R& S
# 将最后一个特征输出替换 序号为【6】的分类器) Y: b5 A7 B# ]6 ~/ `6 i
num_frts = model_ft.classifier[6].in_features # 获得FC层输入. D, |( n3 N/ D2 F( B
model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
, I9 g- L) X- j4 ^ q0 ~) Q input_size = 224. K& A K' w) X- E4 R
& k6 L/ C: o* M' b$ w- G. D
elif model_name == "vgg":
4 b6 l r6 I1 M """. j; f+ E+ i6 _6 ^. J3 ?+ c5 S
VGG11_bn& i, B& e" t- v+ B
"""/ R7 K9 e8 e7 i, m
model_ft = models.vgg16(pretrained = use_pretrained)6 W" P' A j, A9 n' h! r9 h
set_parameter_requires_grad(model_ft, feature_extract)
; s$ k$ V9 F' @: W I* I num_frts = model_ft.classifier[6].in_features& p/ B; I1 {# J+ O& i! P. t
model_ft.classifier[6] = nn.Linear(num_frts, num_classes)8 ?$ E) E8 `8 ^* d0 h% w# m
input_size = 224
+ l, R1 a. W# Q* e) e$ c( k/ A. U( {
elif model_name == "squeezenet":
/ `3 y; N4 K7 q% K3 _0 k """; p" d6 S" N. ^( G r6 z
Squeezenet
1 G7 O7 b( p5 }1 b """( K9 F l- m' r# R; i0 M
model_ft = models.squeezenet1_0(pretrained = use_pretrained)& ]2 a& n7 X- y) b/ V
set_parameter_requires_grad(model_ft, feature_extract)
& u" ^4 \" @' @ model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
; G+ j3 I( J- m3 r: t model_ft.num_classes = num_classes1 ^9 X2 _+ A% O1 V; y
input_size = 224
9 R; |7 k5 o. {/ m- G, k! \2 C- S
& | p1 Z7 U/ X9 m elif model_name == "densenet":- R0 N& V- u- y% A
"""7 s& r) e8 ^/ M& H
Densenet
/ N J' I2 b1 }$ [3 V """. c1 Y/ V3 O% e- Q
model_ft = models.desenet121(pretrained = use_pretrained)+ h" P% g3 r& d5 t1 S
set_parameter_requires_grad(model_ft, feature_extract) E' A% Q4 d# i& y1 h
num_frts = model_ft.classifier.in_features/ i0 ]" Y1 `6 ?3 l
model_ft.classifier = nn.Linear(num_frts, num_classes)
0 k7 b' ^3 Q3 o# R4 v3 G input_size = 224
4 X+ D: c$ H. G, @1 t% E; O7 D. T
elif model_name == "inception":: J- p, l5 U' h1 Q; E6 @/ ?
"""0 T- D0 L: G2 m0 r% e& I
Inception V34 }" J5 Q* Z6 a9 I P
"""% d) u5 @( h7 @2 w/ c! Y. I) Z
model_ft = models.inception_V(pretrained = use_pretrained)
, k, B# n b9 A( M: x q set_parameter_requires_grad(model_ft, feature_extract)
6 e! v) z8 k6 K9 S: @% }/ E7 e
! J- |; o- a* r; |3 x+ H: ^& H num_frts = model_ft.AuxLogits.fc.in_features- ~( M7 ^8 i B0 {& K9 C
model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)
5 E# \ A$ e& f k( |
, e4 K2 @5 O5 E; S num_frts = model_ft.fc.in_features
1 q# ]. Q4 R- @6 O& H( Y, B model_ft.fc = nn.Linear(num_frts, num_classes)7 C3 E: ^3 d2 E6 d. g. u- v
input_size = 299/ r, T3 h$ r* A3 Y c" F
Y3 i- e; ~3 ?- Z$ n+ x+ Z% S
else:& S `3 d2 Z" x# B u+ U: t" j( B
print("Invalid model name, exiting...")
7 i* ]1 c& r3 k! y% B# Q exit()
6 M0 ^- w4 Y7 L/ z; M9 G* }' B$ X1 \: Y0 q- P" S: N
return model_ft, input_size
1 j6 w. n" M5 s# f" x- `3 a9 m( _. e* o$ L, F" P
1& F. N: U) J4 \; I8 [% ~! h
2! @' u X5 e6 p1 S# w I% X
34 Q3 O# M- W; J) g" {
43 d# Q- w, q" Y* D/ u
5
6 B" Z8 z% R" u$ c3 N& t6 E* k66 l L, e2 n+ k( D6 ~+ l* W
72 X7 r* A! A! ]# {9 X4 F3 K
8, m# u3 p$ D/ N1 A' ]4 y% v
91 w! r4 ?2 y S9 U2 r8 m6 [- {
10/ ~; G0 j8 C. L
117 H1 s. V6 k# ~+ N* F: Z& R
12
" L+ c8 | S6 ]. \13+ q% @; ^: E5 S6 a3 S' Q
140 v8 Z t) {' a3 y- X7 X+ C
15
7 @2 S A; b1 V7 L16
7 I4 |9 w4 J. ~7 t# E8 @5 q17
: E3 G! q# Y+ @18) Y5 G, V9 `2 F) d" e% z
19* B: F4 f) r/ s( q
20$ {9 \; I( c3 d) p% M
219 ^& G$ ~! o6 W8 e! P
22
- P0 m- {& U5 d/ {' X: ?23" {" b, w9 d, o; E! U' m
24
5 N+ H- E6 _6 B3 R) g7 s/ T6 O25
9 Z0 u$ }1 P* l; l" H, X2 f `26! `* Z3 M |* T2 o0 u6 b( b, r
27
% R7 Q. ^4 w: |' Z28! z: l3 ?& l5 L* H3 u
29
* N4 U* f: k6 l* l, f300 _5 W5 m. H* N9 K
31
% [; \, E2 B& L9 t; ~! |4 q8 e32; l0 E0 g2 }) ^/ R6 I8 {
33" U# K' ~& O d
349 r( t, X$ g' x# ]5 g9 _
35
5 Y0 e! ` v, x$ c36! p5 |, q, [+ ~& }
378 _0 @) h: J8 P2 E1 o8 C/ L
380 q8 S, ~ [+ p3 i4 c
399 x, ^1 s) T5 j% k: F
40
/ K9 w3 d2 `# t1 c$ o4 K414 \ ^! Y7 K# l( [
42- U* H4 I1 s" ^( j' J% c
43
9 _9 i$ k- f9 a ^44! ^; }7 r5 }/ s3 b- M, g4 R" ?7 \
45
' O( C0 Q" d) ~3 }46
% V( O* }- W9 a3 Q, `+ _47* s1 u# F" ^1 E2 b4 H
48
5 M* T4 b* |3 S6 p) [, z49 L' ~+ k% f' j) U
50
* Z6 i. \3 h! x5 R- d- G& W9 `51
* Y& v; O( s9 M$ @6 }52
& B7 w* |7 q- X( x3 N53
+ d, v: R( N8 ?8 s54' e- ?: W, H8 |; a; C3 F$ P- s
556 G; }! Q; \ \) z8 I" ?
56
: h0 u# t! H) l; [576 p2 U& T( R( E- {# o' k: H
589 I, W9 k) ]% {, C* w
59
+ F2 C# z% ^7 n! Z1 ~! H60
' r4 d- {) |6 V) ?" S( p& P61
( ~7 O$ o0 v7 S4 l" C/ Y1 x621 _/ N% D3 l' q
63
+ e1 Q5 j# e: \) e64* P+ _" O3 w2 e$ C$ _" m5 d6 y
65
( q3 }/ T/ U! v3 Y66( l' H. s* O& @- u! t( e; N
67
) i7 q' C; J( e g6 C+ U$ ~68) B: ^" c; v3 a1 Y! x4 \& f
69
) I$ i$ O' l4 n @70
- O2 o, z/ @2 |3 ^' x) v, {8 c% S71
( L2 f% z7 ^: g2 o# o72' x5 Q: u. g7 v* X4 T
73" T' p& e4 v0 ]* c3 j6 v
74: i; \5 |7 o& _3 `% E* i: G. X
75
9 j. n: w6 M3 |9 Y76! ~% @1 B8 ]6 T0 }0 Z, {: p0 z
77$ J0 b) P4 W4 P3 V) h' z% z
78# ]6 o- R( b; F
795 R# O4 ^% F& c) g6 g0 x. k
80
8 h' N1 \' r- I3 q1 F# r81
; l+ Z2 A. @: p825 a- Q$ W# _: n, d; S" y0 N4 D
83 r/ G6 {! w) k$ X/ k) K+ _
7. 设置需要训练的参数6 G+ p) p, b {3 s
# 设置模型名字、输出分类数/ A; ^% f* }8 |, a0 a1 O' J
model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)
6 k8 x2 W q. t4 ?' K5 ]+ K& P' |7 i# P; @. I
# GPU 计算
- y, E o9 e5 l' o! O$ @/ |model_ft = model_ft.to(device)$ e( k& t6 \) {
" E% r" i4 Q( s+ B; x! b
# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取
$ `) X5 ^) d: U pfilename = 'checkpoint.pth'
3 |0 Y# i" l5 F, W! C$ A }: }
% m# t$ G8 v7 U f# 是否训练所有层
' D$ H! c2 L$ jparams_to_update = model_ft.parameters()
, } J6 i4 g8 [( [2 _& Y# 打印出需要训练的层, A0 W6 H, _$ F" W- \: W
print("Params to learn:")
$ j) ]. u' m2 C( E! k6 g5 F5 Mif feature_extract:; d( V$ C# \; f4 E. C4 r" X' V
params_to_update = []
+ T5 p8 E8 @8 T for name, param in model_ft.named_parameters():
. T; A, }: ~# T- L6 O if param.requires_grad == True:; Z+ P- ]! r7 l) z5 x# D2 k( l
params_to_update.append(param)9 \3 C. l5 P1 x
print("\t", name)5 d* W! |; }; b! e+ [: m p @! {
else:
+ x E/ k. ^1 f for name, param in model_ft.named_parameters():/ Y) b/ ]/ L# v2 f; C/ S
if param.requires_grad ==True:$ @4 T8 J! }" n" p
print("\t", name)
4 \, x0 i) c" Q0 I. B1 y7 W! W- @8 R8 Q6 O+ M W, i$ r' S
1
' ]7 N( _) p% L& P/ Z, x2 v2 B27 v3 S7 N; ]/ T+ O# _% P( `. R
3
: ~9 S/ [9 l$ ~7 K& _9 a% I4
& k) G+ U$ a/ j6 m0 q0 A0 O5 p% p, Z5$ I! B5 h' j$ n9 J
6
8 S" r& J! w. i7: p7 \% x% C) N e" K! A
8% \4 C0 @1 b$ g% y" N- Q
9 f" m8 f" y' t' A$ X$ Q
10
% y) r8 F) m& g11
9 s0 K I) o$ f! f, f9 c12, ?* Q; d6 B* ?- w
13# C. |, u( W: E7 r7 i
14
" j$ p6 @/ O) N" k6 p3 r1 k155 f% d0 M8 z) s7 J% o
161 i v+ K& v; Z0 J" G) o
17
& D7 I5 P7 N5 M' o18. X w! W; r" H) b
19
- B1 C7 G9 T5 R4 a. A6 ] ?+ ~; _20
* `# @# r1 w5 ?216 b# u; }0 w3 w" ]6 }" T
22
" V+ J0 K. f9 ] k. t235 j/ j7 z' [8 _7 ^1 D+ u/ W2 ^- R
Params to learn:
7 V7 q6 B/ n7 z$ O& L% [) h fc.0.weight* y, J' B A; I
fc.0.bias: [7 x3 I) m+ r, K
1$ g- d" z/ V6 {. y
2
2 T w1 {( L$ `2 E3
e8 G( G8 z& H0 C3 b7. 训练与预测
8 j. u) d; C1 m! [" p/ o% k7.1 优化器设置+ Y1 F8 P0 U; z3 w: G/ a. G
# 优化器设置" X3 @ r3 D' \" b: } H, I$ R/ P' d
optimizer_ft = optim.Adam(params_to_update, lr = 1e-2)
- m9 }5 _% l1 F6 _/ W0 @# 学习率衰减策略( v0 z4 S# G# J0 a% ~3 l
scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)/ O" v5 {, h* x) q+ f
# 学习率每7个epoch衰减为原来的1/10
; L5 @/ F0 K @6 E$ n, P9 ~# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算) ]# d( A- j+ Q+ D
5 W5 {7 c8 H X ]' k1 u% I D
criterion = nn.NLLLoss()
: j) G1 h% N: u9 A, Y; W; n% h7 }1
7 E# |) e! F& c" h( d2
6 m% A) Q$ S- ~0 g3
2 Y# }% L _, j5 u, c# h0 M6 D k4
) k# u8 k2 E) _1 I# V, q5
9 p- U) B# s4 a66 \" o# n3 ]& p$ }6 f. i* @( A
7
% J1 c4 d8 C4 k86 G5 x: r& k" {: k4 n& ^# f
# 定义训练函数8 H: d+ B$ a/ ~& K" Y! Y
#is_inception:要不要用其他的网络
- M/ r* Y5 A0 E. K% T, Gdef train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):3 L5 c( p* S8 {6 u
since = time.time()
. @/ I4 @( p) `$ T #保存最好的准确率: |" Y h9 \# U6 V2 P! B/ h# l3 d
best_acc = 0! w( {* L/ o T$ u
"""
1 X# f5 S! K7 H, c3 j7 Q; Z, Z0 _ checkpoint = torch.load(filename)% T7 b9 V1 b% A0 g+ N$ z
best_acc = checkpoint['best_acc']( z7 N7 G E) g8 Y
model.load_state_dict(checkpoint['state_dict'])& ~) m5 b8 J+ b( S
optimizer.load_state_dict(checkpoint['optimizer'])- C! I% I0 I; m: h3 W$ Q8 {. X
model.class_to_idx = checkpoint['mapping']& V1 ?$ I! ^) {3 i
"""
7 z5 X( R {8 K3 V% L7 T #指定用GPU还是CPU# {( x0 k+ L1 Y0 w' x" D2 p
model.to(device)' ?! M6 G% j6 \; f
#下面是为展示做的
7 x; s5 Y; Z' T6 u# B8 m val_acc_history = []* h1 ?0 b. y$ n- ~ M2 ^: O( A
train_acc_history = []
/ P- \" a7 n; y# @* ]' h0 { train_losses = []7 o% M) q2 X5 T
valid_losses = []7 j( [' I2 e2 k3 i
LRs = [optimizer.param_groups[0]['lr']]
) ~9 K- b b3 u1 a #最好的一次存下来& Q% `. h4 K+ h& }* b; x& `
best_model_wts = copy.deepcopy(model.state_dict())$ O/ Y+ N5 q* }
" H$ x4 s7 u: `( S7 {0 F
for epoch in range(num_epochs):
$ d; L/ O& g9 V& E" E print('Epoch {}/{}'.format(epoch, num_epochs - 1))/ D" \6 R; I9 c) P9 ~
print('-' * 10); F9 w/ |- [, k# ?, |3 M9 E/ h+ V
3 z! a3 w2 h% `$ h( c& i
# 训练和验证# P5 k3 L- ]& B3 K& w3 m0 C
for phase in ['train', 'valid']:
. `; ]: ]; b6 Q# c, h: e7 G. O if phase == 'train':( c; R: b( t6 H7 v+ A) H/ |$ `2 \
model.train() # 训练
) W1 g: b: @) |/ }- _5 i else:7 G. M) D1 `2 o; v6 r: [) P# M; J) X- b
model.eval() # 验证3 F# x6 Q0 g* F9 J7 M1 o
! w) B) a$ O. m5 d# L0 j running_loss = 0.0& j9 s4 o0 }1 M1 c
running_corrects = 0
- v' P( l/ R6 Z. e+ A3 D' b& \+ l% W
- ^1 c. |! c$ B* i9 }# X5 U5 x # 把数据都取个遍
6 k+ P7 g; w1 b# Z- x for inputs, labels in dataloaders[phase]:
! b0 G8 `6 B8 E. Q5 }, p% V( { #下面是将inputs,labels传到GPU
" \. p$ r8 v' ~ inputs = inputs.to(device)
& J0 Y- ~1 `8 O) m. O labels = labels.to(device)7 G: s6 y5 Y* ?9 T
6 d) [7 x8 b; `6 r+ `/ W/ x # 清零, S% J2 S2 b7 a( J
optimizer.zero_grad()# Y8 `9 U6 _8 w' d; }
# 只有训练的时候计算和更新梯度
! b j- p$ u& ^3 i, n& v7 X+ [ with torch.set_grad_enabled(phase == 'train'):3 l q' b, P& }
#if这面不需要计算,可忽略
+ o9 f& S+ y1 _ R* n if is_inception and phase == 'train':' b6 J. t$ ^9 S+ {! M7 G$ \
outputs, aux_outputs = model(inputs)
) m r1 D* O1 t, x! _6 r! v( y+ \8 { loss1 = criterion(outputs, labels)
: b. U% J8 O+ d# y, G' i9 ^1 r& H loss2 = criterion(aux_outputs, labels)+ g4 q5 i( E: h' X0 Z: m
loss = loss1 + 0.4*loss2, }, {; U ]# p; S/ Q1 u
else:#resnet执行的是这里4 p1 [2 h% e6 P- |
outputs = model(inputs)
3 N8 P) d: c& g7 p loss = criterion(outputs, labels)
& t1 d/ V$ r& G2 u7 ]/ K' f1 z1 t7 B8 D2 o( C/ }, q2 [2 T
#概率最大的返回preds: f) X& L+ c' [1 F3 O, f* {" Z3 F
_, preds = torch.max(outputs, 1)- b- |/ h7 Q2 l
; o3 m' M+ i: G8 F
# 训练阶段更新权重# y* A6 n. R- u$ Y- W- r
if phase == 'train':
+ M; @! n8 F' ^4 U) k loss.backward()
; B1 _, K3 o! ]4 b- ~ optimizer.step()
- K5 [- n0 e' m; u
- A; o3 O X$ l `4 C+ l# F # 计算损失. s7 b f+ f" f. q4 H
running_loss += loss.item() * inputs.size(0)4 I6 D7 F. c+ S3 s
running_corrects += torch.sum(preds == labels.data)
( P$ T* r3 v2 Y. z0 K7 [5 v# v& e! V
2 {6 h3 J; ]* b$ K- R9 K #打印操作
, ^& }# Q0 t4 A! a H epoch_loss = running_loss / len(dataloaders[phase].dataset)5 o3 ]4 n* _# k- b
epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)
3 \ t+ ~( R6 Z% m* u7 x0 R; j
3 E. z/ N1 x/ o8 V/ B2 t7 H
7 q m1 q. K! o% Y1 U. i. _6 d time_elapsed = time.time() - since/ `& ]+ t- C2 ~
print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))7 n- H1 u1 n* r
print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))2 a0 P, ]" C$ x4 F0 U0 X
0 |5 m6 ~$ |6 z$ }. J% \9 ?
4 o# ^8 [( @/ \8 m" q( ~ L # 得到最好那次的模型1 Z' B- f1 J8 l' z8 f6 d) k/ f
if phase == 'valid' and epoch_acc > best_acc:
1 j/ U1 t; M3 p" m& A8 p best_acc = epoch_acc: R: a9 G$ t6 o+ s1 n8 Q: A" N" d: s
#模型保存
5 E6 s9 ~8 h2 \! ^7 z4 c1 q best_model_wts = copy.deepcopy(model.state_dict()), N: E6 T; w$ a$ }8 W8 j# w8 a: O
state = {9 A6 ~( t& w9 v6 e. x
#tate_dict变量存放训练过程中需要学习的权重和偏执系数
3 i8 K S7 J5 w. [2 Z+ J1 I; \! \ 'state_dict': model.state_dict(),
" [; U; V! Q2 [' D- o 'best_acc': best_acc,
4 _% v4 { V3 ~( T6 N 'optimizer' : optimizer.state_dict(),
' x, C5 c6 A8 h$ w/ j$ R: ^ }
& y4 r+ A5 Y# x3 k' {1 H# q+ v torch.save(state, filename)/ L2 v! `/ G9 T4 p% T
if phase == 'valid':9 U/ N' o: ?# I! F, a5 Y
val_acc_history.append(epoch_acc)
. T% H3 K$ A' T) _. m R* `4 D/ _ valid_losses.append(epoch_loss)1 I1 x0 c# w6 ^# f
scheduler.step(epoch_loss)
0 ?7 \* U2 h0 X4 S2 a- V+ k0 }/ _ if phase == 'train':
% G5 W# N9 n) Z) R8 s% h train_acc_history.append(epoch_acc)5 K7 R( z# S8 [+ }+ Z, t0 R3 A
train_losses.append(epoch_loss)9 C$ Q' i6 D% f. G& p0 q) h
6 c/ v! m' i: B
print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))
+ H% ?3 K3 ^, y LRs.append(optimizer.param_groups[0]['lr'])
- X- o; w7 f$ W print()3 ]9 e) h! a& D% M/ D3 [9 ?' K
" R6 Z. l. }+ S5 o5 l% S
time_elapsed = time.time() - since
- \& J9 s8 D0 O; [ print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
7 d! d3 j/ \2 y$ p9 m+ A print('Best val Acc: {:4f}'.format(best_acc))
5 m* ]# G/ C/ J. S
: C% V# S* N1 @8 o" I% Y1 {% @ # 保存训练完后用最好的一次当做模型最终的结果# \( |2 A, O1 }, W% s, C1 Q
model.load_state_dict(best_model_wts)- B! m+ p0 J+ z$ y2 T! H* B
return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs ; q4 l! f& U. D9 e; b% O& a( d7 l
% i& I9 @9 o: ~4 C) b: X( W) s* `: Q) o7 s$ G& ^
1
2 b* `; z: f: [, O: s5 s8 Y( H6 k2
0 }) G; J4 n4 A' d36 G3 t7 ?# _" Q& o- v, b
4
3 E; N( ]; B4 r, N: Q3 d. s C$ y5
! Z+ B/ C% J6 _% U7 Z6
' b! X4 p( h' H/ K) ^! H7/ Q O! k9 j. E6 j: w" B5 N
84 O/ B3 h. L3 ]2 A2 Y+ x$ D
9
; }- q8 E' h6 y+ k) G+ A10
1 u8 o2 B; o! V B9 a: M/ P0 h11
) x1 f4 U6 B+ Z* `' ?+ X+ Q2 V j12 L5 M: q3 I! Y0 @) q+ u/ v
136 S. S L* u: N) B5 _; p9 c0 L
14" R Y, R+ O4 ^' c9 h
15
8 {. t% ~4 d) Z% k- }( U16
5 _* x. ^7 F- ]. b4 d7 p" R! B0 D% g17
( s! m0 E9 X1 y2 \/ p$ B0 U2 D" w18
0 l9 d, d4 @& U# m$ D19. z8 q0 ?& f" q% y" O
20* r( Z! y8 S$ C6 y
21
/ f- _ w- U5 U% W22
' ^; {) R8 l$ W; A) v23
7 v6 M! g; S) t24; s0 m5 z/ q, ^+ o, |+ |5 E. T
256 x, T% v* [# B: m& A; l
26) B/ Y8 y/ @; p/ W% Q/ i4 ~, E' O- J) U
27, n; M8 U8 F: ^+ x6 f& T9 P, K$ A
28
: ?. f) Y% G( L' U h0 O294 w0 g ^7 b% b' q6 p1 d% P' q
30% G: ~& e' j; ?" A4 w
31
% c. Q! D! i) {32
" q% n8 |: R+ H$ B2 b. d8 L, K7 W/ u33
: |" \1 L3 N+ S/ x+ x3 P34
# Q; P# @0 U0 x( u35
B& D# R8 q: H& A9 T) V% i& T; }366 b* t* C& P# Q- e/ l% r7 i
37
* O4 H1 b3 p7 E! o4 ~5 H383 D' `# b( p) M a: ?5 D
39
' x; M7 f' |- g. U$ u' q. @. }40
" s7 q# ]7 @5 b* v- R41$ s6 }0 L0 d+ X/ a z" R, R7 h" l5 {
42
4 `( ~. I, g! ]* @9 W433 H0 M+ \; u+ p& Q0 G9 }1 A# P
44
0 s: l( W6 ~) i4 e" X454 v, S8 L, {7 T% f* R4 N" w
46
, i3 c7 @/ O7 F47& i( D; G' W- }; ] j
48( f' p' |' x; M* s) H" r
492 x; D( G: u3 v4 R
50+ i9 N( |6 Q8 R8 Q
513 z2 Z7 e e. v' Z. }$ J' x* r
52. N4 U/ x6 i0 h$ D. g9 D
53
9 @) A) s9 O2 V3 s+ c54
j& O9 B5 A; b6 N556 _ t+ I- h1 d5 P# F9 ?
56
0 t8 I( r* N4 r/ M" ~/ O7 _57
7 b2 O$ d6 i: D0 V7 v @2 H0 }% \58
2 A% ~7 ]0 Z5 t# o59* @/ C' m/ U# ^& `
60
" q) C6 W4 C0 h- D; [61; G( u5 K. L! _+ F
62
% ~- _ H3 y6 a; b6 Z. c) O63
& t% T3 a+ \$ k. o/ G64
! k1 D; y" G6 u3 r8 V65
' U( ?. D# D* U. @ @66( z+ `) b4 D/ ?* }
67
0 W6 [6 [* n1 g7 F68) ?" f4 O/ n. ~0 |
690 H% _% }- Y2 F5 j s8 }
70& F- O, Z- F- B) n `, n
71! \8 T& b% P: ?
72
+ E1 P+ T9 a0 F, R3 c73
4 R- i8 [( `( B* y743 o# l1 {! }9 W5 \$ n
75
6 ~' r6 d& c7 @/ [) t6 b76
^3 ^3 E! c. E4 v4 s77
$ g9 r6 t" R6 M. q9 J# r78; k5 |; A4 e4 L( f/ R
79) Q' b, J* f- {9 @0 p
80
4 i/ N$ `7 K) U+ v4 W81
; M* H% j% O) ~7 _828 J; t& o- l/ f, O- i4 {% |* G) K. I
83
& \: f2 f( V9 s" H {, d% u& j847 ?2 z. n8 v; O. e" G: b0 S3 @
85: L" q. `- V6 X1 O
86% i/ _; ^6 q! h; Y# J; O
87 s( X0 }2 l C: \) R6 Y7 u( b7 a
881 ]: p- a. t+ q7 [* y8 R5 G/ m
890 o( O; E {; x' Z/ p6 H1 u7 e
90 t! f4 l, S: ~% e: _/ _6 V
91
+ y, d' ~6 p( r% }* E; u92% E, G D2 P, X/ Z( ~# M
934 G3 }- v- U( K& U/ G; a
94. B! N$ N/ S) o9 w! s" a
95# G; ~/ A3 X' @4 J9 P7 \1 N
966 E# p% U# @* ^
97
& _8 g, l- Z$ u5 Q0 G98
! x+ w) s; N+ o: i1 y8 S3 Y99
' |( P; Z+ f$ ^+ e) p! T' L100
" z! j- W8 _* Y3 Z8 A101
' N% l4 K7 x0 x- @102
0 i# P6 X3 _/ ~, \) ]1037 o4 v5 B- R/ V9 }; I3 t+ P* f! {
104
; s! n$ `" W0 _- k1 c105; b' ~% R. I! d$ l9 D3 Y% A
106
2 A- D+ r. {' V2 q) `& C107
" C- {+ r) Z, L% {( k1088 L* I; b3 C0 l
109
9 h4 x" s; `" W) }110. c5 `: @0 z0 ?* a
111
& @6 P/ O# _( v7 r' `& a1124 g4 o8 m3 X7 k. @8 z) B! V
7.2 开始训练模型5 ~4 @& y1 h- c' R
我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次
$ u& }4 y$ J0 z0 Y* U- J6 X7 J4 Q2 D4 n- S3 w. `
#若太慢,把epoch调低,迭代50次可能好些
1 {+ V& d1 z6 a# i9 m#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合" B6 F' w( ]# e4 p& i [% R- R
model_ft, val_acc_history, train_acc_history, valid_losses, train_losses, LRs = train_model(model_ft, dataloaders, criterion, optimizer_ft, num_epochs=5, is_inception=(model_name=="inception"))3 R# d9 x, n- F" Q" m2 w
h9 W' b4 o' F; l- F: D10 P: u6 k. E8 v
2
3 p/ L0 [. |9 d( P3
_0 K* f( _" ]( c4- g7 s `$ R. i: h- o% Z5 k
Epoch 0/42 |$ ^6 `0 ?. X0 l
----------
J) J3 c# b& c$ cTime elapsed 29m 41s1 {/ @% s0 P$ n% o" u( g
train Loss: 10.4774 Acc: 0.3147* y1 i5 c8 ~* }: B% c" ?9 _3 y4 N7 O
Time elapsed 32m 54s8 J, X* O9 n2 ^" m# n# f
valid Loss: 8.2902 Acc: 0.4719
2 P/ E) `; l2 X1 g$ z d2 Q: ~2 DOptimizer learning rate : 0.0010000; D4 F. j% d! a! D; ^8 J
7 y' p4 _ B. n7 f, N% o0 q7 e" b6 z) v
Epoch 1/4
$ [0 |* H# d! O1 G0 e5 O----------
) {! j) R `& v; pTime elapsed 60m 11s
, @" i0 `, ^7 C- Dtrain Loss: 2.3126 Acc: 0.70539 O9 j, z0 f7 b4 e
Time elapsed 63m 16s* A3 h y+ c5 l0 N% m
valid Loss: 3.2325 Acc: 0.66264 Z, n! s% V8 g( y9 J
Optimizer learning rate : 0.01000007 i4 I" b( M* c8 A' [
7 C0 {$ s/ C h. O# fEpoch 2/4
A8 A+ _& y& q3 h. m----------
$ {& m0 R. T8 MTime elapsed 90m 58s/ b2 C* U/ e8 \: u; A1 p, q. L
train Loss: 9.9720 Acc: 0.4734
1 o2 @' ]0 q. j; M2 E" n" PTime elapsed 94m 4s
- M2 E' ?8 u* X3 I% j* `" {( F3 m6 Xvalid Loss: 14.0426 Acc: 0.4413* Y* r0 }" W1 d, z! J2 i
Optimizer learning rate : 0.00010003 E2 H% ^. e4 ]1 L
7 m% q+ G% q- W: Q3 H% {& UEpoch 3/42 m" D, Y! h& P2 Z! w# w
----------1 r4 q0 e4 g+ a9 H" `6 r& c
Time elapsed 132m 49s" }; y: Y6 X( O; ^+ V4 E( V, b
train Loss: 5.4290 Acc: 0.6548
: v! \. N5 o: \2 g8 }Time elapsed 138m 49s
% W/ E: N1 J/ X& xvalid Loss: 6.4208 Acc: 0.6027
7 I3 Y- m' e8 ?! }- ?Optimizer learning rate : 0.0100000- n; F9 J. B: h
0 {) B' p5 d* ?4 |, D
Epoch 4/4* N8 p1 f2 J. _
----------7 d* A( F; x, m6 C7 Z" ~# x8 U* @+ ^
Time elapsed 195m 56s1 O9 @ p2 h! k3 m
train Loss: 8.8911 Acc: 0.55195 A( t9 G0 I5 a4 a' D% x2 a2 M3 D% `
Time elapsed 199m 16s
, x: X6 b S) A$ P9 B; Cvalid Loss: 13.2221 Acc: 0.4914% j3 C6 x2 g# ?. O7 e. |
Optimizer learning rate : 0.0010000
1 t4 y# h- Z2 U, H( Z( H
`# a( l- a0 L5 y5 Y# ZTraining complete in 199m 16s7 c2 c- @$ q- T+ g" [! K; o+ R
Best val Acc: 0.662592( z. E& d( Y2 {( L0 Z1 A3 _0 S. s
+ R& ^" H5 {/ _0 T" `1 l
1
' h X* Y. e9 z2
5 H& h' K9 o! [* G* @34 f: C' x- w7 @* u& L+ c) M. b; x
4
/ J- S* T. s2 ^+ k1 m51 R) d) F; U& f0 Z/ I
6
U- ^9 }) E, }% a4 k' O4 Y7
4 v; W' X6 Z5 e9 A% T7 M9 I1 Q8
. Z" P& a* _! g6 d9' |& H- k, A' g
10* U1 |' R. J7 v1 R/ a+ Y& g$ P
11
P! q6 X- u( K: E( C, o12
. P" `- Y% g2 D# W" `0 j, w13" u" {9 \8 C: r2 W2 [/ N
14
9 n2 ` U! L8 I/ t0 X5 K8 N" o. q15* _5 p* e6 l& f6 S4 v% e
162 @. J( W7 q8 u3 D! Z+ O0 b
17) r% E$ D; Q' Z4 n8 X3 o
18
& K+ ]% q6 g" t4 k) F, B( Z, v19: Z& w- `: l4 b+ j) D1 k& K* i
20
. c7 K/ `- Z# {' z7 c& Q0 X21
5 ]1 z+ o2 n5 z22
9 e0 I- G8 a3 a23
8 o3 I% w3 s; L3 o+ @# R24/ A8 h# I0 U+ z3 O+ E9 f6 l: H+ S
25; i3 S6 c s: ]7 s; k
26- q4 `: l1 G; b4 ~
27
0 [9 T# _( F) a& ]$ H288 x+ k" h5 v5 q% y
29' v2 T3 t$ H+ v4 y! t) ]! X8 k( p
30! k' x J6 t, C1 ^3 ~5 u2 R9 W
310 i7 z3 y; l' M# j
32
* v" ^: t4 s& D/ p2 G33
9 |# Y+ Z: {( ^2 v342 |0 S6 x+ z! x% t0 U
35
- s; T/ r, C, P0 j5 k9 w& ~1 U368 m1 ]+ i2 O# m. _/ v5 K+ ~5 R
37
3 S, a/ d6 e2 A9 L' Z1 I38
& K) r/ R1 n) x39 N0 [8 s% i a+ `
407 `- _( X- ^0 f1 [* ^$ _
41
# y; S. e9 z( [9 I42; {) F" d# m- a+ p+ A+ q2 ~
7.3 训练所有层4 c9 N! U2 B+ `6 N$ q5 e
# 将全部网络解锁进行训练% ?: }4 \/ Z/ _7 a- N& |; g
for param in model_ft.parameters():
# q" \* G9 Y/ M' n param.requires_grad = True' H' c5 F4 _" U: u2 H1 ^
( }5 p+ }4 Z7 O5 q
# 再继续训练所有的参数,学习率调小一点\: h+ [- \; m$ d, |0 G( M
optimizer = optim.Adam(params_to_update, lr = 1e-4)
6 Y% _6 c" p9 k( o8 `1 J, e' S; Dscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)! X* _7 v" r' U+ c1 T
5 i! I. s' [9 s1 @. {+ j
# 损失函数
' ~' w- s o& K$ P4 ^criterion = nn.NLLLoss()) s5 D; Y$ q: ?) o9 _* g
1
: V7 V0 L9 O. n; n# L/ t; e: Z2; @0 L- U3 e$ }9 p
32 W- S% R# Y. j( p; ^: |
4
. K x' m+ N' M+ |: S9 n4 V; O5/ ?* Z9 y6 h/ ~0 Z) y5 |: b% K1 d
6* ]7 |1 ~7 X( A6 Z- w# Q4 H
7
2 v4 _4 ^/ J- O0 v8 P8
5 S4 ~3 N$ H! A- q9 p2 L" t! l! e9" F; G9 s/ A; A5 H/ M/ W; Q- K) K3 b
103 H; w( s5 W3 t& v m( m# O
# 加载保存的参数
9 Q: b: |' B5 Y1 l* u# 并在原有的模型基础上继续训练- M4 T# [5 u! R# M
# 下面保存的是刚刚训练效果较好的路径) s) W/ k% r/ |
checkpoint = torch.load(filename)5 A0 H6 v- d2 ^, y) a* d8 u! {7 j
best_acc = checkpoint['best_acc']
, h7 g+ }/ g5 M9 l! Z: Imodel_ft.load_state_dict(checkpoint['state_dict'])
; t! \6 ]3 P5 {1 b8 r8 Aoptimizer.load_state_dict(checkpoint['optimizer'])' n/ a; I, p7 z/ v) y
1
$ G$ R+ ?; T9 `: o22 f) H! Z" H9 Q) a
3
# E* I( D' j+ S3 D6 k; ?4) @& P& `' m& x5 T% [
5
, m& @! O, [+ E! O$ V6" I( f' R1 G+ \+ ?3 ~, d
7+ }1 g7 C2 |6 ]6 n; @7 P
开始训练( ?# l/ u! p0 ^( M) N
注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考
( p* ]- n6 y2 T2 }9 n4 p; o- v) \( {
model_ft, val_acc_history, train_acc_history, valid_losses, train_losses, LRs = train_model(model_ft, dataloaders, criterion, optimizer, num_epochs=2, is_inception=(model_name=="inception"))0 J' n) [; J) g! m
1
+ A! \1 e+ {; f. f' JEpoch 0/1! r8 T- f& ~* `0 l q- D1 H
----------
7 _$ m% s2 i9 ^5 N kTime elapsed 35m 22s
! M1 U ?8 f' ^! Itrain Loss: 1.7636 Acc: 0.73464 V2 P" x! p2 t; @: M- t
Time elapsed 38m 42s/ a6 C* W6 O- h) X
valid Loss: 3.6377 Acc: 0.6455; x6 i9 T% y3 }& u0 j7 {
Optimizer learning rate : 0.0010000& Z, U0 Y3 S: n0 h) O3 [
! K/ f! s3 F. Q) o5 s1 g; EEpoch 1/1
# x! M! j3 Y& P N) l2 U----------! ~( y! U8 b: U8 V
Time elapsed 82m 59s- l7 i' u D" z% P! O
train Loss: 1.7543 Acc: 0.7340
1 g* y9 C6 h" Q" h1 m! t" ITime elapsed 86m 11s
5 A! @; O$ F3 u* M& Bvalid Loss: 3.8275 Acc: 0.6137' N% ~7 ]3 b- m4 i7 f8 n
Optimizer learning rate : 0.0010000% A+ }: r; M+ r, h9 K ]
: U p/ }9 l$ K s: rTraining complete in 86m 11s
' ]% M3 y2 L# u& yBest val Acc: 0.645477
: y1 n$ t# J, R3 N ?
3 e, @+ k2 N9 o, V/ V7 K1
i/ s$ X# I" x/ r2# l4 I" s d4 o7 Y
3
9 e* B* c8 v8 x4 z9 n5 ~4
" n( b, n2 H3 p7 L' w3 E6 G( @5
/ F" \: a" J4 }6 j6- M! H+ I; R: S5 U
75 X. A4 D2 D7 G. E& J
8( b% B) l" A5 o2 i }
9
+ W2 R/ c, I, @102 Z' W1 F$ o' M8 U% B: b/ @
114 m& g: z. h. W* O
12) f* Z, c% \2 E5 _
133 W2 }8 X5 r3 J1 l0 p
14
# }5 I7 g' ]% b8 L) r5 E+ {15+ O' Z( g$ y, x& }
166 a9 M0 ^% V3 i9 e: I0 X; u$ k
17
+ C4 P& b. \; B2 s, H5 C% d5 h18! a- a4 ?$ Y, B2 y8 e P
8. 加载已经训练的模型
/ D: ?* c' b6 v E5 K( }相当于做一次简单的前向传播(逻辑推理),不用更新参数
! y) r* K$ O j: w- q
2 p/ T$ w% o3 k! i, a; h) fmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
# r9 l B; K# b9 m
, J: ?" ^7 S4 W0 h8 e; y7 B% e$ m# GPU 模式 `" a5 j% t) G% M( M* P0 T
model_ft = model_ft.to(device) # 扔到GPU中6 D* k* `9 c! |* t" ~
2 q0 |3 s+ D: X" j1 a# 保存文件的名字. @; Q k1 v5 @7 ?. V
filename='checkpoint.pth'
' W7 C3 s! _6 C' P$ K- W
+ F% I& n3 ]. R/ }6 U X# 加载模型' y. N6 O; z* a {
checkpoint = torch.load(filename)
* K8 {0 H* I8 P6 ibest_acc = checkpoint['best_acc']
0 }8 J; {4 Q7 F G t, F& G+ ~( _model_ft.load_state_dict(checkpoint['state_dict'])
j6 v7 n9 }6 \+ Y- [: v1# K# G$ r# s Y6 I2 B
2; X3 ]0 I* v Q i) l- D
36 {$ E- t( Q- \3 Y8 y# k
4, \9 ^8 {' ^8 J2 q/ B# u# l' G
59 M+ f# q7 E* Z2 j$ q) D2 T
6
& C- m' K0 I9 A9 n0 {7( q( h: t3 H/ l0 s
8% A/ m) G9 f3 w- l5 S' P3 w
9+ M1 \& W" T+ D
10
! x; W$ l0 }4 m% X2 Z! {11% n& B* ?- a5 Q( [. ^3 V7 p6 L0 u
120 Q3 A5 m% i" H1 U) b* S/ R
<All keys matched successfully>
. V! Z' T( k' j1 L1! \% l: ]+ B0 K* ?( a
def process_image(image_path):7 G3 C2 r, T$ C C, {& \. S
# 读取测试集数据
1 d) ]/ D/ @; @; y1 @1 K; y v; M img = Image.open(image_path)
- I# x; K3 a# v5 M6 {( L& V # Resize, thumbnail方法只能进行比例缩小,所以进行判断; _0 Y9 Z$ q$ X- Z
# 与Resize不同
- V( S+ F, Z h0 v1 j0 p4 Y$ A # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小3 p w% n3 H) g: @; X
# 而且对象调用方法会直接改变其大小,返回None$ F4 Y* n& ?3 d% `9 x1 ~
if img.size[0] > img.size[1]:
# Q4 x$ H0 j& p* Z" ]% L img.thumbnail((10000, 256))
2 c- q. L! S! l/ J, k else:' a) d L4 r, J7 ` K+ @( N7 ]& U
img.thumbnail((256, 10000))0 l4 e' U/ D: z( _" Q7 e3 X
8 n2 ?: W6 y3 A! _+ Q; Z
# crop操作, 将图像再次裁剪为 224 * 2248 W) ]( S( u' y( Z% {; {
left_margin = (img.width - 224) / 2 # 取中间的部分
7 D% n# [2 K4 s5 ~' B bottom_margin = (img.height - 224) / 2 , z* w$ W" X! y' I* [+ P! X5 Q
right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度8 a. q* `* z7 |. V& Q8 M7 i4 V/ }
top_margin = bottom_margin + 224
2 f3 V' D0 e: s. G B7 Y6 L H. D8 ?0 c0 C5 t" L9 }8 E
img = img.crop((left_margin, bottom_margin, right_margin, top_margin))8 i. [1 v/ W7 Z( ~% b# |
. A& [' m# P% Z
# 相同预处理的方法% Y) b. _/ A* i% Q; i- E2 \
# 归一化+ r& w8 H' F0 L' `7 ~/ n8 f
img = np.array(img) / 255 a2 j8 c; {+ M* _
mean = np.array([0.485, 0.456, 0.406])0 g+ F, O; p; v
std = np.array([0.229, 0.224, 0.225])
# z& O# N- }* b$ J! J img = (img - mean) / std- r8 m& u0 T5 N! [9 Z! @3 E
6 N; \( t) U! V+ n* y
# 注意颜色通道和位置+ ^2 v: y- F5 j0 D* x0 A
img = img.transpose((2, 0, 1))
0 K: I6 A. M. V' k
1 }& {# ]& ]9 z9 J return img
! i% \: F1 h0 O: w( o) {! a1 @; R8 P! I2 u" z8 w& s
def imshow(image, ax = None, title = None):+ V7 X V% y/ D, Z
"""展示数据"""
% j k# m+ W0 l: D9 I, b8 S if ax is None:' V& L! V0 Z" z: D9 T3 u
fig, ax = plt.subplots()
# g+ I! M/ L1 t* y1 L' E4 U! L3 ?1 D1 y+ z2 o3 n9 f( I
# 颜色通道进行还原1 v2 r6 R! ~- ~$ f# A3 T2 C
image = np.array(image).transpose((1, 2, 0))
5 |" |" b4 u+ x8 a) [( }
" Y8 K+ ?. [1 p% d! Q # 预处理还原; Y' m0 U& |8 A! p+ Y5 B* i
mean = np.array([0.485, 0.456, 0.406])" b) H- x' U$ _* ?3 T4 U
std = np.array([0.229, 0.224, 0.225])
: }8 d) ~7 p" \# u image = std * image + mean
- B' m) ~6 W) X" N! {- d4 @ image = np.clip(image, 0, 1)
. C2 J9 j; E d5 Y) s6 @ x2 F" x/ W/ y& l& O7 s2 F
ax.imshow(image)
- O ?8 m9 `( f% m& F2 T ax.set_title(title)+ |1 z+ B8 Q2 r0 S* k% f3 D
# a9 l8 s0 j- g6 u; {" Y# z! Z7 c return ax
1 T1 `: a; \1 B1 J7 s
# b* j4 _. z0 C" [$ O( R2 Bimage_path = r'./flower_data/valid/3/image_06621.jpg'# Q# I; D+ C$ [
img = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理
) w$ c$ y. T$ s$ H& ximshow(img). q$ h8 Y5 x0 M6 q3 ]1 t& u# v
* ?( t$ Y1 i! v' ?6 v7 _
1% H/ p/ J. q1 H0 f5 |( N$ ~
29 M- m9 Z: O5 F2 B
3
2 O' j, }. \# ]- M4 G: ]4+ T5 W% M" O" ?. o1 s# C) i
5
0 t. b9 T4 ^0 k" `3 N6
3 D% Q7 ~) u( ?7 q0 j7( v2 G7 _4 f% c
8
$ K# L! Q$ w; f8 ?5 F' [1 i. @9$ w# z+ V' @. D- |9 x7 o! ~( f% K
10; V# L. R& l; v8 a9 n# h: y
114 Z6 l* Q1 b" i- C8 Y; w9 N# X. F9 `
12
4 E* {# j! l7 B3 D+ H; F13
4 g. l! ~$ ?5 g% m1 `0 H8 H145 N, b5 y/ ^0 e* ]4 T* a- ]) I
15
$ u% l! L# G( Z! Z6 P$ z& K/ v16$ ~ C. O7 M$ i* M5 U! F$ l
17/ o1 {& W& |, l' l( ~
18
3 p1 m* V" e, B9 ?19* h" D6 T0 L3 ~. Y
20# g7 x) T1 C% R8 r* i% Z
21
( X9 k g- E' H4 h' {22
9 Z+ Q- y; O X$ }; I233 G: k$ G6 x9 Y
24* C6 V& u# h- X
25: y' G4 `: s) p; C. f2 p
26
9 o7 l/ d$ f' b27
8 q9 ^5 L) T$ |$ k7 E28
9 m- t" f9 T/ J; H+ a295 q! s" d/ |# U0 S# _( B8 s, Q
30( L- M5 D7 x: \3 T
31
/ m$ _: Z) U# b; G" H+ L+ W32 G2 i" y+ P9 r' o1 ]1 c* s' _$ \
33
+ z- R. S* g6 ^0 t% H$ h7 c7 U$ Y9 J1 q34
- m* v- d; P/ J" q# E4 K35! V# h! J/ H- A' R0 p* O7 F
36
0 M& n4 Q# @2 c8 r; A! K37" f+ C2 `, L. ]6 x/ \
38
4 T( y/ w* L& U* H/ o$ n- U% b0 k9 |8 F39- Z3 A+ ~, U/ x. a7 W0 ^
40( E8 z4 \5 o/ |1 m
41
. X7 B7 e8 h$ l5 e7 g6 q! O42: K& Y* d) k/ f( q; _) f c% }
43% @- ]- c( P$ R3 d8 S; u- C9 e
44
! p( \- H: G$ n% i, j2 q! \45
' J# W( P x* G1 u, \46
6 j3 x6 v2 n6 V1 p* u9 I! ^47) ?% k4 Y/ R1 l; E7 D
48- \, J2 G9 \+ G$ B7 }% v" H2 D
49
/ W5 h% c+ p& x; r+ t/ ?50
9 {; R6 P7 G# K( T51
( K9 \( L5 u0 A1 i% k* S525 n: X" N2 J1 W/ B- \
53
% V1 t! ^, n9 o54; g. P3 K1 b0 m: n8 z' q6 b$ c9 Z8 l
<AxesSubplot:>5 { Q% e/ i. U* E7 }
10 y I9 L8 ]5 }4 I/ k4 V( k7 s/ a
, q5 p/ n+ V! [
上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确 B& O5 \- \2 O4 v* ?
( l0 h7 _( g# [" R9 Aimg.shape1 M8 Y+ k1 n1 W# A$ H0 g
16 \; Z* D6 A' }8 R4 Q/ o
(3, 224, 224)
0 t7 F7 |& ~# }1
5 V1 D' R( o! b, U5 x4 F0 F证明了通道提前了,而且大小没改变( y( ~& B. E/ H6 a7 ^. {7 \2 T$ `
) G1 N+ Q' \! M# d6 {1 E9. 推理 e* R. s& |8 I, z. I/ c
img.shape% g) A+ j, k" O7 |/ S
+ \; h- A' `" A) N4 M# 得到一个batch的测试数据
7 Y2 g7 l$ V) N5 k0 k2 T/ Bdataiter = iter(dataloaders['valid'])' c9 F/ B* Z5 h. _; K* q( x
images, labels = dataiter.next()
; k7 I3 o8 Z- `1 L L" \- j! f: t+ q* _% H" B; p
model_ft.eval()
3 [5 v1 I& ?0 x7 I
, A. i0 m& @: Eif train_on_gpu:# h6 a0 i( M t: I$ q" [ n
# 前向传播跑一次会得到output
+ S! G' h3 S/ @; ?3 e6 ^& Z output = model_ft(images.cuda())0 x. Q0 B8 C ?% ^& E
else: t9 O$ c4 N9 @) x2 ~
output = model_ft(images)
, W) u- X% X# F' }% f$ ?( N6 H0 v! E5 R
# batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值
2 m5 K% U6 q1 C9 doutput.shape
6 ?! W* ^8 y! e; F R8 o V. j5 ?
1
" q# C: r- ?: F, n5 P$ i( O2
T8 c4 O! W& Q' I3
! Z- g" Z6 y# c& |9 D K1 x. x4
3 v0 s3 p/ q5 Q6 k: e( F( ~5' {5 k8 M [1 X, e, a3 P
6
+ @, N1 u0 G/ R& j! w7& n( v2 G5 G; Y1 Q; g
87 B' [9 ^ F) a7 @% A6 Z/ l8 F
95 _4 }' W% s% k- J6 s3 @4 g9 a
10: Y3 J* P; S# M2 Q8 ^3 |
11' T, ]1 a Y/ q" |7 c' I
12
: T: T6 B! Z$ s) Z1 F/ i- A0 U8 N136 v$ f% {$ H2 ~
14
9 @' B" T. w$ p) [; }6 i* K15 A4 F! T; ^+ d' p
16
Y- _5 H- L. Q) g8 [# }8 otorch.Size([8, 102])+ c# h" h A* x/ s; T" `+ t
1
" A& f" W7 ?; r: M9.1 计算得到最大概率
; a% X( I @* U x_, preds_tensor = torch.max(output, 1)( L F, ]# q8 H+ P0 t& w2 w6 @4 f
/ S6 M- M3 u: U1 l
preds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量
+ Q5 W4 p: E% F$ ^; q1
. r) E1 e0 p1 F- { ^$ R; x21 n- _7 q( L! _$ p
3
( N4 a/ ^2 k) _2 [6 z9.2 展示预测结果
4 F: g' ~; M9 M; \' Xfig = plt.figure(figsize = (20, 20))
# n8 q+ |- W5 u- ?. qcolumns = 4
! `# H- V6 H+ V- arows = 2
9 E: T# u; X! J' B7 e6 l7 N0 G
( T, z$ B* f6 j# h" a) o' P$ A: Jfor idx in range(columns * rows):
2 l6 v8 T' d: z4 A& }! M$ p8 v ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])5 K/ R7 T' [) U
plt.imshow(im_convert(images[idx])), n+ z& q7 J' O* {& K% E
ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]),
6 m8 N+ U1 e- w A+ s; D color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))
0 t; E2 i) y% d9 M4 }plt.show()1 ?3 o5 E5 \+ M( o5 w8 C
# 绿色的表示预测是对的,红色表示预测错了
7 v! l, A, ?8 {3 |1 [) Y5 b3 y1% H. j8 Y0 v/ b$ Y0 _" R% {6 Y7 d
2
: E( v2 v5 g8 |' T$ @3& N; K5 K" H, H
4
4 |8 Y/ }0 p6 d# {7 F/ v' e53 T4 i7 t0 B7 [
6
) \% p: ]! ~- C7
x9 H- `1 x- `) @% S0 T! Z( @87 D; Y" _3 ^0 m& `( V( x
9; \3 K" j# n# X! M' ~* B+ X0 C
10
* U: t: h1 C; ]- E O3 L11
' ^' O" R/ Z/ i! o. P% E# ^. n: |! z/ |+ o
% L E, |/ I3 y: q/ O8 l3 ^0 i8 l, n4 l
+ [, r5 [# L0 `! l! w5 R& s! Z/ _ C
————————————————; m7 L& k$ |7 T3 O, x/ j8 x8 w9 b8 X
版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。# A# r" w& |! S+ O9 n% L+ A. a
原文链接:https://blog.csdn.net/LeungSr/article/details/126747940, E7 O4 j( E( O
9 w$ v n3 _1 S9 I9 q
3 v# m0 J% ^% R0 A |
zan
|