- 在线时间
- 1630 小时
- 最后登录
- 2024-1-29
- 注册时间
- 2017-5-16
- 听众数
- 82
- 收听数
- 1
- 能力
- 120 分
- 体力
- 565742 点
- 威望
- 12 点
- 阅读权限
- 255
- 积分
- 174945
- 相册
- 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)实战案例
& J- y# y3 X) M \* Y/ T- b+ M8 p& d, a% n
文章目录
. N3 s. J$ x. I% G卷积网络实战 对花进行分类2 h. l8 j) @9 m7 u& Y* H
数据预处理部分 L& u# z; d" m$ x5 W) n
网络模块设置
' k5 d( e4 S n( B7 M! K' p: D网络模型的保存与测试
6 Q3 S! a/ I* A" H1 C) B数据下载:# t- |3 Y# T& O
1. 导入工具包1 w. C! T# p. E( ?. P! t7 e4 z
2. 数据预处理与操作
& o! O. m. o t: o5 q: e3. 制作好数据源
$ E9 Q1 W7 d5 ~8 F4 \读取标签对应的实际名字
W$ U, X% H2 d) D3 w4.展示一下数据
) K1 ]- ]& i) H+ C5. 加载models提供的模型,并直接用训练好的权重做初始化参数
! Q. S3 d9 J1 ?3 [6.初始化模型架构6 W9 |' k% w3 t+ a' X( _
7. 设置需要训练的参数
& [ v0 {& f$ [6 A+ F7. 训练与预测( v3 z% ^$ f& {# n' o0 e+ A& ^6 R# A
7.1 优化器设置
" c- A4 t6 L h- h7.2 开始训练模型
1 E1 ?5 X7 _3 T; c7.3 训练所有层% }9 A5 B8 a: K; \
开始训练
& D* w& @2 a( H8. 加载已经训练的模型1 N& J/ T9 P7 o8 p
9. 推理 S4 d, @5 T# n! O1 ]
9.1 计算得到最大概率
: U, p, d) u; r" x A9.2 展示预测结果7 i0 \) Q4 n' ~. S6 z0 i
写在最后
4 }' z4 N& b* o' `$ E; M3 q: F卷积网络实战 对花进行分类
7 d& k2 b) X3 O3 ?" P5 @5 C: P- d0 ]7 s本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能
5 t( ]. h6 g1 ?2 ^1 Y( i0 }! `, _- A0 W, S, U
在文件夹中有102种花,我们主要要对这些花进行分类任务6 Q# f* M. J% x5 J; F. \0 J+ C# K5 j
文件夹结构9 g; [7 A& N3 w! f! U, n9 \
) l% z2 ]$ K3 F
flower_data( N: F- H* \8 g6 \+ e2 \# V2 f
0 f0 a) u- w& O7 P
train
. ?+ F) B4 H0 D7 ?& p' V, J! v& K# A. F6 I) m: B
1(类别)+ ] j2 n1 O- P! W4 w
2+ ^9 p: u! P; P0 {% M, R/ d
xxx.png / xxx.jpg
# s( l: W6 E8 e7 K. |valid2 l6 _$ {1 ~5 s; h& b
- O, y0 @7 O( \1 Q' Y; N
主要分为以下几个大模块9 `1 W9 g+ \1 a
) j- V: P% ]! U" [" t2 T0 m2 W! @% l数据预处理部分
% M! D5 [; M+ z' D' [$ g( X数据增强
9 x' F6 A1 }" ?! w数据预处理
' R+ i. o6 {7 j: |8 N/ a网络模块设置& f8 j5 v) M2 T- N* h
加载预训练模型,直接调用torchVision的经典网络架构# A6 v) R! a0 G: E0 K! \
因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务% Z6 |5 N' }3 Z. B% x3 j( t
网络模型的保存与测试
3 _+ f( I/ @. d" ]8 g: I& w模型保存可以带有选择性' g% U. U- d+ q% Q
数据下载:
; q+ U8 D% q Z5 ]& ]7 P5 [3 a% yhttps://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset
+ z! | f# ], [1 P, D8 @1 K, E9 E/ H' ~/ x7 A" ?
改一下文件名,然后将它放到同一根目录就可以了, x- ^, U' v3 e, L( h
6 n9 h$ G. y3 n: \下面是我的数据根目录& [& _) W! W. O( d, A
$ v; l% K. w6 e. b ?% b6 X1 t
& l+ ^- [# c9 h5 N6 f) |* g1. 导入工具包
! n! j' L6 W a8 {: N; t6 himport os
" @ J4 f, i& x( k8 V$ _import matplotlib.pyplot as plt
, ?6 p! E/ b# V* N, {# 内嵌入绘图简去show的句柄
( a9 D/ ~$ f7 }& M2 y& ?) f%matplotlib inline
' C6 y& {+ U# [: `4 z4 y# ]4 Jimport numpy as np2 g/ }- Q5 ]: E- Y
import torch6 x# B# n4 p q' s+ @8 q
from torch import nn
% B8 G# L: G% ?( e+ G) f! U0 ^ B; C- r: L) {4 [ U7 `
import torch.optim as optim
$ u' P9 }3 S' D2 L4 G: Vimport torchvision
( X }3 Q5 d9 Y6 ~% k' Ifrom torchvision import transforms, models, datasets! _. t5 e6 S7 s/ @4 O: p/ k* C
: Z2 ] u/ I9 l' y& E1 M
import imageio
! O4 c0 a& {+ r* r/ Himport time5 D% K0 w7 Q3 P2 E) h- V6 W0 k
import warnings! X2 `9 Q+ ?" ]3 ^* s: b$ I3 l
import random$ S" f$ y& R3 n7 [# A3 Y5 w
import sys: r" a' |9 l' ^" K* L3 p
import copy
- V0 }, R2 ]4 r0 v9 _9 y% J% ]* Jimport json
W- f$ C% [+ E3 Q( |- tfrom PIL import Image; l+ s+ o. e. E
# V- [. l7 @8 f3 c
( [- p; E/ C; H7 w
1
, p% x/ V5 K0 Y% d% f- n) W" R21 L( u c' H& l* K+ ^
3
! w i. Z6 z* |0 o. @4
# J, P0 H* b6 A' E5+ F8 W( n% C n' G$ Y( `
63 L: r& e1 i1 h6 q3 b3 y' E, Z
79 ? {( g" @: |& \% s" t
82 W4 Q# n- @9 [( C
9/ R. T) N; [$ u+ }
10
9 r; k3 r. R, N0 \0 F9 x. _' D117 Z/ n, E8 x* f1 e
12
( _$ X8 k8 L1 t$ k( ]6 j5 @13
/ q3 W: G& \! `* N" b& X14
8 ] Y X! g# T& ]153 K) D$ }3 n6 O& ^) ? `
16
& @% z# w9 C! O: O4 ?2 m17
- m9 i+ m5 f9 B1 F H9 B18
" {' L; H% S5 q19
/ Y# Y- {) }: |3 }/ x206 ^% K0 ^3 V* t4 t
218 C( I' C1 ?# i& Y7 B5 \& V
2. 数据预处理与操作
& b, C) K0 [5 ?, l* R0 X* _#路径设置
8 G0 i7 B' S6 m# W1 A1 Gdata_dir = './flower_data/' # 当前文件夹下的flowerdata目录
" ]# R: _ N' L! q' P& d1 etrain_dir = data_dir + '/train'
5 p6 e [: K* Vvalid_dir = data_dir + '/valid'
) o- |# `) g, N' i. a" U1$ l7 r/ ^' w# l" G5 } d
2! }7 X& Z; f; B
3% y' l9 c1 S, X+ @* v
49 ?) ~% J9 k @/ d8 q5 Z& H
python目录点杠的组合与区别
5 P5 q, G/ l: N* x$ U注: 里面注明了点杠和斜杠的操作9 u: E2 ^+ _/ r6 D- [
7 s# Q9 Q* q* F- W: @% }/ b1 x3. 制作好数据源5 Y6 C1 ~) T$ o9 T5 g' A
data_transforms中制定了所有图像预处理的操作
' K, S" {& }* I( v5 WImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片: ^9 s, K# s& i* V* E h1 y
data_transforms = {
# s2 S! y, L+ I6 y; V # 分成两部分,一部分是训练3 P% z# Y# L9 q4 @
'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间! U7 x, u4 H/ S, d) d
transforms.CenterCrop(224), # 从中心处开始裁剪% X0 {; i5 j7 g, e. f" c& y0 f4 R
# 以某个随机的概率决定是否翻转 55开
^7 g: ?4 t" s* `3 i% z transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转
* M; b1 P! w) d& L3 F" [ transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转4 [3 Y9 p5 J4 w1 j5 n& v6 s( T
# 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相
, o' Q9 I6 u. R1 w transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),9 o; w- o$ G* ]( a9 g
transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB
6 e. |7 w# |% h$ Y # 灰度图转换以后也是三个通道,但是只是RGB是一样的5 f2 p* K2 m7 g
transforms.ToTensor(),) M# t# i7 O. i3 E, z7 g
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差
$ e7 R1 g/ G( \3 y ]),
2 F6 b& s* G: r) m # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
$ c9 B; _. b4 B* h8 Z 'valid': transforms.Compose([transforms.Resize(256),8 F9 W0 Q" \4 v+ ~
transforms.CenterCrop(224),
' X5 Y& n: W" p- t7 j& E transforms.ToTensor(),
, N6 M1 f/ p$ ]0 ~( o, B transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同1 f, ?; H/ L! N) j9 r" n8 r [
]),) |9 W% u; a5 ~ g/ g
}
3 h9 { t5 v6 l {5 P
- t" _/ |7 n( n7 M8 B' F/ O1) u$ S# W# K" I! ]) Y
23 ~! S8 A, P/ a' C
3
2 _; B T3 A2 O3 s4/ t" O0 S, o# g" E: s8 s) W% \
5' B8 v% C n# _4 \
6
' p% E7 ~* `4 c% [7 F7
" F0 B9 k; g6 K+ x5 K4 d8
8 p, x x* N, p/ M/ B# B( G9
/ ]1 @ a+ J, q( C4 u10
: [0 M, r* L6 X+ d }11
- _6 A# k! ]! T' h- k12
/ o# b" t6 C9 U- T4 U13
4 T) C! r1 f" T4 h: ~- v* X142 s9 h/ l6 V# R( u
15
7 e; v* t0 I# l5 L16
' s' H$ m& M6 h9 X17
! A# X: C9 J+ X) L; \% |" _188 K8 W9 V! V H+ s7 s- f2 j& W
19
. c) t& J8 N0 Q, n5 g$ c1 c20
) }5 p5 g# A$ [$ }21* h: Y" u5 l; V: ^4 u3 R% U# V
batch_size = 8
& C) g9 a% ^& t6 e- `4 [9 ?image_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}! Q, @$ A4 I+ i$ Z- H
dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}
! n: s! v1 Q; B1 vdataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']}
. R4 s) L; r3 g! Cclass_names = image_datasets['train'].classes
5 A5 W/ W/ ?3 b6 U8 x4 q% M
5 r! g8 O2 }8 e" l. Z7 r* K#查看数据集合
# B) ~- |0 R! U" a$ J4 J0 z4 yimage_datasets
- K+ h! ?$ _, f% ~2 S. L( a3 K
" W1 K1 \1 s" g7 ]4 @1) E1 C$ ^8 _) g8 C, f. Q, t/ k. o
2" c! p4 V$ q0 q% u V# `0 f
3
+ @$ A( Q3 o. \1 {4/ w2 `+ E5 ?6 I% D0 X& {7 R7 ?; i
5
4 {6 V9 y0 W& k/ t% ]6
8 s0 n; f8 z* i5 R/ _9 m6 ~7
$ _4 ~ s' R! v+ T0 a8
7 G( b+ n4 N" v6 y' v9
8 ]+ a ~5 R/ F2 `{'train': Dataset ImageFolder
$ j) u8 ] ^6 x5 v' f4 B Number of datapoints: 6552" l7 R- {2 b( {8 }* E) O* r
Root location: ./flower_data/train
5 Q; x+ t0 z# o6 ? StandardTransform3 f0 U+ Y# V) g1 q# g5 @
Transform: Compose(
2 s7 v4 Q/ g4 t j$ e RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)
; k2 R# _! }) T, n5 R6 ~ CenterCrop(size=(224, 224))9 Q7 @; P7 m8 j: X7 a2 v- Y
RandomHorizontalFlip(p=0.5)
! u6 ~+ ]: {0 N& ?/ F& m0 o RandomVerticalFlip(p=0.5)( @$ r& [8 D. A2 h# c% R
ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
# Q ~5 ?" y( D% p8 I RandomGrayscale(p=0.025)
; I* \9 b4 d" R. S- `0 G9 ? ToTensor(). J- b S) W: d* I; p0 u
Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])$ @9 F% d6 M; k5 }5 M. o8 S- n
),6 ]8 z# z; J& u4 _8 i5 O! J7 ^
'valid': Dataset ImageFolder
: o x5 O# a' F- {) t* a' X Number of datapoints: 818
2 T4 Y3 S+ d9 L# b( w Root location: ./flower_data/valid) \; ?+ Q$ X% A( O% L4 p6 A
StandardTransform( H7 b9 _4 X" V8 ]3 l
Transform: Compose(. }" Q/ L. M& ?7 n6 z. c
Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)) s! ]8 N2 C" m/ `+ b. h
CenterCrop(size=(224, 224))
+ f& u5 [* H) P% [ ToTensor()
- y1 y$ m' u2 ]" f# G Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])( n7 r+ Q) B, E% ]
)}
7 Q; w, l" m+ z7 U
9 q# e7 t' L+ L# y! d1# r# E0 ]# m! z' S2 }
2
' F/ M1 T: W3 J1 S3( U0 y3 u1 x0 }, g# f
48 b3 z5 M5 X5 Q& j; l
5
. R6 d9 P' w* O% @* t" P5 C, Z# n6% K; {) g6 P2 n6 e) \8 U0 B
7% P7 C. F$ Q. P
8
* _- A; F/ J& S; ^; l3 n9) P7 u* i+ ?# f m% U' ]
10
' ]2 w- O7 @ x' E113 C7 u' M1 E; f3 r& G! n
12" h, k; M; r! y$ z6 V7 Z
137 O+ h8 q4 i% j% M7 D
144 B; @% z- f) E5 T+ i, h# J
15
5 m5 Q1 s) q+ T3 y7 I16
T0 B# e; Q4 ~4 g; {17
4 ^% l2 E( y5 H18
/ V, _! e, u9 ~- O19
, s5 j4 `7 ?% }6 U4 g20
! S" b9 I' I7 g* {+ I213 x" Y$ Q/ i) m$ ^ h
22
- r: ~6 ?4 C8 |; t' C* a0 ]23
% l! v" f" T4 K9 Z5 @9 e' P3 c24
+ b" x7 U L: R6 V/ W/ w# 验证一下数据是否已经被处理完毕
- G/ t Q8 ^9 ^* d5 Xdataloaders/ ]3 H: E: x. D: S
1' q# `: U) K- t! F* o
2
) T% f; N- `( X{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,; A% j: M9 m! \3 ~4 C: \
'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}
& Z4 K( r& U2 J6 S5 U6 o1
" _+ P9 k5 Z8 u2
B) w, s, R+ T8 u5 pdataset_sizes6 W$ \/ b |3 V( j$ M
1
+ r% @5 p* z; W" B# F" a( o{'train': 6552, 'valid': 818}
0 g1 K, i& y5 o1
, {! Q8 u* c' ?) u( t5 }读取标签对应的实际名字
; T% w; D! r/ w& Z. _, v5 m- p% P使用同一目录下的json文件,反向映射出花对应的名字
+ ?' X/ |3 K q( f: o! `, \ K( A: U7 P& I+ p/ k: K& n% J
with open('./flower_data/cat_to_name.json', 'r') as f:6 Z0 c. D z8 [1 N# T. b
cat_to_name = json.load(f)
, S# |& _3 O* S: ^8 K1' t* u" M+ X. p( i% P) _
2
; v0 _, r/ h3 G% z0 wcat_to_name
# d1 b5 s' O: \1 n, ~1
9 S) G/ B+ o. m; u2 |{'21': 'fire lily',! ]' x& ?: _4 W+ W& J2 U
'3': 'canterbury bells',& \7 x, [/ M3 q l6 s: v
'45': 'bolero deep blue',
: o& I7 M2 m: Q$ i0 b6 C '1': 'pink primrose',& K0 z' I: M! g0 a
'34': 'mexican aster',& B/ C6 j2 ]7 b: j [; P. @# J
'27': 'prince of wales feathers',$ }: G8 }" q+ M
'7': 'moon orchid',- G2 i" D+ T- n/ v3 i6 } W
'16': 'globe-flower',
/ b% f$ N# w( y" B' f) Z; r9 X4 \ '25': 'grape hyacinth',# h3 Q6 [- f+ m" K
'26': 'corn poppy',
6 Y ~/ k6 O5 E' G '79': 'toad lily',4 Z3 ~- n! ^" M, N
'39': 'siam tulip',3 g" ?3 [% y3 P( p. d0 F0 x
'24': 'red ginger',
( G! f- i! D" ]! r G' W. T '67': 'spring crocus',: J6 @; G6 k! u: {9 `+ b
'35': 'alpine sea holly',
/ N9 b/ c- T$ ]/ F9 P( s! r '32': 'garden phlox',
; `2 J! p- p* i' k5 B '10': 'globe thistle',$ V8 a! o$ t! C w# i4 X
'6': 'tiger lily',
% c7 _5 r% c& D( j '93': 'ball moss',
, A$ O9 Q( p# s- J7 e% ~. K# d8 h '33': 'love in the mist',7 A5 g8 D( G3 X& N {
'9': 'monkshood',
- b6 T4 v6 c% p3 R '102': 'blackberry lily',
# u8 m& Y8 t8 e1 T+ U. \ '14': 'spear thistle',
6 e, l0 O# H) C6 q' |# u+ Q' _ '19': 'balloon flower',
, e1 L a# J+ c5 h '100': 'blanket flower',* k' O* C$ \/ O
'13': 'king protea',+ X! {3 G# Q6 f
'49': 'oxeye daisy',9 ~, _. R% [) w
'15': 'yellow iris',
* P- v( h: R, o" _* M '61': 'cautleya spicata',
7 D! z% C$ ]0 a$ ~" \; Q' g4 \3 W '31': 'carnation',, N) X) m C& A3 R' B- ]
'64': 'silverbush'," a. M# r0 b1 ~0 Y# d5 ~5 X/ {
'68': 'bearded iris',8 F# d. w- [' g7 ]
'63': 'black-eyed susan',
; Z) j3 g9 \3 _3 e+ a '69': 'windflower',' z4 F8 |$ G& l+ a; Z
'62': 'japanese anemone',
' A- i6 p6 y( V* w1 l5 J: Q0 [! H# | '20': 'giant white arum lily',
! c: o8 c6 m0 s( \# R. a '38': 'great masterwort',! \ |, p# x7 w- [
'4': 'sweet pea',% R1 z1 S( R1 Y" Z
'86': 'tree mallow',
- j! G4 o) @6 w r# S& [ '101': 'trumpet creeper',; R }0 B' i W/ K j& ~5 l3 q
'42': 'daffodil',' J) L0 O! s7 U/ k+ v1 F! Y2 D
'22': 'pincushion flower',* F5 s$ L+ `6 N# L
'2': 'hard-leaved pocket orchid',
7 Z& ^ B: S+ p0 z$ n. N '54': 'sunflower',
, O# p( p) l% j5 S. ^ '66': 'osteospermum'," r9 i# Z8 w8 z# R( J1 z& `
'70': 'tree poppy',
4 f0 O# `" F# V; }2 H" U '85': 'desert-rose',8 ?8 b& S0 G! ?
'99': 'bromelia',$ k, y' Z- }6 r0 S; e5 ^% u
'87': 'magnolia',
$ [* `1 q3 E* }2 i '5': 'english marigold',
% b2 w3 ^. l0 B2 G: ?7 S8 _ '92': 'bee balm',
; }6 y# [/ J5 Q# [# A6 G d '28': 'stemless gentian',( }$ J0 y) U, l# Y
'97': 'mallow',3 X1 ~6 `2 z" m9 f3 j
'57': 'gaura',
5 z+ }: ~0 G9 _+ ]/ B '40': 'lenten rose',# N3 s9 Z4 Q- L
'47': 'marigold', H! n# ]; u& c; f8 q. S5 T
'59': 'orange dahlia',
7 A/ ]7 G1 n6 k. P4 T '48': 'buttercup',
5 e" }( T9 g9 d! S6 E* c '55': 'pelargonium',
0 F$ C4 z9 u7 E '36': 'ruby-lipped cattleya',
( @8 ?% t& P: F7 Q, `1 Z9 B4 K '91': 'hippeastrum',
) A+ x h. q0 O) {4 l '29': 'artichoke',4 h6 p @! X5 u X* n. M% `( H
'71': 'gazania',
! j+ |2 ]! d7 F5 r '90': 'canna lily',
" [' b6 r) I( J, W2 \8 J '18': 'peruvian lily',
) U3 x. o% j6 r '98': 'mexican petunia',
* K' h) d' y4 G- J9 B* I* q '8': 'bird of paradise',
& }* D0 B: X' h. ?/ J: R7 {0 l( e '30': 'sweet william',6 l* D6 c3 `# C; I O& ?6 [+ Q
'17': 'purple coneflower',
/ r$ n" f9 {1 ]- A '52': 'wild pansy',' F3 |/ P9 g3 S+ Z- f8 z7 @: o3 K7 F
'84': 'columbine',
+ Q/ J0 Y: b. z, \1 N- n% S '12': "colt's foot",
# p) q0 N7 F! X: x a '11': 'snapdragon',9 l' M, I( Q, K
'96': 'camellia'," a4 H, B' z9 b2 G4 X0 }+ s5 ]8 T
'23': 'fritillary',
P9 d! k0 Q7 o '50': 'common dandelion',- |3 B! Z' c5 w- q- a
'44': 'poinsettia',
, k0 v, v: G5 u/ n3 [0 Q5 ~ '53': 'primula',$ o" Z! C3 I, s: W% u6 T) ^
'72': 'azalea',- W$ {8 k3 p/ d
'65': 'californian poppy',
( S( s5 W! f4 B. s '80': 'anthurium',0 S. m: \- z8 a
'76': 'morning glory',
2 Q" j0 z/ l2 m# y# w0 E) F '37': 'cape flower',# n4 W( z' m; @+ d x
'56': 'bishop of llandaff',; ^8 m, C* ] i6 G
'60': 'pink-yellow dahlia',
+ N! _' y/ N, G' B4 U5 a7 O '82': 'clematis',; r6 h- X" ]3 }) J A+ p
'58': 'geranium',
( l+ ^- c, W4 w# ^ '75': 'thorn apple',
! H" w- z8 v: r. ~- q '41': 'barbeton daisy',' G Q# F# u$ a
'95': 'bougainvillea',/ t8 \ {% @5 Z, e, J
'43': 'sword lily',
' A) V0 B5 r* c8 ^7 _, v( f; Z" S2 J '83': 'hibiscus',
* G- e0 B4 C% r4 K+ R '78': 'lotus lotus',# ?0 J% P e4 r) r4 T& m4 }
'88': 'cyclamen',
P+ V: F% V# z, ~+ A1 b '94': 'foxglove',
' o5 }# n+ |8 v4 L( v& T '81': 'frangipani',8 ?8 E( C C4 |5 E
'74': 'rose',
/ ~- T5 G! w; F4 K2 C( X '89': 'watercress',
" t* n) w5 F- _8 [4 F( q2 X '73': 'water lily',: f n% l. L8 K% o& m* H$ J5 d2 s
'46': 'wallflower',/ Y1 n6 O; w: K2 v
'77': 'passion flower',' s; {( |/ e( T8 c7 z2 j
'51': 'petunia'}
( ?% d7 ]7 {: T8 e8 y0 g# h0 k s1 m$ ]/ [' \. U& j |9 H1 n
1
$ g5 J$ a2 q4 E6 V6 H& J21 Z+ f: v) f3 \6 s+ d
3- ^/ t' M: N/ j9 Z* r
46 Z: G' z9 c2 q7 P7 D
5$ n$ [. }8 f+ u" N2 d" O6 S
62 F$ a) F1 e* T/ e: X0 }$ l
7) w) F. G$ U4 C0 e+ @0 d
8! i; e5 J2 `4 \- x
9
/ @- l' Q( j; c6 M7 T10" y) a5 h5 E" v3 H
11
, b/ @* w6 l$ y/ M/ n6 M12
( q# A6 p" V5 s13
# A, P' g7 ?/ w& {* |14
% j) M" y8 b8 u) Q/ [15
! i! R# I t: e+ b16* O, O+ {- ?; p' F) E( P, _
17! ?: A. }8 s& M. V
18$ k* @$ B, b/ c9 @) x9 O1 X
19& D; E) q" [# u' }. ^/ w A5 i$ x
20
2 V8 j! z P# U4 S21
' ?4 y- c5 u! Y7 H& _& V: z$ [5 [% q; y22( f/ w+ S& i+ U; Y7 U
23+ A1 J* k K/ a' |. J' q
24, m. `+ O5 a) `( ~( O
25: r' A- B; l/ Y, V' v
266 x. A% F N( v1 p e6 o
27
$ J2 e+ Z" e+ @28
1 q) S- y+ V: g3 o* _) d29
! C6 Y% h8 h# k$ P30 i9 B/ o P1 [+ j6 e. m
31
2 f" f/ k7 |" p! r5 Z+ n! P32
: B' A7 O, _8 V! M( @* D% m339 |6 c' l) Q6 @5 [: X0 y
34
% D, N9 W8 w" p) D1 W% v6 f35
6 j9 L8 R# q7 A8 w# l# [" g$ B36
1 b! Q: @; _7 h/ A37/ t8 y( E" L7 [8 G% i) a
388 j$ Z- g$ w0 D; I& k0 E
39
' W- i$ Z8 {' V/ b; O4 y- K) @40
& W( l. x. a! @/ z4 t: I41: q$ l( j4 J# a* s4 _
42
, R) J0 m$ O3 Y$ o0 @& l ^# p43/ C9 J0 R3 \) r% ]
445 s$ ?/ K7 r; A6 L
45
8 w' y# D, u, ~3 v. w/ W46
0 l0 B" B3 Z9 }1 F% f+ l: o47
3 i4 a% {1 U) ~8 l1 D$ [48$ I6 U% u$ Q8 p. I; W8 A& F0 ?
49. R- l( ~* }4 J, a0 ]
50# F, i+ f- I1 k, ~, M6 C! D' i
51
" M5 O7 w' v2 m8 s- E52
! R- X- ^* _/ D9 c; P536 Q* D, Q) B3 Y U2 C- J
54
/ i% l* n% A( ?, `+ O55; N) k8 Q5 e5 N
56
; B2 l- a9 c: @) e4 d+ P57
! _9 l# m7 s) `58
8 k8 c0 g1 |. t2 o" h: v59! j) ]7 [ p5 V1 g* m$ t
60. l$ F0 m/ u0 p$ i
61
! ]9 k4 {" w' o; b, s62
Y) w) G8 f- Y+ z63
8 ?" R* z9 ]( J2 ~8 L646 J, i3 y [* n+ G& N: b
65
1 ~! ]* H+ J! A# ]5 {66
6 \ H' S6 M$ b% g1 F67
+ F. M" D7 D9 U. }5 \/ R! u68
1 O# F, ]- x) H69
5 f6 \7 k! |* k/ z6 f- ]. `70* ?+ ^! l3 Q0 v8 U/ a. m
718 r% k# v. {1 h5 s/ _4 A
72
' }+ t' K# B& E% `2 \" r7 g73
- L5 E% G6 s% X- G74' p9 e" ?" O* E1 _) x
75( E4 @% ^5 ?- N# n' p
76' m! [3 S# Q1 a) ]& ?$ Y: l
77
( N. o7 X7 w) c0 l' x78
4 f0 S; t' ^) d) u2 h5 C79. U( c% ]3 F8 V8 O
80
4 j& o" `" s9 r7 S' z81
) d$ Q0 X% b4 a/ J* a$ \82
' n, x8 Q3 U6 ~83
* @( X1 G- o. _1 H$ u) w2 |84
5 {6 ? N% N% f/ E8 B858 |! o* G: N2 Q( ]4 o5 r4 t
86
( q; I) C$ A/ y# y! @5 W- z3 p) G2 b87
) o$ {& u5 o( C; A" |- \" \88
, p; d3 {; C0 O89
t5 Q S; g) c4 { s7 w& w90
# t* @0 p& L2 m1 z' b91
! k' |! b0 M9 n92
: \1 j) @( F" u: |6 m93, U7 C& V; r; T' G- Z; w) |7 U1 X: a
94
- }. m4 e: @5 U7 ?95
, n B A3 w0 `) x5 H$ x9 W96
( Z) Y$ R2 \. y97
7 w+ P7 h2 E! x$ F; ]98
' e9 e' |# Y5 k1 p- b99
1 B1 e8 g$ ?7 \( L- N; }100
K% q3 R. R4 n2 R; `, O- i2 k6 v101' a9 k( b s3 G
102- d, K( L9 Q& P7 u& ~( B
4.展示一下数据
3 M) Q F8 f: c: y% U V, ddef im_convert(tensor):
; k- A# l3 q2 [( I0 x """数据展示"""
; a) B- k/ g4 r8 q) F3 U4 _% K image = tensor.to("cpu").clone().detach()% ]8 H' R4 g- d3 \
image = image.numpy().squeeze()
0 r$ S, T- x: i; l3 ]- G) K # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图
4 f7 h- \1 G7 w; q. ^7 H( D1 } # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)
! a, I! Z/ G& h8 h% W# @! F. s, } image = image.transpose(1, 2, 0)$ |: }3 Z/ u; Y0 R& X& A
# 反正则化(反标准化)
" J1 {) S3 u. z& O image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))% n: X) }/ O3 a0 `, y. R" B9 K
8 _ R7 O V6 a% C # 将图像中小于0 的都换成0,大于的都变成1. R: Z# b3 f% }
image = image.clip(0, 1)* d: {' S2 U1 H" j7 I! o2 m
T9 G- k1 s" T return image* }) J; K- J; n' t& Q8 J
1
$ Z9 V2 w8 l7 M5 W1 x2
: ~0 d. H2 O6 |$ Q36 o, X7 K, W3 p
4+ J R* K# ~$ S R9 V c1 i
5- w5 V! `# B F9 U+ j5 Q
61 h, ^- ~, t/ b2 j$ |' B
7
! N6 R: _1 r1 W5 U' L8; \& D+ j9 H7 n2 {9 P
9' u) u m! L! t* z, O
10
6 Q- M3 b1 L! q, ~3 D c11" L: j9 i0 n. W. e
12 l- b2 H) S! q3 g) l/ C0 b
13
) ~5 r" z( ^! b) B2 X14
! f. I+ i+ s- s* w H+ D1 b7 r# 使用上面定义好的类进行画图: C5 q: l) |% q/ Y
fig = plt.figure(figsize = (20, 12))! X! k2 [/ w! b" X
columns = 4
$ Y! k8 T6 l% c' x( h; }7 krows = 2
" I( F) W G5 C; I5 _' g) W! n' c& Y: j
# iter迭代器
% p' e7 w8 j$ G/ m, S3 l# 随便找一个Batch数据进行展示
1 y6 K: V4 M7 C r( N: K; jdataiter = iter(dataloaders['valid'])
+ A0 ^2 S) {- Qinputs, classes = dataiter.next()3 p1 S5 B/ O% Q) X6 e, g5 C/ t
- a7 U: X, Q; m# s% Jfor idx in range(columns * rows):
! U6 x' r4 z. b5 d ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])
' R/ [7 ^9 c9 W3 d7 P # 利用json文件将其对应花的类型打印在图片中
5 I9 r, N/ d' Q. @ u1 F6 q ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])
" v' l0 g4 T4 a$ C plt.imshow(im_convert(inputs[idx])). A6 K3 D. y- }, ^+ R/ A+ h
plt.show() s1 q1 I) D/ r( g! D1 q7 W% `
" u! c q! F8 v( G
1
0 H( M7 S4 d* l7 K% y0 s5 q25 q; e( c; a& \2 v- F3 M
3
0 e9 I) n) y% l' ~+ N. s+ s4
7 Y. _! k: I! U# t0 c2 X: A. z6 X5
& l3 B$ k6 N- f9 v" }& w/ j6! {6 y5 a2 C2 e8 \4 ]
7
% N8 ~- Z4 i3 H6 B) x! j* w83 @1 c$ a4 u* t
9 i; u. ~/ W: [( i6 Y
10
" c% H, w( l: H/ p% H$ U/ G11
5 f9 ` C6 x+ k3 G/ ?( E% ~( C12
! V1 l7 {6 x* T) D) \4 B13
$ c( \3 z4 f: U' _" R14
% s `1 g) G9 O5 Q) ~0 I: i/ @$ |15
% N7 g# A4 N: O$ W. c16/ V! ~: {$ P1 d7 z, y! ^% w" c
% N+ n# ]$ h5 Q) @ {* ]$ z" j: R! ^( @6 |' q. o6 w4 Q. h
5. 加载models提供的模型,并直接用训练好的权重做初始化参数& W& ]3 r0 o1 w* S. W, G
model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
# b" Y% U# v* n& g# h6 t# 主要的图像识别用resnet来做
$ q8 Q& S; K+ _# N. z3 C# 是否用人家训练好的特征# Q# A$ D$ L9 G8 J$ |! p0 T
feature_extract = True9 O4 `$ |: y% f# v
1* H* E' T' p! |
2
: W4 V0 ^& ^6 \" _7 R* I1 d; k30 o/ ~/ B, m& f z9 \
4
' o* e v9 b' `* D+ j# 是否用GPU进行训练) q0 Y) `* {# T2 V
train_on_gpu = torch.cuda.is_available()
& i$ W! S1 X& f6 m8 A* d; d% M, h6 v) U/ M8 H
if not train_on_gpu:
+ b% z& V" p" o) w T! f2 r5 M. _$ v print('CUDA is not available. Training on CPU ...')$ B* a' w% ~! N% ^* v7 H8 L
else:+ z4 S N! X: Z% b6 Z
print('CUDA is available! Training on GPU ...')$ x& A' Y7 f0 `3 l% M& O
$ ^0 b" h3 U! V1 N$ h- E0 qdevice = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')
4 y \' v: n4 @$ ~7 X- A1 T1/ V2 C9 G# ~( ] l$ {
2/ [4 ~% w7 N/ B* [
3# z7 b, g0 V! c3 \1 [9 ?
4
9 E% W [. S' @- {6 s& a& E" I5# c+ T3 l/ O# R, g- r) S
6
/ L: @' p" j) B! [; i7
; y! v/ B1 W" K: v3 C81 R& V& C9 Y( [& I" Z6 m* J
9' `3 {3 k8 W0 L
CUDA is not available. Training on CPU ...
& N6 z/ V- w! F. S1
* b5 ?1 n4 d7 ?& L; Z4 f( n' k8 x# 将一些层定义为false,使其不自动更新- C( }8 M+ x8 @8 l
def set_parameter_requires_grad(model, feature_extracting):" C) ]8 l9 h( ~+ ]
if feature_extracting:
0 T U9 O2 n# g3 U9 L% s for param in model.parameters():0 Y, I* J8 A. Y* v
param.requires_grad = False
- G4 u, m% ]& V% m+ P: {1
1 d3 @" `6 |1 g) ^8 m2" F* t7 I3 y% X+ T1 s8 @
3
6 e4 g8 u w" X, _9 ~3 X& T% K4, N6 O/ z. C% S
54 }0 Q( b* R% r) v4 i' L1 y
# 打印模型架构告知是怎么一步一步去完成的, Z. U# [$ _1 Q, D& A
# 主要是为我们提取特征的0 G7 @8 [) Y" x# e
K% W2 x) j' v( p. r
model_ft = models.resnet152()
" {( W$ ~# w+ C1 g6 H4 imodel_ft( X+ L& k2 M m+ G5 K6 j
1% _/ I9 `" O! `* D# F
2) n0 k& m6 Q, v; b
3
8 I- _" Z" G8 ?) i+ J* l S: C4
7 S* p; ?! q8 q# L8 j5# q2 `4 x' Z8 j( n, a0 o/ k
ResNet(( f6 N& c2 B4 c" D' t% o7 T
(conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)7 _4 n; ]* ?/ O% F2 _# |
(bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
. ]. g" U( b" L ?4 Q- }0 ^ (relu): ReLU(inplace=True)+ q1 P* K" m2 G2 G$ n! L- e
(maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)" V: r: O/ S: M
(layer1): Sequential(
* i% W7 A: L5 r: [1 q (0): Bottleneck(8 Y9 f( Y0 }( n7 j0 P
(conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)
C) h0 n$ @1 Z8 q7 w (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
9 d$ K5 P, C$ u+ e0 P% y$ f (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False) H& c. [$ t, k* t; w' r& _
(bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
3 }9 j9 M* k7 }3 b (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False): j" H4 U7 a5 k& I5 k7 r3 s
(bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)/ V3 L2 `7 Q) o, V2 V
(relu): ReLU(inplace=True)2 b& M( `( ]5 I! n
(downsample): Sequential(& C$ |9 M, T) `
(0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
( N/ ?7 c r( y (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
, f8 `$ N* \& b: O' O )
! I; b8 d8 O8 @" g( k )
- v' l3 L" \' Y& ?中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。7 ]. O7 B+ ]1 t" J& y
(2): Bottleneck(& Y' u# t. ], q7 H0 F7 Z4 _& [
(conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)
3 S: n O8 p- @! Z9 g% r% F* I' @ (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)0 M8 q' P9 P6 J# B1 c$ k- P5 n
(conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)8 i4 l6 J/ E/ |% C( l' d9 i6 z& o
(bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)+ T! l7 G4 k1 t9 h: U# u
(conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)8 v, {9 ^- R+ X6 B7 V3 Q) p
(bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
7 F; @0 E; c: Z* S+ E9 J% Y& t) W (relu): ReLU(inplace=True)8 d% S% j. q3 F
)
9 u1 c# ]1 B9 h* m' s )# d3 X* ~) ^" J1 r
(avgpool): AdaptiveAvgPool2d(output_size=(1, 1))
- D, f' i+ Z0 E5 i: p4 u" j3 o (fc): Linear(in_features=2048, out_features=1000, bias=True) h; u9 @5 h$ B7 R4 {& o$ o
)
/ c& W8 F# }7 ?
$ A& ` A2 n5 O3 b; w9 W1 A& S. O- C1
) [; k. Z" k/ K2! X( O$ R- l: F: `' |
3
+ @/ m; a7 z( Q" Z0 C46 O% |7 x/ n! s, d
5
; V& m6 [; _& X' Q& u& r6
& o% a; t7 ]: A76 O- e n/ c& k* |6 s" h
8
# f# v( p' A6 _" ?9* B/ |$ D+ `" U2 \5 H: @
10
6 t$ J7 Q7 g& Z) K9 ]11' Z8 z% q8 g0 Z% j% f; n! a3 M
123 h" c8 [. o% N, P
13
7 l9 s4 f P, Y: s7 B: E7 s14
7 @7 l* R8 i x- c% X5 s0 e15# i2 I/ } P% l0 E3 s& i
16
# I* ?- v! ~& X, r. C0 W17 T+ |/ u. D9 h" c* z5 j
18
% [: C0 s% _9 q0 F2 B0 a |* e: D19" N! S4 ^& `8 b [( P! {- q
20
8 _$ @, G; t \/ M: C21
, K* U, R& a) R' r6 ~7 A3 p22- _) K: Y* |- ]( W+ G C
23
; E% V6 c5 ~, a! g# i0 r0 Z, I241 @$ j% ]- x+ e( K
25
6 j H+ r$ B$ ]. q26% o' q- c- i" [
27- Z& ] t+ @5 A9 n# t0 a D
28
7 z: i, N1 w) v. O, ^% V' j0 ]29
" _5 _% s* \3 ^! K3 r+ O308 s0 u3 e$ h, l
31% \ d# p& g9 P3 g) H
32, T8 z0 ]7 V [1 ?
33$ Q* {$ X, W; I+ ]% B
最后是1000分类,2048输入,分为1000个分类
0 j( Z. z; v/ b. i4 H1 D9 H而我们需要将我们的任务进行调整,将1000分类改为102输出
0 A( E7 s# s) ^( e
) X- F2 g) l3 m1 A6.初始化模型架构3 ?3 a. r" ^, a9 @8 z
步骤如下:
) B6 E% q! n/ O6 j4 F4 L# g4 `
$ B/ W8 E- x, p( K将训练好的模型拿过来,并pre_train = True 得到他人的权重参数1 o1 L0 d" n: O: v3 F
可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)
3 Z8 b" f3 j7 o; F% {无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数2 @7 ?9 [3 Q+ N) S9 ]
官方文档链接
7 p m: V5 R6 r2 m5 ^7 ^& Z0 yhttps://pytorch.org/vision/stable/models.html
7 u& _+ I7 T0 W n2 g) W. B4 O7 E. a) o h% g
# 将他人的模型加载进来5 }6 a) u) R; u2 y1 U0 P
def initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):) x! V# v, b& V# ^6 `/ }. E
# 选择适合的模型,不同的模型初始化参数不同
# l" C7 d+ }0 i8 s model_ft = None
/ Y" f9 ~6 g& ]* r8 N. m input_size = 01 |. {( p; F7 r# A# ?, Y7 E
; ~- }6 I! H2 O6 u* U
if model_name == "resnet":: R5 ]- f" D1 G* T I, a
"""; u( q9 Y+ z) B2 l, J) s- P
Resnet152- Y, q/ u9 [) e0 n
"""
0 e$ \( |+ E- g# V; {' b' s4 g& v' `6 Y4 @6 R
# 1. 加载与训练网络
. l/ H" v- [* X+ t$ x6 g5 E. O model_ft = models.resnet152(pretrained = use_pretrained)% `" _. m4 ~4 f4 u: u
# 2. 是否将提取特征的模块冻住,只训练FC层
5 n. r/ C, A& _8 l) y3 Y! Y- o set_parameter_requires_grad(model_ft, feature_extract)( O, q; e2 V7 T9 f
# 3. 获得全连接层输入特征% Q% O& v2 G% x' |7 a0 G
num_frts = model_ft.fc.in_features
# p9 S! i3 _. q) g # 4. 重新加载全连接层,设置输出102
/ @. Q; q% _7 A( } model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),$ z* W( d9 }! v1 E
nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1
7 m8 X$ G; x, l/ ]% [6 w6 E input_size = 224
1 l8 u/ A( I Z, Y" @/ E. ^: s# \8 k9 u+ l3 C: [2 n
elif model_name == "alexnet":+ i) H4 j- d0 C! {6 I- _
"""3 P: R7 C) M3 u/ v! i4 r! R( A
Alexnet
" a- Q6 @/ a( [3 [9 V! | """
# n" q$ b% p0 ? model_ft = models.alexnet(pretrained = use_pretrained)
* B, c3 V5 {4 U3 U set_parameter_requires_grad(model_ft, feature_extract)4 I' e0 b' y6 s; V1 [
! Y4 R, ~% O, P$ p) D
# 将最后一个特征输出替换 序号为【6】的分类器) t$ Q( B9 M6 }/ u. `, M7 l
num_frts = model_ft.classifier[6].in_features # 获得FC层输入
8 q9 k, X% V$ u! e n3 j( E$ Q model_ft.classifier[6] = nn.Linear(num_frts, num_classes)% x1 i+ L6 F5 S$ J& v. d c6 E: b" j
input_size = 224
0 m) s: d# F/ J; A" v
5 H" @& F0 F. C/ ]$ Q elif model_name == "vgg":# D( T. e! j# k+ ~# ]4 g$ y
"""
, B3 G+ ~6 a J0 J1 J$ N0 G! |; i8 L1 i VGG11_bn z* E2 b: L$ y/ f. [- n
"""' E' w7 b6 e( r
model_ft = models.vgg16(pretrained = use_pretrained)) i: @( \% U& o" G+ T
set_parameter_requires_grad(model_ft, feature_extract)3 k. e- y" m" s5 J: x3 z
num_frts = model_ft.classifier[6].in_features6 @1 A3 X9 [9 p0 `$ B/ `6 v. ^
model_ft.classifier[6] = nn.Linear(num_frts, num_classes)* N0 `1 C" F; B
input_size = 224& Q8 [9 v. R. x G5 x* z
8 }/ ~( N& J/ M- T elif model_name == "squeezenet":
" s% n4 i R& T2 u' p- C* m8 ? """* T; T% [( C( c+ k# O" p* h2 S; J
Squeezenet. i6 S. W) F- q6 @1 B0 Y
"""& F, P1 q; p4 J0 q# F
model_ft = models.squeezenet1_0(pretrained = use_pretrained)& d* b8 H% i( [6 v2 D
set_parameter_requires_grad(model_ft, feature_extract)& n0 j8 L# \+ v2 p# _
model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
$ P2 S! t, m6 T, J" | model_ft.num_classes = num_classes: J" ?* u! z# H R1 J
input_size = 224
: j% ?* r+ A: _. J9 ?1 a; ?$ X% J' q) Z' h' }' O8 O$ a6 Z
elif model_name == "densenet":. d# ?+ m2 E4 Q/ e
"""
# @. {. i1 f: U4 P& | Densenet
9 Q# E+ `: B4 L: W* r% a, L """
" Z6 R7 [% X6 a) E! L: ^ model_ft = models.desenet121(pretrained = use_pretrained)
' g5 W H3 S7 {2 T' \ e set_parameter_requires_grad(model_ft, feature_extract)
8 `$ Y ~! H$ q0 T num_frts = model_ft.classifier.in_features& G# j" E& J! i
model_ft.classifier = nn.Linear(num_frts, num_classes)
0 D: {. ]: i" v& n. G, r input_size = 2245 G+ z) D" k$ P3 x
* z2 M2 ^( O: {) d) w& G/ d, n! ^
elif model_name == "inception":$ m" p1 X) W a
"""
* k3 F) r4 C" O6 ~ Inception V3
' e) @$ A& W) G# J) u """" m# y, x) U( q3 e1 p. y1 k
model_ft = models.inception_V(pretrained = use_pretrained)# H7 m3 ^2 v3 X& E) k
set_parameter_requires_grad(model_ft, feature_extract)- H5 s0 a2 f5 r
* P9 d0 L5 B4 O/ z! S m
num_frts = model_ft.AuxLogits.fc.in_features2 R4 Z1 Q2 A) s" n
model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes); ^' ^' M. [, X" i5 b# D7 E
/ `: ]9 s" x1 C1 I) d# G! W
num_frts = model_ft.fc.in_features
3 e2 s( k: n1 b! C G model_ft.fc = nn.Linear(num_frts, num_classes); _) r" q! a* k/ F i3 _" j. |
input_size = 299+ P* W, u0 m' V
0 e7 V! z* [4 }7 b
else:, O( W: @3 H6 x/ A
print("Invalid model name, exiting...")- X/ ]! `, i! ^7 h7 L) W
exit()
! g Z. z2 [' s& \! P7 Z3 c r
6 G3 d2 ~7 A% l. J return model_ft, input_size0 R k! A- Y1 k! C" ~
" H& `- H8 v0 }7 S4 c1' J' R0 r! B3 ~0 N' o5 B3 |
2, L! f. F4 H% m% L
3% h2 t' c! g! n; E6 A
4
; c |9 w0 {2 K& l; m2 J5
- {: ~1 i M" J/ j( V& x6/ |- F0 y n+ b, D6 D- c9 F4 Q, A5 ?- V
7; H# c. o+ o7 S' d9 s% w- f; H3 X
8+ v M9 T9 |& v$ ?8 i' o
9
1 ^5 h) L) o9 |0 O: C10
$ ]: A% d4 E. V k' q* y11
, U# x$ c! F) Z) d; C2 T12
: s! r2 i+ h2 W9 O13. ?, x8 w a3 n+ _4 |; M) ~& U$ N
147 n+ a7 h0 R3 b! }* S" f
15+ }- K! ?/ {7 W( [1 K
165 f! E! Z. W5 T! S
17) S/ {2 s4 d8 k' E
18
3 D p' |/ u& [2 }193 w+ }1 y. O1 l7 ]6 v5 [8 N5 v
20
2 G# H6 n z% N! c4 X" O21' z4 _. o/ n! [2 G, Q
22/ y! R8 {; R& l& O1 w+ S- K1 f' [
23# R y" J5 P8 r X8 @- N
24- ?2 v) T% `7 x$ t8 y8 O
25
. j& U6 B" X! a$ }3 x3 w5 c26
* w3 @7 A3 L! l272 E- g- a* l s3 `# K# E4 X- O
28) z1 `8 b' h% r5 D
29 z7 t; \2 Y/ v
30
" ^6 R* I" I: D$ Z+ }31
R, `8 X$ z1 x& E+ F* P32% r" n# \' V5 s( r9 ]! }7 z$ v
33
5 ^ P7 @6 d9 b% b+ J( L' b34
& ?: _( k; z2 X/ t35
9 ~, k8 z, A/ E1 ], q36' x- _1 u* V7 G% m8 X
37 g% \) A6 a6 N8 p+ C) D
38& k+ o+ ]# a+ |9 L4 k) G5 S' j; I
39+ Q7 h9 B8 M7 H& A# o$ _4 H
401 ^" K' F }( q3 \
416 O+ x5 b+ y' n5 e6 c
42+ N6 s( i: m x: x0 n- @. Z
430 _8 P/ p) @: Q; ~' N8 K* \
44
0 } {9 l; s7 v" w0 c3 D9 [ D45
3 K" n8 Y& P/ N46; M. y$ p. G' \* f. C, s/ U3 @+ ~
476 q0 k% |' w+ h6 l. A. o+ c
48
5 U: ]2 e1 {/ A4 m49
8 ^# K: g9 ^5 ]/ G# m# \5 D50, o9 q2 z0 D- v5 ]# A6 W. O6 b
51% r5 Z8 G; {8 v- b3 T+ `0 j$ h9 _
525 t* n, g; G2 O- L: r* `
532 S/ @& W% k- w7 Z- |! `
54
& K" A" |. [5 ?- {, b9 Y% U, Y55
$ s2 h k, B' w6 f" {, s% V9 a56( h' t) |# s; h9 L, L9 i9 c. A
57: j0 {' t3 E9 S2 |6 ?
585 y3 l: O0 g8 W9 ]) z0 X
59" C7 m2 u+ `$ ]# j: }/ i
60
/ s0 n! e" P |) r3 ^612 l& v6 \( c! D5 M e: T& B
62
4 X: g; B! D9 V; o; p63
9 q& p$ o- P- F3 l64
6 o O$ A U& H65
7 D& \- u" U2 [% Y1 z/ G66
0 ?5 \7 K! z1 @) V: \( U672 o, v$ I# H/ A$ P
68
2 J/ f2 N' P7 h: G69
! s2 w' P7 s4 D9 q% c# p70
7 Z" J" n; C* |8 A1 X* W71
- K8 b7 q: ?1 m n& @$ V0 |9 O4 w72
7 I( i' |% e5 ^' D- d. N, B736 S5 D$ R6 H3 B% i5 ~
74
+ J2 X$ h* s/ \75) p8 a& u2 M ^, v: F# t0 z1 j/ B6 ]. G8 P
76
- H( O/ i% [0 m" H. t) v% R77* ^1 _7 T& Q: P! e7 J- [
78
e& h7 \0 W8 T( B& p# l% M796 g* \4 c8 _% b
80
1 M; ^! P0 A0 I7 _8 l+ w+ ~81
' ]- @0 a3 t/ A2 O7 b5 s822 h% b+ H3 s3 F3 `8 R" S
83
5 v3 v3 j! n4 t0 N2 T7. 设置需要训练的参数
3 U: a7 t! P* `9 J" Q W# 设置模型名字、输出分类数
, M' X$ b* Q& X8 o' \- ?. a# X% Smodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)* }: `1 d3 ]! `, ~9 c7 t5 N
2 c6 e A0 \+ P2 o# GPU 计算
( I" m. D7 G3 D$ B. emodel_ft = model_ft.to(device)
6 N# p d9 T/ P4 ?% B5 s ~, t/ u4 ?8 ?4 S, ~) J& g
# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取
) s4 i, k1 \1 B0 Z# f& x! P+ m- K* gfilename = 'checkpoint.pth'+ R0 d0 l* V! `4 u$ J
; H6 e- X0 d \8 I2 R7 N Q& f
# 是否训练所有层( d; k0 t( u- v" U. ^9 ~ @" q, x
params_to_update = model_ft.parameters()
' Y7 ~7 G, E( R, d! T; D4 Y! j# 打印出需要训练的层
+ ?3 @1 M: g/ Z! i/ Vprint("Params to learn:")
+ x' m1 ^) M8 N5 W' ?- y! {4 r( J/ Mif feature_extract:
9 @9 }$ B; l! i* D5 m5 y: n params_to_update = []
. J' `# x7 T- ~ for name, param in model_ft.named_parameters():+ {& @# O6 N4 ]7 [3 d
if param.requires_grad == True:( z3 d/ K3 N; C. P
params_to_update.append(param)8 c: f( N4 `3 G! o6 X8 ~- ?8 w
print("\t", name)) d8 P; \9 R! W: J
else:0 l0 j9 P+ v \5 \
for name, param in model_ft.named_parameters():
& R* X. }* b. |- S if param.requires_grad ==True:* @8 r; W/ t4 X
print("\t", name)# E8 f Y1 A. o+ w3 k) Y+ W/ z
+ q! t3 G, v2 H9 O" e1 s16 Y9 Y, R+ C3 {6 ~
2. F9 W2 A) j6 A5 i& E9 N
3. s R) k% ]) F+ D, a* W
4
& A& r3 T$ E& E2 N1 ~' r5
" P8 Q$ l- d3 P64 K) L, o5 G7 z; Q$ |- d" P
75 T- T/ Y: t+ f# r9 @4 ^7 p
8
& R$ y$ d# r0 o f8 K9; g9 L5 u. f# a1 k: @: b- @9 {1 k
10
! h( K; @% {/ B6 P" ]7 e( k1 C5 a! E11 V, M6 E) L2 D2 Q3 w, N0 O
12" R; x+ `% u7 F0 ?0 m) j
13
+ h. G# J. v" @# p' ^: s, o14: U# a: g1 e& u5 f" d
15
7 n) }$ X! @( `16! J/ _1 Z k4 s: N, I' S
17
8 T! N( C5 E# X6 D( Q18
! h0 W! c! s; K9 X- J9 e" y. D( M19% D* _9 i* P- k" }, A/ N7 x7 n
20
4 z0 `+ o& n$ C& O2 E21
4 `; e" Z$ h1 D$ i+ c- [22
. y/ H9 H6 W( E8 b/ h23
5 z/ Y6 \. v2 e' m7 ^Params to learn:
4 G8 c% j; `. F$ F) p. @ fc.0.weight1 E& H, B2 @3 r) m
fc.0.bias) m) D. K {' S# {' y6 |& |' }
18 ?$ L, f/ a+ x1 ]+ S- K3 ?
28 Q3 ]4 r& R: M9 G7 @
3
1 R1 y: C, r! g0 Y/ i7. 训练与预测
! E3 N i ~& }% F7.1 优化器设置2 D+ r2 ?/ B$ l+ ?
# 优化器设置
# u, S+ z! _1 _optimizer_ft = optim.Adam(params_to_update, lr = 1e-2)
* r% e' M1 c1 A' J# 学习率衰减策略: _4 _: S6 P$ o3 a) W
scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)
/ h# {; c9 J8 d7 @# 学习率每7个epoch衰减为原来的1/10
4 a4 K" m% ~1 u# L# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算 ~5 I2 ~4 ]3 x. C4 K
0 c" }4 p$ ~+ K* r4 Lcriterion = nn.NLLLoss()
3 N+ J+ J5 y* i3 M0 ^' Y& i8 I3 A1
$ f. x* J* f4 n, s4 J5 d% K; t2
3 D9 n6 U; |' q: b. D3
7 v8 g* P2 o5 `% [9 v; h5 M8 X4
2 o ^% h o& E2 O% E2 K* f5+ I( j& v0 i, t2 B
6$ E C( r% U8 j
72 N0 W. @4 H4 q/ T' S
88 ~* s+ M9 }+ G r# e {' D+ r2 P
# 定义训练函数# U' T; W2 Z: i* V k( c
#is_inception:要不要用其他的网络
, z0 d! ^) t: K' n. l5 Udef train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):, C3 k( p4 ^5 V8 |5 w6 `1 X$ m; u- Y9 O
since = time.time()
) ]8 ~' |( K: ?( `- s #保存最好的准确率
* I# ~1 L. Y8 @ B best_acc = 0
: a: f3 K* A8 t2 _4 L! r """
# h* k3 S9 Y: C checkpoint = torch.load(filename)1 z- m$ h0 i0 i# R
best_acc = checkpoint['best_acc']
$ |# m+ c0 F. C2 ? model.load_state_dict(checkpoint['state_dict'])/ H j4 O8 A+ @3 R- S
optimizer.load_state_dict(checkpoint['optimizer'])
9 P6 P4 c3 N3 U model.class_to_idx = checkpoint['mapping']
( f" t' X0 e7 M d- d """
; i! _/ G& v4 H7 R( p #指定用GPU还是CPU
& B7 Z$ t: F/ l( m model.to(device)
( U7 p, @/ O2 h3 ~: B #下面是为展示做的
, d2 ^& J6 X+ s8 F val_acc_history = []3 R5 K" m/ K. d; f. r, u$ Q& m
train_acc_history = []3 p4 M y0 Q! Y$ Z' t0 U
train_losses = []* F$ C9 }( f2 ~4 k) c9 O/ d% E
valid_losses = []
# L- `" C/ g( i- T: c" n6 e LRs = [optimizer.param_groups[0]['lr']]2 B) M) H& ^9 h( ]3 `/ b! \
#最好的一次存下来) C" t( w9 `5 T4 t
best_model_wts = copy.deepcopy(model.state_dict())! v0 V/ M9 @* M" R
7 ]( z3 x! L+ V# m! G- e
for epoch in range(num_epochs):2 ~. I; [7 Z. Q/ b) l8 E5 O+ t. E4 F5 t
print('Epoch {}/{}'.format(epoch, num_epochs - 1))
0 X4 C, R9 s5 q+ F$ y3 W print('-' * 10): E. k, w: |" C: D: Y5 q- r
6 I# l, p3 m W7 d& r" m$ Q: x
# 训练和验证# U8 {* Y1 p) B9 z; j1 c: Q% T4 o; h4 D Q
for phase in ['train', 'valid']:6 o3 A, @! Z. _3 w$ C
if phase == 'train':7 P( f" k- w; [# X9 f6 a3 x2 S
model.train() # 训练
2 F5 Q9 t0 A& r% z else:
6 W! H: x3 _ m/ Y model.eval() # 验证
0 |3 o. H% n, b) e/ ?. y
1 I+ V, k& m- P" [; X3 J9 j( | running_loss = 0.0
" w; X2 |/ S9 C running_corrects = 0( d! G+ n; b) | q. Q) v
# H, A$ Q& h s! u X
# 把数据都取个遍
. {8 q, N+ ^6 x for inputs, labels in dataloaders[phase]:
. M; p3 A% W! S7 b #下面是将inputs,labels传到GPU
# T6 n0 ~. T# F2 j0 D inputs = inputs.to(device)
$ G' ^* L {, I1 P6 k labels = labels.to(device)
6 p2 d* M, A1 z/ j8 N- |* _9 ]! y8 M; D" ?
# 清零
& a9 j6 X: u: ?* M! a9 o optimizer.zero_grad(), \& V7 N. U5 {: X4 B
# 只有训练的时候计算和更新梯度
. Z& _- p+ g8 Z7 q: \ with torch.set_grad_enabled(phase == 'train'):2 M6 ?# s6 ?- d* f- H( I
#if这面不需要计算,可忽略
/ \. \. e5 t @" G' z& E7 o' C, H' ^ if is_inception and phase == 'train':
% K' `( a5 t5 b4 p4 a$ J! P outputs, aux_outputs = model(inputs)+ B) o( Z- G5 E( b3 z6 r- R
loss1 = criterion(outputs, labels)$ K' V9 J0 ^% r# F, `: x
loss2 = criterion(aux_outputs, labels)5 w& _0 G; d1 Y8 L& c
loss = loss1 + 0.4*loss28 Y8 L e' [ j; v- [! g
else:#resnet执行的是这里
7 n7 e; F- w4 u3 v0 F outputs = model(inputs)) k# j3 q" l; G' W$ M
loss = criterion(outputs, labels)
7 s5 }7 d# V/ E% h1 y7 }) N/ w2 A! L
, H( @# N: R Y Z# A$ L) w0 B! m #概率最大的返回preds
/ ]5 h: `7 O( @5 U3 }( o' D% H _, preds = torch.max(outputs, 1)9 w: k* W" |4 x% {: f- e+ b' N0 R
. N V0 y- w7 r6 N
# 训练阶段更新权重2 c Z: v( h- s1 r
if phase == 'train':" ]1 P, T$ q7 z
loss.backward()" ]7 b' B0 w, x" c
optimizer.step()7 p8 n8 e' R* {' J; {% _
# A5 T8 F6 r8 U0 U* R
# 计算损失
j6 P' B2 q: q5 Y9 A) A; B4 t running_loss += loss.item() * inputs.size(0)
4 O0 N8 L* r; H* K% K' K. o running_corrects += torch.sum(preds == labels.data)& n2 e5 V F* v2 q( w
; e% e8 b$ O/ d. ^# O# } #打印操作
8 ?0 a) T& M! Z I* I epoch_loss = running_loss / len(dataloaders[phase].dataset)$ w& b7 g# C/ d7 d
epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)
& X ]) |4 i! y3 P$ ^6 }7 R1 S% B4 `' \, Q
) A' t8 }4 U, S; z time_elapsed = time.time() - since
f3 z! Y% a6 ]5 f( _5 _ print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
& @+ i T) J( m& Q6 D; H print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))+ R' P6 A% F8 M. w4 \ C
# U$ l, |: k B% A/ o
8 {# c- ]2 j+ @# D) i% M # 得到最好那次的模型
4 d( N" w9 E- e: t- ~) l0 Q if phase == 'valid' and epoch_acc > best_acc:2 r, |4 `1 W5 s0 g/ B" U% ]* ^: w
best_acc = epoch_acc+ L, i* M h: l7 f
#模型保存
& m& i# `) j8 M I: l% a best_model_wts = copy.deepcopy(model.state_dict())
) b/ w5 c3 L8 D* g" M/ q state = {
& [7 ~0 U0 ]: i d$ T& ^ #tate_dict变量存放训练过程中需要学习的权重和偏执系数 ]! K6 C) r6 j; R y9 C1 E' G
'state_dict': model.state_dict(),* ?+ u8 H9 l! \
'best_acc': best_acc,
# D6 Q* k2 @& d( T; x 'optimizer' : optimizer.state_dict()," Q" I5 `% P* [ Z( o1 ^. k3 l
}3 P8 S3 _6 I* w4 l) j
torch.save(state, filename)
: K( C( ?8 W( v2 k. F5 h3 w if phase == 'valid':
4 N5 k* ~5 Y3 e4 t& z8 U val_acc_history.append(epoch_acc)
4 L+ ]5 t: w7 f1 J+ M valid_losses.append(epoch_loss)
1 K; w5 u8 v5 \ scheduler.step(epoch_loss)
. c2 q* |( r { if phase == 'train':. s& @4 }4 v8 X# q9 J
train_acc_history.append(epoch_acc), G& o3 C$ @) Z! X) d. o
train_losses.append(epoch_loss)
$ [, X5 V- M4 R v; f
2 Z* s9 V* u: T print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))# b% w. f' N' e6 G, n
LRs.append(optimizer.param_groups[0]['lr'])
( ~/ T" ?5 @- P% ^5 Q9 J% g9 j print()
$ W% H, T- }, j3 F+ E2 ^: ^& O# N/ F+ l" k" J* G0 v
time_elapsed = time.time() - since1 v$ D. R# Q/ J4 Y6 p
print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
6 H% j. G" m, X/ @2 W* F$ u# V print('Best val Acc: {:4f}'.format(best_acc))7 R" n% r# V3 B; A0 k
$ B4 y# Y/ M+ \. J% A # 保存训练完后用最好的一次当做模型最终的结果
/ c5 Q' g/ T- I# l7 A model.load_state_dict(best_model_wts), N5 {# E+ @) G5 F: ]% e
return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs 2 b1 [" d/ Y }) F2 Y
# L9 _" [7 U3 f' J
2 v6 ]' J L% s1
8 ?, d- K1 i+ e! a! B! p# H2/ A' e; k$ F2 P; L) L
3
$ p w' T7 A P4 e% k; V4- S: ~( C, G& s
5) T% A) D; ^6 `- f' q
6
1 u/ `' \7 T, G7
) W4 e# |2 U& X+ S9 ^7 @- f& W8
& |- A2 ?, \" E$ p91 u* O, ?: L' r3 x- X
10
4 B4 ~/ n* W* l* c2 o$ h11
8 P0 H6 ^) q! A S% }12
6 u2 J4 F. `9 h" G13
Y* N1 X( g+ F9 X+ G0 X6 B# L14
- G0 O! T# ^2 a* y( l O9 M15
3 h$ ~9 g. t; O7 n* i+ W2 B160 G0 J3 j( F3 ?) G! V
17
( R( i2 v& l: D3 c6 V188 i9 g( T2 ?5 I$ M2 x4 ^" d! h" ?6 F
19
4 x9 n' S( l- E1 m20. O% G+ v. Q: Q) B* `3 ?4 Z8 X
21, E) Z- |3 [8 Z q6 R2 M6 r
223 ~: x* _/ M1 u/ G
23; k& W, n8 N/ X" q0 ?
24
6 k$ ? s2 H/ d/ a25) u; E6 A+ Y3 q0 m/ R0 G! I
26
! }/ ?! H* L$ i2 N27
/ p. u4 y: z( X# X" G/ X5 {6 E28
* |. Z5 f7 X0 K5 k r- d' G292 I6 ^" E/ A f& o9 S1 Y3 Q
30+ M2 R, F4 X+ A' G7 O' Y! r
31
5 `+ M _! N1 L3 I32
( C! k* M5 m0 A% ~; t33
1 w7 z9 E$ K; B! x% W34; G! Q& t1 B; G% T8 Q' M2 M
35
2 ]' F$ ^5 E( b36
7 ?8 d; w6 @, L) f1 F37
7 N8 p) @0 t$ w( y' @. z! J38. K& j' ?# _) O
392 n, t* t, h9 e) b7 \
40# S( I/ t( ] ]! N
41
8 K% N& A7 P8 ^3 N" y. m5 b42' M2 r B' B r t
43
! a, D* [- y" u i- w9 Y3 u$ G44' S* q3 v; x+ C$ e& p) q' k; s
45- ^) w5 B9 l9 T! D: @) Y3 M# h
467 p" \) p& e3 H5 k( y* a, g, l5 n
47
, A: Y' n8 a3 c, X/ o48: M8 _5 S9 s& y
49 v) U. l! v/ a0 z* f
50
W2 H$ i r. e3 S5 l8 v51
4 m5 j3 Y# O: B526 t% h I( Q+ v/ d- N7 r
53
8 F: e; \7 {) F4 {0 X54* \- z8 w4 ^6 ~
55
+ }' C" C6 q! k3 T56
: m1 y# j# e$ n6 ^3 D3 p K& q57
8 [6 d& E8 f' y58
. M1 B& g- {# n" {8 S59
' J8 C9 H8 q, z$ Z$ j60/ {5 S; q( J/ \
610 R2 q1 `2 O" K0 J3 ^" {
62
. K2 T6 ]/ a% q6 ?63$ g ?% g* P) \- p* f
64
# N- d# Z y2 x" `% l/ g! Z65
/ ]' b3 X- d3 V7 V; K66% T) J! H! q5 [( j, ^% H
67
% P5 a$ l ^1 L- F68( }, w$ \# o8 Q4 c& B. V
69
8 T7 y! R% d( r) t4 |# f70
" p& U& b2 X# |9 i71
# H, g8 M* ^" e. V) C72
$ Z7 N- ?1 ?5 }2 q8 S/ t" G73
$ {, N, b3 P7 f1 c74
" A s% E4 C3 q75
, V7 f/ y: j/ u1 Y* P* U- z8 w76
4 r8 s. L( w2 a; k# j$ q77; h. J: I. [# t+ f
783 X! u& C. G6 Q# I9 Y
79. P$ S; j, t+ W( ~% S
804 b7 o' w- I% T6 u# f$ `' Q; D0 `
81. d+ B. X- l0 Y- ~
82* m R0 P0 }4 o
83$ _) c# b7 Q: k2 Y2 C
84
0 f r, w8 q4 y( w9 p! _/ @: q0 D0 Y850 E) t, F5 U) B, Y9 W
86) A0 Y' m$ w, |& j6 {6 ^
87
# s0 o' _' H d0 z ^" D886 |8 a, a, i% E% R+ a b: @
894 @3 t% G, z- w- s5 {
902 Y4 c, M4 D" X, [8 G/ B- D: |
91
7 L9 D/ ]3 z5 ?* t& L) h92/ N. p- B0 m1 R0 |7 ?; ^
93
5 n+ Q) ?2 K+ O3 g6 t# x( I9 \3 a94
Y Y R: c( U% S( _# n+ X4 P; L95" @2 R# N* K. t4 M* ~# ^
96! B) ?* f& |6 M2 b% W1 b* I
97
" q! Z, J3 s7 k" o# o98
$ N& Z/ G# \# u+ a( n# ]99
' C" h0 k, }9 n100
$ z" c/ L, J! _& J3 c101# F3 s1 i3 P6 f( ^+ z
102$ k5 t* K: [$ x, x& J; V* @% ~
103
C5 K5 B0 C4 V- N7 P, z104
" m) w! F% K! d4 i% P5 e2 N3 D105
* e$ e- e- N! J7 y9 [. T& C106
7 c3 \! c" @* F# ?) w107/ c, g$ Y9 j8 E0 i
108+ ]' v. T, }0 i) z2 r4 U
1091 v# C( r1 J) z$ O' L( b; j
110
8 V9 @! W: H L# |1 A( c+ p111
, u: }6 q9 t9 a. u112$ p0 {% K# }( p) ?3 v( `0 ^# x
7.2 开始训练模型
1 e, o* y p+ P6 c' B: j [我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次- {( ] Z4 L3 v( U9 P( Q/ N
4 _3 R% V+ I5 U1 n8 D: a# O: t#若太慢,把epoch调低,迭代50次可能好些: P3 z W$ f7 n, O5 n* Q
#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合, G* c: W' c3 l: G* o- D
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"))% ~9 o7 F+ C- W" N7 s
! Z! v; Z4 y3 Q# ~# D4 F
1
( |& `+ J( h, \) C2
" u9 D0 }/ i) t( _$ ~# }. @3
6 z7 ]) E0 N8 Z2 d n4
, J# S) E# T% ?" D" \Epoch 0/43 V4 K! _2 z6 E, P" u6 {
----------0 a: p, v! b; h# W; y# H" B
Time elapsed 29m 41s
6 v- K% r* x- }% m v: g7 Ntrain Loss: 10.4774 Acc: 0.3147
3 A# n |7 \5 V+ P- E4 pTime elapsed 32m 54s5 e5 {+ Y% t, g) y+ F
valid Loss: 8.2902 Acc: 0.4719
- h" R" n- S0 m. f! JOptimizer learning rate : 0.0010000
Y" L- W: }. S3 G: m: v: r7 d% b& E+ D v# V- F% o' K
Epoch 1/4# p0 |* L" ^* [8 I4 n# J
----------
! [7 E2 K0 @( ?6 zTime elapsed 60m 11s2 i) _( H. o7 d W
train Loss: 2.3126 Acc: 0.70537 W9 }& L1 d3 u* k: N" r( W
Time elapsed 63m 16s
8 J o/ W+ _; I7 O4 b8 p% vvalid Loss: 3.2325 Acc: 0.6626
6 c# q" ?- t7 z6 NOptimizer learning rate : 0.0100000
) o; b8 Y5 H; ^1 [; ]4 p, C* T# @$ B+ r/ `5 @. j2 C6 s5 U
Epoch 2/4. A1 }- C9 g D& S
----------
& U! B; z! c$ v% D* K5 D) [+ M5 yTime elapsed 90m 58s/ x* v1 [/ [3 b. r% t& B. [
train Loss: 9.9720 Acc: 0.4734
$ _, V* t7 y& u$ n4 t- K9 p2 H; R- hTime elapsed 94m 4s
/ W) O; r& i6 i% l, ?valid Loss: 14.0426 Acc: 0.4413" N( }1 m: N4 Q1 a1 K8 {# X
Optimizer learning rate : 0.0001000/ J. c: O* b9 t+ {) O/ A+ A) G$ Y9 c
3 Y) ~' j( C" b0 y- |, _Epoch 3/4
- S5 r2 L0 m) A8 z+ t; b----------
J- R; E$ N, o5 ~1 _8 y! K$ ]Time elapsed 132m 49s* ]$ V& S U) W1 j& n: c" y# d
train Loss: 5.4290 Acc: 0.6548# A" o4 P- K: i* o, E& y& K
Time elapsed 138m 49s
) m/ @2 R) l4 w) n7 G) r _valid Loss: 6.4208 Acc: 0.60270 h/ {2 D9 L. I
Optimizer learning rate : 0.0100000, B& e/ y3 j- k, E& K) g
) H1 K, g: I: P$ [Epoch 4/4
1 ?0 H+ d2 z2 G/ O: @& T4 v- c4 }9 o----------1 k4 ^4 i% @5 Y! o8 _* p& b
Time elapsed 195m 56s
5 A. H, G- \- |" Ntrain Loss: 8.8911 Acc: 0.5519
. Y9 H4 ^# K R4 S4 C- A: \- iTime elapsed 199m 16s8 G% C' c/ R4 U2 `7 V
valid Loss: 13.2221 Acc: 0.4914( F# J& @* t5 e/ Q
Optimizer learning rate : 0.0010000# ~1 K8 s9 G% H/ w6 ?% \; d+ h
, q: x1 @* F$ s3 ~Training complete in 199m 16s
9 w/ A" R. Q, W! C7 E. h! tBest val Acc: 0.662592
, ]4 W! @ l7 u; X# ?# V# j2 H( N; {1 H; o$ s# F! D7 O
15 J/ p: W8 H9 ?' S( E
2
9 H% ^- A$ x T3
- t6 q' L; B2 \3 q4
1 p. ~: d e( j. ?8 H, g5
: p9 u+ @& t( C1 B- N9 b5 A7 o6, [9 u2 Z3 D' z% ]+ D
79 T! t3 C) H7 Z0 [! H, r" s5 {
8/ H3 t- p ~! U0 l/ j& W! d) E# G
9
3 e: L* ~8 H% K+ A102 r9 s5 k H. `" H
11
2 C7 ]. [) s% e0 p12
+ i. F2 I- `5 b$ R138 Y/ y. y( x8 a; S) ]$ T& [
14
' h- r: u4 O5 |' j i15
& l: @' W$ g* m. E16
3 P0 b& {5 t$ T17" ?; a" `0 D# @
18
. I2 v8 J0 k+ X* U$ m. Z! Q1 s19
" n3 H% @4 Q$ C! J20
; N: Y# _# _) R# k' g r( D2 i" Z) Q21& L& y# B1 g; J5 O- Q
220 P; o* h# I- r* I" @
233 k4 W2 e: T) N7 S3 I
24/ C- F8 |' C7 n& _4 t
25
' I' |4 {0 u) v v26$ P8 ]# |, ~) o8 L/ A0 ^" K! e
279 Q+ O! H; ]; h/ r8 W+ q5 V
28) ]7 Y' x, r5 m4 e
29
! F7 q8 ]3 S6 I& d/ B4 a30% r G+ u5 b& Y, u
31
/ w2 S) R. n$ f$ S& X32% ^) y& ^, \: x B# R
33
3 I& A: u+ r& ]34$ q( K$ {1 ~9 W( U0 k8 x; V+ U5 i
35, `# g6 ^- |1 T2 ^+ ^- [8 z
36
* ~6 |: R- o6 U9 H37
# S/ \: z1 i; O$ }38
) Y* a3 M2 C, \& C) w" u4 ^39" z# J& o+ x9 E5 \# d8 M
40
5 R; m9 e' @5 o1 z3 I, v41
8 c1 V" D" V) G42. W; @2 k8 k, r, X2 H+ s
7.3 训练所有层
b" K7 p r/ D! q# 将全部网络解锁进行训练 B5 R" w$ G5 h% i/ n
for param in model_ft.parameters():* h* p; n% D( k) ]: z
param.requires_grad = True! O7 i) j' r# B; ^
! i, N# @ T5 t- P* R, q
# 再继续训练所有的参数,学习率调小一点\& O% o( U9 y& [+ W+ b5 Y
optimizer = optim.Adam(params_to_update, lr = 1e-4)
' }% T, A( q) I; Hscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)
4 x* s: B9 z# b: V7 {2 a- _$ i% n$ m6 R# D) R+ E: K$ H Q5 L6 x
# 损失函数
$ F8 ]0 G0 ]& b: s1 ~& a3 [criterion = nn.NLLLoss()
8 D: ^4 z- M9 K13 {$ w. q4 @. {) a9 h( Y
2
7 ~3 i5 A7 d! z; B3
U/ }* \, B8 T4 m& X" H' H. d4
; z5 Q0 y% ~8 _56 ]2 O, C u. I+ ]8 X: Y# Q' A/ K& v
6
$ A- b% L6 ^' J# a72 j0 k1 v) d6 g
8
5 h$ i& G: y. T) J4 N$ {9( V4 @. b" M5 \& ^' D
10
- \3 M. h" W, q+ Q9 y; J$ O# 加载保存的参数* B' R n! Z. h
# 并在原有的模型基础上继续训练
" H6 w; A' V" i" f6 z3 z# 下面保存的是刚刚训练效果较好的路径
5 |- M1 f& e) B O l1 Y! o% J, L7 V$ hcheckpoint = torch.load(filename); I: z+ I( l! J. D
best_acc = checkpoint['best_acc']8 }7 N2 d7 W! x$ X$ [3 {
model_ft.load_state_dict(checkpoint['state_dict'])/ }; X* ]; P- o1 F
optimizer.load_state_dict(checkpoint['optimizer'])! T: `2 S8 `, b
1
7 a$ N9 m6 w$ t22 W/ h" d4 w6 ^6 |" o
3
' a7 U' {+ \& B41 J# I3 E5 w0 J! E4 B
59 q3 P% ~+ V2 L! A$ B' l- s: U
6: A. V( Q% k0 N1 m+ Q- M
71 p: Q$ w3 s% f
开始训练( G) M- I$ q5 ?3 |, @, H% w
注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考
7 t4 \( X; R F, a: S7 L x5 l* X i/ Y, s2 g' x6 z
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"))# m) F+ d& j5 T: J3 w
1" b4 u D# Y ]6 P2 z7 C
Epoch 0/1- m* I$ o# b' Q# s
----------5 { g( I# F( H4 _( K' g7 @
Time elapsed 35m 22s
6 }, q' ~/ m# mtrain Loss: 1.7636 Acc: 0.73463 A/ A2 i5 W' K9 z) u0 ~7 X1 U( s# e
Time elapsed 38m 42s7 a9 o1 v: T3 C1 y: q
valid Loss: 3.6377 Acc: 0.64556 Z' c+ g* k B
Optimizer learning rate : 0.00100002 D# d" e w Y: s: R7 P
, Q. v. U6 a6 o
Epoch 1/1
2 P+ Y( ^; v5 H2 w4 s----------8 E& m9 m, A$ @
Time elapsed 82m 59s3 J) T& N7 J, D2 Z' G: V8 y0 b1 L
train Loss: 1.7543 Acc: 0.7340
q$ c! |; ^+ Q7 q v0 o1 qTime elapsed 86m 11s
0 x! z- ]; l2 w3 avalid Loss: 3.8275 Acc: 0.6137. \ w7 ], |0 C3 h. ^5 T7 Y
Optimizer learning rate : 0.0010000
9 |( _$ g D" g( [- m4 b9 U0 A6 {
: k9 C E* F3 ATraining complete in 86m 11s
0 A6 b+ q2 C" C1 M$ |# ~+ }Best val Acc: 0.645477
9 \8 E) o! v" x2 l9 a2 C9 L8 n! w) \# W" q( B5 A
1
7 k6 J- W1 V% V9 u6 b2! J0 N1 j6 k) e' [8 `! J- }
32 d# b! Q- u" I" g+ ?$ o; Z) |; G
4
: c% p8 P4 k Y; O3 Q. r8 a; Z5
4 M4 K$ ~% d, q8 k. L0 W4 v1 {6
' T% _( D2 O! |: R7
; \7 }0 V+ B& \* I- ^2 ?$ |4 k9 e5 c87 l) w- S5 D; c
9
% {" ~0 _ i* J1 u) a2 h5 j/ V108 D4 Q/ a: x$ t% w, n( l/ L
11
( G0 y8 j- s( W1 Q) P) D8 u. q12
7 `0 {1 i$ o: N: E. P1 G& o# Y2 x13
1 U( x. ^: U) y! T7 R& B+ c14
) w$ U! V, d( D15+ x- q9 i9 p$ m* M
16
! j, b2 m! ^! A5 g C, y! V17
; L6 F9 v* ~ x3 {18
; B0 L# E6 z4 J+ @* B+ a, ~6 z8. 加载已经训练的模型
& e- j! K5 \8 y6 ]2 B! L3 S# z- i相当于做一次简单的前向传播(逻辑推理),不用更新参数
/ s0 s! ]% D+ {: _' ?% E( V1 b! G. M) [0 ~% ]+ e7 m8 T6 n$ \
model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
% j- M8 O- E0 M" I) |+ K( x: b2 K; G; H+ i( ?" e t: {
# GPU 模式5 M6 q0 ^# V/ c$ N
model_ft = model_ft.to(device) # 扔到GPU中: [' \( q; E! k. \) f1 ?
" Y! z; R! @: h7 ?0 ?+ X' ~3 X
# 保存文件的名字
/ }' @5 Z0 P7 Z5 t2 Q, f# [; I: d% u) @4 Vfilename='checkpoint.pth'
- J5 \8 E/ S* U2 P+ C0 N( ^) b1 p! |1 O( o
# 加载模型) w* e, i, s# {0 d3 _0 ?0 B
checkpoint = torch.load(filename)9 b' I( _1 a/ P8 o+ @9 S* p
best_acc = checkpoint['best_acc']
9 E. i* U: c2 q) D& \9 O/ Fmodel_ft.load_state_dict(checkpoint['state_dict'])$ L$ ^2 |: H% a
1
/ P. }4 g c) Q$ l; S2* ]$ t, ~1 r+ o [6 d& O! b1 U
3
0 k4 _% Z* n- V: |& [4; p' } F# x: Y, d* k: ]7 p
5. b8 V% k- b. H! m
6& j3 C7 w2 u; D7 C* r/ v
70 [ x% O& B+ M) `$ A6 C! z4 b
8! J0 O( h) \3 g0 N1 [
9 F, x+ G6 N9 \% V2 d
10
! w- e0 |# n: V8 w" d11
5 ^6 Y2 x0 r* |5 z12
3 e1 x9 k' G+ }8 M7 v' E- h5 @( ?<All keys matched successfully>
: ]9 R" g" `( T1 v2 j& K4 G1 w. ?$ r- N
def process_image(image_path):# e; i; a8 |1 T3 e9 x
# 读取测试集数据+ S9 M2 q1 X! z6 ?9 M( \* y# P2 ]
img = Image.open(image_path)
) w: k6 ~; B" D # Resize, thumbnail方法只能进行比例缩小,所以进行判断. _$ r4 P. t* w: D
# 与Resize不同9 u! ]% K6 ?7 x# v0 o' Y5 Z
# resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小 S1 `8 N. Z2 t$ M- T) J$ w
# 而且对象调用方法会直接改变其大小,返回None
6 u/ m- J$ W' w2 Z7 [1 I7 V3 W if img.size[0] > img.size[1]:
' k. E. z/ l$ O img.thumbnail((10000, 256))/ I5 {) o4 k4 M l/ `, {8 h
else:9 W6 |& u$ g; F4 g# ?# t
img.thumbnail((256, 10000)) m1 x6 G% N# [3 F$ ]9 \. J$ j; X
. l/ U4 h# T @* ~+ x9 @$ f8 ^; p # crop操作, 将图像再次裁剪为 224 * 2248 c7 T4 G$ E. |
left_margin = (img.width - 224) / 2 # 取中间的部分
+ }. o' h# d9 M m bottom_margin = (img.height - 224) / 2 I1 X+ p. b8 \: n
right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度
2 o6 ~. E/ u* ]; y top_margin = bottom_margin + 224
$ R; v5 L1 s4 w& D* S7 p4 {( r, F' b7 p5 k$ {1 Q! J
img = img.crop((left_margin, bottom_margin, right_margin, top_margin))5 v* y: |8 W4 @7 _0 p6 a
- Y- |: H$ J0 u$ ^
# 相同预处理的方法
' O3 r3 R1 o4 q, b # 归一化
0 J" B* O8 d- j% O, l img = np.array(img) / 255. y! }) P$ L I6 Y- P
mean = np.array([0.485, 0.456, 0.406])
: j; c* G& q. A+ j2 u std = np.array([0.229, 0.224, 0.225])
8 r3 T% T! w% c k img = (img - mean) / std
7 C' A# M. d% X9 ~& v. ^
- \6 I& {5 U5 J, D/ c: r # 注意颜色通道和位置3 L- H W7 }. q4 v$ U& `3 q5 o
img = img.transpose((2, 0, 1)); V. v6 G8 e; M A; [0 }5 q( j6 v
8 Y% C2 u& I5 X1 \ return img
* I+ `0 p; i1 F1 i$ {3 H0 \( l, j/ S: g6 g4 \6 } F) c. F
def imshow(image, ax = None, title = None):
# M3 Z& f7 y* I6 ^7 r! R3 G """展示数据"""
5 W; p" n1 R, N& D/ v4 g6 d( P if ax is None:
! x. o$ y& E8 i fig, ax = plt.subplots()
+ i3 w7 z$ A8 T, J P- g' h7 Q# d6 @# Q7 e& N
# 颜色通道进行还原
, _6 |) Q1 N% p" z1 r0 J% M: t; D image = np.array(image).transpose((1, 2, 0))0 t& U9 b y. p* C& a4 v
* P! N# A6 d$ c2 F, C, m
# 预处理还原
: w; l: A& r! |$ j9 A mean = np.array([0.485, 0.456, 0.406])% y- ^( |; k- S! E& `
std = np.array([0.229, 0.224, 0.225]). {$ l& G9 _5 A0 p2 i
image = std * image + mean: U' d% O* h$ C# c( L2 a+ [' `* c$ A
image = np.clip(image, 0, 1)" U" ^9 g# g K/ Q
/ a/ X+ W. ?4 ~9 R0 T( i ax.imshow(image)
# Z, L, @! t2 u) ^9 t% L# Z ax.set_title(title)
) y! E- |9 \5 e5 J' \: ?/ ], E* V1 f) o' d0 u6 i. o
return ax
6 Y# a6 m" t0 f, |, y+ u) a% g7 c- y4 _# w( p2 Y$ f
image_path = r'./flower_data/valid/3/image_06621.jpg'$ M0 d4 m& ~+ ^
img = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理
/ Q7 \, V# r3 \4 p0 Y- Gimshow(img)/ U4 ?. g' o$ x( l
9 I2 |8 \" J* G) T, j( ]1 a1 s4 j
1. u; W; d0 n+ C
2
& I6 p7 l# Z, h$ Y3 q3
3 a) C. l6 |7 u$ \4 x7 ?4' D& J# T. B. M: ~) N ?+ C. c ?
58 V: X2 k+ q2 s6 R. u: H1 i( \6 Z$ _
6
; u& E+ W% T" M( z" W7
" T/ n1 L7 Q" Z5 l% S1 o8
* Y: y7 S0 m" r2 v& u, X) p2 D% v( ~9
( ~( c X! T0 {' M+ x10
|4 C- U D( {# v* L9 D- W3 L% C7 X+ l11
, Q8 T. l. i2 i. D: y) ]& \! X1 c8 u12
0 o& i* a8 \& J" C# ~7 ~, k13
% t8 d8 c! O3 K2 }146 V# U( E/ M" n1 z; _
156 k3 M1 Y9 n5 W9 b h0 j
16# x3 s" q6 j+ i/ q4 E5 i
17
3 } e; Y; Q o) G7 b18
9 j# [- S4 q- w9 C19
7 {/ Z: d! ]$ Y8 S g+ g2 p20/ M# h0 t9 O( e3 _+ X4 v5 C
21
1 W& }- `& v$ |! _3 S% N+ t22* X, D: z) k- s: ~6 u. b9 x, K
23
% i% a% O( a* E24. L, U* i$ U2 _' K( g' P2 n
25
+ C6 |( O; v/ N6 P' I4 p+ Z' @% g261 _. _, h7 q7 b7 D! U. e
27
# S* c; |- D: |2 U28
- [. |- g4 J: N! A7 P& x# p29
- j7 I8 v% s( J30
X5 L; N9 a" t- }7 e' X31" Z) ]* R+ f( ^) z6 A( _% T9 Z i1 R; R
32( Q7 O7 C* D8 z9 W
331 b! r7 @5 q0 H7 w- C, L
34) ~4 {/ S) r. a6 J! W1 x6 X, U
35$ c- D' ^. v8 J1 e* G! M
36* R: O& G8 F$ @6 L5 m
37
* I7 [: m2 I% _. s0 `38
4 A" T$ h# J% [: D& S4 Z39: V# n8 h. \9 r4 F6 I" W
40/ {5 l0 D9 w8 f3 F3 t6 k5 r! j' e
41
% e& C) Y1 K2 a, K2 x; Q5 N$ h42/ C' N' C8 k6 |8 p3 s
43( g7 T/ ]! K! |* k: v, g
44# x: H2 q8 u9 E M3 ~1 R
45
2 y/ h, }5 h8 Q46& {5 J; c( g9 R' M/ X. B, z
47
9 R% j4 R( V8 ~$ V48
4 H, [8 L' H8 Q8 o0 c49) T9 ]- h1 X% C, J# ]
508 |/ b9 f$ n. `4 x K) ?9 A
51
! v, p5 l5 B6 C+ t, P52
5 G& @3 g4 Z& `) }, M53
0 {/ \& N# X2 t% E* a54
" @% t5 s+ S* C7 ~8 `6 U+ o0 @2 ~<AxesSubplot:>
U a7 I. }4 @* Z, I. R17 c4 \- F, r8 f
2 N6 R* K* U$ D7 } k2 s0 R8 V上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确& f3 R' [1 _1 Q5 [# y3 U9 A3 k
; p1 s6 ~- I8 ^! C. P# i
img.shape8 q4 z# \& e1 P P
10 ?* @) m3 s. V% ^: C
(3, 224, 224)
1 \, M& C* I j6 l2 }# h11 \6 O% g- ]4 c' l7 _9 t
证明了通道提前了,而且大小没改变
3 y# s+ f5 U, `& `
$ [* {5 _/ l/ W8 a/ \) e; v2 u9. 推理
9 h' R/ W* t4 o' }& G" ~img.shape1 p& r6 x+ w/ b
* k/ O. F9 L& K) u6 M
# 得到一个batch的测试数据& {" m5 e) `& t, Y0 ]2 N
dataiter = iter(dataloaders['valid'])
, Q4 a2 t/ w$ e$ r- e- O: b. l5 Kimages, labels = dataiter.next()
" F8 H/ F2 f: B. m8 h# g
. H& o# s5 |$ A4 ~model_ft.eval()
( u; G9 @3 l& }/ o( u0 M8 f
; b1 J- W2 f p0 S' U, [3 k. @0 {* Jif train_on_gpu:
4 Z, h# s5 e# }% Q, b3 k4 X # 前向传播跑一次会得到output
; O6 M0 S" h. q7 a1 c7 w: J; ` output = model_ft(images.cuda())
9 T) z5 g* [8 Z: x/ ]else:
$ H6 v1 k+ V# m5 N. [) X! G6 E output = model_ft(images)) L% \& X: B) G/ ?0 _
) Q- F0 y# s/ H8 L
# batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值8 ?1 j& [" x- Q0 W1 \
output.shape
: ]! x. ?. u/ u, Q# v; f0 l! M4 \
1
" q9 T$ q# l8 v$ f* Q" F4 q2* c2 Q" @- I' t% B" N& @; m7 Y" o
3
: @9 L5 ]9 l6 `) V4
( M3 q1 J, ^7 {) f! e( ^5
; @: \/ b4 Z# ~6 H, ~) C- i/ C6
5 P! \" B, o) b# C. `( O+ C7% Y$ c$ H% o3 p) S
8+ F6 D/ ?/ ^3 L: A7 {, \7 H5 `. p
9! b9 Y; [- _% ~! @+ ~
10
3 B7 j( X) r1 e6 ?11
- E1 j5 O: n0 c12
& d! h9 S' E7 ]2 e& ^# {130 B/ h! }, t5 r! s7 u
14
3 E, g( R7 T! z9 x- o) A15. z! H. n: ]0 T5 ?: A: w
16
. y3 y8 x% @( J4 [7 Itorch.Size([8, 102])
( W+ P2 z5 L" H' ~1
- h, ], N4 `# p; l4 G. K9.1 计算得到最大概率7 }4 ^0 d" z8 Z* i
_, preds_tensor = torch.max(output, 1)
1 k8 G% b" \% E" h9 o0 o {( S
: v; }3 o3 l5 |1 [8 V( {7 R$ Jpreds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量
) V( W1 K R4 G# k6 [1 A7 I' E; a8 z
24 X8 d1 _/ t1 I' T
3
3 Z8 k7 j2 }8 A7 L a5 a9.2 展示预测结果
\4 f+ z: ~% ?: [5 t. l' Efig = plt.figure(figsize = (20, 20))' a: e" W4 y2 F& I( q
columns = 4, l: J! `+ t% S
rows = 2
6 ?" l, h3 ~( d+ i
$ b: d7 U% k! G: ~3 Kfor idx in range(columns * rows):
' z0 A: c) Y9 r1 e2 H0 F/ y0 p ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])4 }* S5 x- X! W! s' P
plt.imshow(im_convert(images[idx]))" K5 l7 I T; a
ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), 5 K5 s5 {* m% ^: |* B" m% t
color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red")), M/ _0 ~5 ]# D3 n' g
plt.show() H; K2 i1 @. y: ^- V0 P
# 绿色的表示预测是对的,红色表示预测错了6 {0 }( i2 Z0 Z; a/ D
1 J4 Q% D e1 I
2
% a+ d5 u3 u3 g# E4 Y3
1 R! O: S' d0 U) Y9 H4
5 i% J1 ]" j7 A) R" ^5
( ?6 t! T$ N/ U# z8 ?" y7 M6/ j0 \* R: O4 r
7
/ G1 r u; N9 t* U* D: ]8
: X y U2 h* t8 E" q' W9: q7 z4 q3 z, e* O; l3 l8 a
10
/ _: L. G# Y2 U: v11! D. f% H3 l# A# e6 o+ u: O7 x& _1 b
6 V3 o0 {: ? [& N& N* H/ O- {
, R" i( v$ S& D9 g. M% |' S
+ }( m( u8 o- U) {————————————————
" x% _, u/ S: f版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。4 k5 g# g2 W) S
原文链接:https://blog.csdn.net/LeungSr/article/details/126747940
9 |1 ?; x! P% ?6 f. `) f& t1 Z7 M9 ^2 }7 v" Y6 ]0 N8 M! R3 I
9 m! h4 i( O0 Q6 s* L |
zan
|