- 在线时间
- 1630 小时
- 最后登录
- 2024-1-29
- 注册时间
- 2017-5-16
- 听众数
- 82
- 收听数
- 1
- 能力
- 120 分
- 体力
- 565610 点
- 威望
- 12 点
- 阅读权限
- 255
- 积分
- 174906
- 相册
- 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)实战案例% k1 \- y' M( M2 A9 R
+ r1 N3 f O4 m% K
文章目录8 M+ U. S$ ?% Y+ T4 d0 y
卷积网络实战 对花进行分类0 A3 e. Q* U4 w
数据预处理部分" C- ?7 k/ W( i1 C
网络模块设置
5 n& ?* R, M; h3 B' y4 } G: Q网络模型的保存与测试
' K" e' ?+ l2 O数据下载:
2 n$ {+ x: F* V w$ E1. 导入工具包$ h+ O8 ?9 E0 t
2. 数据预处理与操作9 @8 c* D5 i/ |, w- ~7 R) S
3. 制作好数据源
* ]% W6 O, a0 L; |读取标签对应的实际名字
! m. t- p$ E4 r4.展示一下数据0 _- } _, e- S% a* G" K
5. 加载models提供的模型,并直接用训练好的权重做初始化参数
" O4 E& j+ \: c2 o6.初始化模型架构$ l' a5 F9 w, M: Q4 o, l, Q/ ~0 t! o9 T
7. 设置需要训练的参数0 q1 {4 q; a9 }0 k0 z3 K5 h
7. 训练与预测
?3 O# M1 n3 H* n2 O6 G- C7.1 优化器设置& [' E* j0 L' i7 k; R1 H' o
7.2 开始训练模型 e' F$ A* V1 e8 Z' n
7.3 训练所有层
4 | L" f# q- B开始训练
1 a+ n F9 c3 O3 R8. 加载已经训练的模型
* { I5 a8 L" Q, t6 X9. 推理: P5 e. c4 C: `, ^6 o7 y
9.1 计算得到最大概率: Y7 D+ l% }6 B
9.2 展示预测结果
' N" ?5 I8 b0 q0 h- a+ \写在最后* x7 [- i/ S2 i3 w/ K7 f) ?3 L8 u# T
卷积网络实战 对花进行分类5 p( t- _- F- k, N3 G5 E
本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能
`7 O& c9 k2 N; P' _
' D6 R$ L6 m7 [5 E f5 ?在文件夹中有102种花,我们主要要对这些花进行分类任务1 J8 K3 B) ^# \- Y3 F+ ?* R
文件夹结构( m0 B. P% ]1 n5 g& N, J0 l: @' N
$ r/ G# w: ?* W& W
flower_data
) H! _2 N4 h6 c! A3 C$ b2 L. f8 p- K# k- ? X& ]+ j$ `6 m% R* H
train
* S6 n4 y, `# B0 R2 y) C
) F, P# D9 u, c5 q4 Q3 } b. a1(类别)3 r" [3 `, u3 M- P! h8 q' K
2
& m6 r0 C' \% O& mxxx.png / xxx.jpg
4 K5 j. L' _- `. f q2 F/ N& uvalid7 @" H/ Z1 v4 |8 `2 S
# t' q7 I; C9 W$ c$ O6 u0 M
主要分为以下几个大模块
) d0 ~/ J' l. [, I( \; X8 m0 G) }- m6 s5 ]9 @2 u) c
数据预处理部分. [& M8 y8 s- U$ `6 q+ _: U
数据增强% K, a. n5 ^. J' G: a: _
数据预处理
* d+ b0 k) U# f s网络模块设置
4 Y+ T: l4 j4 h9 w: \8 r加载预训练模型,直接调用torchVision的经典网络架构& P% f$ R" P* T3 o
因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务 a- G1 q8 o* v; g
网络模型的保存与测试
9 w9 R. \4 F4 [模型保存可以带有选择性
5 L8 @# o/ \( w( s6 u+ }数据下载:
4 w' Z5 v1 I. F& H1 Ahttps://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset, e7 i. [8 [# y7 z
2 P6 H6 b( W, r
改一下文件名,然后将它放到同一根目录就可以了/ s. x; M. S, \/ r$ d2 y+ z
( h V" d Y( R# w$ V
下面是我的数据根目录
% j, h! r4 |8 J& Q0 a- h& g; S8 T9 D: X# S" ]# m7 k5 i, w
5 H9 N9 O2 P5 ]+ B- D1 i2 x2 x" O
1. 导入工具包
7 O; G r( ]* |" r; Limport os
- Y: i+ y7 I% O9 |+ ximport matplotlib.pyplot as plt! ^! g6 y; F Y$ B" ~
# 内嵌入绘图简去show的句柄
+ s; b3 L6 Z# z& f3 z, g9 X: s q%matplotlib inline
6 x- w3 N6 O0 k! I" ]% Ximport numpy as np) P2 d1 r' J2 j2 A% g5 B0 I
import torch
9 S1 ^8 F1 x9 F2 \# Q) pfrom torch import nn
2 ~& C' G, h3 x+ C. m- z+ R+ G4 Y/ i4 Z# }8 _ b9 j7 S* E }
import torch.optim as optim
7 v" W' p+ {" ?) e$ U7 |, m& limport torchvision# p1 [: [2 v, B* [* p
from torchvision import transforms, models, datasets
" S' A9 e$ a- R; n. t5 t) h
2 f G# M6 a) T4 Qimport imageio
. @; Y2 y1 E. N! Iimport time
: r7 c8 r. D# O7 e6 U: Dimport warnings
, t! \! ^/ E8 y- t4 u" bimport random
) c7 }$ p9 S) t4 @3 l1 simport sys
8 ~+ y7 O( N/ h1 {" Q% Kimport copy$ S' s* t7 h# O% o# z8 e! ^
import json$ ~0 C1 I+ y5 O
from PIL import Image: [! v- J8 C; z# a! c- K2 g+ w+ Q
% t6 O9 v4 p! x9 c* H1 u, O
/ z2 M5 }2 ?& `& n2 R1
# K% e/ [0 Q3 w2
' E& s j8 o: n( |9 Z( R( {* P& f3
! ]. U( ^( k% b( E# m h; H4$ z* v* [; y; \" q
50 v' a+ W5 f& e# O* Z' ^' ~4 T
6# |( K6 Z: ]8 \7 H$ O$ ?1 t
7
! k# c1 c5 V( c H- o8
9 o& Q: _; Y9 c' D4 _9
3 b6 {) T2 O# a5 L7 F& q10
4 v" m7 c4 `3 d9 b4 D11* ~! R" }! o/ R6 _3 n# _8 U. i
12; a/ B1 H* O4 H# E7 V( m
13. P6 a( T- ~, K% T9 i* t5 E5 p
14
/ U8 ?% n! ]* |% R% \15, V( ^4 Q: p# L
16
4 P6 M l4 n9 ^$ l- X6 X, e17& c. N% }$ O5 h2 Q0 ~# K( E9 M. K
18* A* n; n q% H( E: `
19# R: t, S9 Y: N6 K5 m
20
" P+ s4 `/ A5 d' d: Y21
7 ^$ L5 W% q0 s0 s1 H: B" }2. 数据预处理与操作, H- b* C& t5 u6 Y
#路径设置
# |7 Z- [8 R3 T" A8 n( Q" Pdata_dir = './flower_data/' # 当前文件夹下的flowerdata目录! x' w; o+ B& J
train_dir = data_dir + '/train'
: i# c) _8 W0 t& B! L+ Pvalid_dir = data_dir + '/valid'
3 N" O: j e1 {8 `% P0 ^1& I' t8 X# g9 b# ^% r8 E
2
, t a& K! |5 w# q3 W8 t6 K, J3
5 ~% w5 n" L1 I6 x; q4
# f: t3 f, d. H( x$ apython目录点杠的组合与区别$ @* U6 ?8 x: y2 u8 `
注: 里面注明了点杠和斜杠的操作" h5 J: A f! f% O
! \% n G3 R7 A1 D0 B M! ?3. 制作好数据源5 N+ r: Z# u5 L6 v1 d
data_transforms中制定了所有图像预处理的操作6 e7 f! E. z4 L, y+ L
ImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片( |/ f/ H; P# n' e
data_transforms = {
' S% L1 G! s2 M% m # 分成两部分,一部分是训练7 |8 R( Z! w# r6 H
'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间0 J% p4 S8 B# L5 ?
transforms.CenterCrop(224), # 从中心处开始裁剪2 Z, N& e+ D1 q5 t4 }, N
# 以某个随机的概率决定是否翻转 55开
- e7 v9 `- I: i! q, Y; ]2 q transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转
+ K5 N8 q8 x4 t% L. N) F4 r transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转( U/ ^: R* N J) W, Z: O' B1 Y& w7 ]
# 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相2 s+ ~+ i. r( o- H& Z v
transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),) v5 o% Z! P+ |! g
transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB
; z5 N$ I7 B" \/ u) n1 O. O # 灰度图转换以后也是三个通道,但是只是RGB是一样的. I& R' {" q2 k: K: ?6 n) }# e+ x) G
transforms.ToTensor(),! G$ v! i" H4 |, R. v$ ^- @4 n
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差4 K$ _, T4 B+ \' ]
]),
- p) }& E5 j. _+ w # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
; b! o* g7 X* c7 o+ n5 ?- [ 'valid': transforms.Compose([transforms.Resize(256),2 D" T" F& n' h5 `& R9 Z- E D0 Y6 ^
transforms.CenterCrop(224),5 w$ k- ~$ y3 }( h; E; L9 B0 m
transforms.ToTensor(),4 v4 [, o7 d4 I9 U! ~3 a. c5 E4 l
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同8 m: h7 B m% W0 ?( M
]),( P, \# C7 j& ]% X# F2 p4 ~. r
}
4 `3 U# o$ W( G! `. `, O3 D. [6 A9 a4 P+ M3 T" m
1
" \7 s( C1 ]& s- U2
7 K! w; _% ^1 j5 L* }. P/ q( s/ w. J3
: ~" T' a8 L \5 W4* R4 N! W# m+ d" d8 n+ E
5
O) \5 \& l4 r' { m" E( d$ _$ r0 z61 d n4 c6 R1 z5 v: L8 S9 R
7
1 Q7 H* h. c( K: K8
p' k+ U# \8 H9
( t/ g3 L: o1 Z" n2 e' p1 p10
- j1 E; t: {! D) b/ I11* C9 E% a: ?" l( L( C- x
124 e, G8 G9 R6 i5 w8 M
13
' f- |& D$ T+ u- Z/ }' E14
& _0 L% z% |, `6 G& R8 e2 v15. t/ T! p* C# D& t# W$ d: r
16
4 ^* g7 b& v/ \! b175 u# c0 C3 S0 U/ i8 ]& x7 ~$ \
18
# r4 {8 U/ R0 r; P& x$ O19
/ b) x( B6 o4 D/ I, A20
6 s) T; O2 L& Z21
; q5 N; }1 W1 o# }6 {batch_size = 8. H; ]. v( t" @9 l) D
image_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}9 c( B. z8 Q, \
dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}( Y, @: l; \" E# ]3 l- f5 ]5 P0 f' t8 N
dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']} ; ]5 u; G+ R2 a S
class_names = image_datasets['train'].classes
' _/ b) E# G U- d
4 K* G0 r: k" a# L& S4 o. @% W#查看数据集合
! c4 g, b6 a) k3 I4 I" aimage_datasets
+ f; |9 ~ b- |( ]& n" v5 j8 v3 f1 ~( S, Q
1
% W" |! ]# i" N* C* ?) T29 `7 i. O$ |& {+ m9 W' m
3
; Q3 J4 ?& V. Y- W5 n4 q4
4 ?* Z& b) [) ` G& f, ?+ H: x5
5 ^6 O: z$ a- g0 L, j6
0 r/ {( |) y9 l8 \/ |) ^) l9 H7/ r# ^# N4 ^9 r3 G& X5 |- w
8
5 i" ^9 g& ~4 r/ K2 H# B) _2 ?/ z& q9
' d, N9 u4 c. W{'train': Dataset ImageFolder* {& c3 a5 z7 l# ]& M6 K* W
Number of datapoints: 65520 e8 c5 U: z! @
Root location: ./flower_data/train
/ D$ K! @5 u" n Z9 |( m+ s7 B+ s& T StandardTransform
0 W, s* y3 r5 P/ ?: ~ Transform: Compose(
' G) n, D' H1 W/ S RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)" ]% b' u5 x6 G
CenterCrop(size=(224, 224))
0 ?/ i' b' t; O RandomHorizontalFlip(p=0.5)% B6 i& D3 I4 D+ I/ p9 q* g
RandomVerticalFlip(p=0.5)
$ Z$ E3 K! |, w5 J5 @ ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
$ i; T% L5 @( O; u% z8 i7 ] RandomGrayscale(p=0.025)
$ @5 ?2 o5 \- m+ [4 Z ToTensor()
0 g& M! M) R$ Z3 K6 { Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
3 b5 M5 S6 ^) X' O7 y" Z ),- { f% z+ u1 A. n! N6 A* M
'valid': Dataset ImageFolder
* c- x V- [ Z9 O! | Number of datapoints: 818
& @. E) u: R" L' W/ f7 p3 w9 L6 h/ U Root location: ./flower_data/valid7 {- \& }4 b0 A' H; A y& l
StandardTransform/ V* m t& O5 [0 ~: D# k, Z
Transform: Compose(3 [( G" }/ H2 o- M( e) ~
Resize(size=256, interpolation=bilinear, max_size=None, antialias=None), `: \ F4 n* I3 f
CenterCrop(size=(224, 224))
+ n2 v& d7 T t/ P/ k6 @ ToTensor()
1 k T) P2 G, T2 z' r, a; m Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
0 V" F" R, ]( _* {/ T )}
/ ?1 w2 \3 M- W& r+ A1 e& ^. X
1
) S. R0 w) t# l/ A4 A/ k% J25 j' F- F7 R) y
37 F" R8 a) O( Z9 }
4
& s v: K: }4 O5
! }7 g, g0 C7 e, O) s; H69 Q6 Z" W& f- [& a
7% T. F- o$ o4 t- j. R3 W" F$ Z
8
* s. X2 k' V% H1 H/ W# \7 g9
+ m1 M# o3 L A" Y. e4 D10
# b5 P$ e' E9 m) S2 u2 y11
7 V6 D8 k; I' s- }! b6 n126 a% M) _# i, \! l. v$ y
13
" `. G6 [% C3 P+ k) W14
" R8 N2 T/ {0 a {4 ~( l15& Y' ~2 G% I; Q/ y
16: o( F, m* B k5 i, f: j2 R
17
3 Z3 U5 ?2 `$ l18
1 ~. f) J( g, N" r1 S19- _( E6 x ?% Q. |8 w* o9 I
20) m6 J# [0 [ v! q1 e2 J# N+ C/ t9 I
212 l" G) X1 ^$ y
22) r6 a. B \) O( }
238 `! b- \0 e) Y- s) |2 K
24
' G% {1 g( F% k2 K* P+ S& ~/ o# 验证一下数据是否已经被处理完毕3 y3 K, u. l4 q+ f
dataloaders
2 h; L$ \* N7 V1- j# T5 q9 b( F0 r4 [$ R
23 \$ x, w) c% [+ _
{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,
9 s9 T# L, c, W" e* W 'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}5 N9 K$ D5 w5 J% O7 o' q9 y
1
' K! z2 {8 h: j29 O \' g" G* z$ Z
dataset_sizes
" h0 ?0 l* P3 g+ \- Z1
0 c2 h. U1 j0 r: H5 H5 S{'train': 6552, 'valid': 818}
" c7 e7 Z! {! j1 e11 J8 j7 j B1 {+ |' J
读取标签对应的实际名字# N. q3 f E2 U8 l ]
使用同一目录下的json文件,反向映射出花对应的名字- H/ F3 c& O* F) J' H! s5 Y' e
4 I* @1 d4 Z( X+ a5 ]
with open('./flower_data/cat_to_name.json', 'r') as f:
# W/ Q; B* m) c% M" P cat_to_name = json.load(f)0 B T2 } q8 s
1
* u6 l2 ~1 |0 K$ E8 U2
+ V( _/ _7 e- P# r. x9 x6 ^cat_to_name3 {. c) {! z. `8 `: `
1
, R0 n# T1 U$ m H9 a. ^{'21': 'fire lily',# z( h4 u0 |1 Y" ?* P+ N& O0 u2 n
'3': 'canterbury bells'," u. L& {6 H# K; _0 u4 J2 @: ]
'45': 'bolero deep blue',/ Y; R0 I1 e' A8 k6 A
'1': 'pink primrose',
+ S5 G, r& z4 R9 T! @6 _, E# d) f# U9 h '34': 'mexican aster',$ W1 X0 V: v D9 _/ J3 h1 B g
'27': 'prince of wales feathers',
- x2 W$ u. F# k5 q '7': 'moon orchid',! C' G3 X9 f `( w9 T
'16': 'globe-flower',% R' D; r/ u# H# ?* J( ?5 { w* M
'25': 'grape hyacinth',: L' I! ?7 \* x3 ]
'26': 'corn poppy',$ @$ I! O1 @& b5 Z
'79': 'toad lily',7 x4 I/ U8 K9 g2 h
'39': 'siam tulip',% A# v! @; J3 r2 f8 k
'24': 'red ginger',# X# | v! v# _. d% a4 V4 s9 W
'67': 'spring crocus',
/ X/ `5 \& n3 R8 F" p '35': 'alpine sea holly',9 K d0 [- f$ j" j4 ]/ B$ Q
'32': 'garden phlox',' ]; O- d( X9 j
'10': 'globe thistle',
3 } |# f& ^; s+ v/ J0 c '6': 'tiger lily',
8 i) w- X! y& m- O '93': 'ball moss',
# J, [5 t) h, O' ?. a/ u6 _ '33': 'love in the mist',
6 Q. t/ [0 P4 r8 `6 e/ f% v '9': 'monkshood',
6 v3 Z$ V0 t, O) u '102': 'blackberry lily',
5 q4 a# L4 W# a% j- N; @1 I* a% H% K '14': 'spear thistle',2 j, Y' R0 T+ ]" {# {. Z3 X
'19': 'balloon flower',
" q7 x, {: `! h( l6 Y) D '100': 'blanket flower',
, V/ H' g$ F3 j, Z '13': 'king protea',+ r# ~9 S* M6 Q9 v$ \+ k0 s
'49': 'oxeye daisy', e. S: {- J8 u v3 R4 Z& w. J
'15': 'yellow iris',0 J3 d E! }# M/ f% @! V2 u
'61': 'cautleya spicata',
& O$ B1 K+ ?0 J4 K '31': 'carnation',
; H+ h6 u2 r; R' ]. {+ n '64': 'silverbush',
- \; {; I- `% v- A3 G" Z '68': 'bearded iris',
% B# n9 K$ T2 V3 H! F+ @ '63': 'black-eyed susan',
/ u7 ]" V2 T- n4 Q/ c '69': 'windflower',
! i( b+ M! ?# Y' K% O '62': 'japanese anemone',# B H( w! _' s# k/ I$ f. i
'20': 'giant white arum lily',
. B# w2 l/ _6 [ '38': 'great masterwort',6 Z ]5 l" n) u# ^% p' i/ E' {% h
'4': 'sweet pea',: q' n. v; K' g
'86': 'tree mallow',
3 o2 {5 q- W1 {, ] '101': 'trumpet creeper',- p) m: j/ D6 H* j5 b, B
'42': 'daffodil',8 l; J) N& @: S% \5 \1 f0 [3 B
'22': 'pincushion flower', v6 J# v5 m* G/ n
'2': 'hard-leaved pocket orchid',% Y# W( w1 S0 Q" n8 I. N
'54': 'sunflower',
1 f5 R' P) \7 N4 p3 \$ J( j '66': 'osteospermum',
% h3 M$ q; z {- a# a '70': 'tree poppy',! S# L7 w8 q @% e8 L' u$ u5 H
'85': 'desert-rose',
$ ` K b. \( R '99': 'bromelia',5 b: E& x( @/ r" z
'87': 'magnolia',
$ w! G, X0 K1 y& b5 B '5': 'english marigold',
" Q; r* y; f! G3 Z7 r '92': 'bee balm',
# `0 R9 ^: F9 h/ d- P; w '28': 'stemless gentian', i9 U% x5 w# H# K
'97': 'mallow',
" S4 J$ a* V( ]$ r9 y/ O1 M" \ '57': 'gaura',
" G$ A; m9 t* h7 }: U! H '40': 'lenten rose',
# t/ F) b* r) u. k$ s) l '47': 'marigold',$ a3 j% C/ n, p" k/ S
'59': 'orange dahlia', L3 w o1 N+ l5 c& m
'48': 'buttercup',4 a0 r/ w) C& h9 r) q9 {. m
'55': 'pelargonium',/ B) x3 ~2 a* _
'36': 'ruby-lipped cattleya',
5 Z: Y7 q' g, V9 l1 i8 ^: {# U5 i '91': 'hippeastrum',8 _. v) _5 T% E. G' J: K
'29': 'artichoke',/ u9 \! `3 N4 M4 u* M
'71': 'gazania',5 F+ d8 e, S6 W7 O: T# ?* r
'90': 'canna lily',: ~( L; Y9 Z* Y) K) s) j
'18': 'peruvian lily',
' \8 D3 j3 Z6 H/ `+ x: S+ C' S '98': 'mexican petunia',$ R0 N) G4 q6 T- `
'8': 'bird of paradise',
! @0 [( Z4 N$ V3 j '30': 'sweet william',0 h" W/ @" v8 I# [
'17': 'purple coneflower',9 j- `* D) Y7 q$ q. T9 h+ h
'52': 'wild pansy',& s: z0 @: c# r) S0 ~' E
'84': 'columbine',' ]9 D/ {# F7 r, U+ y) | ~/ I
'12': "colt's foot",
4 Y" F& n0 f( M# |. O4 S '11': 'snapdragon',
5 [$ t( C$ a1 m" y: e5 J) |% A7 | '96': 'camellia',2 A. W+ {; v( J. X/ F" E
'23': 'fritillary',2 K V3 o7 i7 d m" d5 L/ ]2 D
'50': 'common dandelion',6 ^6 d" L5 [% }; w
'44': 'poinsettia',/ W l k' L# i$ E @; b
'53': 'primula',
2 f+ b1 E! u& I '72': 'azalea',
! y( m6 \: ]% b '65': 'californian poppy',. g/ Q3 ^' n: K6 J
'80': 'anthurium',
1 I, T) m& U4 d- M! n$ g% M '76': 'morning glory',
a$ { O4 x3 R: g4 D- K, y' _0 ? '37': 'cape flower',
, L) j# x( r9 ]" r+ _1 I' T" Y7 A '56': 'bishop of llandaff',3 X( a. _& G H9 y6 L0 f
'60': 'pink-yellow dahlia',) M3 l' K0 C9 ]$ i9 z; g, O
'82': 'clematis',9 r# E; d3 I# W
'58': 'geranium',0 ?" g: ^* i1 s" y; l
'75': 'thorn apple',
3 Q* l+ B5 n' g8 B, K3 h '41': 'barbeton daisy',
+ b. Q/ W% n7 w/ c0 @ '95': 'bougainvillea',
( [5 M( K0 c! ~' b" v' p, B2 s '43': 'sword lily',
1 J8 ~- ?* e0 ]/ P6 c '83': 'hibiscus',# ?" B8 d- g7 N( z( T; y5 G+ K7 {
'78': 'lotus lotus',
! E2 g$ [. F8 \7 ^# L8 m; q- f '88': 'cyclamen',
' g |! s* b* w6 _1 w' T+ {. v '94': 'foxglove',
7 `8 p" s! I& h0 G+ G) o* P '81': 'frangipani',
5 v2 h2 O1 m* b4 I) m+ }9 ^& q" \ '74': 'rose',1 \. A5 A9 E" m! v" e
'89': 'watercress',4 B8 x% v5 Y# n( [$ z
'73': 'water lily',
8 j- R! p$ l3 H8 M/ J g '46': 'wallflower',
. @( c- R( _. f- j4 f* D& w6 U '77': 'passion flower',
1 F; t" I. q$ O$ K7 G* W6 G' r '51': 'petunia'}
. ^8 a# |6 A' {5 U- A7 J1 q! O# T! k; f* T2 M
14 H* L: b, N, i5 V9 R( I
2' \, [& o% F. N, ^8 `1 z- {+ I
3' |/ ^: Z1 p; L% t4 D7 x' B3 h
4
0 o/ z# O8 }% x5 ]' x5
4 V# R# \2 x/ v2 o6 b2 n5 E6
) \9 U. e. {5 k! U1 a+ C' j* ?7
, B1 F6 y; E3 F* o2 } q- y+ e8. E. x/ Y' X5 e0 Q
9
/ i7 x/ ?- U. l3 |1 X10
& Z, g. |! c% Z# W4 t11* o' u6 d, ^' w
12
$ X7 I; L1 T! |. i, [13
" n& A+ U/ `: T14
* ^8 D$ u: a9 F) w/ C. X15
5 L7 s2 g% m( i6 p; G16
2 u% \( x: _* a; Q173 f* L9 j- c1 e! l. V
182 ]# F& Z5 _" O6 ]& E! f
19: z3 _+ h9 j0 j
20
8 T4 {- f; n3 `) F- M0 Z, h21
5 F% z2 B+ R6 ^7 Q22
* H# k) S, |. ^9 m0 N. B/ @/ M) I23$ R# D. r. b* {. A1 d* T2 d
24- S8 F$ G8 J! f9 O+ M- N
25
& X1 v6 U0 Q( ~' H/ G26; {$ V5 N! s% j) {5 E1 {" K" O i4 E
27- Q0 o7 ]/ O4 U4 D4 b5 y9 Y- R R
28
2 e* L; ~0 T0 u+ N) Y# ]8 F/ |0 e29/ j: j% [( t( z9 I, {; m
30
# i4 b# A$ B3 e3 R2 u* Y& }$ `319 ?" R" `! n! R
32
, [/ Q. k# r/ q33
R' X9 w. w8 f' }6 ?+ h34- u% L7 C$ J3 P7 H3 J
35
' ~6 f- ~0 ^: Q) j; b360 @( p0 L& b h9 n/ {& c+ f
37% }7 Q$ F# E! m! U" y
38/ `! x0 m" K7 d5 s) k! e" [: P! [
39
# I6 f; W' }+ s/ O( j; m40
% E i p+ v1 v4 I6 p414 o: _' Y+ x* g, W: N7 {
42. S& `: ]4 a' S |0 O" `
433 A; [: b, [/ b4 T8 F* J
44. Z- [ s/ i3 l- r4 r5 n: {
45) W6 G: ?) d* \9 \
46% l- } D& u: I! ^$ V; R
47
# v$ t% j R0 m; q: g0 h! i& ]48. R1 u# O: h H) @
49
9 a: m3 @. w$ Z4 N6 |7 P/ K50
, L. b& o, S0 a7 L' h. J1 A518 ~% V2 O+ z! t$ M3 u+ z
52
! E" Z6 _6 F b0 B- T4 r53. i9 C4 E3 |- ~+ \, }
54
# G0 g# W, O: F9 E+ T* G55
3 e. ^9 Q- `% y) d56; M, n/ u; G/ ]* r/ y1 C+ S3 o
570 B3 q) k* ^% N2 ~2 d* R# x8 s
58
/ A k' l* t; H2 q590 T0 Q6 c; b& Z# E, G, J& m! B
601 |4 `1 z( O& B% Z5 s F0 z3 l: Q
61( w2 J" }2 s% g2 I
628 H0 T' d, b9 X9 b" ?. I6 Q( ~: M
63
6 ^1 P: O5 C4 y: _* ?2 q& J5 \, G644 c9 B# n& {* r5 @# g! O
651 \' R8 r$ L. T0 j3 v; Y6 B' x
66
& X; ~% t/ d. g67
6 @: `" O) F2 S; B: B685 u/ }* q' l. |9 E4 D6 U9 \
69/ u% ?' R2 f5 `0 D2 a1 `3 J1 n
70. t) V4 @$ v y Z( |5 N
712 U. V9 T+ h0 h5 g _, f% \; _3 C& L
72# G+ M: v" |: ?! b$ \
73* C' |* [0 J: M, h) C
74# A5 V- Q- q7 N- r! s
754 |* P: U0 f- E( [1 i2 Y
76
0 D, U# w' u/ R/ Z s/ h% \4 W% X# L779 e# @9 `1 y3 F! v7 f' G
78/ T" w4 k% \& U( I, o! _
793 s0 \5 F* v6 y) D3 Q
80' r7 l' o$ m# z/ G. Q' e8 G. I/ M+ U- x
81) e/ s. q/ Q7 U* `' C
82
! l$ \; f, R8 H6 Z0 \' c+ ?835 e! o( a, E3 s! l3 m" O4 Y9 H% {( q
84* ?/ B) s3 T W
85
3 k: W4 C p$ B5 G8 @. a86. s( }& G. o; P( W" K
87
2 Z$ r" m8 N- A88$ L8 A/ T( a: Z% r# H# X0 b9 E
89+ Q! I! a% k8 ]. Y* \* t" D
90) d+ z( j+ J) {8 [
91$ K& Z# G4 m4 G/ R5 k/ }
92/ L }* f' X, c
93
0 I: _/ u* H1 v9 l5 E; K94
0 w3 H+ G. e8 v: A95
; b3 v* S6 r2 B1 \96
0 k/ @9 P" V- Q5 y' j* Z97* T0 K' m& b' g$ [; j) \, B
98
9 {8 w: o+ R) c9 K1 Y99% h9 N0 i2 ]+ ~4 Z5 C
100: q+ Z4 U% e4 l+ d
101
z) s* f/ g1 q9 r, _102
- W) R4 d8 T" F$ \& Z5 l4.展示一下数据. S& r B, }; v( J: o
def im_convert(tensor):
+ T* m% Q7 N7 Y1 w- W! |5 y, w """数据展示"""- f4 ^# r8 Z8 C
image = tensor.to("cpu").clone().detach()
$ W! ~8 a8 u3 ^ image = image.numpy().squeeze()
/ l. j) h. r: C$ o+ n( R1 N. n # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图, w; J V# ^3 P% u8 z$ T
# transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)
: K) _0 y0 C; ?+ E/ J, e image = image.transpose(1, 2, 0)0 l4 a+ O, ?2 Y/ v
# 反正则化(反标准化)9 E P" d; A$ m6 x
image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))) v" ]9 [2 g6 X4 j# f% O9 A$ }$ _
. K, w2 `8 K/ {. q3 }1 l& b # 将图像中小于0 的都换成0,大于的都变成1- ?2 y* Q4 O* j" W! A
image = image.clip(0, 1)
/ L. ~8 U$ _% H1 ~ V. s4 v& |: n$ @$ q: Z! ^
return image
4 Z9 a8 P# T1 ?9 }1
( s3 b; J% n5 E7 n) o0 f2
5 W2 d5 z8 X2 U( y1 X3
% j# Q% c2 Y8 L9 s/ L, o. e40 |1 E- W" A9 P( ]' ]+ f$ d
5
; e" Z+ B8 p; D l7 a* ^6
4 ^! @) M& D; Q9 f2 ?. y7
6 [+ h2 g" @# ]. a/ r# t* K8
: H" r: n7 _% T$ X9* D( w; y- j& ]4 f3 ?" m
107 D, V- u Y. ]( [9 Q$ U
11
! k: z4 ^( _2 ^/ {- Y: t12
' h( [5 l: I" P' w" o2 M9 ]13
" E$ O& {% V3 F- ~/ O! t3 O C3 B14
7 h9 z" t r& b6 T5 t9 G7 L% c/ a# 使用上面定义好的类进行画图
. f9 G% i3 |7 H9 ]fig = plt.figure(figsize = (20, 12))1 q* ?( A0 E W1 m3 t6 `3 [5 w) q* ~
columns = 4, \$ k$ h; G2 o" `5 s5 v
rows = 2
8 q2 R. r: e4 G8 \, E7 d
. W% t9 p; K; M O* M: G3 Q# iter迭代器
! z9 c0 Z% c$ Z3 B) A$ f& \3 Z7 j4 C# 随便找一个Batch数据进行展示
B: B7 Z; \( i! ?7 q: s2 g- F6 Sdataiter = iter(dataloaders['valid'])& h+ k$ C0 X, O% ]" K" d
inputs, classes = dataiter.next()- q# z P% j' @7 @
# N( F! ]; V4 Z
for idx in range(columns * rows):
7 W r2 n, W6 _1 N# s. c ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])
6 x0 I* Q8 ]" P5 H9 {" G6 ` # 利用json文件将其对应花的类型打印在图片中
% X2 F/ S& b# |. e* _( a+ N$ T' @ ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])
! h1 f* g% @# T plt.imshow(im_convert(inputs[idx]))
0 W7 G; ]! \1 [ t3 Dplt.show()# Z4 F; w6 v, Q
5 I/ E- }* ]8 H% d7 |7 c
1, x6 g) G4 U7 G" \; r- m! X
2
0 V- m5 D. k, Z5 _6 j% K v# @( l3. T/ G) m, J% B2 {$ v7 v I
42 Q7 }& t- P4 }$ b! I1 W
5! Y" i: \1 \: K
6
9 h) H R: Y) m! W# d& `7
# n4 T4 T. ]$ h6 F; f# Y- }9 l8
5 c& ~( K& i V9
# x/ H2 ~& P0 N; O. u0 g# Y10
2 L4 m; }" N* } N3 S7 i3 b/ l11
- y& \8 N' |1 S+ a% F$ ]( V12* l1 j7 J/ Q# | O7 `" R+ A+ a
131 Y, h$ L- d9 I8 Q+ p1 ?- ~
14
( Z- o0 [+ s, B" {15
* C) V5 J# S/ k: _16
! q# F. d B. p; N$ W, |) p5 P% \/ z$ M# m
5 f' `0 n1 m3 f7 N2 r% l
5. 加载models提供的模型,并直接用训练好的权重做初始化参数
: U6 O: p5 h5 I5 ]model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
- N2 V, F" C# w9 g7 [& z6 J. G# 主要的图像识别用resnet来做3 x6 h& ~. d5 A9 `
# 是否用人家训练好的特征, ?% t' H! b" R
feature_extract = True1 N9 {) b0 Y6 c& K2 z/ G
1
6 s+ G7 r" Q: Y& m' r$ n5 M w2
- ^2 f7 ^" E& n/ `3+ v9 J, ~# J/ U" e$ ^( \
48 D, E) V" n0 m: X. [
# 是否用GPU进行训练
. h% u9 b: z- x3 u2 c ytrain_on_gpu = torch.cuda.is_available()
a/ g2 W/ {+ {) q1 x* J9 r) A4 N3 w
if not train_on_gpu:" o& C0 w* n3 y% T# a" F
print('CUDA is not available. Training on CPU ...')
& d* ?7 F8 T& L+ q% D1 A- {0 _$ b" yelse:' V# z. L3 s1 L) ~* H% [) _
print('CUDA is available! Training on GPU ...'), ^5 X6 e) W& g! S- s5 c
1 U7 z: B" s' g. l# H# W
device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')& l# j3 R/ Z! M. n9 n" M
16 a) k* L5 j A) F3 O
2
0 X. U9 o+ H/ J1 K! X# ]33 f3 W5 @. B l2 {7 P7 b
4
9 Z# l/ c; w1 ?. n. X5 N; P. r: K& p56 B( U- W G/ T7 h" K8 W' C1 }
6
0 [8 E: t; ~. f7
1 p, o/ l; ?! @% I' L" W8' F; O) f* Y k# M
9
3 f; ]; @( G2 V4 @. e8 A: xCUDA is not available. Training on CPU ...$ x- f& G/ n, W) S& l% G% o
18 Z3 r* l7 N3 d9 N6 L( g# y6 [2 |
# 将一些层定义为false,使其不自动更新
1 z! i- i; J/ a: Y4 W: ^def set_parameter_requires_grad(model, feature_extracting):
6 b2 V/ {# M- ]( g x) q1 i* c: F if feature_extracting:5 l7 D S' `6 j! `: e
for param in model.parameters():
9 u9 _1 E9 K$ d# U) q& ~ param.requires_grad = False
R% k! N- t- V |5 U5 X1
, Q. z! a# b" ^5 }2 f& x' q0 X- d24 D$ d# M6 M5 l4 m1 T
3% Z; I/ X K w% }# e" @
4" f6 A0 D0 U! L+ B% _; f, F
5. d9 f/ `/ S% m
# 打印模型架构告知是怎么一步一步去完成的
+ P9 @ c; w- c. g# 主要是为我们提取特征的, A! w/ \/ O: K0 O
7 s: t/ N4 S6 h8 h2 E4 } cmodel_ft = models.resnet152()
4 k9 G) N& O3 \& d% ^+ Jmodel_ft d' G( ?: v" l) F
1
: L% h d$ c, ~) ?. A2' H. y% d. b) { T$ {
31 P3 Q/ U" ~) G8 ?3 K
42 i5 A2 Q, J" ?9 B1 `
5* ?: U3 Q5 X# H& P
ResNet(
( o- z& r3 D& _' J1 A3 B (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False); P5 R$ F4 ^# ]9 \% n
(bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
) _) M7 k6 y: m# O( e: J (relu): ReLU(inplace=True)
& Y3 D; I `. q% E3 K6 Q (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)6 V: v' O- U+ ~/ Z5 ^$ }4 f
(layer1): Sequential(
! A4 V8 \. o7 _8 X9 ] (0): Bottleneck(6 N3 c8 ^% |2 \+ `+ R0 |
(conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)
1 F: J2 x: V4 z6 }/ f (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True), u8 }* {+ C( {1 Q# Y) r
(conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
0 ~2 o9 ^" o( H2 c3 t7 w (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)0 t) M$ ?, ~+ \; V
(conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)6 E) R! X2 @' H- [1 t
(bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
" ~. G B; `8 ?# X' p (relu): ReLU(inplace=True)
/ h( ^9 y7 q/ i. v (downsample): Sequential(, P# p! F. ~' v; W
(0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
9 E! ?7 h) m( t4 Z) m3 H$ l' `6 _ (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)7 O$ g! @1 v2 q% {/ Z h
)+ I0 ^" H* N2 R* l
)4 S1 \& I/ `9 o/ k! ?4 z8 c4 |
中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。" ]! c' c, w2 A: e( n* A/ C
(2): Bottleneck($ d% I( X" y+ m1 m
(conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)
- n6 G) @4 w) i W8 C8 ^; h (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)4 w4 Y1 R- J& e: n
(conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
) }$ H n/ Z4 v3 v1 \1 e2 Y& B (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
5 ]7 B+ W* ^* J, e8 W3 T7 @4 i (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)( ]( P, N2 @# Z5 J+ g9 g+ A
(bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
6 \8 n* X' ~0 X1 P) Z- M4 a (relu): ReLU(inplace=True)& t& l- n5 Q1 p4 W: U
)
! {% l% Q g7 e) s" ?) |& x )) m7 e4 Y0 S. r( _% e9 {& C
(avgpool): AdaptiveAvgPool2d(output_size=(1, 1))
0 z6 U" v2 {' z- x6 e7 I (fc): Linear(in_features=2048, out_features=1000, bias=True); u- }1 g4 E+ R4 L- M2 {
)
* `# u! T- i6 ^6 m0 }% }) M1 K$ l2 H
$ N; L( K) A3 |$ T15 w# ~3 `+ O: }
2% i, C4 D4 t- w! @+ G" |
30 u) Y7 w; |, C2 i
4) T3 v0 C$ i' U; ?
5
6 S6 |" g: O# P+ Z6+ {( k: J+ k3 F- I0 x
7
+ J: R7 f% V& u, T! q8
- e; U; `6 R& k4 q/ o- c- E) j95 N8 `% }. p5 H
10, ~. B$ u6 _9 H+ K% Z
11! O# i7 P. n) d- u0 m' Q% }- \5 N
12: b+ V: `5 }9 \' {' W& E3 V9 N
133 X- U0 p z9 T6 W* P u: b5 |* x
14
8 {8 U8 K8 D( G1 H6 j q$ k15
4 P0 J! [- d# n16
7 K1 o E( V3 V175 G/ [8 C9 p3 U- i! S/ Q2 f, z
18 G/ E5 p8 u) `
19. ^( e/ s2 v: U9 e% V
20$ l9 S% w/ r2 `/ t7 P% O8 d6 l
21! h4 {9 X. ?* g& J3 S V6 a# l7 @3 }7 h& W
22
q5 b+ Q6 r. p* [0 K236 V; N8 i6 y. _7 \ ]+ W6 z/ \% C3 e
24
% _5 w3 A. a! `! I" J9 e: M25( R& U% ]' O# ]' m" ], d& B
26) I' E; q: x/ |9 @- n
27
& e! i; P1 ^7 P28
1 e4 S* `2 @( Z29
& r1 w: B7 R7 |, k& v) ?% d* b305 d' [# \* b9 G0 h5 P* }
31: w5 }0 i' e5 m, m3 x2 q) b9 ]4 x* U
32
4 n4 W+ {% x# Y7 y! l* N2 a9 Q33
& ?: c6 S& r, L1 z最后是1000分类,2048输入,分为1000个分类, W( ^; L8 g$ Q, m
而我们需要将我们的任务进行调整,将1000分类改为102输出! ~ {; x+ H6 W8 f' f. D
+ J4 q& I; S# a# L# C
6.初始化模型架构
A) N! G- W* k) t步骤如下:
1 f( G s2 X/ e$ z; ~/ w; ~3 D
5 X, ]+ h. T5 t- G+ t将训练好的模型拿过来,并pre_train = True 得到他人的权重参数. Q) {' S7 i; ^( i) G5 F
可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)* G/ W/ e- y' s; X* z- k' @' [
无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数7 {3 b; `0 }6 C9 R! a
官方文档链接8 r4 k0 t, U: B5 }, B
https://pytorch.org/vision/stable/models.html
6 H: {0 r; i" a4 @% @* G
* t' E0 E1 R, v( g# E; o# 将他人的模型加载进来
+ _2 {1 |. h, _+ F' s. I/ L- [def initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):) J, N/ J( i: a1 Z1 Z: T( a
# 选择适合的模型,不同的模型初始化参数不同* f$ F& s2 ~7 ~# B+ @
model_ft = None
8 w# j0 S2 A/ a1 H* F- K! k. v input_size = 0
: D+ ]4 V3 i4 i8 l" o
$ m- N! o/ y" ~, V' \ if model_name == "resnet":5 q" H: b% _: R6 c! c9 L3 }% G
"""% X, `7 n" p9 H" _: n: m4 M7 c
Resnet1522 U+ B- Y4 ]- k- ?5 G
""". o$ d+ C& n+ I1 Y8 X
1 n, ?9 n l: Z
# 1. 加载与训练网络( `: E5 ]. T' o3 {/ J, r
model_ft = models.resnet152(pretrained = use_pretrained)( O8 j* a- H6 j' V" O0 G$ L* d
# 2. 是否将提取特征的模块冻住,只训练FC层
7 t! k8 H! V' R set_parameter_requires_grad(model_ft, feature_extract)
6 h( q6 Y9 E" V # 3. 获得全连接层输入特征" T7 Z6 u$ E3 b) c6 _9 U; o
num_frts = model_ft.fc.in_features
+ [" I, _2 f& a3 e1 t: X # 4. 重新加载全连接层,设置输出102+ P9 v# B' r# p/ t" }( H
model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),
- V4 }& A& z8 P/ y% A nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1" s* r6 X0 Q8 P* ]4 J& u7 h# k2 f
input_size = 224
% W. x* T# }8 T. s, n3 U" H& s' D. I. Z$ k1 i3 u0 W# ?" F7 c" T
elif model_name == "alexnet":6 s; D, |( K7 `$ S6 n4 v
"""
, a4 i8 O: P" q8 _2 P0 r Alexnet: S' v7 p5 }3 O) b7 b
"""
2 _' _3 L( a# r9 [& b8 I6 F model_ft = models.alexnet(pretrained = use_pretrained)5 n' q: t2 @6 O6 M. A- V6 ~, m: c
set_parameter_requires_grad(model_ft, feature_extract)
; w9 H- G' _" d8 r7 O. E
) ~( x0 Y8 {) h) B! ]/ B: } # 将最后一个特征输出替换 序号为【6】的分类器
* V' j3 m# `8 b* A: z num_frts = model_ft.classifier[6].in_features # 获得FC层输入) Q) T* o, d* ?4 @( y3 o+ T
model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
4 o( k" q! ?9 l, A; n input_size = 224
/ Y1 v$ a6 J0 T1 n+ |. ?% a6 T" m, N& S6 L0 H6 Z- P* @
elif model_name == "vgg":, @# X- _7 C( {. @$ q, M E
"""% i+ ], m4 k7 m
VGG11_bn
3 ^9 N5 L* f3 h5 i Q """
$ u8 b! O; I- D& r model_ft = models.vgg16(pretrained = use_pretrained); k/ b% a8 j5 i9 G% `
set_parameter_requires_grad(model_ft, feature_extract)
' O2 A! f9 U. E2 y num_frts = model_ft.classifier[6].in_features$ b( h/ X; D' {+ @* k# e. p8 D
model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
6 c9 W7 V; U6 D2 i0 n2 ? input_size = 224
7 `+ G6 A7 L3 g+ Q: P, K; C# _3 I# t8 b" b3 P4 E7 N% S" u
elif model_name == "squeezenet":7 _8 h8 ] g n- }) J
"""
! d! |2 `: C& g k% ?3 o Squeezenet: D0 M3 r( s |- T) c
"""
& ]) x0 P: ^4 c) [ model_ft = models.squeezenet1_0(pretrained = use_pretrained)
: R' a9 |5 t. x+ z6 T3 Y. |1 i e set_parameter_requires_grad(model_ft, feature_extract)* | `, y! J: D; j" t! A
model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
: ]- E! O; X' d* b$ b2 p model_ft.num_classes = num_classes5 \+ b6 f) Q% b3 G
input_size = 224
% a0 W, u, A2 A* D9 _. e' p& Q
9 f) H/ K% G+ b6 i: O/ S elif model_name == "densenet":
' _+ Z% n5 c7 C: V( y+ H """, o/ O- ~* ^# w5 B' B/ K% A/ W5 k
Densenet' L2 ` U- R/ o7 Y5 N! x& B: M
"""" T. f9 z+ ] `6 A
model_ft = models.desenet121(pretrained = use_pretrained) f2 s. n2 e9 W& r% e5 Z) _0 g9 P
set_parameter_requires_grad(model_ft, feature_extract)
2 S T8 R* o5 X5 Y num_frts = model_ft.classifier.in_features
: W( s2 y4 V8 o- g2 ] model_ft.classifier = nn.Linear(num_frts, num_classes)) q- S4 p9 S$ W( S' c a" q
input_size = 224
9 o8 y$ n p6 O7 c' |- b6 |9 ?, w/ i2 y9 I3 c& Z
elif model_name == "inception": ?; S* }0 s- g$ z/ R
"""+ x" M+ h; X( P, S9 `
Inception V3# X" n% |) `: [. _0 [ g$ e* h
"""
0 e$ E4 p$ y' |+ ~- [3 C model_ft = models.inception_V(pretrained = use_pretrained)
! T# w/ p* F, i9 u5 K set_parameter_requires_grad(model_ft, feature_extract)% [) B0 X G2 Z& n# N
' W& h. |, Q* V0 M7 y0 [ num_frts = model_ft.AuxLogits.fc.in_features' J4 l) Y' F( E& P
model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)4 I+ g) }0 ]. e2 ~1 d, H
7 Y: k7 J7 d5 I' {% ~0 C4 [ num_frts = model_ft.fc.in_features1 S( v1 U/ v; e! ^. N7 Q$ C6 ~" j
model_ft.fc = nn.Linear(num_frts, num_classes)
0 n4 x+ V% G* b9 R input_size = 299
, W1 i ?6 r; W7 |- v8 Z) A9 D1 M7 f0 [6 |
else:
1 y+ N: }, N* ^% r print("Invalid model name, exiting...")
$ M* K) g, G% k7 ? exit() { O2 L6 {$ p! s! E
( j& e, n, H3 {8 Y
return model_ft, input_size3 h6 W+ X0 f5 f4 U0 Y
- _# w+ t. y+ l* W6 w& C) \
1
" ^. N9 [* I6 ]6 B. C2
9 b+ M/ F/ h: C8 H" c3
' q' a( L+ U" a4+ h+ z; b: C1 e( {
5! } i2 }2 J& G& C% l. B. E/ x9 ~
6
* W0 x: h+ e k, A7
) r! K/ `/ M& F4 q1 T8% d4 X) _. i! Q1 ?" i7 j- F
99 T6 d/ W6 A3 w! J1 ^
10
0 d8 e" ?4 p* M. m5 O d11! z1 Q8 c, @* S7 a3 |- W: ?# [/ h
12
% l) X# b- R' t132 F+ J# D* D `- o0 j8 m: J- n4 ?
145 r" _. |/ X" O. k6 M! U; Z, {* z
15/ x1 P9 |- D2 X% L; ]
163 s$ y s! w0 I
177 J/ y2 G3 Q$ g& C# n; z" A9 J9 O# [
18
2 |5 E N( N0 L( Q/ u, c$ u8 L19
& [: M. g' w+ x7 m20
9 y! }; t' c f3 g& g+ C" u217 Q8 _3 p! \9 x3 E! k7 ?$ O
22
3 X) o$ a* @( p$ s( c* V23
+ T3 E2 h% F) s5 _241 V3 w) d, t1 n3 {) n. D
25! B& |1 @! H/ L6 E! u/ f
26. J1 I% ^/ ]4 ~+ S. \
27) M* d6 n. _1 }, ^7 F, A) @
28
( O, T- E4 z- H29" p. m/ L- M9 r" \3 ^9 d D# T
305 ]" i. L7 M- k# X5 l( k
316 \, r3 n# T8 f9 h( a
32
- A/ f% \& T/ j/ \$ H33
# R7 b) E4 e5 N" q7 F# [; N34
$ W, N% U1 g; G, d$ a) o357 B. v \. W* @
36
8 S1 L) a: [5 t& U37
& C! N8 ~5 L: p* }* z) R2 }% L0 h384 G9 h- @6 Y) |: Z9 ~
39
M7 X# o$ H! P* m+ \40
2 V E& u I3 U41
$ _# Z. t; w) f; A# A1 x42# g5 W5 Q) r/ D
43
8 J7 T' \9 ~5 e0 ]8 A- f444 ]- R0 [2 i- R$ y! q! @
459 A( `( f6 x. L! [+ E$ a
46) R3 V! |! H8 ?6 F& g
47
[" M1 r' t+ }. ~6 `% c1 G48
1 ]$ B% M1 h$ O% ?+ \3 q49
1 A5 p' `) F; c4 i0 j50+ k, ?) \3 B7 E1 ^& u$ p4 n9 a4 Z+ t
51
( e& W" I( j8 U$ M5 ~0 E9 F6 S* H52/ C* P* S u* S9 {' O9 K, I
53
- k$ E2 H6 C- a" G* I' e0 Y540 s; B, q8 w+ [. g' y
554 I8 k/ k7 F3 Y. R
56$ ]( v2 f+ w, `0 h, Z: k2 ]( C
57: t5 b% N" z0 x. q9 r: A
58
) o. {6 p4 n$ l( D6 `& C" {. g599 M- E; J7 i' M( A
60
6 h: T8 G0 t; s7 M1 n9 i61
( f- _; F _5 c3 `62$ n5 @0 z# E$ @& x" b
63: m; K4 R+ T8 V I7 U- \* D' K! b
64
# b+ z( M/ ]8 O' o4 k1 ^65
0 Z# B* p! ?+ I J66
2 D9 h: S x R% y67
5 q" Q3 d/ C+ J M68
0 b) Y9 u0 ^1 J$ Z! ?1 I: }. |, m69$ ]- m' z! C5 C2 i2 r
70/ d# S4 d# ?1 `$ o, v1 I
71( ~/ ]; i, D1 j* N7 V" |
72
0 [3 Y, q: M/ J& x$ q# g. Z i73
" q! f$ f9 z# J' S74' @, l7 f6 l1 L1 C+ A
75
$ X \: B4 I$ O' \9 E* F9 f$ J/ C1 u76 l7 f) y% @, }% b- w7 `0 J
776 `3 ]2 J! p( E% O$ @7 k2 Y" a
78: c% }* X3 f# }8 \' I5 g, c
79
/ H* ^4 m: w1 Q: q% F' [( }3 Y- l80
+ H4 {- Q" N' q m4 l81* L4 L/ N) ?6 i Q0 h$ K% \5 Q1 D
822 f4 w9 S- Z' e( M
83
" p* S/ [+ [2 p; G7 A7. 设置需要训练的参数" }# _' U& B$ a4 C, H @5 R
# 设置模型名字、输出分类数
# k' x8 [* D0 wmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)
& }2 d) a; X4 y; y" E$ Y: H5 S! b3 u# V5 A5 Z+ F% K
# GPU 计算5 V& k7 N& `, A; B& v" N9 o' K( J8 K
model_ft = model_ft.to(device)( {5 s- \9 I) |2 r) S6 }4 W
+ j' S. b) @: @' v
# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取4 A) h: I$ Q2 i0 O6 _& Q
filename = 'checkpoint.pth'% B7 L+ w( \2 ~( I1 _6 a+ R5 U
/ S u# s' y: Z# z# 是否训练所有层% F" B: D# @. Q3 ?4 p- Y
params_to_update = model_ft.parameters()7 z: A, | u& I: P1 v! e3 K0 l$ @
# 打印出需要训练的层
' t* ]: c- }* w( d, L/ o" Tprint("Params to learn:"). A- D! J% w* x) ^
if feature_extract:
& O, w" R9 K3 D- ~4 K" F% z params_to_update = []& O% O6 W, [: t3 g0 n
for name, param in model_ft.named_parameters():0 a$ w/ p' y: T
if param.requires_grad == True:6 [- W$ W! I0 V9 B" t# g
params_to_update.append(param)0 R. x2 g+ V _% N. F% E0 c
print("\t", name)2 F7 e4 v( h B! m) x
else:
- \' k3 k# T: h* g5 ] for name, param in model_ft.named_parameters():
8 ~+ r9 Y; K* R+ d if param.requires_grad ==True:, H" m" N M, [7 V2 k$ U* K
print("\t", name)% M# i- J# {! I1 o
4 [6 z! J9 L+ Q
1
/ @3 Q* f3 T2 P- ^# @2& }$ [1 ]# a2 w
3$ G& c8 x5 T* U- b" y6 F1 ^+ `. D
4% y, x2 `% P, L7 R/ n2 R. M
5
( I, S. r/ s s( M2 ~2 q9 b0 ~6
f3 ^$ [) V2 ?, ^6 j1 ^8 E7
* E) _( \; V+ ?8 K$ F3 |' x# k! J- f
9
. x$ b- ~( T- g+ Z10
( f( ~! a$ S- L5 s' D11
! u# a7 X8 r& \# R$ m- L12
* q a5 t" N& G: }3 Y- N13
/ u; ^( E8 R$ j f: ~140 f+ a, c1 |7 R0 t0 L
15
7 Y0 M* C: v- V) K! g16
& M; q3 |% ?$ F1 R: A174 ]% a8 w0 g+ b _5 C; Y) U! f7 q
18) X ^5 P m6 ?0 V0 b/ J6 N1 Q( N
19
. F4 ~, h# _7 _) @) J3 f203 X; K7 X9 Z V2 P% C# j5 D
21! K7 C) `! `" x3 ]1 t" A! {* V" ~- D
22
) V) K0 J9 J/ ]" c$ _23
+ t8 T2 W8 M- V3 T$ u A' NParams to learn:
# `4 u, G$ P) g. M3 y7 `& m' c R fc.0.weight `% d! L5 W( V4 r
fc.0.bias& k5 I5 _& p% ^$ T+ f& m, R x1 ~$ X
1
7 M' C! P s1 R2% R$ f; [2 z3 U4 A- w: f3 |% W
3
2 p8 N# Z, _2 e8 D7. 训练与预测# Q) O* k# o3 d2 S2 o
7.1 优化器设置
$ N6 ^* Y4 N: u; N! \, J# 优化器设置
/ `! d3 i# \3 o( Aoptimizer_ft = optim.Adam(params_to_update, lr = 1e-2)
, A( K5 [, ^2 j, H! w# o# 学习率衰减策略# n5 Z6 ^# M$ }: f+ E5 q
scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)
! t4 b4 Z, p/ e# 学习率每7个epoch衰减为原来的1/105 x: Q8 P0 I; ^3 ~% q
# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算, Y2 [# U1 B) g. D% u' o8 |1 W% k7 D
2 u# _# L9 `$ B v0 w* H% U
criterion = nn.NLLLoss()
; @& t# q1 g+ x! M/ _/ X1
; t1 f& X2 C. ]- P2 E; G, L2/ u0 w7 q- l! @ X& L; H* C' _
3
) h+ [: V# f% V& d# T1 H z4( [% b/ R$ q# X: R) z9 \0 f5 s0 {
5: O. J; h- N4 }9 o. f6 k+ ~0 v/ O
6, N! B( `" B8 j
7
4 H6 |6 h: }- w1 T7 Q8
- `2 m; R8 B( o1 K9 y5 F$ t+ F$ d# 定义训练函数, @& Z, U9 x* c4 ~6 J, h# u" z
#is_inception:要不要用其他的网络
1 o6 o5 O4 W* S. X5 xdef train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):8 r& z0 F8 t( m7 N
since = time.time()1 y7 C7 j& `7 f1 O
#保存最好的准确率
5 q6 m* R6 H. }9 [ best_acc = 0
4 _* G" X% w h. d """
7 o. B' c2 W2 [ _/ a% f- a2 Z checkpoint = torch.load(filename)- M' B4 a2 R+ w
best_acc = checkpoint['best_acc']7 t2 A5 v7 c: N
model.load_state_dict(checkpoint['state_dict']). t& k @* S3 m$ a! D" K
optimizer.load_state_dict(checkpoint['optimizer'])$ L% }; C, _3 a/ O0 X- p0 s( b
model.class_to_idx = checkpoint['mapping']
6 B# G& s; t( [ """
. }! B4 J; ~) t5 e/ _1 a, f #指定用GPU还是CPU& r# g) @9 y" R" O) q
model.to(device). v q( Z/ J! c; O9 O" g3 i& n
#下面是为展示做的
; y# F* n- L8 s6 v( N" S5 o: C val_acc_history = []
$ y2 F, g. A& t; { train_acc_history = []( j! f1 f8 ~8 Z" Y
train_losses = []/ V6 O& F# j4 U* i% R4 y
valid_losses = []
4 q* c, |) s: t) |4 e$ J j; D$ I8 R6 v LRs = [optimizer.param_groups[0]['lr']]9 k/ r' r% O7 C p. s
#最好的一次存下来+ i& a0 t: I! G4 ^
best_model_wts = copy.deepcopy(model.state_dict())% s. ^( J- Z+ h/ I# [7 K7 W- p
; u4 b' z) D F# g; A6 B; [ for epoch in range(num_epochs):
1 j" S3 N+ Z/ r, K- t print('Epoch {}/{}'.format(epoch, num_epochs - 1))1 f( m) l0 p, l! t9 p" q: {
print('-' * 10)
3 c& N6 y1 C4 _) D! ^
6 f; A4 i% J i# e$ O8 n # 训练和验证
2 p) \9 P' x/ H- [ for phase in ['train', 'valid']:
; {+ J* U+ a6 i y$ E+ v1 P if phase == 'train':- J& W8 q& [+ [" Q% {% A
model.train() # 训练0 J, C& ~: C! b. Q d; A
else:% J: M( t7 R/ {+ \0 K4 T0 b
model.eval() # 验证
$ W4 T, Z8 d: M8 p; ^, l6 F
k( a8 c5 p; w running_loss = 0.0: k8 |; y0 H5 G% x J3 r+ [: Q0 U
running_corrects = 0
* i/ ]3 E, ^+ u0 n) \0 `, k
% h G) l: u3 h k/ i/ ]. d # 把数据都取个遍6 b; h, d. L/ f9 _ V* `
for inputs, labels in dataloaders[phase]:
5 w: J! e, u& |- m3 R8 Y0 F6 [ #下面是将inputs,labels传到GPU
0 N; Y8 I/ V& G+ q6 h( _1 H, c+ Z inputs = inputs.to(device)
7 i7 M4 b4 }) M% S8 E9 W! s labels = labels.to(device)0 {: V/ F1 v- s" |, u# A4 Z
3 C. X8 v l! o0 M, M3 a
# 清零+ g! }2 W l- d8 W5 G6 z
optimizer.zero_grad()2 a' M$ R1 a3 P" f; C
# 只有训练的时候计算和更新梯度
' U0 d* G% K- D with torch.set_grad_enabled(phase == 'train'):. t4 o2 W; s& J* Y# Z7 Q
#if这面不需要计算,可忽略+ a9 Y; o& v+ E! f
if is_inception and phase == 'train':
3 u8 ]8 \2 @6 A; a$ q7 [ outputs, aux_outputs = model(inputs)/ P& ^; K, ]; P7 l7 h# Q! U# K
loss1 = criterion(outputs, labels)& t; M; S: {! R7 S% q( }
loss2 = criterion(aux_outputs, labels)- d4 [- e) f: W( F ?
loss = loss1 + 0.4*loss2 W. z, K6 z. m6 A- a w; v3 O" b( C* y
else:#resnet执行的是这里1 L3 X: [$ s# f0 Q( U. u
outputs = model(inputs)" F9 b% _- S @, e- p! F
loss = criterion(outputs, labels)
% y2 d! Z5 C' f6 i" x+ O6 _1 u, ?: r+ [5 a8 J" W8 e
#概率最大的返回preds5 x' o$ G, r2 U' \& b
_, preds = torch.max(outputs, 1), T& U$ V5 u/ x9 L+ I% G2 F
9 h. d, O& v! G- i; X+ A # 训练阶段更新权重
6 Q7 g$ i! ]4 W2 j+ ~ if phase == 'train':0 ?1 `# P, B: a' g# ^
loss.backward()
! ~: C+ X! t* D$ g; T optimizer.step()
/ g' s( w/ Y ^' B- I
. L3 k+ R0 I- j # 计算损失# }$ n. \9 D- V X
running_loss += loss.item() * inputs.size(0)
& I+ _8 R/ w1 e6 R' ^# }& {9 T3 Y- c running_corrects += torch.sum(preds == labels.data)
9 T4 X) w9 k" [' x2 R* ^
5 m7 W$ I& ?$ n9 |% q #打印操作- B) V" ~, j& e) s
epoch_loss = running_loss / len(dataloaders[phase].dataset)( O7 r8 |; a4 o9 t7 a
epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)$ H' D& T" e- ~( c) A s# J5 x
& G, g. ~, D' @7 P( K% j
. @2 A3 V/ A9 ~: V time_elapsed = time.time() - since! h0 a5 l h0 E, K C, x; q) h ^
print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
, m% i' |) g5 a. V. C; I1 b print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))
: _ f% R2 [" H# F: r7 C/ t" \+ N; T- P6 m3 O
+ D J4 b2 h6 d* s9 F
# 得到最好那次的模型
) j+ D/ ~, {/ H: w if phase == 'valid' and epoch_acc > best_acc:5 m5 H. ^8 \6 V+ l& L A
best_acc = epoch_acc* c9 V1 \: @6 \) R
#模型保存! v8 V3 u7 h, w" F. u
best_model_wts = copy.deepcopy(model.state_dict())2 }+ a7 L8 |* w( `$ X" K
state = {
3 p& K- y9 M! | #tate_dict变量存放训练过程中需要学习的权重和偏执系数8 k. i& o( ]+ G5 H) _$ i3 ^" U
'state_dict': model.state_dict(),* }' R* _3 p1 d: G3 s
'best_acc': best_acc,
2 U: y: l/ r5 R, g 'optimizer' : optimizer.state_dict(),' R) B, W' w# A, _
}- O3 G9 O# w, u# b3 I
torch.save(state, filename)
( I8 `2 P. K- F! m0 c5 e- M4 h% S if phase == 'valid':
$ S. n* n) o& y0 i- Z% {2 D val_acc_history.append(epoch_acc)
# X$ `6 P% o: u" u valid_losses.append(epoch_loss)
- [* p C9 d2 m" ]. F2 G scheduler.step(epoch_loss)4 H% k% S( a* K1 v2 c; B5 U+ h
if phase == 'train':* R7 m Z. a% P) z
train_acc_history.append(epoch_acc)
5 ?$ x1 b- t Y D) s2 m9 i' f6 d train_losses.append(epoch_loss)
6 _/ M$ ~6 |$ ~. y# w6 y
+ Q3 G" m* J+ a! o print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))
* @ l; e2 O/ D7 d" Y LRs.append(optimizer.param_groups[0]['lr'])
3 O- ~: M; J5 i( e+ l print()$ n( ^+ F8 q" O
9 V. p0 k8 M( w7 a# s
time_elapsed = time.time() - since
5 ~4 E# T2 i9 f M4 A; t& ~. c print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))3 X( `7 _ a E+ A* d1 T) e+ A0 _' A! ~
print('Best val Acc: {:4f}'.format(best_acc))
! Y+ |; J5 U ?- ?2 I. G: T
' K' J4 x% f5 u' E$ f' l8 M # 保存训练完后用最好的一次当做模型最终的结果, o5 x/ p$ C$ N" |# U
model.load_state_dict(best_model_wts)
* i1 ]% h' C; e$ ~2 _4 X O8 F return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs
( e! J# T# f1 K" d# J( I8 j7 n/ r' t
8 A+ ]' q( V; A d$ J
1' S: p/ R( j, C$ r* I/ K' H5 d5 U: W
25 c9 Z+ G* |1 p' Q
3; U$ t1 e) \: A7 T( s- C1 W
4
% t4 e/ w" Q8 b" ], Y4 y6 R7 ^: C7 _5
m/ O. g0 {. n: L& Y% {0 _) U: J% h6
3 F8 v# G% D. ?. V4 y6 K- Y" D7+ Z5 x4 M3 \# @3 }' `
86 f- B$ i8 G; q& o
9# p) K4 e, c. U ~$ _
10
2 z' l' [7 V: @) Q11$ P3 d! D4 V4 N- c5 j4 M+ Q- O
12
1 n, V1 a \1 @; `13
4 @& k; W+ ?0 t0 q) L- }- t14: K. g: I, j- X2 s5 D( d5 }
15
$ L/ p& W9 g6 k% f* y& R9 I' w. w, B16' U3 l ^' a, q3 J' K# @
17! H; Z. o V' Z& ?& d2 f$ P E
18% e M. f$ B5 t+ v B5 z/ S3 p
193 H* ]5 i( q- v% R/ \
209 J [% \% X9 I8 k0 I L% t
21
9 k! ?: r4 f- |; \& X+ j3 l22
; Z# {; g, V' g" w/ c; Q23' |; I' O9 U i# g$ P' e# ~
246 ~8 d9 b6 o! {* O! ?% T
255 S0 q; [1 T: P
269 [2 F0 G1 |0 Y. E8 Q. n+ T
27
* v5 W1 Z+ u* g5 g0 J28
3 z4 u+ h) V( M29
8 U5 a* ]8 V ~% c( F307 [( j+ X' A. e- y' Y6 ^
31: O- O$ n& B3 N# b' ] D% {+ Z
32' \' j* s9 [+ P8 Z" t% ]9 ^! X
33 r. n8 f( h$ R; M9 r- }
34' a2 `4 S; }& _
35/ @ L2 ~- K) m/ e: t( M
36
6 \ \$ Z! L% W, R37
1 f( I& b/ d6 g9 o8 f+ H* w B' ^38- k5 Y6 @. }' F
39) i( p* H5 U0 h0 ^% v! F
40; X4 `$ G3 V: w& E, s8 N3 h" Z
413 _2 f6 I) k# s6 H" W& v2 U$ g
42
( H* }9 a, L6 c43
: `. ?# Z; T; s& n" H44
$ ], C5 B( I# m0 K7 I45+ s, T G; y, l/ A0 M$ T4 }
467 Q& [7 K$ {6 N+ r5 a
47+ N9 J2 H, k4 A- I- m
488 m' d4 @6 F! [9 k) x
49# q* h3 k5 `* s. Z
50
r. O4 e: u, a0 c) G3 b" p51& a% u" B2 u3 \2 K) ? p$ G
52$ {/ } h8 p0 [, t
53$ S' N* g5 `# }
54! `$ b' @2 T! O) M
55# p( {8 f/ ]3 o' n
56+ N% S- g' O7 L
57. |) }3 K( S* x) {0 G, z1 y+ p
58
' p7 S: s" O9 y1 S' ]" N, }59
% l; t) E4 p# |( b606 }8 v' d/ X0 o: d, u+ \6 |' Q0 o
614 B8 l3 O7 E$ a6 q1 S
62& |! Q9 \0 E( B9 E5 a( C' [4 I7 M
63% U) H6 Y4 q6 H+ t+ r4 t4 S' d; q. H
646 t: f# \8 U# y* A
65% v* f0 @" ~' [& p
669 B& \5 V7 ?: G' Z. A
67
# @' s/ s6 ?* w4 [( m0 e: Y4 X* h68% p( _0 P( @2 _% `! X
69- o1 o$ g* ]" R9 u0 }8 `! O' H7 `% d, ~
70
3 h, m8 J: F" x5 U3 T0 x* K71& C7 l4 T/ w% K# V) B" T" ?+ {
72
+ d3 W" E9 ]! j4 S' y. S73
7 {2 V$ E" ?; M' V* ~74( l3 g P W V( x6 M8 l; {% o( u
75
. ?# o& I: ^" t/ A76
. d4 J! ~ |: g7 G q77
" j) E4 q' N( D: V% M78
. \ _) Z! w" m5 e) V! w1 P9 p" \8 }79
& a1 n; V9 {3 r6 P80
' N; t5 b* q; U. D2 }81+ K5 A* `4 u/ d" s: T) y0 x
82
+ H4 b1 M' z8 _; k6 `/ Y5 `, ]! d83
6 N. o+ c* b: D0 Y6 O2 B: [2 E4 k842 r" D. q" X7 q" f9 X6 G2 i7 b* x
85
* T2 K3 ~3 E$ m y+ d86
u2 ^2 Q! D. F% P; \( X87
$ `7 d4 A- Q0 }& }$ K883 t- e. L! W5 a$ ^
89
/ Y' M" I% i O \90
" ~+ E7 O% a2 P4 i919 D/ o) I3 T. X( J( A1 \
92: u% T: _8 K$ ?$ [, m" O
93
) G# \8 i8 }5 w" ^2 |( ]( C) r8 d94( Q( S; n" O. p
95
C$ b) ]& \/ K9 h96
; w# T1 @; }: [972 \1 V+ p! ~- d/ X
98; R1 H8 f& x2 ]6 ]9 m7 Y1 n8 [. k7 x
99
+ g" S$ ?9 n8 V- u) h# x0 \100
* T0 W$ n# J# Q8 h1013 j; T8 D8 n! |& l
102
$ s. c/ Z8 r1 b: A' w- y103
; q4 U4 D0 v, F; F9 p% s) y* [' W/ v104& h; W2 r! |$ C. p- f& d
105 T9 @( S+ \+ N; [) B
1063 G \. I. o5 Y9 a4 X
107
4 `* _- E1 a' ], T108- b+ h2 Z+ ?+ G& Q% o% y
1090 j- X/ R. q: s1 Y! M+ F& a
1108 d& f6 d2 S9 K1 y
111
" i/ W; N+ v' I: C. D) f# n112
7 y8 s( c2 m: l7 k/ f- a7.2 开始训练模型
, s v4 K( l" M" z% K# v5 J2 y我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次$ g$ o' d. ^" R0 Q
9 C$ N+ ]# M& Y9 I#若太慢,把epoch调低,迭代50次可能好些4 b7 j" |% i6 f- K
#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合 z' |, _9 S3 H# F1 d* d' l
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"))
2 t1 j( ~$ J7 W; z$ _- v* j8 f7 n }( v* y. K
1
' P% b# H! f ^$ k6 q: c2
) S$ L2 J6 Y# _3 t2 f39 E5 Y' W+ E* Z8 D
45 i9 i3 D# F# e' ], Q" F; z
Epoch 0/4
$ e X' }. t; {5 g" H7 j, J" o$ e----------
$ p) A" D! e) E6 d, f4 pTime elapsed 29m 41s; k2 K+ P2 I4 p* V% Z; p3 x
train Loss: 10.4774 Acc: 0.3147$ d& i& V/ K/ y5 G
Time elapsed 32m 54s
, i0 m8 w* Y- I8 [9 Svalid Loss: 8.2902 Acc: 0.4719- a# v$ k+ f; d: k7 b
Optimizer learning rate : 0.00100001 B% d( ?- X' ^* _
/ o9 t0 h6 D: G1 B L; nEpoch 1/4* L+ v4 I" m$ O/ }3 {3 C
----------) ]& l9 R- k3 K$ @7 K/ Q$ \5 ~
Time elapsed 60m 11s
6 c4 x$ A5 @- k Ftrain Loss: 2.3126 Acc: 0.7053
! Z& @9 G* ]0 k8 DTime elapsed 63m 16s8 q1 S# q& q' c- W( m4 j7 [9 E# ?4 v
valid Loss: 3.2325 Acc: 0.66260 y6 E4 r; ~# L4 s: u- s+ |) j
Optimizer learning rate : 0.0100000
7 @2 i( T ?# W- I A: s. D
5 J8 e. a8 x$ O& e9 d5 B, cEpoch 2/4
I* y1 s- c3 z2 X4 X) `5 z$ U----------* z; Y1 W8 {5 k; B2 }
Time elapsed 90m 58s- P, [; a* y( A5 J3 W
train Loss: 9.9720 Acc: 0.4734
) c3 j0 h+ v1 f7 xTime elapsed 94m 4s
6 d/ G7 q* J$ @% @valid Loss: 14.0426 Acc: 0.44133 I% X# b& N# m% p
Optimizer learning rate : 0.0001000
" V9 _! A4 |$ d. S
$ p; n7 D$ C' ]Epoch 3/4+ k# x6 J) Q& b
----------
4 O9 @0 e' F& Z' t# yTime elapsed 132m 49s
. H: l& T0 S0 c7 }train Loss: 5.4290 Acc: 0.6548* A# P, M7 q# O K$ P
Time elapsed 138m 49s, H$ E0 V' \; B& a
valid Loss: 6.4208 Acc: 0.6027
O4 e+ M9 i' e' XOptimizer learning rate : 0.0100000
& A) W! Y; V5 b
% }3 ]1 T7 @1 lEpoch 4/4
$ W8 ~4 G+ A6 W# p! \$ \& ]----------- k6 H8 {6 R) M9 t9 Q* ?$ z0 W
Time elapsed 195m 56s
3 u% R4 ~' I. t2 ktrain Loss: 8.8911 Acc: 0.55198 c) r. y) I ?8 c' t
Time elapsed 199m 16s! c# m J1 k* K- s
valid Loss: 13.2221 Acc: 0.4914
( d" r0 V( L! a W8 G1 YOptimizer learning rate : 0.0010000. k& W/ u+ s$ i+ o- t' y
/ H6 N) u O0 X: M
Training complete in 199m 16s7 r) @% { b2 V7 E9 V
Best val Acc: 0.662592
, O1 l0 z( L; B" n" q3 z( X0 |1 S5 K
1- U6 L, u E" e. i6 p; }7 z
2
7 G5 W) S" N& t% A- B; X3
: K. T' ?8 i6 p) O4" ~5 A$ H. ]! m, a6 O
5
; i" x, \2 I$ {/ G- Z6
( {9 ]3 r" K2 [* l+ }3 P: x7
: V( U2 y* ?2 T) d) B5 A8
2 c( u2 J+ s! ~) [( V j* ^" N o! `9 a9( d1 ^* o) N4 |9 |5 j5 \" X4 R) r
10/ X% {8 L1 g9 ^+ _$ O4 I4 P: ]! m0 k* S
116 q& j% K6 j. A; L
12( ^+ k( V7 K4 u' |$ I1 u0 b
13! I. L0 R2 x8 g& E+ w# J
14
% x- }; p4 w: T% x! Y; t15
! P& r0 K% L9 P0 I16* F; @: V- Y" V/ |9 U
17
" ?6 E$ C f+ s O* I4 W) G18
, q6 \2 `. H: W, H3 J6 w19
% M2 X. M$ s) Z6 N% k0 V7 X20
9 e3 c/ |) t7 g. C2 k6 ?' f21
; G2 f9 B1 }! y6 p" N22, `7 z6 u9 p* b4 }/ \# p+ N& x
23
9 i0 f" ~; g, a241 P6 M. V# T# ]& E3 P v
251 }$ B, f' A0 |# r, V% P
26
4 B! Y+ p2 I9 j27+ h4 y4 u2 C- t
28; @' i& P; T$ I' A) v2 S5 M6 Y% K0 ?$ ~
291 Z ?# X" H( K H/ v8 f
30
; [5 T7 @+ b0 z D& H* C314 V+ J) ?" d* ^$ J7 s
32& U* V5 t7 A: A+ Y
33
0 c% {6 c7 Z8 a, ^+ n& t34. w: n, b1 z N5 k
35# d7 F" N/ S) a7 h
36
$ o6 }. g" c' P37/ d7 Z" |/ G# W/ D5 S0 v6 O5 }
38
: J, X% w$ O5 K3 H3 p39
7 w4 l0 ^; O" [, e40& m5 ]6 ]) c) \! o- k" W3 E. }
417 f+ A) |1 U# m7 G4 v0 [
42
5 w- c% D; D7 F* F7.3 训练所有层
4 h! p" i9 B4 s0 F) E @# 将全部网络解锁进行训练( Z+ j6 \) [) I, M# H( u' E5 \; p1 |
for param in model_ft.parameters():
3 l# ]; F ^1 j$ F7 q param.requires_grad = True
C: y) d( n! k5 l1 g9 D# a0 F. W6 B' h9 c0 m! l
# 再继续训练所有的参数,学习率调小一点\
2 E. O. h3 _( u7 qoptimizer = optim.Adam(params_to_update, lr = 1e-4)+ k3 C; ?. z+ ^2 D
scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)3 f; t. D5 Q# G& C7 z; ]
# A9 E! ?. u* g
# 损失函数% i* y2 y& p ?) X
criterion = nn.NLLLoss(). }/ Z& V9 k3 x# B3 I. K
1$ A4 r0 F+ F E8 ^. s
2
/ K' z: ?& C" v3
5 H; u' Y2 V- }5 [& e7 H4
3 H0 j7 ]. O" `5% A6 G2 v! k! e: n+ d. A; F, t
6
7 Z/ q- l' A/ p0 H' C+ p8 q; ]$ P! D7
1 D7 a7 _8 w' Q' }8, D2 X9 U% L% i) O7 w
9
$ i" Y8 u t: g) w" ]5 i10
$ Y, ]+ j' ^6 G3 m1 m W3 n. A4 \& _# 加载保存的参数/ |; T; D9 k! n" f# \
# 并在原有的模型基础上继续训练
- T: N0 @: C3 Y, ^3 {2 a. L# 下面保存的是刚刚训练效果较好的路径
* j; [3 v8 T* S' p( f* tcheckpoint = torch.load(filename)
( H' X% q3 l P2 y0 s3 z5 ebest_acc = checkpoint['best_acc']$ z; ?. ?" M% R8 W7 q$ Q
model_ft.load_state_dict(checkpoint['state_dict'])
8 \3 K; | w" foptimizer.load_state_dict(checkpoint['optimizer'])+ R1 n6 m0 _$ j6 Y
14 T+ {0 s0 }' P
28 r2 m" a5 @, T! q+ l* v
30 c" U: ~2 d& F$ N/ J' Q( X
4
% w& m2 H( t4 B. N+ J: X56 W+ z) T# |' ~/ O0 |: _% [( H" _
63 z4 S$ H; x5 w+ m/ x
7; I H( v) Q7 ?3 S1 _
开始训练
, C& O- W, u$ S) `0 m4 a9 t- T! [! [: q' Y注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考5 u+ Q: N: f# {: j7 W
8 b; \7 T# D" Z7 ^1 z
model_ft, val_acc_history, train_acc_history, valid_losses, train_losses, LRs = train_model(model_ft, dataloaders, criterion, optimizer, num_epochs=2, is_inception=(model_name=="inception")) o# u' y+ W- V7 o4 |
1
' }8 P# J' n( K M6 u7 jEpoch 0/1
0 i8 [& M4 w& ?( D----------$ r8 K/ @# E6 p% ?4 `, }
Time elapsed 35m 22s
4 b! C! m. a$ u2 S3 Mtrain Loss: 1.7636 Acc: 0.7346 a) B+ `0 m+ w3 _- I8 C7 V+ J
Time elapsed 38m 42s
1 O/ g7 y* g' H+ {9 rvalid Loss: 3.6377 Acc: 0.6455. N5 e! r- p- S( {
Optimizer learning rate : 0.0010000
) }! a( d' u$ z0 { v7 ]# T) Z) P' g
Epoch 1/1
4 V, r" P' W- o# }1 }3 ^& z0 [----------
- E A! G! G# h- D- b9 kTime elapsed 82m 59s
: w* ~( [$ M2 q% P: O; q+ ^* b ytrain Loss: 1.7543 Acc: 0.7340. V* ?* n! j: N: D5 U
Time elapsed 86m 11s
: R. s% Q1 }) e0 [0 }7 ^+ Q- q0 Xvalid Loss: 3.8275 Acc: 0.6137# M+ n- v$ ~4 }1 b, ~( x' K* E
Optimizer learning rate : 0.0010000
) R& b* A0 d. Y2 q R0 ?- ?9 l4 G7 ~5 [% J r E
Training complete in 86m 11s9 A* P3 a. @& {5 w$ G A* |5 ?
Best val Acc: 0.645477
- n3 |. e7 `9 C" M: p! B
% z5 ~& l( D* t& T% Y16 o0 u) |6 q4 y* D
2
) \; H0 e% u8 l$ W+ p' \3# L* j. C, J! D( I! I' c, ~3 W6 f
4
2 k, y5 R( C+ Y7 j5
9 o# t3 i$ ] ^, b6 R/ a. E4 |$ v# d6
. ]' A( T( ^# [* e7
# k* R# M! S3 ?) Y1 R0 v5 g1 p8
1 M/ _( e- _2 ~* z1 K# B9% d! a* r# z6 z/ K i; T
10$ g6 F( c. o+ W
11
0 L" b4 ]. u- @* r12
! L3 V2 r3 [0 A$ {9 J- }13; V# M1 b# E. {
140 q+ x i1 T. V1 g( ^6 f# n8 J
153 k9 C" R h% z: ~3 j
16& ~1 N) h( n- e8 {3 `% T7 Y, |3 B, F1 ~
17
" r* @+ I8 b1 x18
0 v7 \- c9 U* A) u; b9 Z" f8. 加载已经训练的模型: p; a2 t% q* F: Y! G+ q/ P
相当于做一次简单的前向传播(逻辑推理),不用更新参数8 o" q, r; K( F) S
% j/ t! N) `- ~. Xmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
0 i" ]$ Q5 Z+ _$ [2 p' `; {3 o6 s! o: R, W5 r8 F7 n
# GPU 模式- B( e H! C, o: D/ l1 _1 x& M- W
model_ft = model_ft.to(device) # 扔到GPU中
7 K: {1 g7 f' T* ]( F
$ }& W6 b6 @: ^& B3 D B! a# 保存文件的名字
1 k- s2 G9 T# \0 G5 C" nfilename='checkpoint.pth'
# v! i6 {5 k5 _' j5 T% t& ^
. Q( }2 z' w9 x6 e# 加载模型/ q: W& d+ o3 v
checkpoint = torch.load(filename)
2 ~( V; a' a# O9 kbest_acc = checkpoint['best_acc']
3 l" h& g/ z& X' F/ o( A+ e Imodel_ft.load_state_dict(checkpoint['state_dict'])7 F3 ]4 [+ U `2 \2 f# \
1" F9 d& f N% j X0 n; _
22 @6 K8 E. b T/ p9 w, j
3$ V7 e/ H5 g; p3 `7 V) Z$ r
4: F9 o$ T; V r) g
5
% y, U4 b6 R, _% H! v* s1 E' H6. F% N2 ]4 c' I, f5 j
7' d7 G7 g4 R0 q+ I* e, Z8 \7 n$ ?" g
86 c' b5 _; w0 O6 }
9& ~7 Y1 r3 P( e: W& S* D
102 h2 M$ F4 c( m/ l5 _1 d
11
% J' x Y5 I( q0 }, n3 c2 N5 Z12
3 F- f- V8 u6 Z9 L- a+ s$ ~<All keys matched successfully>( C1 z! ]1 x3 h9 k% r/ N6 h" y% q, y
18 j2 h' C. a" f3 S! K$ J
def process_image(image_path):) h; |- }. n4 L' [4 S) U# a$ ^
# 读取测试集数据6 P" R' ]5 Y$ [9 @, N! r9 n
img = Image.open(image_path): u- h) d1 Z" H
# Resize, thumbnail方法只能进行比例缩小,所以进行判断
$ e+ l9 K/ r; Y # 与Resize不同3 E1 E* W; {# o" z# d
# resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小5 o6 A' f) I& ]! e: |! k
# 而且对象调用方法会直接改变其大小,返回None
& T' z4 _4 I# \8 ~+ o# [ if img.size[0] > img.size[1]:, c3 U5 c/ A' w7 ~2 J6 a8 S
img.thumbnail((10000, 256))) a0 e. O2 h j0 g2 a T; S$ b
else:% y( X! I( k, }7 d. q% }! L4 e7 L' `9 e
img.thumbnail((256, 10000))# t3 X" B& T; U! A8 I* i
* c- ~9 l5 _4 I" ?
# crop操作, 将图像再次裁剪为 224 * 2247 H7 {& }/ E, G: {( P, |2 T
left_margin = (img.width - 224) / 2 # 取中间的部分
. X9 G5 v5 Y$ W, b7 A5 l bottom_margin = (img.height - 224) / 2
0 b4 |1 J! H( X- W0 R right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度
+ r' w) a! ~' ~, a! f top_margin = bottom_margin + 224
/ N: F5 [5 R9 f$ C( e: K6 u W4 z" R. J+ U/ o L
img = img.crop((left_margin, bottom_margin, right_margin, top_margin))
6 y0 ?( V1 K! {$ y) b+ x9 F3 ?8 l% F* o, s. | y; ?* x$ o
# 相同预处理的方法& Z( d X, e, |9 b0 O3 K
# 归一化, ?$ g! ^+ B, |5 {1 T
img = np.array(img) / 255
% F+ w( M5 F$ Y7 L$ W mean = np.array([0.485, 0.456, 0.406])! a2 p; G) H% B( a3 o# m
std = np.array([0.229, 0.224, 0.225])! Q$ L1 L+ W. c$ `
img = (img - mean) / std" s9 G8 k" `, x! a5 T" {3 Z. O
% w+ i6 n$ A& r+ ?3 r" H( p. \ # 注意颜色通道和位置9 Y0 T+ [8 _$ u3 Y
img = img.transpose((2, 0, 1))
8 e( W5 F; z* r6 |
& v, Q* s; s) k+ g return img
# y" e* i/ ~# Y4 Z5 x# N' A4 `! d: R/ E/ }; Z" O
def imshow(image, ax = None, title = None):
: I; D( e; p- @% e """展示数据"""
0 |/ S) i4 h4 H# K1 |9 q; q4 m if ax is None:( c" A# U$ `1 _% c0 x. _% n
fig, ax = plt.subplots()! l) Q' [6 M7 T3 i
: }2 \& S5 l" c* x, T
# 颜色通道进行还原
5 L+ P, c% ^8 d; V% _ image = np.array(image).transpose((1, 2, 0))1 \% Z/ u9 J8 v3 t! h& t
+ m6 F$ h, n. H9 k+ c$ _) \$ A
# 预处理还原/ M `) {: X# ^* l) \
mean = np.array([0.485, 0.456, 0.406])) @9 W: u0 S' B$ P' f/ j+ z
std = np.array([0.229, 0.224, 0.225])
( i( S. k9 ]$ l* }, G9 P image = std * image + mean; u' D! Y+ b" y% [
image = np.clip(image, 0, 1)" ~9 U/ S! B; z4 H
) u' a, b$ D7 ?
ax.imshow(image). q) l" b) S0 @9 d" h
ax.set_title(title)8 c. i" o; r' n! ~1 S4 |' H
0 b( U* q% `9 j1 P" s
return ax
1 [- l; s/ G' a
+ G9 u, e3 |/ E% S6 g' @( t8 ximage_path = r'./flower_data/valid/3/image_06621.jpg'
; _+ _. X* p( s. u" q& G! S: L' \img = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理
8 q0 m8 |/ F5 Y! ?; G8 e( Vimshow(img)2 M( H1 z, J5 |$ ~, P% w2 u, m
! T6 G4 K0 r9 Y! k( ~1% g; p/ h& T3 p' b) E( J% H5 `
2
" G( H, P" d' i3" M! P0 Q. I! a' X
47 P9 t* A' V, p$ S+ o6 r
58 f$ W, [- M2 D- E
6
" G( r' R+ a) O" n/ ?5 _; b77 D- P j/ h( Z- x- V, V
80 B, ` I0 |% j x, ~1 ~+ Q
9& i W/ C: \: N
10+ h! j+ c; Y! M6 E# d7 W# G
11; p, g3 J/ q* j9 T2 G6 E& ?2 y
12
( d# O$ C2 `& }! h% j* P8 g132 Y& }$ z y0 w
14
; a) k3 |4 p( @, }' ]% B/ x9 n157 K* s! T9 J2 A4 G, ]4 A5 y
16
z( d+ ?- m; r6 h- |7 X( H! x17
( Y) {0 h0 V- C, _1 H/ [8 G7 _18
0 R( m+ Y/ L+ h19- C, t/ v" `) u. _; u' _
20
9 [6 N' d" ?6 D2 K {& I21
: M- l# y# [$ g E4 I22' z8 \: Y1 F, h2 D8 W+ c- V" j
234 X. J5 W$ g; ]& l/ Z0 g6 _
240 K. e! Q$ Z2 t# J+ E
25* _- M- c, |. p8 u. S
26
7 {4 W# W- _4 l ?27
# }- {2 o7 p! z$ z28% F& Y2 s8 B; s8 W3 a$ X
29
* y* s' k% Z7 H* J8 }- d% [+ U30
3 x' M: f' K$ y' x2 j319 d* @; j; _; H' F* ~8 m
32
, X& c( Z8 o' l33
* r# V( [, }( k, @( G34
9 S; U0 X$ C2 Y35
1 A" `/ O" n# c/ G" T7 l366 s' K" \3 ~* n3 ^4 Z& j
37; d h7 t* u# k+ h& N+ V! D
38
4 y2 p. M% r5 ^ a39
. c9 {* m- X1 t- S7 ]40
1 P) E0 A; x# \* i1 g6 B# b" z$ H41. [7 h) ]: G' K
42
5 A: B, q4 ^' r2 }$ a2 x; a43
' c. k9 T9 c6 M2 E, [44
# z/ Z) g$ _7 I45" n9 L9 e8 Q( h5 E- w* F
46) |9 f, Q, @: N" S* O& b; F( L
47
& K1 @% ?* v/ D+ X3 e. p0 B/ ]8 Y488 J; {$ v% l$ w# L8 \
493 }& [) x0 o+ f$ @, ]4 K4 i
50
" n0 Y1 |( ?" l, ?, Z8 S3 A51
" {3 v; Z4 \# Y4 R, w52
3 c; k5 g$ f" Y- n! U# G3 D535 k- s& C( ~9 {3 X
540 g( Q) B' d2 `, f" s" j# j
<AxesSubplot:>
2 P9 j/ S* W' x1. h& N/ B" ^+ C5 P9 ~
" b5 @4 q( u0 F( S' H
上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确6 \ D. D( W/ K) u: w
- w/ C; W5 ~2 Y9 v+ s2 H$ t
img.shape" W. E- P! I/ F4 B
1
/ j- \6 K" f* c(3, 224, 224)* q d7 V# c. p% k4 T8 X1 F. G
1
; F: K- \* `) V$ O5 ^! A( F证明了通道提前了,而且大小没改变
5 b7 Q0 t1 B4 q+ s, b8 n5 A/ T& @8 V; `- ~
9. 推理4 c# c% O: i; o/ g4 r
img.shape
$ A, p4 R+ w0 m" g& ^/ Y( w
& K' n0 @& ~% D) o4 N/ k1 }0 c& K/ T# 得到一个batch的测试数据
+ g% p4 l4 b) d+ m7 _dataiter = iter(dataloaders['valid'])# {! a* T9 Z# Q' c+ ^" B# ]
images, labels = dataiter.next()
6 J* U/ a5 }2 q
, l% h# }) N) i# l# ~, V1 t8 R3 Cmodel_ft.eval()$ x8 }* I0 X. \) X$ m/ h, X9 r: x
2 A3 \& V* U; U: T jif train_on_gpu:
7 T* J) S( p0 Q2 H3 S$ h+ ^ # 前向传播跑一次会得到output1 C2 h. N; e J2 g( `# {; p
output = model_ft(images.cuda())) z) M2 v4 F X8 |- \9 ]) r
else:* B8 L1 Z: T) G& ~6 u
output = model_ft(images)
2 U6 Q6 `" i: M; z/ ?& C& R5 c W2 w+ E) m0 J9 L) I; W t6 Z# J
# batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值
! l: F& i- u* s5 p. C5 H3 @output.shape
* i8 t- n& @5 _7 a" V; A Y: n8 X1 p
1
. `9 K. D) \. K Z6 E a- r3 [9 R2+ z# I. `/ q, L! D9 K
3
3 j. H4 ~. G: h$ i2 A, u/ c4
W- g: B3 j8 {2 M5
! L5 y" d7 T0 I4 r1 h6
3 Q5 e& p2 j+ ]7
: [# ~" _2 M9 I, r. @* J8, H. B! f. @2 H! _1 X6 R" S
9$ `$ q6 T/ S* b- b
109 C2 n! \) }0 ]0 I
11
4 o9 y3 E* ^6 O. o& O12: B, x. n& T w7 U& F
13
9 ]9 c. F: x( @- p$ J0 o14
/ e3 b- g5 T4 u15 [% m5 n7 w& }* c" i9 `8 R
161 B9 e3 M5 j6 m4 q
torch.Size([8, 102])
& h& Z# U" Q8 Q+ F1
* N0 b( E/ M( K4 K0 _9.1 计算得到最大概率! ]% v9 L4 r/ w( }; f5 A
_, preds_tensor = torch.max(output, 1)* _6 P- o) Q4 y, C D! u# E) U. F
8 K) Y: t; u& Wpreds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量
/ l$ T5 d! ^- J% b" m- v8 P% {5 `1 J; }1
. H; S6 o9 v4 n- H/ A3 P2
% R0 d% x9 Z; l2 z3! ~, k5 z" F0 G+ U$ x& p1 U
9.2 展示预测结果! u3 I8 k- |9 P7 m5 V
fig = plt.figure(figsize = (20, 20))) r3 |% T) {" s- [* X x
columns = 4
7 C. U+ w% }, [. _' prows = 2
9 v0 }6 _7 ?& A; b) E- h( x) m
) T6 f+ F4 E# Qfor idx in range(columns * rows):+ ?! c% R- a3 H+ l3 g/ C
ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])) q2 F! k* M1 E' R
plt.imshow(im_convert(images[idx]))
" f0 s: V3 |9 j1 E0 U ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), ) u5 Y+ K. m% B0 p
color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))) V! q+ P" c& P9 p" ]% W! ?" ]
plt.show()
" \5 g& @9 b# o* k9 v# 绿色的表示预测是对的,红色表示预测错了
- ~$ r4 C0 f) u* r8 W1
- v7 g$ p# `8 D* t' x I2+ y" M6 n4 I8 C; E+ p, l
3( z# T3 S' q0 P% M
4( t* x1 I- a. {. W8 ?5 e4 b2 y9 ~- s
5. v! G! [: ^4 y! `/ d
6; w* C H; |* ^6 g" f9 ]
7
, v+ Y' m* p* W5 s8
/ B+ f" U2 W! N( F- z' G9
+ r( ~( X" K; K) k104 A5 ?9 ]8 @2 w# `
11# j9 S: |+ N2 U/ ^7 [2 m* x
/ K% K* ~6 I. n& \! D
; q" P3 ?' O& `9 T0 B5 z1 N* p5 [ V1 Q- @# G
————————————————
. R3 l( w0 w$ b9 f' M5 V版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。 R5 M7 M% n% s2 v" e& I
原文链接:https://blog.csdn.net/LeungSr/article/details/126747940$ A. X2 ^( u. ~. |3 o. b& X
# P2 s( J) a$ s
4 y: |9 L& |3 z4 R6 p6 k |
zan
|