- 在线时间
- 1630 小时
- 最后登录
- 2024-1-29
- 注册时间
- 2017-5-16
- 听众数
- 82
- 收听数
- 1
- 能力
- 120 分
- 体力
- 565692 点
- 威望
- 12 点
- 阅读权限
- 255
- 积分
- 174930
- 相册
- 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)实战案例
' F4 Q* Z: S1 p, ?' J$ |1 u$ z4 B& | k$ G8 s
文章目录
1 F N/ s. d$ u5 d$ i1 P卷积网络实战 对花进行分类
t P4 g2 o! B' L4 z数据预处理部分
! p8 U. r2 ]" Y5 A网络模块设置
4 J1 }0 I" Y/ K* R6 H网络模型的保存与测试
! n4 ]( | q- w5 ?. }8 f数据下载:
4 O* a3 r% L- R4 x$ g9 E9 F" Q) R5 o1. 导入工具包
: `2 K9 O' J- d; r1 _- x2. 数据预处理与操作
+ T* M, [7 i; y; D& H5 K# k3. 制作好数据源9 i* _+ F/ V0 o5 Y. X
读取标签对应的实际名字
" Z& X* U+ Y6 }: L3 I) y/ O4.展示一下数据 j0 |* F9 k* C# o; a
5. 加载models提供的模型,并直接用训练好的权重做初始化参数- ^& @5 M& K: F
6.初始化模型架构
. w0 c @! U3 s7. 设置需要训练的参数% a, g" M7 n+ t: t1 F
7. 训练与预测
7 d, k. v, S9 _% t) \: O- ?' c7.1 优化器设置, a: W9 _& D& t4 r
7.2 开始训练模型
" z) u' j% r- v2 z4 l) ~6 x, b2 m$ W7.3 训练所有层
3 w8 c7 _* w" _/ Y开始训练: D# {# t5 v" j; G8 a+ U
8. 加载已经训练的模型
: c2 J S/ o- N9. 推理
2 X8 I" I; R6 u: z9.1 计算得到最大概率 `# B3 ]5 G3 d. }
9.2 展示预测结果
. i( V" S# [& f2 O写在最后
- q. C# R' c7 K2 _6 H' H卷积网络实战 对花进行分类
4 w* \) R6 Y/ @本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能 }! ^: ?. D. v5 j: K) j
* E* j. a1 U! v8 j在文件夹中有102种花,我们主要要对这些花进行分类任务8 X. V3 r/ V2 }8 T/ h- e
文件夹结构. L: m1 k8 B- ]
3 t0 [7 d" b3 Z, ?, D* g0 ]
flower_data+ f$ x1 ~1 L/ _+ S( e
+ p: |2 @7 [ _% Ctrain
7 c; {+ c8 o( d1 g, I% m. M' n4 |/ ^! r. \, E
1(类别)" P; M$ i) y; k# \- v3 ?, N
2
' x( m: u) K' exxx.png / xxx.jpg
" I! |" q+ C0 z, x, d( R2 W$ Jvalid; z5 S2 s, |* w" K8 A; e0 P4 y
0 e9 P. c% M1 w7 I6 b! x- j2 J2 B9 n$ }
主要分为以下几个大模块
9 @$ z/ X: k* _# a/ F( @+ S8 U
数据预处理部分/ s3 B R0 v8 z3 S$ H
数据增强' q6 u% S; b( g; m
数据预处理
3 D; _. h& ?8 b( J v( N网络模块设置8 I( P3 Q( k1 ~6 V
加载预训练模型,直接调用torchVision的经典网络架构" l- ~: F9 h! b2 I+ ~8 g/ ~# r( y
因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务
- D W v& J: f9 I* l网络模型的保存与测试
$ L' C H+ r {+ [2 |: R- U模型保存可以带有选择性& O/ f" B8 U8 o) M% m8 @
数据下载:4 i/ |8 `( H Y& V& {1 C. V2 P! u. Z
https://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset
2 w; h/ [, H) Z% D9 O/ \. S" H7 @
5 \5 \. S) Z3 `, c% Z$ l9 o" t+ b改一下文件名,然后将它放到同一根目录就可以了5 o: q$ [5 M$ R
- p# w( `( p/ S9 p# `( ~% `9 L4 _ C下面是我的数据根目录
, L7 V+ {5 a/ p$ V( O' ~, t( n' Z8 k8 P5 z! W: e s8 F! o: K. t
/ q d# e4 J% `, ^* j/ v: d
1. 导入工具包% v( \: d' z8 {1 ?5 j) q' C
import os# C# M. ^: _- ?
import matplotlib.pyplot as plt
6 H( C" H1 l' n# q: W7 C# 内嵌入绘图简去show的句柄" c% g2 b, J2 s O
%matplotlib inline - f& A+ O4 `4 |& _* R; g4 B
import numpy as np
* ~; J" O8 K3 ^import torch$ s) z4 d: N' d- ?6 v: s1 b
from torch import nn
8 F$ [1 I" B# T' A4 {0 G6 q0 I
, z- l7 w- |) M- L ~9 C+ yimport torch.optim as optim
" W, _% ] v; f6 J! l% U+ Z; Rimport torchvision4 z9 j) v. I2 W
from torchvision import transforms, models, datasets
/ V1 D/ ^, ~+ P) b9 X' r- I8 _" a5 \4 J/ e) Y
import imageio V, d. k" q2 C
import time( u' Z8 m* a9 ~+ \) b; A2 L) Y
import warnings! j* O( J \0 F( F, I
import random
7 ?; q8 x* O% Dimport sys
& l$ p/ i- S( N! q, Gimport copy
! a# k9 W+ B) d4 Dimport json% M7 L7 q1 ~& j
from PIL import Image
: T% J, h1 G7 E9 _
8 X9 l6 C- m2 R8 ^, b' Z! V& v
; U2 T% q" p7 ?" ^' `13 c$ [3 q5 X) O
21 I; R0 E7 G' u2 r
3
% R. w& J0 [3 Q' o4
+ F: i4 X; h. \57 g4 E/ s8 H7 h" K& X3 D# Y4 Z' W
6
9 r3 H* {. T8 m* H4 M7
& l2 G" {7 W8 R8 a4 N8 ~1 W6 q- v9 q/ R4 d
9) G: s* `- ~+ V2 D' S& k
10
) S# ?4 v. E9 C. L' u4 F11
" l: S, J9 {3 F" h12+ N& L- Y! q$ B5 N; _. D
13 {0 u$ V3 X" o8 \# \- W1 j
14
/ x* k: w; Y \8 }15
: g& r; c H. ]/ Z16
3 f9 H" G1 Q& T6 X8 i177 F3 l3 H, S: S( f7 z
18
8 o( L3 N: M* j* d19) L8 [1 z( X/ `: I
20: y4 i9 U& X; u2 j
212 S" S$ y0 T, d! y! Y% C2 o' R
2. 数据预处理与操作! Z4 u! ]+ N; y& r3 f; x3 B3 m
#路径设置: w% x+ M6 _2 ]& `! P, w8 z
data_dir = './flower_data/' # 当前文件夹下的flowerdata目录) V" _1 J, v5 n; p
train_dir = data_dir + '/train'! a* r( i" B, z; O j
valid_dir = data_dir + '/valid'( v9 \, W/ M' O0 ~
1
% ^ r. m2 ]& L2' O6 R. t1 S& L$ Q! J! ~% T2 x
3, h! l) ?3 W* y* [
4
' \" j! O0 Y. z+ Epython目录点杠的组合与区别
3 I, O! x9 D* J6 s注: 里面注明了点杠和斜杠的操作3 ~# }8 @" y: J, F
, _) @5 R& x- k3 ]' f, ]: i3. 制作好数据源
$ B" e) A% ?; ^$ A( H; Mdata_transforms中制定了所有图像预处理的操作
[/ s+ G5 w3 O* B" pImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片
% \0 f. a, O, U0 j. Z7 Idata_transforms = {9 G5 J5 g9 V# d3 g5 I9 y" p! J/ @ `
# 分成两部分,一部分是训练; b. T/ J; l. n' P
'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间! |, G# x. X9 g$ P1 m. f
transforms.CenterCrop(224), # 从中心处开始裁剪
' n) o6 _3 t) s' j0 ] # 以某个随机的概率决定是否翻转 55开2 D/ M1 A7 w$ L5 r
transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转
! B0 L* F' j; [+ ? transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转
# p4 S D5 n3 f/ O9 w& [6 T4 K # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相' B- \6 {$ {4 [: e2 J, W0 J2 K {
transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),0 b+ h9 U% z( A3 Y( ^/ V, x# @
transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB
1 g) K3 x8 @3 i X$ |. y # 灰度图转换以后也是三个通道,但是只是RGB是一样的
/ H9 u, p7 a* T# m. K, s transforms.ToTensor(),9 R6 L$ }" D, }: |' C6 l) M
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差
( e r, A8 f: _ p0 U4 ` ]),5 f9 z5 n+ {+ K" D; h0 T1 P
# resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
& y8 ^: ^/ o+ ~ 'valid': transforms.Compose([transforms.Resize(256),* P0 [* I! u& V
transforms.CenterCrop(224),
( U0 Y7 z; V- A u% i4 C8 _1 M; f transforms.ToTensor(),; t. @, S$ j# D( U+ m! s
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同
: D5 V/ I/ e4 ]3 z ]),( ~" b$ F. Z* {8 s1 B6 e; A
}
7 L/ R+ d8 P5 ^3 b; T: c5 \5 e% j7 d, }0 u5 o: V- l( {6 z
1* `3 P& F/ b( S8 t3 s2 t
2; U& J+ i* `9 D
3- f1 O/ H% d" k$ R& `( s
4/ w$ F1 T" v+ S" G( x& U7 [" B
5
: ]6 ]& Q( y6 F* h6
: K+ }) _* H& b0 @- J e% _1 N8 F7
* a' _$ Q6 X- X& _( h. }0 `5 R; I) c8
# o2 a, z1 ` r3 G' ^9
7 e6 \# E& l+ b& t! n0 }10
+ U+ a4 x: [; d) h: O1 W11
) D3 o4 m$ a3 A12
) X" v9 q7 F8 e' }- v" n9 }& S136 s. U7 z0 L8 k
14; ^! v! d }, p& m
15
. r; j0 b* n |- _0 A7 p16# w- W1 A! V- U& `
172 {% N: k- b0 e$ H5 M$ {
18* _4 t/ b7 Y0 N
19" ~6 I% h4 {5 F' r# V+ r z5 a
20
. Q- n5 L1 t( ^8 l# D21
0 c$ i; ^" b) X) T+ X: ?batch_size = 8
5 p- K3 k _- a3 I1 j9 Timage_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}/ C$ ]" j. h. n
dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}
B' Y! }: d6 t; E) B" i8 j! {1 ~dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']}
" ~$ P$ M" V) bclass_names = image_datasets['train'].classes, v( P" s# d) V: P9 O' ~
; U* A) ~8 c" J% n, B
#查看数据集合
& T2 d* E0 a7 {# j% x# uimage_datasets' L" M1 m3 {! p- @; S
2 E5 f8 t5 N+ r$ l9 c2 z
1
) u& e/ a* O5 j* |; j2+ U/ Z. }5 D1 U" b& q5 I
3
) _, j, V: H' l8 g& a; o- M44 R* \9 H @$ a/ U- m- ~
5* ]# b, f/ N! I: j% h+ K
6
9 w, g3 ^9 p) n3 {- S! @7
! O8 u3 ? s8 j85 z# D5 b5 W/ N( F' x
9# T# T, Y4 X! x" q
{'train': Dataset ImageFolder
" X7 D! S& |$ ? Number of datapoints: 6552
3 ]. @& T$ \! N0 M( x2 ? [* S$ r Root location: ./flower_data/train' B. o3 V0 D/ v: X: ]( c6 N
StandardTransform
$ I4 ]+ I$ I( [2 a; D% k- R Transform: Compose() P7 U8 U& B% i" j2 R
RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)
& g' T. n6 u! z' v) c CenterCrop(size=(224, 224))/ j3 C/ k* o" L/ l
RandomHorizontalFlip(p=0.5)
, ], s# L5 C6 }+ |# ?3 G* H/ V RandomVerticalFlip(p=0.5)1 U2 p, z- Y! @2 b- Z
ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])/ v: @3 {. S9 \% R3 p ~
RandomGrayscale(p=0.025)
2 I. }! d: u7 _ ToTensor()9 [- }( \" k' Q- T4 t" ~! c; t* d
Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]); A4 m! _0 H7 ~" W. j8 `
), d P) i5 {: J9 x
'valid': Dataset ImageFolder
- `3 D+ M, l" o Number of datapoints: 818: `: [) @7 Z$ f0 Q, a) y
Root location: ./flower_data/valid6 K" e8 C. y0 S
StandardTransform2 h/ b4 |4 N3 ], N4 J' ~* f
Transform: Compose(0 G) j# I( W* A& C
Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)3 w/ S4 p5 Q: K" i
CenterCrop(size=(224, 224))3 y( s- s. L* P2 D
ToTensor(), Q+ ^, Y; E, M: s& z$ Y3 F. O
Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
, g1 {, K' P# f )}8 c) k4 [8 `- P3 U& K
; O0 c" g! |2 G
1
# F' o5 e) O" z! R! s8 r2
1 ^8 o5 k( S# V" Z# P! F, c: s3) A( u' ], q; y- i$ z
4
, ]1 X+ s! U8 Q( P53 x+ \6 u+ B3 z- ^4 g# E0 C1 G0 f
6& N! S' y# E2 |3 s; e/ K0 i
7
1 u0 D- d- ~! `) Q8
% f1 _2 H! s. Y) H; C. U) W1 g# Y91 r. w8 C$ f9 j+ G( r' i0 G" R
10! b& H- M; l0 v! {% i6 u1 S; |
11
: z# `# R% X9 d& B: }: I' }5 O122 d% f# M2 X" f$ j. I
13
6 i% \: F, j( p3 b: Z8 p9 H$ U# B14' P8 p' s5 v7 R4 {: a1 i( I
15
6 ]% J5 Q% q3 x' K! B" a: @16
8 k3 e$ a6 I$ r/ [17
/ d* c+ c& u9 V; w( L; ?# `18/ ?7 B9 v p$ z/ H. z. ~
19( ^9 c% n5 i8 K" b
20+ {& i; @% O1 K) \& x w( O
21
$ V' `# W" A& C4 t6 L$ Y8 e2 [22' |4 I) {7 ]. h& ?
23
5 t$ p/ F+ j' H% n' E/ [% ?24- O0 `4 f( Z% l
# 验证一下数据是否已经被处理完毕
3 u0 v L4 s* b! A C9 @& _8 Udataloaders0 _* {0 u$ |. d& d
1
4 N9 A" s4 {% c7 u3 g+ K2
( V9 G, a/ V. N2 k' l3 n3 @. ^{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,
* u: C1 g4 C" w5 D" a( H+ P 'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}
9 T6 o" ` B D$ n1
2 X, ^4 E8 G: y1 |' L2
\, P* R" W# k. D( m. @) wdataset_sizes- r8 Z" h' q% s9 @/ ]
1& e+ T: |) d+ G5 K; {
{'train': 6552, 'valid': 818}2 p5 ?! B/ b* }5 @. g
1. H! F$ {5 e% Y
读取标签对应的实际名字* P8 D2 @4 U+ U- F: G
使用同一目录下的json文件,反向映射出花对应的名字3 c4 W6 _) l! {3 i& t/ f
0 f# Z' H3 L# `5 R2 `0 \, i4 p. qwith open('./flower_data/cat_to_name.json', 'r') as f:
' X8 Y! |" W( }& u! Y cat_to_name = json.load(f)
" ~ P) r5 H. c3 b( c1
) }! P) ~# W$ F: e; d: R2' C- A# L y$ [. L, |% E* M9 U
cat_to_name
- W) @0 d* E6 F6 ~, U% r$ T1
, u d. z# Q$ D. o% W& I{'21': 'fire lily',
' l' {" Q" T; ^8 Q '3': 'canterbury bells',. y( R/ x) w U, E+ ?
'45': 'bolero deep blue',
! Z5 \2 K$ m: O '1': 'pink primrose',
* p; L# q% k0 b; f4 j8 ` d2 z2 o '34': 'mexican aster',4 M1 m! l" ^/ J* a. @
'27': 'prince of wales feathers',
9 Y2 m! o6 A( B1 E9 @ '7': 'moon orchid',8 u# _. ?. N9 Z, P5 ~
'16': 'globe-flower',% T* V6 U N: Q! j; b- N* a
'25': 'grape hyacinth',# b) f0 D7 A. l- a# x
'26': 'corn poppy',+ P/ R0 O4 g% A# s8 r6 m
'79': 'toad lily', l# N. i: C- g% m4 Z$ U5 O( U9 l
'39': 'siam tulip',9 U5 @; ` j3 U. `6 j2 r3 w. E& @) y
'24': 'red ginger',
$ Q/ \0 l9 G, x4 e% a '67': 'spring crocus',0 Y* f& I: p3 p2 X- l
'35': 'alpine sea holly',
/ M) J$ N- n: ~$ ~# v '32': 'garden phlox',
7 P+ E" Y, P5 |) M: ?3 F '10': 'globe thistle',) \: {4 y- q, Q9 @( ], d
'6': 'tiger lily',
" a0 O9 u( ~) b( o6 @ '93': 'ball moss',
5 U9 Q- L# k9 {- [4 M '33': 'love in the mist',
3 V: `# F2 m+ ]7 _ '9': 'monkshood',
. h6 @. X' ]7 L4 s# c '102': 'blackberry lily',& {+ ?6 Z5 e7 F p
'14': 'spear thistle',! J* Q @6 \) M5 O
'19': 'balloon flower',
3 a0 V2 J: _0 x4 W '100': 'blanket flower',
# N/ {' n7 ?9 A1 D( o0 A5 X; H, F# ] '13': 'king protea',
. B: g d5 N7 G '49': 'oxeye daisy',
& p" J) L8 v: q: f5 ~# o# r7 J# _ '15': 'yellow iris',
' X. c( b* { r5 } m, n '61': 'cautleya spicata',' O& h7 ~/ W! Z% _ D
'31': 'carnation',
4 r0 @2 `% t7 G! x' l '64': 'silverbush',
8 r9 X) i: c8 H. k `& \5 g& o '68': 'bearded iris',
, B/ Z& N7 C) \9 x; |* k4 J4 k '63': 'black-eyed susan',
* X" m# \- V/ m; @ '69': 'windflower', \5 z2 n; c. x) c, _
'62': 'japanese anemone',! W5 G0 Y0 s4 ~( f
'20': 'giant white arum lily',
u- ~% B1 s; |$ H '38': 'great masterwort',
# d; P8 |" M7 g6 p '4': 'sweet pea',
* W- J. T% ~1 Y& ] '86': 'tree mallow',
" j+ l5 S+ x2 Q/ J5 b9 t8 R" r; R '101': 'trumpet creeper'," \0 m* T( v# u( I- w& Y
'42': 'daffodil',
7 p9 P; r9 F @3 W '22': 'pincushion flower',/ F+ ?" C: R4 r# ]: q5 h
'2': 'hard-leaved pocket orchid',
$ z! W0 `( f/ k9 l* O '54': 'sunflower',$ A4 a. J2 \1 P5 j i1 U- g# g
'66': 'osteospermum',
: n! e+ V3 R& L9 g6 ] '70': 'tree poppy',
~' f" c+ O2 M0 [ '85': 'desert-rose',( \# O1 R- t7 P& f% Z& Y
'99': 'bromelia',
5 a5 i) ~1 f/ P$ f4 _ '87': 'magnolia',% T1 U8 ?1 ~, `3 }. |
'5': 'english marigold',' b, ]% S. U1 s' a& \. o
'92': 'bee balm',7 x1 Z# v. v* m6 q
'28': 'stemless gentian',# d& P# Z0 Y, c9 ~9 [1 j. w& Q* m
'97': 'mallow',
9 c* w, I2 t3 i3 | '57': 'gaura',& P: d4 ?0 n6 S5 n0 j+ y5 H
'40': 'lenten rose',
! P! }; E" ]: C) t, i* X3 A: p '47': 'marigold',
5 j' o: ^# ?2 ^- G; ]! r0 _ '59': 'orange dahlia',& b. s1 `% Y( ]0 u
'48': 'buttercup',
' j8 Z/ L, s( A! E+ L2 Y( O8 I '55': 'pelargonium',, i2 P: g' D, l0 a+ n6 o0 \
'36': 'ruby-lipped cattleya', ~) r& l) j. s0 r, z
'91': 'hippeastrum',
+ C! K1 i; E/ u' o. H9 Q '29': 'artichoke',: K" F1 Z/ B+ ]% B- n: C9 Y
'71': 'gazania',5 y/ D7 ~+ d1 L8 R0 _0 o
'90': 'canna lily',
. Q) L+ l, O* @0 @& P8 R- v. L '18': 'peruvian lily',
6 c- `6 U: }" F4 U" l '98': 'mexican petunia',) L8 |& ?7 |( ]8 z9 t% i& q( N, A
'8': 'bird of paradise',7 B7 H6 L( e0 s$ M$ a5 E
'30': 'sweet william',
, `$ {* f" `" z. w4 x '17': 'purple coneflower',: w8 @# W+ [0 ^; G+ W( {
'52': 'wild pansy',( ~6 m. C* f" H3 x' Q$ m
'84': 'columbine',
+ u( q! |4 S% w, b% V '12': "colt's foot",
& N. p) W/ H9 v" `/ l$ p4 B '11': 'snapdragon',
# v9 M! y- ?* B d8 `0 c% O '96': 'camellia',
! ~7 @! j U5 W '23': 'fritillary',
4 H9 w( f! v2 d! v9 S# F8 W '50': 'common dandelion',
/ S. |+ _. K' i! C: Y0 q '44': 'poinsettia',
8 h: M3 d5 w, o! n$ r+ O* ] '53': 'primula',
' q) N9 Z, ]/ z6 z+ [ '72': 'azalea',
~7 y# c1 Z; I7 Y( M+ N '65': 'californian poppy',; c6 k, `% Z: p2 F$ R9 O& V7 }
'80': 'anthurium',
" s8 x5 }- A$ ^6 { '76': 'morning glory',# o. @8 W8 m0 _+ k
'37': 'cape flower',: N# ^: @* |: z1 W* h8 N
'56': 'bishop of llandaff',- B8 K6 e# Y3 R
'60': 'pink-yellow dahlia',; k, s" H2 y% P. S" F0 r8 w* p
'82': 'clematis',+ o& u- d& |& I$ C2 _/ @
'58': 'geranium',' t# J, o$ t( R. K/ J) V5 j$ n
'75': 'thorn apple',3 Y3 x- M4 [# M8 d' W
'41': 'barbeton daisy',; s8 b n+ R, ~
'95': 'bougainvillea',
5 ^ m4 |0 F' L' }% g '43': 'sword lily',3 m! B! G: T2 h; B
'83': 'hibiscus',
6 U, V' ^) m% H4 ^3 d9 g. z '78': 'lotus lotus',
' c w3 F6 I! N9 Q% Q '88': 'cyclamen',
- ^7 ~4 Z4 q9 h' _3 e' m- m '94': 'foxglove',
5 x I; L; S! m- O, X, i4 R& d '81': 'frangipani',; j+ O$ y$ q# |9 T: y
'74': 'rose',2 K) d) L5 }; r! K" M$ } F
'89': 'watercress',
% A3 }5 Y5 H+ Z b '73': 'water lily',
$ U8 @! E3 [% z '46': 'wallflower',
6 n: m- |+ s. b) X7 o- D '77': 'passion flower',
3 N7 F( _2 {+ w9 f" p( | '51': 'petunia'}2 _" Y% l) s6 B5 q( ~1 n4 Q- b E$ _5 d
+ A& Q5 f" z9 ~8 q6 z. ?1
{* l; e% X3 X2 d- I, {2 v1 n) a5 J/ I l
3
3 Y% k. @' Y( k# a5 d; W- G; A# O4! a& S4 p! d$ T
57 Y, {# z- M, U8 g
6
$ J1 F" T: y& ?4 P4 V7; T8 n6 a4 O! ]# g+ ^2 w: O
8
$ A- Z" a! R1 y; `7 c. d( C99 F' c; [5 W; H; W: H, g# \
10; L$ J w Z* e, @1 z2 _" i+ l
11; i! Z+ z: D/ n# q3 G2 f- q% d
12
, _" B3 g& v, c3 M, d13
$ D9 M8 Y- X* }! x0 Y: o# S14
. D3 }8 {8 |8 S1 S' m. n15
' V9 K) }" n+ c2 \" y) S163 G; _. Z: V) _, ~
171 ~0 G J/ j- @) k8 f0 A0 v
18
3 N) \/ D9 N1 }' ?, {19
, ]% |2 k9 e! h) h20, e8 g* H; C# a8 P- B
21
+ L4 L2 ] g, w2 O22
* `2 x6 |' \: J+ L+ [. E' x23' t, ` J' K) R1 m) h
24
& G; k" ^& S, ~# W! Y25
& W9 I& U8 p* y+ q0 o26$ o* Y" u$ Z& o' }- q
27+ A$ }" t/ ]5 e$ F
28
4 Y8 f4 T0 l5 T4 M, p29 ?( q; ?) U5 X I; C) C; V$ S
30
( l5 \! ]# S! g+ b31
( A: j; [9 W/ o" N: f: \2 g( ~9 b32
" I* D& Z- I6 v33
$ h0 o1 j+ i0 f$ @0 e: u/ \34
( F" h* q- @9 O: f: F- S2 Y35
4 G$ p- J9 D& _- a6 z365 y3 Y+ f' H% c- \ p2 A
37
! \. b4 }6 q2 ~38
( q9 f8 K$ T/ u1 O39! {& n3 B$ S1 r7 r0 M
40+ r0 C. @, F# F. M
41/ K5 c/ G$ l6 _6 C( o; P; T
42
& x K/ y9 W# u, L& ]# F" a* o43
3 |8 o6 }3 ?' O! a2 P441 ?' e. @& i, v8 u* Z6 c
457 v) h- A+ K2 r7 q' m
467 ^4 n- T, q, j1 R
47
6 V! b5 ~7 V5 q+ \& l48
, J& N: \, L7 l/ V! {496 }* m5 z/ L; ]4 ^' L0 k0 q
50% j- ]/ Y0 B) w% s6 n2 Z$ n
51, Y; A- [( H* A0 E0 @8 e0 q+ n
52. u( g* x' s3 m8 a4 M
53
% {$ V8 u+ q: b3 Y( O6 g543 f) `/ [# B. j! i$ `
55
' k1 F8 n8 P* q9 y. W% {2 J56
6 d% S: R, S* L' S57
( H2 N* H7 U$ F58/ l! x( ^: O" a) V/ j
59
+ \7 e+ U3 H1 u60
: d+ f8 ~$ W) d$ `' `$ J61" d. b' \$ s! W5 z6 r
62
; ?' z Z' q. f63
9 T: S# q: @/ m: h8 z* _& u- _7 B64; N0 |3 ^7 Y" w& ~2 |. |" N
65" M j2 j" k. ^ m" g
66. G2 @- i6 ?) S: Y l7 T
67
2 J/ s! S7 ]. _/ w$ V r; r1 R68
0 D' S# y8 p7 o69' h" P& r5 S, D- b+ r$ A) x
70
9 h5 R8 D! U3 K! x' r8 C71: _1 H6 u5 q2 f# F5 o$ ?
721 G) l! q, u1 t4 K$ D9 d2 \- u
73
. o; I/ M) q* x" c& }$ z74+ U) K1 E* u4 ~/ E0 ]7 f8 `3 } w0 ]
75
) _0 M. z5 ^* i4 n, u( N765 h& @0 [, ]& d' Q" z
77
+ D5 J6 S" b% _9 J78# ~7 k2 }% a- _9 S0 B5 }
79$ q3 f9 q& W+ y) J4 z( n) l
80
: g5 {7 U4 ~# C2 @( O6 ^7 z813 Z9 h' N0 L' o7 s Q
82
9 ~$ H' h9 F j! g0 z83
8 a5 S r/ t7 P1 m7 ^ G" Y5 @84% f# D/ ^% G& }9 H! n4 d
85
; {9 G, Y5 |3 `9 ?" j. [/ W86) `( Q+ x9 k ], W5 ]6 n0 f8 B
87
% A# {" w) s/ p3 A88/ C: N/ W6 W/ l8 m a
89) y j6 n" R5 v' a+ z1 v; ]
907 y- ^, h% F5 i& P8 B! B. i
91
+ r. |# @9 O0 h$ C/ E- l2 ^8 ~/ t2 k92/ C( U4 b6 Z& @8 M/ S
93
" V3 ]5 z( r8 ~8 a94- j u( \7 W+ w" l; M# f7 f; E
95# L' x# j/ R2 ]$ c1 g
96/ ]4 Q# D( ?! R8 G) r
97
$ Q' K# ^5 Q e: C8 L; \( Z% E98
- H; I" ] p9 R7 ~99
W4 G9 ?& L0 }: F# l( u0 E100& b% l( e p5 F; G0 }3 U: I8 j! M/ U
101
# _5 P; i# R, c9 G. K102
- C1 X* W @, I4 U- \4.展示一下数据8 k8 B; H5 A5 a' E: V
def im_convert(tensor):1 F; m. `' Z- W) r6 z, o
"""数据展示"""
! v; D2 x& Q1 Q' g# x: ]* E image = tensor.to("cpu").clone().detach(): n" F# ?& \' ]* K1 f$ F
image = image.numpy().squeeze()
9 A, I& w# d# |$ S2 p! e; m, n # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图
5 U0 c. _) B: `6 I1 G # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)8 ]$ r$ r6 B [; X7 w
image = image.transpose(1, 2, 0)
+ ]# b. F S' @+ S # 反正则化(反标准化)
1 L4 H$ l) x$ g/ d image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406)) W7 I: G) M6 k+ O V! X4 o$ V; r
( F& h2 h3 E4 R, v7 a. t
# 将图像中小于0 的都换成0,大于的都变成1
/ h2 h+ U/ u" B/ C9 h5 v- P: ` image = image.clip(0, 1) a! M1 _* H C! N
4 l+ ]- e! U( z* X1 ^8 D: O: [ return image
) _+ i1 Y# U8 Y1
8 x6 m( }# N' E& O0 o" m2* W1 f# {- Y( k/ t0 ~8 R
3! `9 i) ~# T5 t2 N [' _# n
4
/ M* P4 C/ [0 S' L' }5
7 E# C E5 d. ^: Z) c6 e1 L8 Q7 j6
0 b# j% x! A# g" V8 s% ^72 q ^, }3 k& W
8' x+ ?( E0 J5 _' c. K% g- g0 r
9( x2 M g7 S/ o% i' y+ z2 N+ e" _3 c
10) h) |4 {; r8 E# T$ q
11
* ^! N) @, f+ [' n6 f; T3 Z" c12
% O' H! K* {/ H* M' C13
; ~7 `$ x2 M5 r! W14: O6 N3 z$ J1 l2 k: m/ J$ X% Y
# 使用上面定义好的类进行画图! \4 ]1 K! K) O9 w+ a$ e
fig = plt.figure(figsize = (20, 12))
4 s: O7 H8 ~; Xcolumns = 48 s! {0 ]$ Y) Q: t( f
rows = 2
( |- e, Q- P) e* `. U ?7 E
7 d3 o7 _9 C/ E: l" K) c# iter迭代器) y4 I( ^6 w7 c' m# V# A+ l" U C5 f
# 随便找一个Batch数据进行展示" { U m* C3 |5 K9 ^5 Y& l. y
dataiter = iter(dataloaders['valid'])
1 R/ L) a0 F4 i- I0 n# i" [inputs, classes = dataiter.next()$ z! @8 W6 L7 F+ v; A; F0 i. o
: V( a2 K3 Y A: F- K
for idx in range(columns * rows):
/ j) t1 l- L( J! b$ F ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])
8 r, G9 u3 P" s # 利用json文件将其对应花的类型打印在图片中) N1 D' F/ d$ w4 x! I1 X& M$ k% @0 S* B
ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))]). c/ d2 X) T8 ?
plt.imshow(im_convert(inputs[idx]))' t) ^) g! I' L5 m' {0 O3 u+ J$ i
plt.show()1 d" h1 X. u6 a. E
2 w0 A/ m' p0 c; Q1 d2 i! K2 t
14 M8 l/ q' ~/ u! X* B+ }$ b, z
2, d# g+ E- G6 `7 n
3
9 d) P+ i. v" X2 g! P# p0 K4, O" Q. z# B+ F4 ]; @$ S% t8 B' ?
56 j$ \5 o# i+ B' u
6
: {/ H* z9 D! ?. E. A: f4 ]6 n, q7
: B$ ], J! C* z6 D' O h9 h8& p' j. |' M3 l. K
9
1 q& u9 }# X2 r% V7 _- Z10
; d' _! |7 C, r; s3 a8 I11* A8 @' P( {" R* x; O+ a/ `: m
12
# d. ]& N! X( M& N0 }13
) o' u* h0 m- R; a( s2 p148 n4 ~4 Y9 o# C7 z6 d
15, a1 A; x; M$ J, R
16
M( s7 }' D8 z5 B9 m9 J8 a- g$ F. y _0 S
. E5 [/ O# g( m" m
5. 加载models提供的模型,并直接用训练好的权重做初始化参数6 D e0 t0 ^0 t# N v2 S# l
model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
$ e9 J# h# V7 O+ B+ [9 q5 ]; Z# 主要的图像识别用resnet来做
% z7 |" i9 s% s% A% K7 D# 是否用人家训练好的特征
, H \# k6 b( Y# gfeature_extract = True
- R a3 S! |2 B) o2 p0 \, m( F1 L" ^# _3 [& E! a4 z
2
6 R/ @% H& w' z* F. P3
* F1 O7 u3 O7 r4/ M4 N8 u9 m* S/ y
# 是否用GPU进行训练" j5 Y" S9 T8 H a0 V1 R+ ?3 R
train_on_gpu = torch.cuda.is_available()
- z* S3 l' J4 r, I
7 P/ g; f/ s4 e- C5 f: J7 }* Z$ @( Yif not train_on_gpu:* w2 e& h- z8 c
print('CUDA is not available. Training on CPU ...')
3 J$ h6 l/ b5 [7 X1 Xelse:
: s8 y+ i# L% \8 H: \. x+ q print('CUDA is available! Training on GPU ...')
3 d. w/ f4 O; `4 P+ ~, P* p6 _! r. ~3 z# M+ i
device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')- e( m+ V8 \4 c; ]& u/ l
1; ] J0 u; S: i& P
2% ^8 O5 l: Z) B; I
3
+ u8 n0 j4 |) A/ ^. e& U; J46 R$ M7 u* {/ H. Z6 n6 O
5
+ `0 D2 E. c( C$ x* Q9 J4 ^6$ a8 ]- F1 |' f \1 u) B6 i! j
7/ N+ M$ a9 `1 y" _: Q- G( O5 {+ A
8
" E6 q! E, Y* ^& [9
$ D, ]( c1 W% R# RCUDA is not available. Training on CPU ...
" H* w# {# J: ?' O. Q: B1
3 R u. s! \5 h! n. D# 将一些层定义为false,使其不自动更新
0 z3 k8 r' r0 c/ P2 @. d( s* adef set_parameter_requires_grad(model, feature_extracting):. }- g. E; p! V0 f$ Y
if feature_extracting:' u l# m' H3 [, g0 B' e
for param in model.parameters():. m% T9 H' ~0 u G' p
param.requires_grad = False
" r& i5 `5 t, b3 V9 @. } I19 N! D. s& H" O* D" P( x0 N
2
0 z6 J) A d, t$ W: b _7 K J34 ?9 ?2 w: G! l& [
4( z( F: c3 Y/ Z9 y! q
50 R6 i0 W+ c0 M
# 打印模型架构告知是怎么一步一步去完成的. |' U5 q6 D+ q3 U% f
# 主要是为我们提取特征的+ F" G& R: z! K; f
" K- k% n( @/ [( W2 @
model_ft = models.resnet152()" u" \$ n+ _$ e9 B
model_ft
$ K5 d V4 i# {, v% p' G! B1# l+ K' Q5 q8 ^* W% w+ q' H$ k
2
6 a( p. K% t/ g+ |$ J) y- z4 a( e3: j' h6 e. {. I' `0 V
4
8 Q2 X3 k! m3 t: R1 o5
6 G+ {, Y: y o) d( l1 hResNet(% L$ y" n7 `( `9 |/ S X
(conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
; h0 T0 _$ ~' A* n6 Z0 } (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True); m9 Y; C. u) r/ N) b# u5 ?5 w
(relu): ReLU(inplace=True)
: f' ~. M1 x2 C) Y9 p (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)3 K+ T( v4 h1 l9 k0 n# K
(layer1): Sequential(
! f. H) t) J" b5 P7 p (0): Bottleneck(3 P$ x _1 D% {$ b$ y
(conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)! h$ j( A9 `& U5 z4 m1 _* S
(bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
/ P+ ?/ I4 @: h9 X. `1 A3 \ (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
& I( C! o# a1 |6 ~* C6 b/ b (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
4 P/ S. H2 g4 T8 e( ]% c (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
- V6 ~3 P4 }! H- x9 c$ h (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
/ _8 S- @* y8 a* R: s) ~7 T0 r (relu): ReLU(inplace=True)2 k. N( _) Y, ^" d4 K2 |
(downsample): Sequential(9 o& n- U* R F, j# Q; Q/ q
(0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)( y* O! p1 F" P) B% v; a; ^& V# b' l
(1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)1 c% V, Z1 D1 W L' o4 u& F; K3 m
)' x, r8 ` C2 H+ \" G7 z" ^
)
; Q& d3 l4 ?7 D1 ~; @( T6 x中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。3 P9 `. H" _) q" W
(2): Bottleneck(
$ @/ u5 b: e6 D$ } (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False) s* ]7 a3 v* \5 F
(bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)+ C0 _! b9 P5 s
(conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
4 [# n- E9 e4 c/ B( }: y9 d- }7 D (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)& L9 n: B- X4 B' j, Q& Y* I
(conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)) t2 o, v2 u. d3 y% R! D% U
(bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)* u- @$ I, [3 H% x8 ?, O9 S) `7 v
(relu): ReLU(inplace=True)9 @6 c7 W6 U6 I% Z( I; d- a( [
)
0 i% W+ }: R* l$ ~% ] U4 B )
& c" _' O1 p$ a2 w( b' O (avgpool): AdaptiveAvgPool2d(output_size=(1, 1)). g' a) S! h/ Z5 Y
(fc): Linear(in_features=2048, out_features=1000, bias=True); }, U( l5 Q& e$ J
)
, n) @- \6 u; J( h
0 N4 M8 A% ?/ e9 ^8 c1
, @3 M# l& T' M% x) ?2
) Z# X4 M% `' w3) |8 ?1 t- W' P7 D" U2 k% \% b: V
4. F8 r5 B4 R. J5 O' C5 n! p9 x' }
5, d9 w- l- T9 x( _
6
2 B4 X* r/ C" t+ V7 D& s3 e78 F( ^0 J% Y4 ] R3 c' E; B
8
8 q1 c8 r9 u, o2 q- F: H) h9% Q4 u# J% V( d7 ~6 n+ u
10
' c' q, `0 A O( [# _11
Z& z! ?& Q! X2 L) f8 `12/ ]" q2 r/ J5 [0 c7 G
130 w/ e! F8 M+ [8 v4 _
147 i! G9 P) K" x/ B3 ~) g
15. ~7 v2 T' ?8 Y! E9 P% e
16( l1 }2 Q# a3 R5 Q* O; \
17$ R' K7 h; | k( @8 W+ M: ?
18 A* q) a8 o) b) h
19
z9 L/ n3 ^8 S) ~& r3 a* X20
/ Q$ B# D9 w9 c" U& f21
5 K- X! C! Y' ~6 R1 F% G2 L22
# ~0 Z0 V+ p K/ u6 U2 u23
& n8 B( W/ I$ R; }+ ]/ W& i p24
9 O3 l+ B% \- P6 U& j. M0 |25
/ g. }1 M+ R! p% x$ L) ~$ @265 u' h1 S0 F4 ^7 a- [' s+ l7 m
27
- S: q1 R/ m3 i' q+ G/ g. c. \286 t5 m& [6 v8 [$ }# X7 ^
29" L$ _! p" h: a+ p7 M
30
! E8 G8 P& e8 T31
4 e+ Y( a! P/ D) s1 b32
' I$ j8 c% g) p7 J. x( Z% r' G33* w3 _' j* Y) |0 G( u$ j
最后是1000分类,2048输入,分为1000个分类 C4 l. h B" x" M3 P( J
而我们需要将我们的任务进行调整,将1000分类改为102输出
" V8 {8 f: E+ ~: N9 J8 ~% `. q5 w8 \ O% r' v0 G; X" g% V
6.初始化模型架构2 H0 n7 H6 i& M- {! h. f7 f/ a3 J
步骤如下: P9 [$ C+ Z5 R9 x' E6 x' H
0 X) J9 H. j7 f% J' l. i将训练好的模型拿过来,并pre_train = True 得到他人的权重参数
+ h0 K7 m+ L# V# K% {可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)
! r" ]0 I( G3 O1 P1 r: Z* Y: [无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数
8 x' |' b0 v; t& Y官方文档链接$ S$ Y$ V |" V9 X0 z/ s# }
https://pytorch.org/vision/stable/models.html' o5 u" s' X, d7 R" r0 b8 {9 z
* q' M" N+ z- j4 Q4 c# 将他人的模型加载进来: P8 h5 x) f* [/ R
def initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):& n% H: C$ z4 E7 |6 i, a8 g
# 选择适合的模型,不同的模型初始化参数不同
3 r1 A0 N9 Y2 I8 ^ model_ft = None: y& D: x7 R# r) m4 U. [7 ]
input_size = 0
! E; E# I1 ]& a3 y! p' M3 @' O2 ~
( E% R @1 ]. }: h if model_name == "resnet":
6 B+ F7 z; i5 Y: S( U. ~% ? U """& B$ r0 ?* `6 \# u
Resnet152( i- D8 C9 q2 h0 q7 j K' l
"""0 G2 _; x' B8 U7 ~" @$ @( B' p
R3 P- ]& t& h. b8 L) U # 1. 加载与训练网络
3 n2 O: Y" X4 S _ model_ft = models.resnet152(pretrained = use_pretrained)
* v+ Z6 W$ m" F) D* c, g1 \4 ]/ W # 2. 是否将提取特征的模块冻住,只训练FC层" u# e6 e4 t+ T5 l% g2 t
set_parameter_requires_grad(model_ft, feature_extract)7 i3 r G) r2 f6 y* `6 D+ y
# 3. 获得全连接层输入特征) `) g* Z/ I9 y4 t, K" b) w
num_frts = model_ft.fc.in_features
9 f& m) T" a# S, o! N # 4. 重新加载全连接层,设置输出102
# H5 L. o' J# n' \: b% L1 K' j( x" d model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),
) x+ o J% v @2 P) s V2 n nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1
1 D- p6 n) F8 O. Y input_size = 2244 V9 [; T9 w1 x; k. i S4 h5 L7 @
! C7 u, _5 q7 w
elif model_name == "alexnet":
3 C! f- g. v% ^# }3 }# t """
C! f" j: U3 a4 n7 B# C Alexnet
B" ^9 `' m/ I0 A """
6 L2 }! T4 F/ ]2 l: u! _% _! ~& } model_ft = models.alexnet(pretrained = use_pretrained)* U3 V; h& S/ h# t" v& r
set_parameter_requires_grad(model_ft, feature_extract)
5 T: c% k; X! k8 Z) k+ C/ O3 N2 _! X+ H& b9 B- b' o
# 将最后一个特征输出替换 序号为【6】的分类器
4 I- i. d: [7 @/ P5 | num_frts = model_ft.classifier[6].in_features # 获得FC层输入
, m* g7 N0 G' E: C( V( K* C model_ft.classifier[6] = nn.Linear(num_frts, num_classes)( Z1 U/ z8 M3 ?. l( a) B2 q9 w' U
input_size = 224, q2 l: A$ v: m. d, h5 Y
, P- `7 r# }8 h! R elif model_name == "vgg":7 D3 h% ]0 v' \% Q$ P6 k
"""4 G0 I- p4 C3 L9 u* y
VGG11_bn, _9 Q# E, _3 D* B: A! |
"""
# l4 p2 M! R: G Z1 B model_ft = models.vgg16(pretrained = use_pretrained)
! o4 F. R2 Q( e f% f set_parameter_requires_grad(model_ft, feature_extract)1 ^# I; {( h6 ~4 {- c1 ~: f
num_frts = model_ft.classifier[6].in_features
# q3 g6 B ~4 N" D E. v5 \ model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
* p, m P7 _+ B/ C input_size = 2247 Z# M& ~2 ~7 [2 b2 o( b+ a. a
, S, ^: `7 u! T5 m2 ~. L9 |# j
elif model_name == "squeezenet":
7 K) v0 V7 p3 ?4 j7 x6 M """4 Z3 V" ]. C$ O" P
Squeezenet
4 I" d+ I# C- U P """% B0 B U$ Z/ @% @2 \
model_ft = models.squeezenet1_0(pretrained = use_pretrained)
5 B. ] w U5 u- s1 a set_parameter_requires_grad(model_ft, feature_extract)6 c6 ~/ |& Q+ O; t
model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
' R' M: F( I4 ~1 O$ l1 U+ }" N# I+ }: d model_ft.num_classes = num_classes
# j h% a9 I! W, w' B) {) D input_size = 224
# H6 }' B& X6 b; h& b; O
% h/ w8 G% T; k+ ~: k7 ^ elif model_name == "densenet":
2 l s1 R6 |" Z7 x& F- f """( f" x: Q$ d* Y" k& ^6 O
Densenet9 ?8 S! ~" g! A+ B& {4 y
"""/ _; F4 Z$ w: G' m x
model_ft = models.desenet121(pretrained = use_pretrained)# T& w' H c$ ?; n7 \ s# L
set_parameter_requires_grad(model_ft, feature_extract)
8 d" |3 e4 L5 ^& w, j7 G num_frts = model_ft.classifier.in_features/ `) @. U5 ] g8 w' t
model_ft.classifier = nn.Linear(num_frts, num_classes)1 J* `3 }- T, |3 Z
input_size = 224; [# V% j2 ?! _! K0 d. H0 M! L- L
% s3 {' s& O! C4 O# B6 i
elif model_name == "inception":
7 P4 r- Z6 v. s; R """# }! e2 o! E1 P5 u$ s
Inception V3) o2 R, Q6 i0 x
"""
) J$ f; r. ]8 y% V4 I model_ft = models.inception_V(pretrained = use_pretrained)
: c: S" ]6 Y* N- \* t+ {( M; H set_parameter_requires_grad(model_ft, feature_extract)
b i# x$ r% }& O) j s4 I" Q) v5 r( @6 `0 m
num_frts = model_ft.AuxLogits.fc.in_features4 r# z1 [) r9 \
model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)
6 _0 H+ k) Q4 {# I3 p, I7 m& l; \! F$ m. }5 B; ?: w2 ?* I
num_frts = model_ft.fc.in_features, F* P1 {) j9 `! x7 |* ]' S" g
model_ft.fc = nn.Linear(num_frts, num_classes)& y4 }* t! ^9 M8 g
input_size = 299- L% \3 y) ]% j$ Z6 p
7 K6 u Z/ X5 C, D" t4 A
else:% v$ X0 b* _: k
print("Invalid model name, exiting...")3 r" g( N/ k+ D
exit()
; C4 V6 N; x P* Y8 O5 U2 w ~- q, Y h" D! t/ ~" O, X7 D' r
return model_ft, input_size5 X, P( m5 X5 c8 m, ^$ s0 z
; m: b* T! C" r5 z/ c1 Z; w
1
6 R' A) C- t( E9 q1 X2% V# I7 a" Q! q3 S5 Q# p
3* w; T. x8 C6 u7 B& H
4
5 u+ `" E4 E5 t* d- k5. N& n( {6 ^0 \
61 k% J7 _$ A/ w/ ^8 z
7$ e/ P3 V) k* ]. p
85 H% M$ u# u$ U9 {6 y3 M- M
94 |* L0 @( ?( a) \
10
. A5 w) Z# h9 ~. h( ^" e11 x M" G' Z( s! }# ~
124 X) V% g L% W3 s
13
% K4 j/ A7 j3 w( e" E4 N14
: z& @ m8 I5 K5 K+ P2 \15
0 d2 v! X3 @% _0 N- [* _16
$ u; [, O+ J' r* B8 S8 {8 s6 V17( B* y! E: }2 m7 u, Z, X; f
18+ o5 f3 D* f U$ ^8 w) z
19. e# T) E4 r- q
20! ?: P. i: a/ P) J9 g+ n& R
21: C- {: g' C; P
22
+ |1 J2 g! j+ x* g23
% f2 s, b. f3 F0 P: x, G0 }! w24( h! i' v$ P6 n* S, K. O& F
25
: v: T9 F' `$ W8 V! C. Z2 |( j26
1 K* f. v6 G' C( z+ n- J$ |& Y( H277 h5 k7 f( o; x; T% b- ?' M
28
; B+ X8 ]& T9 V) [) P! @29, V3 P8 [4 W0 C w+ w- t1 j8 z# n
30
$ p( F. C5 e! U% |3 y$ F31
! _* v3 B3 B9 k _; {! ^329 a& M: F8 T. E) A, x9 h: A
33
7 U0 v! H& i; f ^/ J/ q; X3 V( C+ ]349 X/ ?, V) X/ m! F! }/ r
35
+ J4 k) L3 P9 S8 V36' g( Y: l# c! M" i& Z; t7 `
37
1 X# w& A% E, @& V5 u38
" Y/ f7 l$ S8 ?1 E39, h/ R3 R" A% U8 ]
40
6 Z5 L3 ~: s$ E& Z; e41, z+ O% ]" F2 p1 S' B5 Z1 b
42
0 u! X7 |# K/ a. r( G/ y43
! d7 o# D3 h7 F- b+ s44
" g! q( C ?& \0 W: x45; y6 L. n$ E3 t. o: k4 a& c
46
. F1 F& P" F2 O" \+ U47" w9 c' L8 ]& w
480 T; K; ^& f$ c2 i
490 H9 c8 W* p7 y
50
! Y- J( ?) {: K' h" W: `511 _* i; I; n- R3 m, C/ S- R* q
52( K0 R6 }4 G3 y8 F
53" h& @4 Y8 v4 E9 l$ e
54! V5 b8 [8 T+ A. y( c! _9 {
55
% ?" o, F% i' f" h# N3 k563 \% ?9 `! a, A/ k& T6 _
571 f- U: {( F$ J! A. f
58+ y5 ^8 E, B! w8 j6 Q, ^
595 m0 ]0 L5 _3 h+ f6 N
60: w8 `- b1 w% t* a, Y7 Q
615 d* c; O( }2 e5 N [
62
+ V1 I4 M# ^% b/ G! [* g63
! |% V) z5 l& L" m$ g& l2 V, V7 V640 u4 I: i$ j# }, P/ |
65
: w) K/ j: D4 x. ?66
. a( C9 p0 j" G" @ ^67# Y# }; t2 A* t% R( @+ z/ f
68
& c. F' e4 q# `! F4 z- A- p699 i$ x. ^1 [6 ~/ `0 ?
70, @% E; e( Q1 Q& s
71+ p& z- E1 p$ J
72
% r& b# z$ f7 t- C: k& z4 v$ u73( t, h1 h% D4 c! d, j# g
74: f- a( h# Y: u# f
754 R, V% s7 L) J% g# _& V: e7 Y# d
76
5 |9 y' s, C, Y v/ z& e77! r1 x, @' a$ c9 ]
78
. b; }5 L, L, f791 `+ y% B! H* a/ B0 M
80
% }4 _" {' z8 f81- E: @. j+ ~7 N: w) ?& d% V: n
82
, d6 ^- {. |/ Q+ `83/ u$ L$ x2 B* T- t
7. 设置需要训练的参数
$ R, K4 `# r; R* I8 |# 设置模型名字、输出分类数
d7 l! u6 S7 }" f; V, A- }% Fmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)
% H- r/ u a7 m4 q1 R4 C) `( y7 s6 y; K5 b* a# E
# GPU 计算
9 B& Y& z6 X: d1 W7 \9 J$ wmodel_ft = model_ft.to(device)
2 a& i6 T3 g% k: c7 F* n6 ?3 h/ y( V7 o( W$ N7 V
# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取4 J( V. u" k6 }( `
filename = 'checkpoint.pth'; m# h9 Y( I7 R5 d
$ W4 x6 G4 B* j$ D Q* G# 是否训练所有层# k6 W, a p* e, H5 u
params_to_update = model_ft.parameters()
4 `: a x* T8 |, ~# 打印出需要训练的层4 U0 Q1 _, H6 [$ D0 ?# m0 ?
print("Params to learn:")
' s5 ^4 h/ F" B' z- s) Eif feature_extract:" C! v7 o; ^ P; `+ u. J
params_to_update = []/ B: J6 A- }& |3 T/ j3 l' o3 r
for name, param in model_ft.named_parameters():
+ Y. @9 L7 e9 G0 u k if param.requires_grad == True:- g* s/ y( z( W& ?
params_to_update.append(param)
8 H" ]6 v8 F6 E" ] T3 B1 P8 i print("\t", name); B6 x/ z! s9 j6 `7 I9 v- ?
else:& l# w6 l* h: j! z Z8 Y% ?7 u. U
for name, param in model_ft.named_parameters():
( l# z2 T! e3 w3 k if param.requires_grad ==True:6 [( a) O0 @# C; g
print("\t", name)+ W% y) Y( i2 `8 T, z+ W
" q. v. |8 E, O+ R4 y
1: s. }+ Y8 Z* E U1 @
2( B4 n8 a# B% h/ {! S8 F
3, p6 r2 q$ B) W( }
47 J7 }9 X; ]$ |: o) P# Q# a" ]2 P8 S
5- E# Q9 r! Z( K% \2 A9 P
60 k: R7 j, I( }9 \
71 j2 V$ Z7 k* K
8
# Y( R: ^: G7 N8 ?3 D. p9 b9* P' W {% B$ I
10
* [# `( V' }. h5 w4 f2 B11
p |; n$ F: x# c9 _8 \12; U, D) \9 S6 k
13
" w: p) G {1 M* r% b4 ^146 `+ ^4 A! x: p" C! k
15! @7 b2 q/ i2 \, t" I' ^
16 ^: Y6 t4 r& V! @* c
17' T- U! }. Y" V8 {5 Y
18
7 }' {5 u& I9 e% a8 U192 K( I. q( h& y8 o8 H: ]
20" J5 Q* O( T8 c5 l/ y
21- f- J$ T1 Z) ~) V2 b0 u
22% p. Q, t; J" p' @9 p$ u& f
230 M: S. t$ I! r
Params to learn:
) [( @3 P# f9 t c% W/ E" M fc.0.weight
9 q) b" v/ ~; l; L fc.0.bias* N$ f& t% H! ]2 o& X6 j& Z
1
3 m' c* P$ `. S- J* Y2
/ ?# [7 g+ o9 A4 ]1 }39 C1 q o7 Y% u2 M7 _7 I- v
7. 训练与预测
( c* h0 D( \% p8 J7.1 优化器设置
4 F! E; j1 Q1 A$ [" s6 G) h# 优化器设置
& F) w U' t, Z+ x$ ]+ `+ [optimizer_ft = optim.Adam(params_to_update, lr = 1e-2)
4 l- o$ N9 a9 j. B! t a- z# 学习率衰减策略
6 c' m& K* |+ }! I: T4 C# y+ nscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)# b1 a4 E! w* p: y6 @
# 学习率每7个epoch衰减为原来的1/101 c8 b1 U8 @ l( ]: u* T- \* ^
# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算
. W; P- j& \6 Z4 h4 X- M/ \1 W
+ }5 b5 K4 ]7 {8 d8 xcriterion = nn.NLLLoss()- [) f9 s" E. E9 @3 ^0 q0 K
13 V% \9 B/ s+ i- ~( f) L- x) q2 s) f
23 p5 Z" e+ F9 F$ H( t+ m' E. H
31 W8 e c4 R: n) A2 r+ C
4' M# i6 n# p& K. U
5" A: G9 a$ D, f {# d7 \
6: @5 _0 X* y( T4 i5 b7 S
7
$ A$ ?/ X# _$ m$ y4 Q8
! V+ a) K! M! T) s& a3 G R% S# 定义训练函数, m. H) z! N3 b9 k0 ]
#is_inception:要不要用其他的网络2 L. \8 D0 e z$ C
def train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):6 m4 D' q1 S% \& h2 n
since = time.time()
; {1 x/ S. X8 }9 D; ^ #保存最好的准确率
% \+ b8 }. F4 e best_acc = 0
; c" P4 e* m- L """; f$ |: X8 W; \. |
checkpoint = torch.load(filename)
, D1 M# M( K' q7 m- Q- L' [, D best_acc = checkpoint['best_acc']; g' r3 t( W, r! g' l8 r
model.load_state_dict(checkpoint['state_dict'])
7 Q% s H( ]" r. l0 [ optimizer.load_state_dict(checkpoint['optimizer'])
, Z( e, b8 Y$ J- l! V model.class_to_idx = checkpoint['mapping']) m3 Q/ |$ }4 Y
"""9 ^, |- C" w1 `6 s- ]# y
#指定用GPU还是CPU0 D' [1 D9 T9 K, i; t# b
model.to(device)
/ K1 ~7 S+ N) Y #下面是为展示做的
( x. x9 y* ?5 Y" [7 J: t/ O* T) f val_acc_history = []
7 ]: G! X9 S5 o train_acc_history = []* C% k$ l1 D9 S' i1 {- m
train_losses = []5 }/ @1 Q+ z. Q& E8 Y
valid_losses = []/ N" n: Y5 N+ H* q5 C4 [
LRs = [optimizer.param_groups[0]['lr']]! b' a: d; t& R c y: P- a
#最好的一次存下来0 s' |, Z' @: Z
best_model_wts = copy.deepcopy(model.state_dict())
' l6 B) q) V( L5 ?& [# ]$ c$ @2 W% r6 t+ n
for epoch in range(num_epochs):
6 z) N( ^* f; J7 Z' w print('Epoch {}/{}'.format(epoch, num_epochs - 1))/ ^5 d4 @+ n* }8 y t" M
print('-' * 10)3 Y' o% q1 c( e1 `+ A( Z
& v1 \9 H7 j8 j1 q # 训练和验证+ Y8 z6 e4 e; b
for phase in ['train', 'valid']:
3 d1 f4 l" }0 r* ]9 Y$ i2 e) b8 S if phase == 'train':# D* r) M* h5 F9 Q2 q A
model.train() # 训练
4 j. E5 `: w4 G1 ~ else:
" R! _: i3 [, c8 [: H3 m# g model.eval() # 验证
& j' \3 y/ G1 R4 H, r. d: W7 \
2 u, y7 J) L3 t1 Q1 d* w running_loss = 0.0% L* X8 M% E9 Z: a
running_corrects = 07 E4 ?+ k, \! V- q4 f+ @
: y8 t4 e4 V! @2 |% Z" ]3 Z3 B n # 把数据都取个遍
" A0 f% j# k# Z7 F t$ A4 n for inputs, labels in dataloaders[phase]:
4 v6 ~1 w6 L: b, i g* @6 ~0 Z& f+ s #下面是将inputs,labels传到GPU
* p' T) u3 m7 Q) y, D# z inputs = inputs.to(device)
, R" S4 D+ {: o9 Z labels = labels.to(device)5 v* q- J) k, G' x% j% C S- Z
+ J* ^: d0 [: `# F3 A. w
# 清零
: ~7 k# r8 W$ Y3 n( t, R: [( K. q3 j optimizer.zero_grad()
6 l4 ]0 Y! c8 c2 G* A' I# _9 n # 只有训练的时候计算和更新梯度
0 e6 o( z! }+ @ with torch.set_grad_enabled(phase == 'train'):
' |& F8 Y6 v+ @+ K9 [ #if这面不需要计算,可忽略3 b2 e5 d( R5 {$ s
if is_inception and phase == 'train':
) ^$ F0 N8 Y3 p- p6 A& I outputs, aux_outputs = model(inputs)' c1 M3 x0 M+ V! n3 G# a4 b2 B
loss1 = criterion(outputs, labels)
: _; v8 \+ m; [0 t1 O: x! O+ z6 c! t loss2 = criterion(aux_outputs, labels)2 `6 }) H! |' e; Y! N
loss = loss1 + 0.4*loss2/ ^0 e# c+ n2 x# K( K* O! p
else:#resnet执行的是这里 j' [8 a0 e1 b4 F
outputs = model(inputs)
* L3 K& E4 v0 G: X4 } loss = criterion(outputs, labels)
+ O' |2 K, g3 y1 N3 s0 R5 P2 I
+ G+ B8 ~+ }" I( d #概率最大的返回preds
" j& y+ @/ A H2 _$ t: Q( n y/ x& I _, preds = torch.max(outputs, 1)
5 G, X6 @# C8 v1 v9 _2 x" i% N! E1 ~2 |
# 训练阶段更新权重
5 y* ~, D8 c# @0 o% i! T5 k if phase == 'train':
# A$ r0 ~! R& \ loss.backward()9 W( R, y# K7 Z# }; }9 M' r
optimizer.step()
0 |! v2 }; l( u9 s9 E8 @, L( P7 W* v. g4 B
# 计算损失
. ^( a. d' Z- B& Y running_loss += loss.item() * inputs.size(0)( b5 i( v6 O# i& L& E) K2 q
running_corrects += torch.sum(preds == labels.data)/ `3 U% G0 a/ k- Q. T5 \1 `" g Q6 p
, T, I1 p; F' i( d #打印操作
/ f0 C; O' @2 h" F; Z epoch_loss = running_loss / len(dataloaders[phase].dataset)* ]+ Q3 X7 C4 N( r5 S
epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)( v0 E( |* s2 z7 A+ Y
, \5 P8 V' Y; @9 D. j/ ~2 Z
- K0 M: N/ h7 d. n7 b time_elapsed = time.time() - since
4 q$ O$ f3 N1 k7 B) ]# d print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
- M" y: p' f4 P% l print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))3 e% l/ q+ [/ @/ n1 C* c
+ q8 R! ~* q* m
* F9 D+ @7 d7 `1 a8 @. f # 得到最好那次的模型
9 X1 i Q: Z$ M: w* T if phase == 'valid' and epoch_acc > best_acc:
, i/ b: m; l t, B" e best_acc = epoch_acc
0 E1 i0 h- f0 K x #模型保存
" a; |/ [% L6 @) [2 B best_model_wts = copy.deepcopy(model.state_dict())
) O9 N! |8 C! } state = {- {+ {% V: ^, T" b
#tate_dict变量存放训练过程中需要学习的权重和偏执系数
" C- ?+ U X F2 @. _ 'state_dict': model.state_dict(),# s6 @) A9 ^8 P( {$ W
'best_acc': best_acc,8 p: o4 ^& ]+ `' b. u6 d9 t8 K
'optimizer' : optimizer.state_dict(),+ F+ p. ?$ x( T, p! c$ O
}- j [; @7 Y! ]4 C9 T! j
torch.save(state, filename)
, \, h+ r1 Y3 E if phase == 'valid':
* `9 ]5 T6 q+ C$ ~# Y val_acc_history.append(epoch_acc)# i) w- ^- ^0 h
valid_losses.append(epoch_loss)
@- S! {9 K+ _ scheduler.step(epoch_loss): ~( v. b4 i! L/ ]
if phase == 'train':
- _* K% y. x+ w: Q; I0 D$ Z6 F train_acc_history.append(epoch_acc)
. m' |8 f1 H1 [1 f1 H. ? train_losses.append(epoch_loss)' x* D2 l. Y+ F3 G8 j
' a0 T$ g) j' W& V. \
print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))" X9 \! E9 K. q7 y1 V; d0 N
LRs.append(optimizer.param_groups[0]['lr'])0 m7 d0 ^; ^# g. i/ L2 K/ u
print()
8 ^ g; `1 o, g/ }% D) M, P' P4 p8 ?
time_elapsed = time.time() - since
* e7 K4 I/ o# ~3 W9 Q" ~" m* Y print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
: Z: }9 V* G" ?9 C/ h2 j print('Best val Acc: {:4f}'.format(best_acc))' y. w# L% E* |+ F
" k2 |& u/ \ T$ [) m# ]# T # 保存训练完后用最好的一次当做模型最终的结果
# H* Y! o/ K9 w# F2 ^; R4 A model.load_state_dict(best_model_wts)6 i' f$ a' T- H Z
return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs
J, d; t/ j) a8 C) ?( v
) G, e2 V8 O0 q, E; m2 T+ b$ b I+ G9 Q
1. f( P9 K5 ]& D" w
2
$ f, ?" h$ i$ i5 ]/ j; j$ g: d3
, m0 T2 w' b: S& J0 I$ @& {4
) g0 H- x' ^& d- T( r( m' r5
1 X* E: u. ~$ v- X6) c2 \/ D# D1 D3 K5 K
7' j8 u( Z6 k+ X8 \
8
! D$ E$ T8 j' V9: \7 i& H* ~% r* x3 s/ t. m
10
# B& c( A$ [) U W) J( R& a p119 h1 K: U/ P. l( ^: o% E
12
2 |6 S# g2 q6 e. ?1 ?. D3 X13
2 P/ Q9 r. K8 Q; [, d' A14$ x! X; ~; q7 _6 Y( r* n
15* g) x. X$ ?& o/ j& W
16
! [4 ~/ L( S: k2 x( {174 B o1 ?# v: Z2 n7 o% z
18
+ k) ]% X# l3 G0 W, j; R \% O5 Y19
; S6 s, _$ a+ N# o20
" }! z" M O5 J. {4 o7 Y21
+ z" j9 w, H7 X7 H# h1 i. W- O22
' U/ s6 S0 |; f$ T23# d" k9 A- v6 A- z1 q. W3 z! i# ?2 O
24) M6 @/ n0 b0 c7 o$ Y& C" g
255 q# F p$ s0 v# c, m+ z4 F: B
26, N5 p; B* i+ a9 V: j9 g
27+ f' k0 {* K; d, V5 w9 l9 l6 g
285 V# {& {* ]' w$ \# T7 V" o
29
7 o/ @ W4 p+ h; X1 z30% @4 w6 `. O, y. F. y' _/ J
31- M4 t& U6 a+ R% I) ?
32
0 {4 U) T) Y6 f8 O* Y6 c0 ^* t33
8 D( z# h# t) v( _! t; a ~34
) n5 R0 N" [5 Z35
' g( y4 N* W8 D$ P9 v36
; m, Y( O1 x: r. m* o$ L37
' P4 c; t% D+ S/ R0 ^6 p38
% S4 d4 q5 S) M: Z, n& {39
5 T) J4 P& _( I0 e% |# P8 {( k% {409 o3 r6 _2 _4 z! d% z
418 M Q5 X/ z( W; M
42* f- n4 v v( c% n, m$ ~$ Q
43' N% Q9 S5 G7 w
44
: c; X1 h: ~! q# N; g45' q$ N9 V2 ~$ i1 |* p
46. k; f4 {: o2 C4 r5 ^9 `
47
; J. J5 J H, K( ?* I- v1 D( j; C" f48
" @" d4 s- s, b, f3 R7 C; E$ _49
8 \# G9 v2 X" ^3 g! {4 d3 \50
. ]2 }0 `% L/ n- E; e/ F# \2 W! i518 {8 H R J) v( B/ N( A
522 Q9 G# g8 t6 k: P' {8 e! A5 b8 ?9 C' C1 [
53
3 A, K# Y7 u# w" h4 R54# {: U v8 J, m, v! x' B
55
4 i4 S! X9 c( U+ b56
9 @, ~5 ~( [: u7 H9 W# G# p6 ~% a5 i, }57
6 o0 E& P9 C" F58+ x) f* K, r8 k+ W; Q! @& I6 I
59
- K+ A8 ~, J" {; x60. |* O! Y7 @0 ~# G; ` `& Z
61/ X: x3 c, m# E, u
62) K F- s% K% q; [- c# r
63
' ~- X) q" I( I3 v' g- h2 ?0 `648 F6 [4 \1 J' l6 u7 M/ C2 `
65% ?" z3 H- y6 { W
663 f' g4 D, n4 Q' v' m. N
67, n- m. o9 T2 y- T% i% C
68/ T- Z; f$ w. F( [
69
2 N0 u3 Y, K# Z700 M3 o+ t0 ]7 l4 X. ~6 i# E" I+ x( B
71: J8 M& m6 w" H# y d
72" t' Q; d6 v" j9 }" l
73
2 j7 m" S+ U5 ~6 |74
" |$ J- l% s C: x# Q* v! Y2 F% ~2 y9 ~75
: i! j8 V' K i. W% l+ J6 o% Z76
" J2 B& t& b- N: |6 O77
5 L+ f/ {, i9 M$ f78
( g% B- ^; a( _3 @* _0 X79
7 x9 G: M7 D" C7 \9 _6 K/ M802 Q" N: I0 p# h( ]% T, v
811 b2 f( `7 ~$ }5 }& w
82
, P; ^! X7 ^6 S7 w" |8 @. G83
: W! e2 ~, N0 K) E' a5 _84# }& B) d7 A9 A
85
* n! C. Z- }/ n4 U" K' x- |( S86
1 U: G3 c% V# z7 g9 R$ J5 l87
3 x- t6 v w! D8 q1 v. c9 V* C88
. Q0 x, |. y' w1 u$ `7 v4 Q: K892 L5 Y. \) | V( B* N( }0 ~) Y0 l& S, _
90
* J2 i& h* z& t! r* P8 p* u) P3 }. o8 E917 \' Q) s& }& Q6 m- ^& X( {
92
) J% L$ Y4 R: G" i0 v0 l9 q/ O93
$ n& L) U2 ~0 U2 }# s) P, j8 u* U94
+ @' _% L- l* G+ m& R6 ]. ^5 B) J' M95
: D/ o# I0 Q8 k. ?# T; C" y96
' Z T* L$ E1 p9 z% u" W: @97& X( Z" O" ^6 g' l- a1 X. \
98+ p: E9 L- }9 F4 @1 T
99
0 m* [7 Z& z( a100
6 m2 e: G0 z& t: K0 M101
- {2 {6 d5 q. N" e1021 @& R* q/ u- Q, m
103
7 U1 ~! g. K) f& ?$ p104
% Z6 \7 E3 Y) k: B; C2 J1 j6 h105
! j$ j, ]6 x9 i; E106
& q3 s# x ~2 u/ D! {' n107
6 S; N" p9 U: S( n, X- m% W9 N108; g5 m; g$ X6 d' ^
109( v) E* v, \0 r" j
110
" @, t; l$ z; p. F111) d( i- H \9 u: H
1128 d5 a2 d# Y+ Y( h; K6 I/ j
7.2 开始训练模型
7 j9 f" p$ l# d! K8 o" T我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次) D* B/ }8 E' n) E% [& P
) l6 Q3 t% _1 B2 q- W
#若太慢,把epoch调低,迭代50次可能好些2 T# M) |7 u6 y( G7 t
#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合
`$ d1 I# t: ^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"))$ _1 F) C Z$ ~4 Z, B2 S7 G! V
8 o6 ^+ B, U9 h' U0 k6 X
14 B( f6 |$ o8 I& X& |& s6 K" s2 A7 c
2
, I" N7 E/ ` b5 O3" P3 f! f( ^+ ~* a* I
4
& \8 A) @) E+ h1 \Epoch 0/4$ O! J# A6 ?. A: t7 ^; U
----------& o% G) E+ X9 O5 I" D
Time elapsed 29m 41s
) A. @% Y. K% e' T! [, I7 D1 J7 w: J% L0 Gtrain Loss: 10.4774 Acc: 0.3147
~$ ^4 h+ C% M3 z! T; UTime elapsed 32m 54s5 ]8 W) K* w6 y2 w
valid Loss: 8.2902 Acc: 0.47194 K2 U+ [4 l% N$ Q Z
Optimizer learning rate : 0.0010000
9 b) t( }, a" \3 b% b7 L) D) ~" G
" N' P) j3 n6 l: XEpoch 1/4
1 Z, b: Q: U# w( j0 i- w----------- f0 R: g1 o9 ?, k
Time elapsed 60m 11s* E! g2 }1 M& q, u5 r" P
train Loss: 2.3126 Acc: 0.7053
. e- Z; N( R! W2 ^$ qTime elapsed 63m 16s. c! h% a$ l2 Z, I/ g4 \
valid Loss: 3.2325 Acc: 0.6626
: c! x9 k) _0 h0 h. XOptimizer learning rate : 0.0100000
. I& m3 W4 F( o2 G$ S2 u- ?7 S) m- w; z6 M( Y+ Z9 r
Epoch 2/4
5 m+ i. E* B$ V----------( g7 T/ ?7 K, D5 i' \' x
Time elapsed 90m 58s+ c: N' w6 u" N0 `* V. k$ X
train Loss: 9.9720 Acc: 0.4734) I; s o! B4 |; T. ?. Q9 @ Y6 X
Time elapsed 94m 4s
( D5 m: S8 j7 S3 {valid Loss: 14.0426 Acc: 0.4413
; g7 d/ [" {7 }" W+ iOptimizer learning rate : 0.0001000
! N: A, T6 q. Q
. A j1 i7 g: yEpoch 3/4
8 G N' {/ }7 H. ?& D6 J- o$ g----------+ X* e4 |) O4 d. y( c
Time elapsed 132m 49s# j: z+ } g9 L; N9 b4 q" ^
train Loss: 5.4290 Acc: 0.6548" ~6 |: Q4 v3 T( ^; B) b; L! j8 n2 `# _
Time elapsed 138m 49s; V0 H1 i5 M j* Q% x+ J; S
valid Loss: 6.4208 Acc: 0.6027
. }4 @& [/ D7 ?Optimizer learning rate : 0.0100000
+ n2 d# {$ w) b- `1 @* F2 T3 c' ^- x7 ~7 `/ v& x% q
Epoch 4/4
1 q% j' ~$ [, a# b7 x& o----------. {4 `& W" B( L" Y, j- E7 c
Time elapsed 195m 56s
- d8 e/ h1 G* |' N' G& z, Dtrain Loss: 8.8911 Acc: 0.5519
& D. a; N4 t( \! {" PTime elapsed 199m 16s
8 n6 M& r' }7 [3 s/ g1 M6 ?7 Ivalid Loss: 13.2221 Acc: 0.4914; x2 I L7 W8 N7 k! w
Optimizer learning rate : 0.0010000
5 G" ]: i0 m8 h$ A) S4 i
* ~* M- m( q$ W0 q5 LTraining complete in 199m 16s0 x1 c1 n4 w& J, N" v% q
Best val Acc: 0.6625923 p! r m- Y/ B& U7 u
2 O3 V/ [1 `4 u: W$ b( W1
5 y, S* S' [ L, [: h3 O2 P2
" F5 r' ?0 Y- M' I, t3
/ x( \( ^8 Y. ]) E* C4+ T' Y3 _* x) N9 N! _9 s
5' L) ` O" f# ?# f- F7 w) B
6
6 q4 O% g9 J% E8 K7
) J# z) o+ Q& t0 a$ M$ e& u# |8+ o# L' b7 C) [9 w# k
9
. a' Q& t. ` k x5 S1 H3 }1 d* m6 K10
0 y2 n" F5 x2 y$ E11# {7 ^' r6 f0 ~! ^
12
% U& b. K1 o" w( p13* O3 s. B6 o [$ ~8 D1 ]
14
# p3 \# }! r, o8 ~% N151 \, H! @0 d# q0 W5 A
16" O6 U5 q# n3 v& d; x
17
( U$ A3 c& }2 W, f* U6 T G187 {% f; P) p b* L
196 {: C( e$ ~) T7 Z9 b
20' L. H4 K0 @+ a' ^
21
7 S' [1 m' Y# U0 q4 z N/ y! C0 T22, Q; X3 L8 {3 k5 l$ ~
23
8 H1 p1 w( a, r5 C7 Q# s24* A% }& S" e L- n" W
25
8 E/ X% z& I& u% a; r/ s266 x$ n$ U5 O! o) ]& q; B
275 y# O/ a4 D" f: Q
28* \2 A9 c3 T; l' S" j2 Q' R( @
29
/ U, M% r" ]8 N9 F7 e301 w$ F O: c* z( k% ?! p9 r3 x
31, w) ]2 b& X4 D' z
322 r4 r/ n7 G' e$ h. Q0 Q
33* n! D) a9 _% \
343 _ ~+ c4 I* ^; |+ N
352 g" N! E5 A2 T2 [- |
36
- l1 i$ w) A1 q3 e& @37
& T; n4 Y, |* Y7 n. i4 [4 ?3 X+ t38. s4 ]2 x2 `, c- @1 ?* T
39
% y; ?& f* k/ w40
6 B8 ^ j9 [: l' {" _" f41
1 E2 R+ n( y6 h4 ?9 w$ }9 E/ e42 r F3 r! {. T7 q& [. {
7.3 训练所有层; `1 b2 A. J# i7 z. O
# 将全部网络解锁进行训练
* v6 r3 c' F Y& Mfor param in model_ft.parameters(): `. \, j% e- c9 j6 p, j
param.requires_grad = True
+ W& p7 o N5 @ T' G4 u# \: ^2 G8 u1 J0 P6 ^! a4 [
# 再继续训练所有的参数,学习率调小一点\0 w& H0 B |7 e+ p& _; b1 p
optimizer = optim.Adam(params_to_update, lr = 1e-4): @' Z/ Y. C! ?% n$ S
scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)2 B! H, {7 I+ k+ p( U: _0 x) W
( a9 d/ h) ?- i' [- M# 损失函数0 ?/ A* C' W, z% I& _
criterion = nn.NLLLoss()
7 P$ A# x) p0 X- T& Y1/ K2 U/ n6 v Q
2/ a* {; A# C$ p& G( w8 P& M9 T5 ?
3. o5 X# l' Z- r/ n1 p) a
45 y2 f! v) {$ I- \( V
5
# V: u6 {& B2 {& e" Q6
( `6 G+ w; g" |& R" ]" S; }9 P: @( p7: N/ `* y, t! }7 e% K
8/ l! i2 L9 I# D; X/ b
9
6 J6 u8 Y; J, F3 R: o5 @; {7 Y: x10
+ a2 ?) s' ?; y9 N) [8 d) y9 U: Y" s) Z# 加载保存的参数
4 q- l2 E& t- R# A/ W0 ?2 u: I# 并在原有的模型基础上继续训练
8 z [7 V. h( L: S' ]& R; x3 i# 下面保存的是刚刚训练效果较好的路径- Y$ k) \! }3 Z4 k) q& U
checkpoint = torch.load(filename)
( k6 I2 H5 {' R2 ]best_acc = checkpoint['best_acc']
% X) J. L1 \1 c& fmodel_ft.load_state_dict(checkpoint['state_dict'])
& F" g ?& Q( m8 V0 v% E( Y6 Roptimizer.load_state_dict(checkpoint['optimizer'])
& I& q6 _) j( ]& V; p10 P* \ @! Q8 q0 e
2
" o- q" z' V; @3
2 Q5 o" _+ W! N4
# I5 R0 x8 Z% V6 D/ q, m! P6 q5
! d8 F4 g( }1 y7 i+ k" {6) N5 ^. q5 A) b8 H
7+ V) T" a, g* V% E
开始训练- [" `" V2 ], ?$ E" C
注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考8 i& A7 J& j. d9 F
2 P: q5 Q' P8 Kmodel_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"))/ v6 f) N7 c3 k7 C$ p# F
15 X4 z$ r( [) B% J+ H
Epoch 0/1& R, n. V' }9 M/ v# C
----------
, O; D) U; K3 G7 O" b4 m6 y/ b4 [Time elapsed 35m 22s
0 O; O) U3 ^! F+ @* h$ k: \9 G: e7 htrain Loss: 1.7636 Acc: 0.73461 ], j" e3 U/ {1 t
Time elapsed 38m 42s8 \9 p4 L8 b6 l9 h
valid Loss: 3.6377 Acc: 0.6455
+ P6 P- i; [4 w' j( J, h nOptimizer learning rate : 0.00100006 t5 n+ q, a" f Q0 _; r
, H( w7 K6 ^, M8 c+ h; q; wEpoch 1/15 K9 A2 n" \! o0 w) u! Q4 s8 G
----------8 {' d% C% l+ i9 n( |
Time elapsed 82m 59s
, ?8 y( k* G+ d3 |) _' Gtrain Loss: 1.7543 Acc: 0.7340
: S5 X" W" d/ d5 ^1 J3 nTime elapsed 86m 11s! a; w* W% I8 Y6 z2 }. z
valid Loss: 3.8275 Acc: 0.6137' f9 z5 Q+ m2 ]# r
Optimizer learning rate : 0.0010000
: U+ P% F9 e" I( `
- n; c; X' Z8 f4 }9 f# RTraining complete in 86m 11s, F T7 U, W* R! w
Best val Acc: 0.645477, {7 Z2 p8 V: [) E
2 G6 n# H$ O) y6 T) t1
) s' O$ K: u& ~8 S2 D9 n1 @2: v% \. I0 @/ Z% R% m3 p
3- @7 {" L) B& }4 D) f! `* I
4
8 J4 K9 x9 N4 C5; }& s4 [; n, G2 a
6
% v: S! Q) T0 e6 I7
2 w5 e" x/ f6 F( ^$ e6 { @8
/ n. p/ i: h, G3 O9 A1 p; Y; j2 m90 p) e& w5 Y- v- ~
10& P+ v5 c- n I% s& I1 P9 L
11
. c7 B* p* Y/ T8 v12: X9 y5 w( @ x' h3 s7 q$ D
13
4 X8 m P8 U8 K) ~" T. `, D14( ]7 U, d9 {' c: c+ \
15
+ U6 o! d' R7 [( l16
* v; P% k3 ^! [1 S7 r17. A9 J) D) E% p" G W
182 o. U1 d3 L' W d; Z: T
8. 加载已经训练的模型
) o- H+ S. e w4 ]5 P( H6 K相当于做一次简单的前向传播(逻辑推理),不用更新参数
% @& d. ]& M1 o( }( b+ Z) T- F5 y7 z# a, p8 {( f
model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
- V9 V% w; u! w+ U# \" A7 @
$ C, ]# P% W4 i- p+ R6 [# GPU 模式9 c E, k S; f2 ~
model_ft = model_ft.to(device) # 扔到GPU中+ u6 q2 u! u" [( I) v' O
/ G1 l! l5 T8 E B/ E6 f- V3 H
# 保存文件的名字8 V( x ~+ U' V+ S
filename='checkpoint.pth': G- X' ~3 U$ u" j
4 C3 T3 \7 c* s. y4 o d
# 加载模型1 d1 J$ A: U4 S R# D
checkpoint = torch.load(filename)( _' z$ C; d. J. o
best_acc = checkpoint['best_acc']& D5 W: J: E i6 [" r& ?
model_ft.load_state_dict(checkpoint['state_dict'])# u+ z8 N- L7 b8 i9 X
1
% D' k" ^5 [' n2 i0 L, L28 w# T# N, P/ h! W
3
. T" U/ T, Z& q$ L4, [3 n/ H) O2 x4 I
5
$ J& |9 @8 [- D6; P+ ]1 m3 Z& |+ I
7
9 I6 \/ l3 f0 N% z! t' x8
' I" Y( N% k9 W' ^0 b: M/ ~91 v4 O% b3 {6 M- {/ q
10+ K/ c& ]) x$ Y9 K* L
11) u& ?3 E; t: q# f5 e& \
12* p0 B B# G+ Y6 _( }
<All keys matched successfully>
1 L8 s. e4 L+ M" j1
: v/ I" n' g3 X$ n7 B ^0 @def process_image(image_path):
/ A3 o/ m. L8 k. \; h # 读取测试集数据9 K, N: j1 k/ F$ @5 Y3 c9 N: Q
img = Image.open(image_path)
1 B: p8 S \' { x4 C # Resize, thumbnail方法只能进行比例缩小,所以进行判断; p% ^# L; K: v' J# O+ p# h
# 与Resize不同$ @* V( V0 V# M* }) h# U) W1 i9 N; r
# resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小# g# S$ n7 n( @4 C
# 而且对象调用方法会直接改变其大小,返回None2 |# V t" F3 j
if img.size[0] > img.size[1]:
" |; t! z: C L$ ^ img.thumbnail((10000, 256))
3 l& Z1 H; _& S S6 t else:1 U2 A9 h1 b2 n( |2 X1 ]9 S
img.thumbnail((256, 10000))- Q7 k' c8 w9 d5 B, C
' J; P- y) F, [8 U+ P/ `# Y& N
# crop操作, 将图像再次裁剪为 224 * 224 K9 g2 Y( S! x% {1 P" s" Q
left_margin = (img.width - 224) / 2 # 取中间的部分% Z% w9 t" K- `8 s/ I/ ]% x
bottom_margin = (img.height - 224) / 2 1 v4 J: T9 @2 M
right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度! l. z. P, w. ], w
top_margin = bottom_margin + 224$ I' U. H6 g, [; u3 y
3 R9 o7 J9 a8 B! f
img = img.crop((left_margin, bottom_margin, right_margin, top_margin))
+ p- P5 Z4 v. @" K, k! F- y: N( f' G& \9 K' T5 j
# 相同预处理的方法8 l3 Z% l4 y# `( D& a% z5 Y
# 归一化
4 I, P1 f6 @4 h' ?! z1 z+ p img = np.array(img) / 255# g- B7 O i6 X( s R+ z* D% Z
mean = np.array([0.485, 0.456, 0.406])
2 v+ v' |4 {: s+ \8 D9 L1 h2 ]1 | std = np.array([0.229, 0.224, 0.225])
( Y( Q, e4 x0 k) m img = (img - mean) / std* Q+ F5 ~- n/ c; [5 [0 k" Z
7 V$ s" k7 R+ B8 h, | # 注意颜色通道和位置: x5 D" v& J7 O+ j0 C4 `
img = img.transpose((2, 0, 1))
$ `. N, A$ l. s
+ q; r. W: _0 o. J2 D return img0 V3 x2 \3 q6 S/ }* l' S
! j* Q) {- A- S) V
def imshow(image, ax = None, title = None):
8 @$ s9 n; j" b- F """展示数据"""
* i- I1 H/ e/ N- z, Y" Z0 e if ax is None:# {2 G7 s! Z/ v6 A
fig, ax = plt.subplots()" Z, Y( b. y4 s+ i
: N0 U0 |! L* h
# 颜色通道进行还原
2 l0 r, L& |2 s image = np.array(image).transpose((1, 2, 0)), w* Y. t1 P7 d
, r: V& b: n' l
# 预处理还原& G/ w ~& t( N5 N
mean = np.array([0.485, 0.456, 0.406])
2 J+ e) Y9 A2 d; @% S std = np.array([0.229, 0.224, 0.225])
9 D+ W; F* g4 G, \% h& | image = std * image + mean
c: J# i( @5 b9 f8 l+ F image = np.clip(image, 0, 1)
: u8 M2 v2 [* ?6 e. n$ H" m, i
$ R+ F# Z+ ~" A6 `8 m/ \ ax.imshow(image)
3 Z+ s& i9 G& g1 x ax.set_title(title)
2 C K2 e" s/ u2 q' L1 w5 p% ~6 R) r( d1 v
return ax
) }( a M1 |' |' N8 F% Z) f4 d' ^; i9 U4 ]4 T, o* \1 ^! b
image_path = r'./flower_data/valid/3/image_06621.jpg'' @/ i' d% p/ T# t; M
img = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理
; x& U' k0 Q+ @% m4 f1 E# s4 nimshow(img)# S' x$ D( y. G0 j& E+ i
6 e, ^# j# G) R: P) r& r8 Z
1
Z4 ]7 ?/ L l5 q8 S( y, Q/ S2" e+ y1 H& W8 @$ \& ^
3
% X8 A* a6 b7 T5 V. I- m4
9 M! X" p$ y; e9 |# E5; q" t! f% m- {; E2 h0 k
6
% c; Z$ X8 W" F x7
. v$ {5 [) A4 J0 d. ?8 w8- G# \7 s9 M! A; V; e( @
9: P6 W. `& k( _! Y, V
10% X0 x% w% @5 E( Q: t
112 {" B+ ?( F9 M9 b" F
129 m5 P/ t. Q. S2 V1 q
13
" r [# k" r, J; u14
( |8 w) w% l1 F. }" ^. K15
v* ]5 T8 R. L8 ^, f* Z16
0 ]4 F3 d* n; X$ N/ [$ `173 B+ `" V& y+ @! o* j
18
N ?3 Z% Y0 p- v/ V7 N8 h19: u* _0 P3 V; K1 y: D
20
& t# y4 d! D9 N+ ^$ d21
; N% o6 k/ p3 m6 }22# M: r0 g# j' p
23
/ q$ K$ X g& ?240 v) y) k" ?, x' U* U( Z
25, g( O) M( ?5 R
26
6 V, F4 E3 g/ v6 ^- U3 Z27
! F0 M0 q/ d, H% u28/ a. w, _3 `. q; W. z: U- D* \
294 O+ i" l8 K x: p @- u! u
30
; }4 G, O6 ^! A$ A! R" |% D31
: d/ ` X. p3 E9 |32( K7 i! ]8 `: }# p: g. F5 J m, K
33- r$ f b* f9 Z8 A4 b. t; \
34
( j. u' Z2 }6 C6 v( r! k35
4 r* n; `4 g3 k. d36, J( L0 v) r# {' j0 f5 S
37
+ i6 r: ~: _ @: M1 g7 q38
' [& i# ? h- Q9 g6 G1 }39
2 A r- Q* P. u+ j40
. W) {3 y# F. L V. J41' x) ]* K) ~0 x: U
42
& [. `/ Z# }. E: E% V+ A; B2 E8 u43
* P: y; Q: p+ t' u1 F44/ l, w1 c/ s' m# I/ L
45
1 l: F. `$ [6 d$ X( T( J1 n46
) v# d+ F. \/ L5 f8 q. Y4 K3 K47! t& @3 c1 n+ l" i6 n
48& ?" V3 U2 [' r+ w! \6 L! K
49
; H, \! w: w1 A, [7 c9 z9 i# S' q50
' b7 M; O! M4 C51$ W: l5 y& f7 K: e4 }" V4 X
52- K$ U; g5 k T" X' ?
53
! b3 s$ E% j- H2 Z549 @; Z+ Q" F7 N# W; d
<AxesSubplot:>3 h3 Z# M1 n9 `. |9 U
1& t0 H! f2 @# c6 d: y' ^
* U3 T( W. H3 c: ~& Z+ g W上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确; |; a/ U7 R, f8 \8 E" q
h3 \; d8 t0 L" X4 L
img.shape
7 a8 A. }; j: z* F% S# H! ~! H) N1/ E$ M5 y' }. {% q
(3, 224, 224)5 E% D. W; f$ k) L6 E2 y! A
1
1 _4 O+ X( L& P- W% z证明了通道提前了,而且大小没改变( f; g# K+ L9 d) S+ E- t; v$ O
6 p/ \5 S3 n4 b3 G. x1 r0 [9. 推理% J' _1 u# O; L( O) }7 m6 V- M% @
img.shape( c5 v6 J2 H: O/ i2 G) U6 T" u
& @* [, G: z+ e- O
# 得到一个batch的测试数据0 z4 v9 `3 G$ e1 \) y; ^! f9 h* {
dataiter = iter(dataloaders['valid'])/ X* v* v4 M! @
images, labels = dataiter.next()) ]0 z6 q# o" ]1 m h
9 E* X* r! ]1 \& l+ M( Tmodel_ft.eval()& u7 A+ E9 Q4 V- l/ M
7 }* ~" G9 V3 Z2 q2 x
if train_on_gpu:2 Z$ h' C' n( g% l
# 前向传播跑一次会得到output
' U1 M2 s1 A/ N, s6 n% |8 a- O output = model_ft(images.cuda())
' ]! x; H7 J3 T7 I- ?else:- i. s! o7 r4 K2 @' Q
output = model_ft(images)' Q A* W$ T7 S0 a8 b( @4 q2 i0 ]
' ~; v* L" a8 C0 C# H# batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值
q& |9 ? p$ S; x! _9 Aoutput.shape
0 o; D3 b3 i9 ?. s6 Q+ A7 U( c n$ |1 R2 H& D
19 Q5 `! L) i/ D$ ]
2
# q$ t4 q0 t- ^ `3# r/ X6 a; [. e7 h; C1 Y
48 ^& q0 G; M2 `! T, i. N
5
7 I$ b* w+ t7 k. s% d+ s& K5 d6
5 l) }: v# T+ ^ m0 I7
) b3 q' G6 e. f8
6 b- }' v+ w- R- U h7 h7 y9* i" b5 T* ]- z
10
# T; D: \1 ~6 t* g$ e9 k118 p1 A0 O: e6 n( D
12" H7 M9 W% Z* ^% L, T* `! D0 ^
135 i/ ^' Q' k& U5 S! |8 v
14 }$ z% j8 Q# W8 R5 i a7 S% R4 e$ p/ N
15
1 y6 E, t/ r& N) e16
/ r, I- U+ \+ S4 |) Ytorch.Size([8, 102])
! y& ^' A# p2 k% T: s1
( T& V) ^ z, b: U6 b+ R! B% P9.1 计算得到最大概率' s8 B& E+ } T- z; }" p$ d0 X4 c
_, preds_tensor = torch.max(output, 1)
3 |- H4 e# e- G' w0 d5 Q' ]! {8 W1 c5 C6 }5 M7 ]' O# h
preds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量8 U7 j9 c: x6 F
1
; G# b0 Z+ o$ U! w5 m( a2$ G& w4 E \: p; X" n
3
0 W4 ]6 Y& [7 \$ p% o9.2 展示预测结果
( t2 `( Z9 C' H1 q' rfig = plt.figure(figsize = (20, 20))
6 ]' X" }, M) {* b2 X* |1 ~1 jcolumns = 4
) {/ V+ I8 |8 [6 X0 j" Mrows = 2
" r* S6 C1 |9 X; H# Z& ]0 t' B7 o0 G& T% W
for idx in range(columns * rows):/ L5 q! ^9 {8 M1 c$ C7 d0 F" \6 F5 u
ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])
) i; l" J/ b/ k, b plt.imshow(im_convert(images[idx]))
2 J$ B9 L1 T9 H! @& ]1 ^7 A ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]),
. j$ V8 r A. L' y, E color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))9 n& A( S7 p4 e9 R
plt.show()) w" y# ^* Y8 J S
# 绿色的表示预测是对的,红色表示预测错了
; H; E) O/ F3 L. y+ v( ?+ G U1% U, ^+ I( Z" f0 ]3 d+ I
25 y/ m& [2 w+ l" y
3
* r0 |/ S9 V' F ?8 o4
) A0 D5 |; ^8 t- C8 L ~- i$ |) f: E5; k& ]# a8 m" z( ~
6% z7 N0 ? l4 O+ e
7$ o, b- j& W0 C. S: T* X q& {
8
: u* Z. ]5 I* N. m9
" Z1 y" q, C+ B, W+ f10
) E$ ~4 H, j* c2 q! R11- b& E5 g& ]% S. ?/ J% A& D
, W2 _& ^! R# |# E% W
* n1 I5 i5 k8 Z% D
4 J' F4 v! v& `% V1 X
————————————————
D6 I, _6 s5 @& |+ W9 m9 M版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。
! p1 @ Z- V' l% d z& V原文链接:https://blog.csdn.net/LeungSr/article/details/126747940
6 t; y" h- U+ \7 k T& [/ ~9 l' {: f- a- b
" W0 h( W2 A2 @) P( S% [' b
|
zan
|