- 在线时间
- 1630 小时
- 最后登录
- 2024-1-29
- 注册时间
- 2017-5-16
- 听众数
- 82
- 收听数
- 1
- 能力
- 120 分
- 体力
- 568716 点
- 威望
- 12 点
- 阅读权限
- 255
- 积分
- 175838
- 相册
- 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)实战案例
0 \# M( }9 k& c Y8 ]2 N3 G
) h9 q, s6 F3 v文章目录* x. ]3 Q5 _: j- A6 r5 |
卷积网络实战 对花进行分类6 k* v, s! \$ h6 E6 d
数据预处理部分; W+ A8 P8 F3 h& W. i3 _
网络模块设置
8 ~- E) N3 i! `; o: G9 a4 S# ~/ a网络模型的保存与测试
5 `1 D4 b* H8 |0 F* i数据下载:9 I1 J9 Y, S6 s7 Z" _' r
1. 导入工具包$ Z( c, I: }! N8 ?" q2 }
2. 数据预处理与操作5 B7 a% y( ^/ ^; }. H8 G
3. 制作好数据源! |' W' e5 E( @6 u
读取标签对应的实际名字
# D, A6 F' ]4 a# \ U4.展示一下数据
' G0 N: c; n) ^/ [9 `$ e! Q' `( s5. 加载models提供的模型,并直接用训练好的权重做初始化参数
/ Z2 V* h: D* B' L8 a' a; h' @! U1 z6.初始化模型架构
+ |5 |" ^: O& L* i$ t7. 设置需要训练的参数8 m5 y% q5 b1 s
7. 训练与预测0 t7 E8 |& @% E( x O" x
7.1 优化器设置2 b2 y7 `4 e, R
7.2 开始训练模型
! T! E2 x& I# V/ j* Z- S' @! ^8 j7.3 训练所有层1 b/ i4 G1 O; D6 u$ ^3 N
开始训练
. n' q- x3 f3 j* I* s K+ H8. 加载已经训练的模型
/ v! J3 O# G4 {. c9 j9. 推理' e: }7 Y. G' A+ V% l2 t7 K
9.1 计算得到最大概率3 [: y4 a/ }9 G* Y3 g/ t4 y
9.2 展示预测结果
' b j& ? x1 e M9 a写在最后( i* l; i- y0 q' U6 N
卷积网络实战 对花进行分类
% D7 w f7 Q$ x. j( Q& K, t2 a本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能: ]% b0 j: ]1 E2 d1 _
6 a. o8 H- N$ ?7 r在文件夹中有102种花,我们主要要对这些花进行分类任务+ @! R% ~. h: [, @& i
文件夹结构
6 c( E* D' P* d% Z' G1 R) o; v D9 n) v: o7 `# P7 `. ]
flower_data
$ V- l' I) R0 `8 _$ W
9 A5 X a2 X. }: W1 v* P4 D2 Strain Q/ h6 I, D1 N! }
4 G' O' T& n9 \8 N
1(类别)
8 }, ? o, ~! D+ V6 c9 {21 Z: l X3 @& H6 x" _: @9 l
xxx.png / xxx.jpg
- Z) O) [& Y, F M2 Ovalid! }$ u$ k9 U9 X
6 e! S1 T/ u: n主要分为以下几个大模块 f+ N8 V7 R' I
, r, ]- P5 I! o% f/ G! M2 F数据预处理部分
2 l! f4 C0 Z* b) p+ h" z4 O数据增强
& d! q0 b7 o5 k数据预处理, f/ K4 l* c+ I1 Z8 m9 j4 B0 H: E
网络模块设置
- P2 F$ a0 U% G9 A2 K8 W- q( R. S加载预训练模型,直接调用torchVision的经典网络架构
0 h; y( w, r* Y/ f因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务
3 w. [# l1 Q; r网络模型的保存与测试& Z3 T( J5 B, f, P. B
模型保存可以带有选择性2 @' ]# f& |" Z$ e6 m$ u
数据下载:
" [! G+ }( B3 y9 V1 M* khttps://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset
; W3 g' ^' u$ p+ C2 {2 \ J* B3 Q1 `& i
改一下文件名,然后将它放到同一根目录就可以了8 e$ I/ { Z) h0 D6 c' Z
@# H: B' L ]2 {3 }下面是我的数据根目录 H" @$ X+ p8 x/ l
- z/ i/ A( Q) P4 `
; u, f) B$ ]( s, \5 p1. 导入工具包1 }2 P: I7 ?' l: i: B
import os
# N' W7 j: x& G0 y8 {2 v! bimport matplotlib.pyplot as plt9 k$ z6 ^2 D% \9 S$ B5 p: b( K, [0 e
# 内嵌入绘图简去show的句柄
5 @! _9 ` y8 K: O' L: ^% b%matplotlib inline
! o/ ~+ x2 Y) u7 ?import numpy as np W# U4 @+ o+ C4 X; ?- M
import torch
4 |' H! v+ h1 o4 A6 c9 \ Pfrom torch import nn
" P8 |0 \3 H0 G0 V" l) \5 t' k
2 B; @0 I+ @6 w0 s! L4 jimport torch.optim as optim' L T5 C' k, T( o
import torchvision
" E, u+ ]6 _: [7 N: zfrom torchvision import transforms, models, datasets: [+ j/ Y* z/ i5 n* w* @
5 \1 _7 c G/ I
import imageio4 ?& M- l3 w, U0 k6 s
import time6 v9 G( ?( s1 W$ T0 K' J& ^
import warnings
h' z7 e0 h! t3 H9 himport random
1 Y: G2 P7 ~; x( Eimport sys
9 ?- w2 Y; Y8 ~: [import copy
6 N" e5 A7 }6 t1 G; ?9 vimport json
9 W. l4 f O1 l; t: Z7 C, [$ dfrom PIL import Image
; B( W4 C/ X+ u+ K m8 G* \+ u1 p
. @1 W6 |: h/ {# O1 V7 r
1
6 x0 R* U) D. j3 Z8 r& k2- z2 y! M. r% K6 b
3* V2 s7 q: P( T4 p9 G f
4
' N$ w& |) G' F3 I+ V. ?% f, M" u& ]5
+ Z4 }9 m9 _- Y5 Q5 g! E( |3 q6
7 W1 K# r2 W' V9 n7
' n+ D' d# t4 a8 f( I8 Y9 h86 t* }/ e, B# O; E& t# o1 `
9
' M$ j, C. _/ m10
2 ]; f; w) n* `11; \" C0 j* D( |7 m9 c
12
4 N. ], h2 W6 X7 @; z6 Q6 z: N13
4 E) d8 R4 ]( g) u5 E145 p) B% G- B _: P0 |/ u9 o8 i7 X$ J
154 Z3 H7 f7 k. r. |3 R4 N
16
- k/ C4 m1 L7 R; J17 ]: Q( _% R- v- G: H* D! L: m
18' ]: @) M/ H0 y
19. m% H8 z8 r& B- E- y
20- l5 A6 w+ s- n
21
9 q/ k6 s2 e6 o$ n: A1 W2. 数据预处理与操作
) f. F( C3 q& H) L) C ^3 x" c( e#路径设置
% c8 u, o1 t2 {3 }( ?9 vdata_dir = './flower_data/' # 当前文件夹下的flowerdata目录8 a- A% \! E4 ]: V# N8 c
train_dir = data_dir + '/train') ]& b& C" Q( [
valid_dir = data_dir + '/valid'1 q/ N# @8 B: Y7 |# q% Y) C
1
9 G3 P: ? U2 |4 T5 A& Z* [2
: a* q3 x( k1 Q31 ?1 z4 o/ v" h( Q! [
4
; _( W( q J$ c) `python目录点杠的组合与区别
. a4 d' l9 ~ _1 u- |! D注: 里面注明了点杠和斜杠的操作
T7 k- z) A% u- G# V
. @" H1 ^" w1 k, s6 V$ O3. 制作好数据源
( {0 t6 v( K9 q8 C# ]# ~3 Xdata_transforms中制定了所有图像预处理的操作
9 ?' n h) r; b3 d+ x. wImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片/ n% B( x) e& [0 j* e. ?4 |( e
data_transforms = {
7 Y, |4 K3 Y, q4 w' ~$ |0 o' ?$ ?1 F # 分成两部分,一部分是训练
+ v. s) G1 U. Z 'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间
6 r6 D3 f1 U4 c/ c* P transforms.CenterCrop(224), # 从中心处开始裁剪: o6 k( l1 O1 U$ ~5 v: U" Q
# 以某个随机的概率决定是否翻转 55开
. m/ {# R: C6 |4 t& @ transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转
: q5 s$ P/ {/ [6 n$ r transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转
+ I- i) i/ j- ?. u # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相
7 [; X1 W+ U. d; q X- M8 X transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),
/ \& t5 G" o* v1 u" H9 e transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB% f% M9 R% t/ P- n
# 灰度图转换以后也是三个通道,但是只是RGB是一样的
8 x& L4 q, j+ F, j" x* O. L& I" X transforms.ToTensor(),5 R6 c* Q! k* F: o6 n- v
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差: O5 R8 h! T6 X. B; j
]),
8 f& V$ j8 `4 C # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
% r* U( ^: Y4 k, R. x 'valid': transforms.Compose([transforms.Resize(256)," l* ?( `0 B" z1 {. U) B5 D; V! z& d* z
transforms.CenterCrop(224), P, o3 u0 F. G7 s$ W. n4 ?# b
transforms.ToTensor(),* l- t$ B; @' _) }+ Z4 H# z
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同
* n6 l" h! r' \5 \) K! ` ]),
, X) v7 {) ^1 J9 t* Z' T4 z( c# L}; Q- A7 T( Z U, r) `
( F+ V' N4 K) z! \1
) ?7 P8 g o" @/ r. W/ v q5 s2
5 P( T, ~* Y, V3
" G( Y4 F: b* p' E41 d5 ]3 l# @& u1 l) P A
5
+ c' d! ?9 l$ p+ e2 S9 G! p6
* F9 D$ Z H l4 `7 _* D7- h, x9 C) A5 m9 }. |" T- o
8
8 X4 ~7 T; Y2 y4 Q96 s, E) p2 @- P3 a0 A5 m
10
/ B- q- L2 g& O4 P( P4 o' D. c" O11$ ]$ e9 m+ ~1 B9 {
12* Y( S0 e; s% s8 v2 c: ?, N: O
137 ]) [9 ?, x9 T0 p( x
14( ?' O9 R/ e5 V: [) u
15
; v4 X2 t( D" v9 q* T/ |, L3 G$ v16
! ^# J$ k/ ?( b2 F3 B4 j3 |17) G! o* T) D% ]7 U) C5 z: u$ _, U
18
8 O8 B# b S, d. n& v; }6 H19
# P. Q- G0 P: M; M6 j" m8 n20, { L- J( ~, j6 b: C
21
1 h& O, y7 v" \# T7 E! R! nbatch_size = 8* ?5 x* p9 f7 M u
image_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}2 i. l- N5 ?# N) ?
dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}
6 h. M& f0 J1 V+ f& C6 n5 [dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']} " ?$ U2 R Q2 I+ M Q
class_names = image_datasets['train'].classes# _% d( ^$ {6 Y6 D5 f9 H( s
4 g0 V3 X* C+ l5 Q! I' q! N1 M
#查看数据集合
. b! u6 f; R; b' \' Z6 C! s2 J% Qimage_datasets
9 _% s* m& u ~% N \3 E
4 k( ?0 M2 H! y: y12 n7 E2 _" g. `& }
23 m. J4 ^' G2 T2 N2 ^
37 ?! w0 m( `' @( E8 j, ]* O
4
$ K" c/ E O3 @; h+ q$ w5) g# M Y# k. ?
6& x* N) P2 Y8 e0 c* o- L/ a' q
7
5 A: q9 f" \6 m3 H) I8( [! O: o3 u. s2 p6 N
9
/ M# c" `1 f# i; X' B{'train': Dataset ImageFolder s7 p0 a* v7 P: ?% Q
Number of datapoints: 6552
6 l) s5 _1 E1 D1 {( F Root location: ./flower_data/train; P' W( G( W9 J" I7 \& r3 _# w" B
StandardTransform
3 P- t l( s+ A Transform: Compose(: T; X$ D0 r: Z. F
RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)( A7 [' }: t2 g
CenterCrop(size=(224, 224))
; P( g; R" X8 c7 }7 L RandomHorizontalFlip(p=0.5)$ L; ~9 ?0 W9 }. u/ z2 C
RandomVerticalFlip(p=0.5)3 T8 V$ _6 C7 D2 M" D, S3 ~- Z
ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])+ ^% @3 q" t$ c' K
RandomGrayscale(p=0.025)
4 C3 n0 x5 Y$ R! h ToTensor()
/ I" E, F2 A9 g; s% S- H) m Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])* O) w8 v0 W. ?* B2 v
),$ b, R! d7 E9 Y2 [* V0 C1 w( U
'valid': Dataset ImageFolder
) w2 x& g, j/ j2 J0 `# M Number of datapoints: 818: U- R# @6 ~. R! w( {
Root location: ./flower_data/valid7 K. U4 m& ^6 _( e" e
StandardTransform
+ O- `4 b7 L9 S: @ Transform: Compose(
9 d5 V4 Q; p5 m/ a Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)" Z0 c; P9 {9 x; L: G
CenterCrop(size=(224, 224))
5 f. c! Y3 g7 K, t1 _ ToTensor()
0 A; z" d: N9 r" o7 P4 l7 p Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])( N* p) z I' m! N9 @" g/ s
)}- S% [6 R1 Y x3 r
$ N1 h. ~4 \$ A+ w
1
, h* c2 b; C2 o2
/ Z1 e F+ X2 R( ]% q! T: U39 ~" r: W# Q: J
4
% W* O Z* f# G. F5, _$ a# V! S% ]/ B
6
% p* F+ h( v ~$ Y6 W( B7
# k# j6 p5 _$ j; J: S8
2 N% S5 v7 ~2 e- T& N9& J9 V6 x2 U3 s& Y& S3 ?6 J
10
) J, X4 r: y% s# n) ?! S5 E% @11
. p ~ H6 S, z12; s) a1 |9 \$ D( z8 D6 E' ^
133 H; e: {& c+ k+ \% a- d
14
. d r& S m9 i8 h15! [# M3 v+ @' @# A3 X
16 F7 q* Z# J) x7 N0 [& u. G
17- Q- S w( n1 A! g' W8 f: H: }! y. U
18
8 k3 l n$ V1 v& W19
, K) |( t) G* {20
$ D$ b. \( h3 ]5 A$ z$ _$ J21
- b# ~" V; {% T# |$ h, o22
% ~- l9 B" _ X }23
9 C' v* e) G, Z& b6 k. A& S( T6 M249 l$ b$ ~, h4 z" {
# 验证一下数据是否已经被处理完毕
7 t! u U- V& y! O! edataloaders" l, } a" B; n; ]0 N/ T
1! L5 n6 [( V% O( t$ @; L) D
21 b0 ~, `$ b2 y2 H, k' q5 d2 O/ V
{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,. q7 k+ O; k" o/ }/ Q
'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}3 w% i3 A" F7 G1 Q, Q0 l
1
4 d. ^) h( E) K, M( Z6 N2- {) [( N5 B T7 x+ Q/ b" Q) [2 I
dataset_sizes
+ v7 S- U& |6 S3 G4 ^& H7 w16 m* l2 p0 u% n; E
{'train': 6552, 'valid': 818}; |0 I @8 z7 k- _" H
1
4 L" B" x1 y# M% \4 G读取标签对应的实际名字( A6 N Q- q E2 s* j g1 B0 Q
使用同一目录下的json文件,反向映射出花对应的名字: s3 F8 \3 r7 L6 V% p7 r0 I% @
0 E" A7 c7 i+ R. v. Nwith open('./flower_data/cat_to_name.json', 'r') as f:
# ~5 i0 f* H* S& Q C cat_to_name = json.load(f)
. z" Z& a4 z& i H' C+ {" \1
8 @; a4 e; A; {* |4 Y21 X% x) w6 F' K9 {. M, w, S4 ?
cat_to_name" Z7 a" a8 n1 Q: l
11 D* e* `. S3 O" Z
{'21': 'fire lily',
; q, g8 w8 u8 T7 y '3': 'canterbury bells',' S! e+ N g, q& w& y# b: L9 Q" I
'45': 'bolero deep blue',7 Q9 P D- d: V) ?$ w: B, j% Z: e" \
'1': 'pink primrose',. E# o9 p; s0 W2 t- t
'34': 'mexican aster',
3 z/ [9 `+ F: a1 x0 f6 V+ z% t. \ '27': 'prince of wales feathers',
, \. p/ g" J) N '7': 'moon orchid',+ x5 L3 x" Z s* W3 h
'16': 'globe-flower',
0 C1 C. T- \( @. g A: T '25': 'grape hyacinth',8 V8 G& D" u4 \
'26': 'corn poppy',5 }+ v- Y; M, c( ~
'79': 'toad lily',. i. s2 M! f! l* E' q
'39': 'siam tulip',- @7 m X9 ]* h( e `
'24': 'red ginger',
8 L, R0 D% x6 D$ [6 y/ ~ '67': 'spring crocus',0 [0 I% c8 E' n+ Z8 n; c5 K3 B, u
'35': 'alpine sea holly',
: N3 |8 o: E s0 x '32': 'garden phlox',6 v* k% p L, A) X3 D2 o* B2 ~' `
'10': 'globe thistle',
- k! G# ]# j' S1 e _9 n '6': 'tiger lily',7 I0 @5 Y) r, U0 r {
'93': 'ball moss',
& Q' }* G, n! P* ]$ O/ w- ?- o '33': 'love in the mist',
2 l5 h9 }! Q& Q# I2 k '9': 'monkshood',- G3 x; P/ R* N
'102': 'blackberry lily',1 F9 f5 q- W) w* h& n1 E" [$ H
'14': 'spear thistle',+ B e* B& _# Z
'19': 'balloon flower',
1 I! L, P- e, ^$ J" r" Y% V '100': 'blanket flower',
! {( U- c0 H( X) f. f '13': 'king protea',
, k" i: |, i* ~3 y8 {6 t& | '49': 'oxeye daisy',, |6 l4 c8 R4 ]) J% u
'15': 'yellow iris',
0 ]0 }, N% j9 A3 } V( W '61': 'cautleya spicata',1 @- m8 w4 S" ~5 X+ B4 {* X
'31': 'carnation',
# o; d7 V8 z/ y: X '64': 'silverbush',
' P |% [& B8 \ '68': 'bearded iris',$ w3 N; e/ Q+ P T0 l
'63': 'black-eyed susan',; ]% S' |: s# K+ z
'69': 'windflower',7 }1 y+ c7 h) D7 c- a" \) t! @# P
'62': 'japanese anemone',
9 i+ Q$ M$ R3 \- j) h7 ~5 p' \ '20': 'giant white arum lily',
2 E- c% b F" @4 z( ? '38': 'great masterwort'," c8 }9 S9 i. r
'4': 'sweet pea',$ @! n8 j7 z" N& S5 P0 F
'86': 'tree mallow',8 s( w2 m$ Y. C1 X5 D" x) Z, a
'101': 'trumpet creeper',, P0 L& A+ O/ q* q) C+ W' Z
'42': 'daffodil',7 }7 b5 V8 d$ Y5 E$ d+ R
'22': 'pincushion flower',& N( E* F J7 @- ?
'2': 'hard-leaved pocket orchid',
2 _% H; v; m# L) L$ C( z: B8 d: D3 } '54': 'sunflower',. ?! N! q% Y' _$ U( ]
'66': 'osteospermum',
' L. V; C: a" ^$ J '70': 'tree poppy',+ ?- j( g( }" Y% E1 i" _
'85': 'desert-rose',
5 p$ g5 `) T! W: V$ _4 T '99': 'bromelia',6 |( H; c' y M* l: L& t$ G6 U
'87': 'magnolia',
( b. M* G& u0 n) @- h" S4 F- L. d8 U '5': 'english marigold'," E( @7 i" Y, H) D. c0 b( l8 w+ k
'92': 'bee balm',. z, i8 G: j2 K' a" `3 J- p
'28': 'stemless gentian',
6 R) b& p) ^* F+ b '97': 'mallow',
i: i) y$ O% Y- w. Y. j' Q '57': 'gaura', G$ F9 @9 j$ a/ J4 D+ X
'40': 'lenten rose',* F2 S# i# u( D0 @2 S, ]
'47': 'marigold',
7 u, `7 T& o8 t- A& f& w0 v '59': 'orange dahlia',' U1 R# q9 O4 I& D
'48': 'buttercup',
# ]% B% k7 Y0 U+ A8 |. f9 d3 { '55': 'pelargonium',; G6 n( L" {; s4 L1 e( u9 u
'36': 'ruby-lipped cattleya',
) \- E K( q9 _; g/ Q '91': 'hippeastrum',9 E( J7 e7 x2 L, A$ n. Y4 u0 Q% {
'29': 'artichoke',' O2 d- o$ ^3 ?& T( {9 E+ K0 Z
'71': 'gazania',: b. O" D2 K, N9 g' D
'90': 'canna lily',* A! b" X) e- |, N+ s- I
'18': 'peruvian lily',6 V" {+ x+ d- e( D
'98': 'mexican petunia',
5 R' ?) J# i! E4 p/ |, F '8': 'bird of paradise',8 w. G# i% K: W+ W8 q+ K
'30': 'sweet william',
" ~4 {: m1 i' ^9 m$ u '17': 'purple coneflower',. L$ n/ |3 E- @* K/ G
'52': 'wild pansy',: a5 z# c( g9 r- _
'84': 'columbine',
$ [ Q+ Y# b3 ` '12': "colt's foot",
- N3 N8 S) [ Q- S& T '11': 'snapdragon',
U% P0 _- u8 l ^( U8 ^ '96': 'camellia',
3 N9 r+ S! r# u '23': 'fritillary',
- c$ S2 Y' S' V$ z( e$ r) m '50': 'common dandelion',
6 {; N- w; k2 }, `1 G* ` '44': 'poinsettia',0 n, n" T+ S$ v
'53': 'primula',
5 @- N ]+ z9 r) u '72': 'azalea',
1 v$ b& e$ l, i* F* @8 \ '65': 'californian poppy',
/ P: ^7 C" }. F '80': 'anthurium',! O' {5 j6 U$ [: l2 I+ |
'76': 'morning glory',; o5 N% H- K. y
'37': 'cape flower',/ I3 K& u8 t: m$ M9 K
'56': 'bishop of llandaff',
( z' J& ?! V0 X% Q- [+ C '60': 'pink-yellow dahlia',1 ~, F8 Z! _2 J) t. g, x4 D
'82': 'clematis',: d7 _( w5 a. Z9 i; D2 e
'58': 'geranium',
! t3 }. }4 L% s7 r '75': 'thorn apple',
+ n# x/ r7 y0 S# W6 p '41': 'barbeton daisy',
1 o2 P" b% u+ ]7 V! s '95': 'bougainvillea',' p6 j; a+ A3 z1 F9 c( c+ C E
'43': 'sword lily',
! ?. f! ~6 z% Z* U% x/ X5 @ '83': 'hibiscus',
( a! a3 L1 q" l/ S6 n' _ '78': 'lotus lotus',
" s: E3 N6 |2 B' ~# [1 f '88': 'cyclamen',
; e: ^; p* d% s6 ]+ f8 L0 O3 m% t '94': 'foxglove'," G3 P2 P" {! n! P
'81': 'frangipani',3 }/ W; i! \$ w3 S. f7 i7 L' m
'74': 'rose',3 A7 m: d9 O: w: s# s# w8 _
'89': 'watercress',
# i6 h5 _* K. A% L '73': 'water lily',
# M8 |$ A* f& _; A& O/ q '46': 'wallflower',
1 `, N( _0 f+ v/ E# _ '77': 'passion flower',
1 `/ l, p f! t6 O '51': 'petunia'}
$ ?6 s. E7 l& P7 v0 w; u: ?* C( y7 d+ L3 ?8 ~" ^3 s H
1
A, \3 r1 m! q6 h2- S: L* ?9 ^/ v, C8 J, S4 b& f2 }
3
$ \: f1 b3 J) \" Y; `# A2 p& g48 s' P6 @, Q2 `/ C4 D8 E4 ?
5- \" {& B- v- t7 W. }
6
: v+ T) k: ^1 ]7) c! ]% M) O. h2 z; y9 A
8: ^% O) g9 G+ \7 A
9" g8 h9 P' p0 B/ @- O# w0 {/ n, Y
10. ~% ]; T: ^ J
11$ j; Z8 e+ q& a$ r- e; A
12* E* D9 m- r3 _
13
1 A$ G* i' ^( x9 v& Q14
) _# p/ k+ k1 a: n; k) X; W9 Y15
/ P# P. Q9 N: Z( Y, N5 v- r16; B( e& c0 l/ C0 L& e4 I4 t
17
* A* [4 g3 O0 {$ D" `% \# `' I18
2 J5 A+ f+ a. J+ _( N( y19
* m( L' {# ^% q20
! v" ] a, O# U7 K7 j: B' O1 p21
2 }7 h9 e0 t2 g# \. D. E22/ q0 i0 {" y5 ?9 t- B u9 c
23( E, [% [3 W3 b6 I
24
' s+ r. M0 N: T# [25& z+ x. }, A2 b2 p ]1 L
26* [7 j6 y: b8 ?, r" x; E
27
6 e! E' M0 C2 T0 h3 e- _28
. l/ c; O. ^0 N ]) k5 ~4 T29: O& q. J# [0 m7 D0 I5 E; |
302 Q8 R* A4 b1 }% a; n0 r
31$ X" U& e) j1 H8 Q( n
327 ]8 M( b& F2 G7 T) W
33* n6 x: I' {0 M6 [4 ?: ^7 }$ E
341 }6 U" t( z( W+ ~' \3 ?
35: F+ D6 L; Z: y5 q* |4 n
36
! N/ I! V- j, S5 W' Y0 k' Y5 a1 O+ p37
1 Z! Z/ r6 m9 S2 Y5 m38
X( \& q' h; @39
8 M, W M: J" h40
6 n; O( t- z: _- d, E u4 x" g41
( U, e" Y# }% V+ l8 F42
0 V+ f' |9 U' a+ x9 X) r4 g43/ r, b# u7 m9 H. ^+ e
44" Z8 u) ?2 W K, \7 j" |" U
45% t$ e1 \0 ?; H
46
H7 q- x/ ?# B" q47+ K% k/ B5 p% L# w" s& x n
48
+ o0 L3 v" z6 K+ O* o' {; [49& ^4 O3 s: A1 ]2 S/ |0 P
506 B6 D: h1 \+ n+ a5 n; A
51
6 x9 l/ z7 n9 n- }6 `520 Q& W! t' S* t
53- W& D" o7 I5 p1 H1 T
54, r" }& e3 M: \) w/ ^
55
" }! p; ^8 d# f! R m) G; k56
/ S. P( v0 e) ~ O57* Q* b2 ~8 z7 l+ ]3 E/ b
58
- M9 t/ T' b4 h6 C59
) i/ c* f3 a5 ?6 ^, i' X60
4 l5 I5 Z) g; A! m; G! E8 L+ E- h& N' `61
$ y5 X9 @0 Q( ?( h; }5 y62
0 W, D: c; r1 @5 ?63. N3 ~" i# Q' r. ?5 z4 a
64
0 B) O' E. f7 T/ ^) A652 y7 f8 ~" g- o8 d8 O1 h
66& s2 h2 G |# z, E1 S( G4 I
67. d1 |8 K' E7 U, W2 ]4 r5 \$ W& y
68
6 P2 U0 X! o- D( f8 F9 a69) j/ e' ?' `* |" e8 |
70
1 }4 `2 t9 [# I: X m3 h! x71
8 R6 s( y e4 z' n* \2 P. l72
1 a0 v& o! u2 ^! F% C# ?73
$ i0 U; ` A6 j9 }% C74
& D2 l, t' {) L75
" k8 K: l. d) @5 o q9 U* c76
# J V( V7 O3 e* S+ z77
7 W9 l5 `7 B7 \" t. l, d% f$ G4 D+ V8 b78: _5 E Z9 y( [) r7 [
79; ^; x3 g3 o( I5 j$ D# S
80- I. L8 \( M8 \% N* j$ S) U
81
' f; X# {4 P8 Q: m% v$ l* Q, t' F# t82
( j/ b& @5 P; \7 D83
, t r( i& U2 U" o84+ Y" b9 B' X9 b! Z% n" ]( l
85! V- K8 }: ~1 V3 t
86$ m6 D' @" ~# x% L1 I! ^6 I
87
3 d( ?* c& w* u* a a( m88- ?* r, x0 c, H h6 k
89
7 c0 v, W% m# A8 {3 D+ p0 M90
( Y" F1 H- w& i) n: y91
/ P1 |3 e1 `1 O) d92
6 c6 c8 }! X q/ U* U93
* H% Q2 D( N% d2 _; k- l$ n! j94
! r! R) M2 _) l7 e7 ~" J95
9 W' D# u7 N8 N# L96
! v. B2 w2 L1 x97
, N' f" \; V* T. i3 K3 G98
* H( i- D2 j3 ~1 Y' G/ F99! i/ p/ C4 C D: s$ o! x& j
100
+ N4 m# j T4 w7 h1 G& k7 Q, {$ `101; l% r! e6 l k
102
1 t5 [" N2 t3 @4.展示一下数据
5 \; p4 }0 D- M, y) G U1 y; ldef im_convert(tensor):4 q' x9 ~! V; `" j1 L
"""数据展示""": [% P1 Q& y% W( M* p
image = tensor.to("cpu").clone().detach()/ |6 ]8 {4 S! [9 z% }% l
image = image.numpy().squeeze()
' f6 d7 \) G- f( s- U # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图
+ t) F7 R9 Z3 n1 o # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)
5 C! W$ J4 b" p3 V: A image = image.transpose(1, 2, 0)9 o/ n2 M0 J B* V
# 反正则化(反标准化). O I3 h" u$ i% @# f S/ u
image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))3 a0 ^& G! |3 X; L# u; o: a
2 T! K- V& {' j6 [( ?
# 将图像中小于0 的都换成0,大于的都变成1; g/ e) ^' b8 Z0 n: l7 d( x
image = image.clip(0, 1)$ a: ^ O( P0 q) m) r
! n5 l9 a) U _8 Q3 K return image. q6 g. O9 W5 d# l5 `
13 K, @1 w+ G4 d
2% d ^5 \6 J9 p5 Y( H5 k
3
9 k1 o6 h3 P; P& A, l) i4
0 Y% q' w- e/ F. C5) l. W# I% P( ~- z+ a- b% w$ ]
6" u' l1 @ U$ K) V+ r3 q
7
* s! O2 m: C9 Z E; G7 h89 i% X' J5 ]6 }! B9 V
9" w! H/ }. P' D6 \3 a$ K
10
0 [( ~7 j8 J( r# c& b# b8 x111 p5 p) z z8 K
12
- c3 U* H& ~9 T7 s% W* ~- D& [13( p: O2 R/ n9 L; k& k# c
14
2 Z% F! X; p) q1 I( l# 使用上面定义好的类进行画图
5 F+ ~* _" E. q ifig = plt.figure(figsize = (20, 12))
# B( h+ o' p( n4 K5 l( t$ i ecolumns = 4
: m. i0 e) x3 x: F9 k! }rows = 2
9 n+ j% t* p8 H; S# K7 G3 f/ N( W
# iter迭代器
5 A% T- ^3 R# ` B" W# 随便找一个Batch数据进行展示' w% P. |# P3 X- ~/ D( G" a
dataiter = iter(dataloaders['valid'])
4 w0 x/ ^( S) `. yinputs, classes = dataiter.next()6 R! _$ v6 ` S) v
7 N3 H" w$ h) y' Z6 b$ Z# h6 m
for idx in range(columns * rows):
$ k7 u* c% j. R Y8 V ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])" X: X/ B, t% s# J% h8 p4 c
# 利用json文件将其对应花的类型打印在图片中- J$ J* [ z! g" ]
ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])7 c' Z- M6 p+ {& H M
plt.imshow(im_convert(inputs[idx]))7 E* [ |0 B1 t/ H% B0 a
plt.show()
3 s3 l6 S1 @! E2 H: d9 A/ j; g M5 S1 m5 Z1 f) L& U
1
7 `' @, p# W6 ?$ L+ |2 }. K2" `! Y+ g7 Q4 w0 e
3
1 F- k' s& c( c6 x& R9 ~+ M47 b# q, ], k5 d" g, O8 a9 I, E+ ]
5
3 d8 O+ v# L4 ^0 S5 x2 F6
, X8 h$ E' s" @( ]$ h7
- C* ]- C3 {# q y( p5 c8% V( W7 l& V" x8 k8 h
9
|3 t: V$ i$ Y D4 o# x10' B. I, ~% M4 G
11
- r) K+ k% c) F% m2 W1 G123 j) p4 w7 p; e3 v( L2 @# V
13
: B3 Q; C8 M3 Q i0 ?$ @0 m1 v14. E1 y. v; A% a O
15+ V5 J4 O6 q1 {. u5 Z/ B
168 b o. o4 p9 S9 f, Y
# t4 w. W+ T4 P5 i& ?; m# t2 B
5 X7 a/ I# ?4 Z5 d9 U% m% @+ j+ f( t5. 加载models提供的模型,并直接用训练好的权重做初始化参数) S& K% k9 I$ T- J0 P
model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
5 H) t! a! N' j1 I3 B1 y7 v+ L k5 j- K# 主要的图像识别用resnet来做. p: d$ G& }1 a! j: H
# 是否用人家训练好的特征
* n& ?$ u) Z% k- Ifeature_extract = True
$ ] s4 P2 r# |: m. [8 O1$ T- f, w: A+ s- h" k' r
2! Y0 }; Q% I: }% ^ J$ ?( y9 j: A6 l
3" B. k: I8 f5 n) k
41 q3 Y R5 J: m) D6 h6 y
# 是否用GPU进行训练
# r3 |/ \( Q+ x3 s/ p' ?train_on_gpu = torch.cuda.is_available()8 d& m, I w, M; W" V: @" y0 r+ P( L
1 i |' g/ M; M6 A# Y y0 eif not train_on_gpu:0 z$ y, L" ~4 r, l \
print('CUDA is not available. Training on CPU ...')
; v: W* Z9 l( d* F4 V0 pelse:
9 y7 k: N: [ v2 _2 @5 z! r8 Z, l- h print('CUDA is available! Training on GPU ...')$ R2 w" v7 s4 R
0 R% y+ }+ b* B. ?3 F7 _4 `" W
device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')
- B8 f* j4 Y# q$ {+ q: y* K; b: n1. a" p& u, ~6 A( ?2 C' y+ e
2
) F5 `" ]+ G( q" q1 H3% |& L* Z* n4 ?+ e: q( }6 H
4
$ b1 [1 P5 E: S3 Z6 H5
3 [. Y' |4 ^+ l6
1 W# \% P& M9 t7
; ~( V8 r- W- b8 i4 @( E! j8
7 V6 u' j$ K; \7 K" X9, q, x$ t5 Y) j( a& F/ M. R; e& n8 O
CUDA is not available. Training on CPU ...9 W: Z" z/ ^, o. ?/ M
10 W* g0 @8 q2 E( N. x" T8 A
# 将一些层定义为false,使其不自动更新
* [- _3 x( J0 L2 r- V% r+ ~3 udef set_parameter_requires_grad(model, feature_extracting):. y% O2 H) r9 B$ E% }1 D
if feature_extracting:
# w" l3 G& i' J7 I for param in model.parameters():1 K: x0 ^: r9 L$ ~1 y4 X. Y
param.requires_grad = False# s% Y5 w: a2 G
1
/ ?5 u4 ^" A% s m$ W20 b) ~; M8 g; z, Y, ^
30 R* s3 o s2 O3 X
4
# R$ ]6 A: ~: i" Q55 B. L# E% c" A% W
# 打印模型架构告知是怎么一步一步去完成的
1 O1 N7 @- y( R8 }, N# 主要是为我们提取特征的
9 a$ \* b: f8 ?5 T$ B& r% k5 [9 g/ u0 \' h8 j5 d: J
model_ft = models.resnet152(), F6 @5 l' O: u) ?3 D1 U: ~0 D
model_ft1 r4 @- q6 ?# Q6 @% G7 | z
1
+ B4 X9 N2 |2 j; p28 z1 C7 D! z! R
3
# f2 j5 b+ k. J& [& F4
0 N& }5 g7 ~% g5 \5
. K. r. t: U2 {4 K1 rResNet(" `. h0 L. T# V1 `5 ?9 F
(conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
" t4 F& v8 e" Q7 \! Z' H9 n% E (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)0 x+ y# Y) j- ^) H; ^
(relu): ReLU(inplace=True): W! F/ @5 j0 A
(maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)- W9 J. h* u6 G0 r, I) d, K) E
(layer1): Sequential(
& o; s: ~' T, W1 `) r2 b (0): Bottleneck(* B0 c. ?( s8 n' P1 v$ t- }
(conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)
; Q2 b* u0 @! G- s& c/ A: [4 A2 | (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)+ m" R# r) S4 y4 L4 E4 `
(conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)3 `! ^2 l' `$ E3 F
(bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
; {. i* x' d8 g: O (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)% o6 P! c: D$ O9 \6 D- g* u; |6 A
(bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
# ?5 \! [& r/ D! Z8 R" @" e (relu): ReLU(inplace=True)
' @7 B2 l, l! S K4 n6 o (downsample): Sequential(
3 j) E) O* y8 r; P (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)8 j$ X. D& e9 h8 g6 _! j5 V$ G+ V
(1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)2 B. v6 g$ ]) k& h8 X
) l9 R+ H) v. C* F" j
)
9 ?2 W) p) B2 J5 u2 O中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。
% H" `0 S) B+ j$ |+ o (2): Bottleneck(
% c: s: r$ o* t! a5 x& L (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False). s9 K- `. r5 I9 W* Q0 {
(bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
9 f+ F, p/ ~4 j1 l0 I& c (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False): d# ]0 N7 i E+ T5 ^7 v( O# Z
(bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
, B$ r' L8 s, J: } (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
8 u9 h1 ^3 A' U1 E x6 k (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
( d3 i. @( G# t. n& V- F (relu): ReLU(inplace=True)# B5 J% [! `# U% B D6 Q
)& i/ @* L j! h
) [- e; A0 p1 A2 W
(avgpool): AdaptiveAvgPool2d(output_size=(1, 1))
. [0 f. @- b0 P (fc): Linear(in_features=2048, out_features=1000, bias=True)" a5 O+ z, x6 V1 X! k
)2 M, h$ L. t: r5 c2 a
, f0 j! f, B, K+ h1, g2 T h/ k7 N/ s2 R
2
4 _: k, {( \' {, O- k34 Y: y8 i ~( T O( D7 h- f; x
4
8 [6 ]# m- y" q5
$ u3 s2 [ d: f4 D* G* I; L1 \6 t5 T, N7 w$ f, c
7
/ S! t7 V& o4 L- j2 U; r4 i86 B5 x8 V# E9 O
9* c' v$ Y2 S, P' J+ v. a; z* P
107 ~# Q4 _5 B' f8 E! ]( d
11
; x2 \7 k2 E/ q* Z+ G R, N12
* d9 r/ x# K/ H' h4 \131 J' `" E6 H- G( Y
14
6 ^. j% W) n4 b15
' [1 d t1 h" o, N- c4 W164 k1 {8 P: Z* _, |- r3 n: ]
17
8 c5 m! ?1 T, G18. s4 j; r( z8 j3 k2 s
19
% l# q. I' d# o( l20( z7 j% g5 C( i: I/ Z- I
21" ` x. i% u/ ?
22
' t. j J8 x; }0 B- O4 t& P23
4 b' f9 a/ b! X! c& n3 J24
% T5 }2 a* }0 y+ r) X25
: M0 w# A7 m( h3 D5 h( e+ g8 ~26/ [3 @3 G' P" f, h
27% e- K5 w3 l9 `
28
, n j7 [% S/ G. F0 Y293 P' {, j5 A0 Z& p/ Z
309 x/ T* y1 U- h5 {2 D0 B$ m% W
31
' [" X5 f, x! i+ U, s4 w& u32
. c+ O f- r. [33% h4 Z4 l- ]7 L# b$ d9 s
最后是1000分类,2048输入,分为1000个分类' L4 V" P/ V( G4 G9 F z; ?/ u6 _
而我们需要将我们的任务进行调整,将1000分类改为102输出) o: C& R' f7 q) W6 U1 n
1 v& w/ ^- y) K5 D/ w# A; ?- L5 d6.初始化模型架构
- E! Z6 w: \3 P3 n/ u8 E步骤如下:
9 J% y3 @, c* }0 j' P- L* O6 h0 f" |* ~6 g
将训练好的模型拿过来,并pre_train = True 得到他人的权重参数$ t- }) [4 E m; J
可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False): V# ~" C7 ^* V5 R7 _
无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数
9 _( J7 f2 P; w官方文档链接$ i2 _! s0 T& T8 A
https://pytorch.org/vision/stable/models.html, V2 q0 C$ q+ W
4 ^# A; `- V6 ?) J6 G0 w& O/ N( s# W
# 将他人的模型加载进来, `# i6 ?. C* p% v1 H3 k
def initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):
, J+ X5 V' A- o- ?. p/ A0 I # 选择适合的模型,不同的模型初始化参数不同( R! X x5 O% P1 W8 B5 T: M
model_ft = None, `( x- b P! O. `# r
input_size = 0
% C6 l0 A" U. P- t1 c5 o1 N. Y- s1 f7 b
if model_name == "resnet":
( o/ H% p2 M) p# K6 f """$ V& y5 e5 [/ `2 J/ e
Resnet152
# Z. w- C" U( `$ K: F """
/ s1 c3 u' ?* f/ E
" |) v$ h$ L1 i+ P: y: d/ u+ } # 1. 加载与训练网络1 I0 p4 N$ Z: W& B
model_ft = models.resnet152(pretrained = use_pretrained)
( a2 b! B# ~5 Y4 @% |+ A # 2. 是否将提取特征的模块冻住,只训练FC层$ G: \% K& t: R' u* l
set_parameter_requires_grad(model_ft, feature_extract)
% I; G. q" j$ p7 ?- t9 s # 3. 获得全连接层输入特征% a6 h# @* O. L L. e7 x u
num_frts = model_ft.fc.in_features
" h% w. x4 E# d) C. V% p3 }, V* H3 V # 4. 重新加载全连接层,设置输出102
9 W. V: y1 |$ B8 b M ^ model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),
* k$ i, N2 e) j3 V- b0 E6 m nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1
' M, _+ p2 O* _8 C1 p# p E input_size = 224( i7 ]; D& U u1 a' c
2 |/ k& b' X/ c4 ~; @5 o+ w' p& e elif model_name == "alexnet":! T) |1 [; u5 q
"""
3 V9 b* ?* B: q- d Alexnet
7 W* i* m$ _; R' K """6 r( t' ], m2 K: R6 E' U
model_ft = models.alexnet(pretrained = use_pretrained)
& N3 M& ^' b8 v6 I0 o set_parameter_requires_grad(model_ft, feature_extract)
4 K. j9 L' N2 S3 d
5 Y+ w/ o, o% R& z6 ?( T # 将最后一个特征输出替换 序号为【6】的分类器
- C9 S3 \% l F0 Q6 i6 | num_frts = model_ft.classifier[6].in_features # 获得FC层输入
( g0 U7 l; m) E1 J model_ft.classifier[6] = nn.Linear(num_frts, num_classes)2 M( r9 U0 B7 W/ u/ C8 a
input_size = 224% {+ Q, q4 L: R/ g% n3 x8 o7 G! q
. K) a6 e( c5 I! K
elif model_name == "vgg":
" n3 B% r) s; \' C6 x) b7 i """
# l: i% S2 r! |: Y0 \7 F VGG11_bn9 f. V8 A, @ l* ]' ]2 n+ T
"""
6 W; R& n- i; `% D model_ft = models.vgg16(pretrained = use_pretrained)3 U# ]7 e7 o( z/ v$ \, b
set_parameter_requires_grad(model_ft, feature_extract)
# C( V+ s, z( r& v8 S num_frts = model_ft.classifier[6].in_features; M! U0 B# u8 }
model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
7 O E: A9 h* K- Q5 K% I3 z: ] input_size = 2242 r$ m! r& p v: F6 I+ |
% j$ N2 z0 Y4 ?* Q1 ]( a8 Y* b5 `1 f. [% P
elif model_name == "squeezenet":
9 n- N' l; s$ l """" O- y F% a. o" F3 t3 R- C* F
Squeezenet3 t: k& v* B7 F5 D
"""9 F P" ^4 D) E, V& N1 m7 v
model_ft = models.squeezenet1_0(pretrained = use_pretrained)
) ^ o6 J, e. \3 r$ d) a; \: J set_parameter_requires_grad(model_ft, feature_extract)+ g1 z: G+ p5 g. Q
model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
2 x* T/ R. R) N% x( _ r: V5 @ model_ft.num_classes = num_classes% |9 y/ f, {7 B* L: Y
input_size = 224
9 m8 G) |& b1 p. Q0 v) X l
( b n+ c7 _+ p7 L5 q; M/ B elif model_name == "densenet":3 g k/ [( E5 L! b& L$ W% b
"""
# T7 {* [2 L6 Q6 Y0 @ ` Densenet
5 a: f; s9 @2 r """
2 q4 B! l8 ~; I" j model_ft = models.desenet121(pretrained = use_pretrained)
3 h! j2 N3 |2 f7 [6 Z7 ] set_parameter_requires_grad(model_ft, feature_extract)
- `3 `8 {" r( l num_frts = model_ft.classifier.in_features+ K# S" z c/ l1 @
model_ft.classifier = nn.Linear(num_frts, num_classes)- R; ?! F& Q( d" x
input_size = 2240 ?" R) Q& @( ~; O& X" |: \5 r
, A( H) y6 @" x& y- U% P5 e% o# Y
elif model_name == "inception":
/ }: d# e2 y9 d. `4 l, q. n! _% o """
: n# f- n" e% m) N) C7 `7 s, z) O. B Inception V3
9 t5 p; P# Q% ?' [. L# B9 Z """7 R( x" P) h6 E! c
model_ft = models.inception_V(pretrained = use_pretrained)
& H+ A( R- z/ S2 t4 n set_parameter_requires_grad(model_ft, feature_extract)
/ E& q. `& G* _8 t3 R
- p: ?! C3 F5 h( y num_frts = model_ft.AuxLogits.fc.in_features
6 w8 o# G* i" M' j3 C- h6 b model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)
! C- i. [( ]; \. E: e0 x' Y( \ _+ f, L2 N% Z
num_frts = model_ft.fc.in_features
" n: Z2 ]) v: b5 i6 E8 `* B9 S model_ft.fc = nn.Linear(num_frts, num_classes)
" a6 E' {4 R, _4 z& x0 L input_size = 299, j5 L U2 p3 O5 h6 _$ N( k, ^
9 L$ H: _4 v2 U) X; w5 F
else:/ a1 E" e7 G4 m, E( h
print("Invalid model name, exiting...")4 K# _. E6 C) ?
exit()
# ^. i8 A& U% ?# {2 t9 o& m2 R+ m! Z9 u* I1 R& I
return model_ft, input_size
3 q& W6 y& R3 P& T$ u, Y- Y1 l+ i# @" A9 ^" G% f# x. V+ l7 l
16 M4 w- M6 u7 W/ ]4 l5 q
2
+ w7 k. K/ ^6 W' Z7 }3
& X* y' t& _6 r& w: U: S; Y4$ f ^' S. u! @
5
E' l$ z9 V+ p6 q7 ~' p' U) j. [6
4 S4 F4 \+ p) d7
7 [% x* y$ ]: J' {8 q* V0 v% B8
2 u, @. M# }5 m3 B8 g9) b! ^' h% \0 D( U1 S3 P
104 T! m, j K; H' F, }
11* } Z, i4 ^) d/ K. ^ {, N8 y
12; s( |( Z( B. p8 I9 j7 e+ \
13
( X, o' o( }! I% [2 D2 Z/ }; v0 ]14- r; t) F# g2 m
15* I0 |2 |; P8 T) P: a- I
16/ M- p; R/ w. o I% Q
17
; w3 w0 C$ h: S' G0 D18% u8 P& U+ W @
19
$ Y; Z& P0 M: V/ B7 i% L( O20
" u/ _: ?( Q0 L) U ]21) q4 C' a. g5 M! U8 j2 o* c
22. B7 K7 n( }8 I
23, U& b$ p$ D$ o1 d
242 |0 k7 J* |% `
25/ U/ S3 S5 T# \) h
26, A% c+ l5 m7 U9 E0 f5 A
27
1 r7 J- O6 u2 R8 t( i28
6 |. g& ^2 i# s! T) g29
1 c7 I( ]* y0 j& m30
- J- T3 H) s) |4 R" F2 f3 g319 @) p4 v5 D' e) n
32
( V( q* f H( E5 f8 y" L8 k5 x8 Q333 F5 ]8 K4 p1 y0 v
34
( ^- V; ~& [$ A5 u( e35: |( b% a: ~! u+ g
36
' D5 N0 M/ R @ e37
: F6 ]" p. S) G. a* Z4 F38
3 ~4 j$ W) i4 H7 c) `, o39' w; q6 n; y- }& t
408 C' P$ u8 |: ^
41! q* B# }! `# D# @
42
, l3 r( N- N% O3 A" e# S% G43( `0 ~- {, r1 V ], C3 Q# C
44
9 g0 D! Q4 n7 @. H% ]3 ^6 d45& i, ^6 ^/ i+ `5 `+ n! i
46: X9 `2 L2 W- r1 \
472 d& H9 D2 ?4 s0 P
48- ^6 K& k: E2 C- S% K5 l
491 {( Y6 N2 ^; j. ~: u
50
3 z6 z& y* F5 f+ B51
/ P- l: y, n1 r! |3 R2 a, Y52
' L8 ^- W- U# W) Y2 S53" }/ U* X' T7 A4 N" J7 u- V
54
( ]7 W0 o) \& I& l* x+ k1 [+ _55
4 q9 B, o& ?" {8 k- a3 J562 d, g Q2 I" J
57
6 I' u! y( k8 ?5 J, l# _" H2 y& I58
. o4 c# Z. K* V9 N' L59
, m1 N1 O1 g: c3 { ^$ T5 G60$ x2 {) \/ v3 W% o/ ^4 n; J
61& f5 K( s% Z' b1 z9 O, S
62
5 h o8 j6 e% c2 \5 ]63
# G! a9 H+ @- ^( q1 x64
, p; z# y* _2 ?( B9 L$ x653 J/ S- u4 h: W9 ~
66: k' m9 P5 J; j9 ~" R8 p& Z/ p
679 Y, @2 t$ G0 R) W# {/ f: d
68. ]8 y. J/ y; n& J" x. D2 o$ B
69/ h( J: k, K' G1 Y r9 F
70) ~: t0 @0 Q. D2 j$ {9 k3 C: ]3 }( \9 h
71
H* B# C4 A: _/ G* \8 j72
6 X( u/ B7 r- s73
. R: u! p6 t5 r, c8 _/ y( Y: a74: S: J7 Z; I5 j; i( Q f% Z5 P
75
S/ U9 q# F/ O2 B' c! ]76' U( [ _: X+ }, R
773 O k" O2 U) B3 o: Q1 Y
78
: d* m9 s, a3 \5 i9 r; [/ H79
d% C+ k5 m+ ]" m) r/ N( A80
2 }) ^( w: j% K8 j- h0 ^, D81) h( L% g$ ]# m
82
9 ~1 ` K6 G# r) L83
4 H' P) u" |5 P$ ~7. 设置需要训练的参数+ S5 ^3 v0 o! X) g5 z
# 设置模型名字、输出分类数
% S' p( u" W$ v& B" Amodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)
7 N- w0 k7 |0 M& k3 B, w
* {7 P9 X, J3 @7 G% {# GPU 计算
; w. k$ n" S& G: e% ]$ f$ R$ `model_ft = model_ft.to(device)
0 J" P- Z9 p |9 D2 ~4 A
8 u; b1 T* U8 v" C% V1 B4 n* I# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取
; M6 r: n' J) T6 w* e4 w: Gfilename = 'checkpoint.pth'
3 O) G# b" R! |. d% v$ M- Q
1 b4 C8 U/ E9 R1 h- }! W4 k# 是否训练所有层 x: m( J a! D k2 n: x$ o x
params_to_update = model_ft.parameters()
! o/ f" v+ F+ U8 s# B# 打印出需要训练的层6 s4 w$ T, Z3 h L. j& ~
print("Params to learn:")
7 b6 \+ A' _& s/ _: q$ z" Sif feature_extract:
2 L6 W7 O; o: I, V! r- U; H1 X params_to_update = []
, P/ e7 {4 M5 I for name, param in model_ft.named_parameters():; c! L; `; `) w
if param.requires_grad == True:6 {6 r( m) p+ n: T
params_to_update.append(param)5 ?& Z( |3 ]. r
print("\t", name)2 R/ P5 ~: S4 l8 t
else:
- C# k5 R1 P" p& y/ _7 N D7 X for name, param in model_ft.named_parameters():
: Z* Y& T: O$ _7 Q& O3 k' n if param.requires_grad ==True:+ T; x& }; }: q1 w7 x
print("\t", name)
( E9 Q( X/ ~' Y I' _# ^& `- e- p" ?! m* M
17 S: O7 ? C3 s0 e
2
% j: p8 h; D' T2 S2 y3
1 u, r, C$ t8 m4 J1 N& y- H! B4
( M- g$ N; G9 y' p; e; B0 W. ]2 i51 B5 u3 b/ W; A2 D. P2 ?
61 {$ P$ y8 }0 P# J2 d) k
7
+ L& Z/ E' n0 [5 a% F( O0 b% q8
; L9 J) b1 a; g7 p; e9
" X* z$ m9 U# L10/ g! A1 G8 W+ H2 k8 O3 I
11
( p2 p6 Z' U& ^& }+ E% U12
. |' K5 [. G" [% p13
$ i. }5 U% I& Q& s4 q* f14
- F6 }4 [) [# g2 q6 H2 m15) o# t( s* z- O- H' X. a
16
v5 ^( ?" W- h172 S) Z4 E4 {3 u: }! j6 Q$ `/ t; m( R% M
18+ h, D/ ?# ^( I0 N- N6 z8 k
19
7 h3 j" U. ~% _( {8 A20
: d+ G$ o% X" X' F& \- V( y) q21. `+ B. G. B) g8 T) f* X
221 c9 p' r% ]+ _; \0 B/ [( z
23
3 l, B6 h1 o1 ?/ vParams to learn:
2 L3 e8 I. { D% x/ C. o/ ] fc.0.weight
4 c" Z" c) j) e% X fc.0.bias
( E5 u% |2 w- k; O' i1
3 ~& p* H/ u# @0 {+ a2 J9 s( `2& b& Q/ k: d4 K" \8 x2 |- b
3) C- H, t4 X) C; N
7. 训练与预测2 C* ~+ P, s$ ?, ^2 M6 x: L; M; c
7.1 优化器设置 W ?" r- O4 g! ^0 f ]0 M; o2 {, X
# 优化器设置
; r3 Q" j7 `; D9 @5 ~optimizer_ft = optim.Adam(params_to_update, lr = 1e-2)* {% N) G1 g7 F. M
# 学习率衰减策略
& p' i4 S+ K. m0 F- {, n+ O9 ~scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)/ K# P0 C% `8 Z+ l0 d/ A! p
# 学习率每7个epoch衰减为原来的1/10
% L2 i# B7 Y3 {* s# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算
- L5 a6 [8 f& N& |' t6 d- j5 O% O6 r9 F0 M3 V
criterion = nn.NLLLoss(). ^5 U0 T \. g1 ]( h
10 z# k/ g3 o1 W5 v1 H, T& V) h1 M
2( m( I9 K# h5 P3 s
3
6 {& H" f' T* g# L" u `7 d7 e4
1 d9 o. r" p" u. C! ~5
" e( _* g/ x+ T6 V& s6
, \3 C* C+ @. s5 K- B/ V- u L7/ ?- l3 S. b2 ?
8
+ G$ v8 }! N, b$ i" M2 E/ S7 Y5 ]# 定义训练函数& R. W7 {7 h$ I
#is_inception:要不要用其他的网络5 l- R; ?7 e1 u3 b) o' K: q. N$ c6 j& n
def train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):5 ~9 J: H: w8 y$ U1 |
since = time.time()
1 M1 ~; `9 @. G #保存最好的准确率& {% o) F0 @1 N1 ?" i& o6 U; `, D9 |
best_acc = 0/ C2 q8 s- {+ X: s/ H7 v
"""3 u. m) E; u8 r, [7 e7 |
checkpoint = torch.load(filename)
. q0 @' u. \. K best_acc = checkpoint['best_acc'] ~3 @( ?( d8 q0 l- Z3 k
model.load_state_dict(checkpoint['state_dict'])7 I$ l# v5 c% y' i! O) l+ X
optimizer.load_state_dict(checkpoint['optimizer'])1 @, o( O$ h" V6 I( {4 a0 Y
model.class_to_idx = checkpoint['mapping']
# r/ }: N y$ Y+ E* i. ?% F """
: O" {; W; J+ w0 S+ z #指定用GPU还是CPU6 a8 W; g4 H: m# E+ g
model.to(device)
3 R# o% M/ k9 P$ G% g4 W* W #下面是为展示做的' I1 q1 t: N" t9 D: g( z( h- i% O/ H: q3 [
val_acc_history = []/ d q4 }5 I, v+ q7 [+ ]6 x" D! o
train_acc_history = []1 r9 Y& K* e4 b7 }* { u$ R1 C5 B
train_losses = [], h& r! r8 v$ J& n
valid_losses = []
" b$ F( [9 ]; n& Y b! C! U, c% U LRs = [optimizer.param_groups[0]['lr']]
7 a4 L% g! d1 I #最好的一次存下来% a. Z' T! `. ]. ~( |) O% p. `( `
best_model_wts = copy.deepcopy(model.state_dict())3 o, {* ^& a2 M
. B$ w- U" a6 k% y3 S8 L& E
for epoch in range(num_epochs):
, c: I, i% V' F& o) l: @# j print('Epoch {}/{}'.format(epoch, num_epochs - 1))
8 T: k* S4 X/ s9 N1 c. ?) q print('-' * 10)
. ?) D5 s% {1 q0 c1 ^" x( q
" w7 Y" p( B. E # 训练和验证- j9 o) `' Y# V+ R- f: Y! A& ~) _
for phase in ['train', 'valid']:
+ r$ @. Q+ T X3 ` if phase == 'train':
7 q; s* L. A' }/ S2 x; Z model.train() # 训练8 x2 Q& c" R: H+ c9 ^% B8 T
else:
( U: }& Y, U2 Y% N4 D2 V model.eval() # 验证5 P# H2 h. B4 E+ b, z* g
9 b* S) t" \5 P# |
running_loss = 0.0
& O5 k1 c8 D) ]6 p* ` running_corrects = 08 X0 J# |/ M& D& `
+ g3 X3 G# M" n: P: @. u # 把数据都取个遍4 H/ t0 Z- ?9 [# D4 r, o& d' D
for inputs, labels in dataloaders[phase]:
, j5 w, q( J D" C; s }6 f- c #下面是将inputs,labels传到GPU
+ @/ k, Z/ {; N inputs = inputs.to(device)8 {, W6 V, r( O% P; ]- V
labels = labels.to(device)
) F4 `7 @0 G p0 a- Q% R$ I9 Z+ e. B# @8 A5 O
# 清零
/ {/ m3 a! g) I# e optimizer.zero_grad()
- X5 L# [( |. Y3 |# r' n # 只有训练的时候计算和更新梯度, I. G1 P/ ]. n7 N* h+ Y
with torch.set_grad_enabled(phase == 'train'):
! c; x, u/ N7 f! m6 R1 O #if这面不需要计算,可忽略, r U: o. L9 D9 y& z8 `0 o
if is_inception and phase == 'train':/ F+ ~( H/ X4 M# j
outputs, aux_outputs = model(inputs) n1 }$ O7 U1 Q% i3 F7 R: n
loss1 = criterion(outputs, labels): i2 Y) x* _9 V' ?- y
loss2 = criterion(aux_outputs, labels)
. d. t l5 D* u) R loss = loss1 + 0.4*loss2
* T6 M9 E% W6 N0 O. }8 F# n' u+ O else:#resnet执行的是这里, H8 r1 H% Z- s \3 [
outputs = model(inputs)
3 S( M8 W: a/ h+ Z# E7 o) c loss = criterion(outputs, labels)' s6 R+ S/ L) Z" D6 v p8 t$ A9 X
4 y/ U( a# P6 }. h #概率最大的返回preds" c* S3 C" P0 ~. Z3 K1 l+ @
_, preds = torch.max(outputs, 1); T5 g+ p" y3 z7 R, x/ [* l
G/ F, c6 G, l) N2 U0 `2 \* { # 训练阶段更新权重
5 t" y m9 L/ k+ N2 r+ a" _1 q if phase == 'train': W4 w& ]- y0 d2 }- j9 R6 P
loss.backward()* C( K9 G k7 y, u2 n
optimizer.step()
A1 j8 {) k( L# b, L8 c9 T- v+ h+ } l! C" k0 b) ?9 E
# 计算损失8 }, p6 p2 U7 T% p" ^& |
running_loss += loss.item() * inputs.size(0)
' L. R! N( u: Z running_corrects += torch.sum(preds == labels.data)
8 P% ~: V3 E7 ^; c4 M) P8 ^0 g) q& n
#打印操作
" Q( ^5 [8 l# a4 Q epoch_loss = running_loss / len(dataloaders[phase].dataset)
/ Y1 _1 C5 g" i7 q( e; B epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)
( C) A4 R+ Q/ p, k! C, V* I c9 b& Y: D9 y9 O& r
/ N2 r r) V8 d
time_elapsed = time.time() - since0 \4 H- z* B# a; ~6 G6 q
print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))4 _' b$ E( C# i; b7 E: j
print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))
" q: s( l9 S, R5 [9 v
* e/ P6 M" p( {% X, w1 I* ^7 \; n {0 z
# 得到最好那次的模型9 K4 o% q. o4 l2 V7 T
if phase == 'valid' and epoch_acc > best_acc:
1 f3 x& y8 _1 i% n7 G- [ best_acc = epoch_acc
% _( z0 X1 N( r9 j' O #模型保存
% e" i0 w* W6 x best_model_wts = copy.deepcopy(model.state_dict())7 P: F* ?3 G0 y8 n* J2 d# H
state = {
3 A A6 u4 A6 R9 j9 `$ z, r; h #tate_dict变量存放训练过程中需要学习的权重和偏执系数
4 p% u2 L6 P, w$ k' r3 w- B& W 'state_dict': model.state_dict(),3 \# k: T7 U6 g
'best_acc': best_acc,/ d: S4 c3 u' n- P' @. v
'optimizer' : optimizer.state_dict(),
' _, d' O- F8 Q+ o" B0 t+ K. C }, Q7 p) v( f6 [
torch.save(state, filename)) j+ Z# W x7 P, D9 ^
if phase == 'valid':8 \9 E) w& b4 i& Z8 h
val_acc_history.append(epoch_acc)
5 T' d) B D/ v valid_losses.append(epoch_loss)
& H9 N9 a/ x B& P6 J" L9 l scheduler.step(epoch_loss)
* Q4 [1 Z' R$ R( o0 C" A if phase == 'train':
, m0 C4 ^ c0 \/ U/ Y* T/ U train_acc_history.append(epoch_acc) r2 F# {. J, Y8 X& Z4 N6 q
train_losses.append(epoch_loss)
; o! T6 q, t5 U# `5 Z# A1 y* e+ E& [
print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))
3 M+ G, T8 A- q" D LRs.append(optimizer.param_groups[0]['lr']); X0 K: k/ n! W) ?" i$ P
print()9 }. P9 ~! w2 J; h
3 R( Q( X) ^! m0 E/ U; `0 |/ } time_elapsed = time.time() - since
0 P% b4 y8 I8 ~6 u( k: G print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
9 V) R- k2 ?4 P print('Best val Acc: {:4f}'.format(best_acc))
6 x1 i" p0 M: ^2 k* ?' p) x/ R) N1 L3 L0 e
# 保存训练完后用最好的一次当做模型最终的结果" O: [7 M( R" Y s7 c
model.load_state_dict(best_model_wts)
2 v" J7 R( l8 k return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs F6 X3 b& g W9 c. K0 X
3 I0 w% h* w8 G2 E; k( K0 }: z3 m, c) n7 W/ z9 F) S0 W$ T0 r( r7 K
1
0 O# X- [2 Z* ^ u( z2
) A1 Z5 X* Y4 F8 ]3
! `. F1 X+ |; c& f) N4+ n6 k0 y( t5 Z1 R+ M
50 z- z; w6 X! }4 [5 P% w8 C
60 C V: O( P' N' Y2 \
7, R* c1 H- T( ]5 W
8
, ~) y# ]- b1 k97 ^) n- u' [' t& w- ~8 s9 e
10. W$ q0 k7 M' c; g. O( R
11
4 P/ s* o, A, d B8 t! N12, c! d% q5 c9 ^$ x- n% c
135 @, t9 ]9 v! _" ^5 D# T! h" O
143 B) G9 q, r0 g# `
15
9 D9 h+ ?! r3 M! _167 D+ [, n- ]0 {; Y0 O( V: A" N- E
17, G3 g; ]3 l# T L" a
18
' U& P% s' Z8 X$ X- m3 b/ f192 I5 @' B# M$ w/ X& K$ f. C$ Y% h/ u
20
, _) d+ U2 q8 R21' e/ e8 B0 v- ]/ N( Q
22
6 P& R! {% Q2 c- f- [' i8 @23
/ s+ u, w" _0 |8 ?* O24
( f( C: @' l$ i7 j' g255 }4 B1 X! p% Q6 y: U" A9 V8 c
26
0 C0 {; U. h* y% w) T2 S1 N k; a27
& @9 E( F2 [: u/ o# B- X( }28' o9 Y3 f. ~/ V' U
29
! v$ u9 ~$ c+ W309 |9 W$ w- [' y& O8 ~0 X/ S
314 y' } }( _4 B5 F& g* |
32
3 E8 T6 |( l# h" ^4 X33
7 G. Z7 m" q, u- z; y: `3 [4 w/ A" z347 M" m( q3 S8 b' m3 ?
35
- n6 |* G+ i9 S4 F36. {2 s9 _; M; h7 G- r
37
* t' e4 U ^) ^+ U38, q+ i* [/ \" J5 J2 b
39; }. a0 l1 t( w9 r+ ?' a/ l
405 O: f+ ~7 k, F4 t
41
2 ?9 L, X" R5 p" ^% \42
% r2 v4 O' P+ d; N43! l3 [& y. X" m( V" ] v
44' o) A2 D2 x; X8 c5 c: s' `. h
459 }, ]2 z2 R4 o/ L
460 G. u& V, W0 ~8 ]+ {; g
47; ~& n+ Z- `, V
48 L1 @$ i h5 D- U! d$ u
49
( A4 N; |) {7 B j; `50
% o/ _, `; X; O: I* P; z7 C51
, Z1 b: Q+ i7 K, h: F52
4 r& H/ q6 H8 T; }2 t F% J" f; `53
- {8 V9 _6 g0 f54
8 e$ X: o7 S% e0 h9 A" I |556 v$ l5 p# N" \4 Y. p# P4 v
56
9 L" e2 F1 U9 v2 p5 p8 c57
. ` o" C+ L- @* H58
; k2 _% N" e/ L59
I1 u$ E5 `# E/ s# e" j/ [600 E, R p# H4 z9 p. F( b, |! K
616 X6 Z" x; r6 L' d' z
62- L! Z; G/ s6 f! p0 s6 O! a
639 X% ?( P) n/ [1 ^; e7 J
64- P) _* r, I( K5 ?
650 k7 h" Y F4 r& a+ P$ o7 O+ q6 I* p
66
8 K0 Q. v2 b9 J! f6 ?4 C676 E+ x7 M1 o2 R
68- T! g! E. X6 o% R2 ]
69
' Z; z& k% G0 O; I( G70( O2 W/ G% O1 m) x, D+ Q/ s$ m
712 n! b, P& k- t8 [* I5 d9 |
72
7 o' Q% {8 e. u5 c, s73" F! D5 p2 L3 F) t$ s6 C1 o( b
74: ~0 }, _5 \* z) F
75* x; }. P2 W7 o/ h
76
# B- L* x9 b* K2 v77% C# F, O0 P! ^0 P+ f4 ?
78
3 U1 r. O" m& a: p9 ?$ I( A4 L( t79
$ o$ B( C: Q6 g% ?80
: E, i& P4 c% t6 {0 @$ E: V2 c81 _; X, T6 d$ S' I( C* q
82 b2 |1 [) @1 Q
83
% J; \5 C0 ^. |+ |& }" `& Y84: ~! b3 K* C8 Q- i; D0 [% \
85
' {; o% Y9 {8 V" A, q H86& i2 H" p% ]' o X" \. b
87
! Y8 F( j+ B7 s88
8 O" D* B9 O; p2 H+ B89
- U( x" R4 |$ Y8 C! b; M' s/ m90. v) j( Q, d+ U$ g4 V" V6 y
916 u4 d2 U. S4 j5 ^0 D( r1 k% K) Q1 T
92) @( d i) _) g* u; h+ e
93+ s9 _ Y' d3 V7 A" |- f: L c
94; |5 u" q4 A! G: C5 M2 }3 Q7 s) ^
95
& F4 l# B) y$ Q1 f: K5 j0 D96, z4 m- u) I8 F7 f+ z* H5 }
97
$ ]* n! M! r. t. q& r4 k982 d, L5 {% \9 t; n. N( l; Z' c
99
$ I8 o+ }, ^: B: i# a1 E100
4 L( {( g/ Q# f101
4 ~- H5 A: d2 V) N$ q8 \$ f102
' _- Y) b U# Y8 n+ |* M9 S( Z3 W# p103
* t' {6 M% k k+ G# y& A104
; A7 p5 E# K' o; Y* |9 U3 p105
; }) j, W7 r, t106+ `( e+ H b1 C! f# s4 }2 N
1077 X% Y( `/ A3 K5 \6 V
108
" \0 G- w: R% I109* w4 H* N. h& d1 I6 b0 B
1100 @2 m: G( L/ d+ N! z
111
1 W' T' h" l5 I( ^% A9 W7 `112* j" R) ~; H1 Z
7.2 开始训练模型6 K9 P; O ?, } ]. Q5 X: w
我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次
' J* j- J: w6 d
/ }8 A a& e- m% h5 O8 M4 V- [#若太慢,把epoch调低,迭代50次可能好些
. Y. {) s. Y2 b0 ]* ?8 ]#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合
# }! d' q% p3 M% ~9 M! Zmodel_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 g% F A3 `. y" }. H9 E* X: W1 e9 A4 J2 I) C5 x0 c9 L" q
17 `9 ^- H- w/ \% Z ~9 Y2 x
25 T9 }/ S8 ~( O) Q. i
3" |2 @5 w ? c) r5 ~
4
2 C( k" Q8 ^ z6 `4 k* ~Epoch 0/45 L# p4 ]3 z. u$ i: r
----------8 i2 x8 g1 |# f
Time elapsed 29m 41s' b3 Y; w- o5 |* q. X
train Loss: 10.4774 Acc: 0.31478 O+ \" e& o6 e! t% F2 Q7 j( q( X
Time elapsed 32m 54s
$ f4 J7 \- y5 Y3 F0 A( Dvalid Loss: 8.2902 Acc: 0.47190 o$ P3 K$ [' C7 } {4 _9 c
Optimizer learning rate : 0.0010000
4 _) A# F& m9 j2 f4 z$ g; Z
0 `3 F- k) c0 GEpoch 1/4) t2 n: D( L( ]7 D2 H$ p1 R7 o
----------
0 _ Y; X" X" S* V7 D4 H7 S% hTime elapsed 60m 11s1 y5 Q3 `' }3 y$ N# h& k. d) f3 q! s, u
train Loss: 2.3126 Acc: 0.7053
4 @' P. _3 F/ S" E. GTime elapsed 63m 16s
9 L; V. r" C9 zvalid Loss: 3.2325 Acc: 0.6626
\( R/ o% M. p- P( m! NOptimizer learning rate : 0.0100000
9 R: p2 `, N; J
7 _+ [" b/ J- i7 YEpoch 2/4
/ m) A3 X- q) i y, K0 W----------
% c3 I( i0 ^# \) YTime elapsed 90m 58s' |1 b: |7 [8 T6 v+ S, A8 Y9 R# L/ E3 E
train Loss: 9.9720 Acc: 0.4734# l: ?) e9 s, _- u/ d! L8 f
Time elapsed 94m 4s7 J; t6 H/ B( F: h2 F+ H
valid Loss: 14.0426 Acc: 0.4413; h+ S2 C; N/ x3 h& k
Optimizer learning rate : 0.0001000. z, a2 H% t% T/ B' S
5 v& J& L5 {# _: d3 REpoch 3/4+ {; q5 D6 w2 l/ A) W
----------4 G& `1 y6 t7 r, V' g0 A, ^( T
Time elapsed 132m 49s
' j- n6 l8 _; Xtrain Loss: 5.4290 Acc: 0.6548
; c2 k* j. R0 i r( L. C) NTime elapsed 138m 49s8 P: y" z" D# E( |7 @, l+ n
valid Loss: 6.4208 Acc: 0.6027: A5 m+ j9 y" f! f6 L* H% w7 y* b) d
Optimizer learning rate : 0.01000003 f: D3 ~' s: S' R2 t- r( ~
5 s% v( D+ N0 H) s# K) E( b& C7 \
Epoch 4/4
) r, t: p4 T& e5 O& E1 |----------
/ h7 g! P' m, v2 ], {* C3 g6 LTime elapsed 195m 56s
$ n2 r) l% `- p4 A( m8 q9 Ltrain Loss: 8.8911 Acc: 0.5519
1 J; V. ^' ~( `0 q* E% A/ hTime elapsed 199m 16s% I$ O3 D" t! K; b: U) [
valid Loss: 13.2221 Acc: 0.4914
4 k1 T6 L6 Z- v% V: G9 o" `% _Optimizer learning rate : 0.0010000+ t2 {0 W# U0 W' y m
) e; `8 p P7 A/ `
Training complete in 199m 16s
1 C& _6 o- z nBest val Acc: 0.662592* @% o3 G4 y+ F. a0 ~: `: l& q
! V7 N& V; B2 B) L
1( V# u$ j- [9 s6 a
2
' h: u: u& o; f3- ~' s9 H2 D& x- B
42 I1 v6 T* t' T9 k' p; m* O
5
7 L# s3 c% {' ^7 t69 U6 X# N+ ]( m8 n; T1 B/ p3 `
70 s* v# I% E" G3 m
8
# a& }7 J5 D9 d- Q& d' n9
3 J: z; V6 C) G N' ~ @, |10
+ H* P' R- W2 H" a% q: S" v11
1 I2 ~- A2 U3 a. Q129 C) [1 ]$ d X, v3 W3 L" G
131 e9 E0 {- j+ j+ H& R- J
143 J* X. O/ e n5 s% `: ]
15
. k. F( j/ T' P& B# c) N163 J y' s! R" Q2 c! G) y! k( u
17* _$ I) ~3 f4 w% t0 ^+ F }/ I
18
& ^; w) F& N' t7 D19
5 s( {: ]6 |0 n4 Y! i20
0 J2 a# d. C& `4 o) a7 \- g21
* D$ u( v& l9 T8 r8 C9 L' l% y7 Q22
$ i0 T. p# C- w1 K- {9 z23
5 [; F5 e0 t" q245 V6 q& K; `7 y
25
0 q; w' L, C/ W& i26) f- q. Y2 F- z
27: E9 u3 A, O, h' a% @2 f
28
' D9 Y: k: |4 p29
2 b5 g0 U5 z g/ r: }/ O6 y30
: o* i% O9 `: p3 O! d" O- I31& C5 H8 ]- H$ g
32
8 H# S* O$ R# k2 u' x9 X( K& {33
2 f$ Y8 F) J+ ^+ D. X349 k! \; K9 [" ^" X
35$ F/ w# y* n; h3 k2 w
36& I* d) b# @8 W8 p8 i/ Z E: P
37& r$ C7 n9 P/ `" L% ?0 S: m; p+ j' a
38
3 C( e) G& w" r4 y! j+ W39
1 w' g- _7 T2 _! m, g+ s40/ g8 l% m1 j$ v" d) s
41
& o& k; \+ E, e423 V8 z9 S) i6 e2 ?; P- H
7.3 训练所有层
6 T5 S& ]% Z4 j# a; S4 Y/ T! A+ z2 L# 将全部网络解锁进行训练
5 V% J) P, {! w7 d2 o* |for param in model_ft.parameters():& u; J5 @5 o; U( B) S/ h& t: j
param.requires_grad = True
7 S5 a8 B( K% _3 J2 R# k# f9 `1 n5 G+ W( r
# 再继续训练所有的参数,学习率调小一点\7 Y0 G& p. ~- W7 M- y! `
optimizer = optim.Adam(params_to_update, lr = 1e-4)
6 ~8 O) H) G$ Wscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)
+ e! Q+ p. D5 d& i2 |' z; k1 Z1 M. Z8 W h* {
# 损失函数
; I$ c/ K& M1 o- P6 ~criterion = nn.NLLLoss()
4 Q" D/ f+ Q0 W# b# i+ _( D1 K( P3 |% ? j7 B
29 r& i# g$ x2 z0 K; T0 h
3
" h" q( N4 s- ?$ C48 b, k- |( \+ G
5
9 u# L1 v8 N* k# y/ k9 ]6
+ z' ~2 K7 z: o1 |* K' H* d9 _4 ]: E/ M7) ?. S' W9 t- e! F s
8
8 B3 P/ k0 q3 f! o `9
5 _& W0 M' X2 F9 T; M: O. \10- T, @3 g. s7 F6 I: N7 L/ R
# 加载保存的参数4 b$ y# K4 U' ~4 q7 @; x! Q- t
# 并在原有的模型基础上继续训练, P* x8 t5 X+ K1 S$ l- X7 p
# 下面保存的是刚刚训练效果较好的路径
- n' d" _2 n4 _9 k; q, P2 U. r1 Ycheckpoint = torch.load(filename)
7 u5 |) z$ ?$ U2 Ubest_acc = checkpoint['best_acc']
6 T9 {. t6 f! Z1 qmodel_ft.load_state_dict(checkpoint['state_dict'])! R, O g6 y: U
optimizer.load_state_dict(checkpoint['optimizer'])7 Y, E8 f# M2 i5 L# p3 o& m
1
; ? R. Q0 y3 Y9 K2
: h: y0 Q4 L; z9 s% C3
6 V. r! @# W; v6 X2 i! R9 P3 `! {4
[/ p$ Q/ @$ b5 m8 p- A1 e5
9 ]# K# j/ |8 B" J( Y+ E6# }2 n% O! ~" @6 y" y
7; }4 k7 o$ v$ u$ ~" t/ ?
开始训练
B; y2 Y3 `) c注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考
$ z4 B# O+ B3 d# X8 i5 C9 h3 U
1 {+ n9 `5 X( Nmodel_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"))
: K" T* x% t# {! r* d: @% a. |2 b1# o( \+ t, D+ O0 A
Epoch 0/16 t# g% l# [$ S4 I# O
----------
* y6 N+ P6 H7 S: a4 pTime elapsed 35m 22s- ~4 U, q- G! |% i# Z1 R1 v& o
train Loss: 1.7636 Acc: 0.7346) N/ |9 I3 j9 m& [
Time elapsed 38m 42s
1 G( P+ u- H. M* S2 _, N' d) Fvalid Loss: 3.6377 Acc: 0.6455
9 R9 }- k$ J0 Y, E( YOptimizer learning rate : 0.0010000- |+ v$ g+ q. X7 p* [% ^
; Z3 \/ f9 n3 w# ~# k
Epoch 1/1
( M0 R' S+ P# `+ j----------' q- c. N% }1 p7 Q
Time elapsed 82m 59s: i" ~, b; b. d9 K, h
train Loss: 1.7543 Acc: 0.7340
* r5 ^( }. E1 M' E2 b3 ?1 tTime elapsed 86m 11s
) ?$ d. g5 p$ L# ?valid Loss: 3.8275 Acc: 0.6137# P! D) y6 X! U3 G) J8 j
Optimizer learning rate : 0.00100005 D0 U! e, c" S7 v6 V$ c( z0 X
8 `5 H9 l. L! ?+ _) z7 S( f
Training complete in 86m 11s
) M$ }* O! t; n. d/ b7 QBest val Acc: 0.645477' w+ y8 y6 `& d6 k, a
$ }9 J* g @( [8 _1 n4 ?/ a1: E! T# L6 |$ i: c
2
+ C+ |$ j# R _7 B7 x: o1 o0 P3
7 Y" K O+ J+ [/ B3 F4: Y& B0 C& z7 Y3 a' r4 |" ]) o
5
' h% z. {+ R* Z% \6
7 \) Q" R# Q9 N; A6 D0 V5 Z7; P2 r M7 t2 n) R* J: p
83 I3 ^" {) d: e8 H6 \
9
7 ^4 a O( C% g% }7 H. u108 D8 f: Y5 [4 [; w8 }
11
; Y. E6 ?/ D. C+ l O; j12, F& J% G4 N$ T9 s
13
- t% M( b1 H. [3 ?5 n; ?14" {; f& O+ V3 I* i. A" h+ P& _: H& o
15
2 }( `! v6 O2 o% t$ O5 z1 S16
4 q( g% {6 I8 U" n1 L- T+ D170 \9 d) t5 G. Y4 t; B! T9 T
18. Y1 q! e; K5 C+ s
8. 加载已经训练的模型! Q# z) R$ l3 H8 P: c
相当于做一次简单的前向传播(逻辑推理),不用更新参数7 Q6 S! z5 \/ s& H0 \4 ?
5 L+ s: Y9 {+ V5 }* o! Nmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
8 z* c: C. P l% V
. z$ X4 E( @- K# GPU 模式) x' Z% d$ _- `$ P+ R
model_ft = model_ft.to(device) # 扔到GPU中6 K7 f% C4 r8 N2 C" N( k0 R( K
" Q6 L7 `( I5 u$ C# 保存文件的名字 C; R- F6 E) x
filename='checkpoint.pth'
) X5 h+ q' u, w+ l: [
# T+ O' E+ I- u. Y5 E4 }: E# 加载模型; r' m" ~# |* e; |7 S
checkpoint = torch.load(filename)
/ |2 Q5 B3 x9 Y( k+ k$ T; `best_acc = checkpoint['best_acc']$ j4 P# E9 B H, G
model_ft.load_state_dict(checkpoint['state_dict'])
: w8 V! J) N: |7 O: p1
& B+ q+ e$ p" ^: ~2" o/ X& e5 ?% M- k$ }2 U6 S3 V
3! Z# Z0 v4 B7 _1 C; T
4
) J- j6 e6 g2 @' K- h1 v5 G' r5 o" y! }; o' T- l6 c( [) Q
6
) V4 u- q5 @) d( Y, K7; U: n# w& ]! O7 M; @
8
& x% E& Y, T7 J+ b9* \3 z2 T1 I. C, q; ?$ c
10$ q6 g' X4 k2 [; c
11
6 q( k2 y z% u12
e$ r8 c+ ]5 m<All keys matched successfully>$ i! J; h8 N5 m# o+ s5 T
17 C; G: p1 r0 k3 d$ H# S2 m' y
def process_image(image_path):8 h# w7 s0 Y$ R8 Y2 C
# 读取测试集数据3 d! a5 {8 x5 C) s
img = Image.open(image_path)
8 B; ?, `- Z) k9 C # Resize, thumbnail方法只能进行比例缩小,所以进行判断. `8 B. n% t, `, Y5 ^
# 与Resize不同
2 ^. w( C7 e- x # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小" B' ^; R: [! U7 R
# 而且对象调用方法会直接改变其大小,返回None2 ^3 }! d. n' B- H) V5 @7 u' a {
if img.size[0] > img.size[1]:8 I* l. }1 @" d; f3 l
img.thumbnail((10000, 256))' R2 `0 r! `% `8 M
else:" A7 I+ s7 D& v4 e( x4 K9 a
img.thumbnail((256, 10000))
5 `3 P/ x' u4 i" m( q! G5 }4 t* b% ~, k2 J' Y( C3 M W5 r
# crop操作, 将图像再次裁剪为 224 * 224
; \# ~, s1 P- i b* e8 b left_margin = (img.width - 224) / 2 # 取中间的部分2 y. T8 C" {; A8 q* O/ b8 [
bottom_margin = (img.height - 224) / 2
, ?5 G) t0 Z0 u/ H2 x right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度
7 D& A# t4 g2 q- d2 J% { top_margin = bottom_margin + 224& x; u/ X- {( t/ \* [! M
9 @" [$ {# ~/ n# q" a
img = img.crop((left_margin, bottom_margin, right_margin, top_margin))- t0 U0 a; E' O: I5 W; e
/ l/ e/ L8 F, h9 X0 l # 相同预处理的方法9 @2 I; E0 j8 p! ~3 L! ~
# 归一化. m, J9 I9 ^( d9 i) {# o
img = np.array(img) / 2556 H- \( \/ P$ H5 U+ o; z
mean = np.array([0.485, 0.456, 0.406])
+ r& r! B$ \5 Z- m std = np.array([0.229, 0.224, 0.225])
" U8 v E0 M* J8 B img = (img - mean) / std% b4 H- _ `1 K
2 w* ^5 |' L1 y% Y; ]2 Q0 Q, `
# 注意颜色通道和位置
+ }! x3 F( K7 o) H img = img.transpose((2, 0, 1)), S$ t9 A& f+ M( x7 v
! V5 ?, A# w# p' @% F: }% U2 c# X' V
return img
9 k3 m; k- V# ~: Y4 j( g: f7 [# M
8 ^: s: g& Q/ }: |$ xdef imshow(image, ax = None, title = None):
3 T6 ?/ O6 w8 t( F """展示数据"""- u, D7 R' p) O8 D
if ax is None:5 U6 ]( L6 z( V9 d0 N9 y
fig, ax = plt.subplots()
9 |' s3 n# L; \! ^& [4 B# ]& a0 h
# 颜色通道进行还原
/ b1 }- l0 C) p/ r1 }- c1 o image = np.array(image).transpose((1, 2, 0))" F+ D( H( j6 Z( V( Z( m
" M' r& t" ?+ U: p: c # 预处理还原+ G/ s# q$ s' h
mean = np.array([0.485, 0.456, 0.406])
$ n. `/ a1 G# q0 `" P, G std = np.array([0.229, 0.224, 0.225])
' ~$ w3 B1 b' o6 M" f/ `: O! N image = std * image + mean* J2 ^9 D( e& Q
image = np.clip(image, 0, 1): |7 C; m: `% Q
+ @" I2 a) P1 t, B0 J ax.imshow(image)
1 v. I- w4 e+ @* p5 s! g" C! l ax.set_title(title)
5 U# I% v* _" x" r' m- E) b
9 k$ L# \- B4 _2 X# Q2 j return ax. u; C7 f# ] U( @7 @
, K# w% X+ l7 X2 }/ X, Yimage_path = r'./flower_data/valid/3/image_06621.jpg'
* Y4 u: _8 W8 l1 O; Gimg = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理2 M4 K- {% ~8 [- @: w1 t' U
imshow(img)
( P' X1 k: Z9 d6 g, S% ?& [& P% M3 o
1
8 d: T% _) {1 j) s0 m2
6 s6 J3 n7 A8 Q3
" l* p, q: S/ y* a9 W% T( H4
! S9 f6 A* y& z3 [% }1 f4 e$ j# N5
& A. z1 N3 y7 h' f$ V7 i$ h. B6
! u' z# b! \3 I7
# @) @, {7 r3 X, p# x* w8
6 S+ V7 I/ q, ?- |/ Z7 | |9( k) j$ O3 T( \+ v5 [) Y5 s
10" n& W4 w2 r, N6 k9 k6 J4 _ f/ x1 f/ t6 _
11
1 j. a; o2 M1 Q3 ~12
7 `6 u3 J) Y* a7 B6 F( c13
" Z* K: Z( n2 \* T. H( ^14! \0 |: \& }. v9 | O
15
" n. n) i" |% V; l2 q2 J; y, I16* X% b: e! I4 v4 L; V' ~
175 T! m/ v. o! z5 u+ Y* P$ N
18
6 F1 d5 m: F" E7 y) ~+ E" h, f19
+ e, L; h$ J: C( X6 M20
% U8 I. R' ^2 r. p- T21" j$ A" d- l+ i) @+ B
22 N; O2 w: h3 u( ~
234 t d3 g$ i8 x! Q: Z' ^
24
# M8 f. [1 Q5 H8 X$ t25
# n$ L2 ?1 u* z3 ^26
+ G x/ X* n! Q* p) H1 |27. K) M% e f) U& H
282 \$ a+ J2 h7 M3 C
29
z1 v0 z+ l$ N! b7 [30# I9 \' w( K! l7 {, _
31
4 B1 C6 N( w# g# j3 g# W2 d32* F7 _3 m: K/ W% J+ I
33' J9 R, B) f6 n% ?9 G4 q p' C
34
g" u- ]* i$ _+ K35
3 W9 |. F! a5 M" Z, L$ _3 n36) c6 d! u' Y2 y+ ^/ U) z: G
376 Q8 z. G2 `. s3 q0 `: G- O
38 F3 O8 I4 S( Q, {) g" ?
39 {+ a) _: O9 n3 J% u6 @* u; I" \+ f
40
/ N% [( _# i% _$ _* t% O" p9 D41
: a. M) W- I0 B42
) q( z8 R1 d! ~9 |43- ~% K4 g$ ~# e, P. \1 i* u
44$ h1 e9 M8 i. c5 }
45/ A7 {3 v3 `0 {1 o. J: o1 }9 \
46
5 Q! h! _7 W' X0 q6 `4 Y8 Z% B+ J47
% A' U6 @6 a! Z4 ~8 n7 y! D# s2 }48
( L" Q% h$ O. S! u$ t5 l49
$ l) M* m1 d( p* d! C50. J2 Q: w5 U' I# W9 |/ a$ E
51/ w) n3 p% E% h( {) S3 I
52 F# h6 w2 e6 }" j
53
, Z+ j6 u& B% z54; b9 c- }; @3 d: ?
<AxesSubplot:>
6 | Q: h- P' [7 p/ G1- X: q" O$ V X
7 o4 u1 Y* J. y; X' U' [' A
上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确
0 K& u [ f. v3 F- U1 h
" M2 g, i- I1 ^9 ]& p) V7 K! Y; Qimg.shape
! j; p) u$ o8 W) k9 W19 A& B: o8 Z4 m& o! [& ~
(3, 224, 224)4 @# K" G% n/ |+ q; {( Q) F8 {1 v! a8 O4 t
1
' C% [: C( |& R; V5 x/ `3 K4 l3 |证明了通道提前了,而且大小没改变7 a! r& c: ~( {( m
E5 Y) s. t( K( L
9. 推理
. M: |. f/ ?1 c$ p8 N& C1 o3 u; C1 ximg.shape# A8 n5 a. U* Z/ _% n1 M
% q- r/ `# Q; n% Z% ~# 得到一个batch的测试数据
0 s: t0 e9 b( R( ]8 L( udataiter = iter(dataloaders['valid'])
' b/ B$ Q! F0 [4 Pimages, labels = dataiter.next()
m; M! u0 K$ I7 U) x. K1 N, E1 x3 ]% v# I0 `: o) _
model_ft.eval()
$ a- H4 Y# S/ `' S# _
0 H( }$ X7 E1 Y& w2 ]- t& J/ [if train_on_gpu:
, n4 L% _% B1 m2 y) V* O$ j # 前向传播跑一次会得到output% z0 y& k l% e$ [/ x# j* q( D
output = model_ft(images.cuda())
# K: ~8 O, G* ?- G# U- Oelse:! k, P# W' J# n+ g' L
output = model_ft(images); u) m4 \+ G9 P
. A& ^- i$ J& s2 v0 U) s
# batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值2 x) u+ Q7 x" D) V/ Z( H: F& W
output.shape% f; u2 i v8 l' [* D0 K
3 y- X0 F! H) Q( p3 ?
1, n! Z8 a @ h* c, X3 d
23 G* Y B0 t1 l' E8 f _
36 h# y2 p* M6 n
48 V2 \" I# ] i6 C# o9 Z+ c: M
5 I6 V& K/ X' J' B
61 G! B- t$ l9 n+ L
7
2 W0 v- y4 Y4 W* e$ d8
8 F1 X8 U/ L7 X0 j: a0 j9
3 _. J1 ]+ C8 m10
) I$ {; g; j# L: W0 I3 j* o11
2 _$ Z! b+ I- D$ Y( A$ D+ H127 d) k1 h+ P6 o
13
( ]( z/ o3 C6 ^1 X2 w% g3 E& V2 f145 x y; t: D( [- V5 p
15
1 R U) H: Z4 Q6 D1 O6 |16
/ y. E/ a! n+ `. I7 vtorch.Size([8, 102])$ P4 R& v( ~. @ }5 Q: m8 h9 K
1
B* |1 c0 T, ^9 t! x9.1 计算得到最大概率
. v2 K- T9 D2 O, I# u_, preds_tensor = torch.max(output, 1)
1 x( M# R% g+ f* y2 [4 _/ K. p$ M; }( E( L. h3 S
preds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量+ W% }. z' f+ ]. x# \& `& S* N/ ]
1
" w: i7 ^" [1 f2) L" q3 G; Q9 {# ]- Q3 A5 \4 ~" K( E% P
3
# ~- j1 q- e0 V% } [9 K3 d/ A9.2 展示预测结果5 _) M8 d' d+ o4 m9 F% y. T
fig = plt.figure(figsize = (20, 20))6 m( c. Y! ?6 z. D ?, P
columns = 4
5 ?5 Y; [( e0 D6 \ K2 \rows = 28 o! n; S# c# a. d
8 f* @8 N/ @" F4 G4 R4 `$ m7 M
for idx in range(columns * rows):
- b$ y8 x& D8 B3 n" O# U ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])1 g4 S4 p1 d3 ?$ h+ z
plt.imshow(im_convert(images[idx]))
( R. f6 Z$ \" n! h ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), & I, R+ t* a$ G2 j5 ?, S
color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))# o) y6 l4 t6 U
plt.show()' a; X5 a( E* _) v* z
# 绿色的表示预测是对的,红色表示预测错了 J3 v* n- W$ y! p
1# p4 [' Y! w& m$ A
27 h& \% E8 A/ f) V- F0 e/ a- r
3% w8 ?' a% I" }5 R& _( B9 N$ W0 |/ i
4
; m$ J3 M7 R' m- Z3 q5 C0 N5( r9 w9 t1 T: |' K$ W9 k( h$ v( s
6
: ]- G, z0 ]9 s5 \. a7
- }; l+ r. o n) N I/ U$ v0 V( P! \8
: H; j& ] ]% d8 r9 C' h* `/ W. A4 J9" ]/ z4 A% |* v+ [% h0 V
10
4 V _/ W( D+ A/ \4 O11
0 H" _0 D! u- S5 I% |6 C7 l8 k8 a+ ?' b; o
# B1 c! P* v" [
: D3 V# G; v, G5 J! X————————————————. |* ~% p x2 u4 T( G
版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。
1 w A! R! J8 ^+ m& w原文链接:https://blog.csdn.net/LeungSr/article/details/126747940/ p l4 b/ L: c1 n# R5 x6 C1 \* `# J
& {1 L" k$ R4 F: I7 w7 C7 A0 q* e8 A
|
zan
|