数学建模社区-数学中国

标题: 【深度学习】 图像识别实战 102鲜花分类(flower 102)实战案例 [打印本页]

作者: 杨利霞    时间: 2022-9-8 10:41
标题: 【深度学习】 图像识别实战 102鲜花分类(flower 102)实战案例
【深度学习】 图像识别实战 102鲜花分类(flower 102)实战案例! L" R, _( {; b8 {4 }
9 l5 \/ z% P* g
文章目录
0 X0 w% @1 q4 v% x) ^9 ~9 D: w$ m卷积网络实战 对花进行分类
, s* V6 J  E. p) ~0 f. R2 Y) n0 [6 W数据预处理部分
+ L2 w5 }' a! Z6 }2 H网络模块设置
% P( ~% ~& z  a, H+ F网络模型的保存与测试
0 ~, Y% O' m/ U1 e数据下载:% q) c" _; Y  T6 g
1. 导入工具包
3 ?! Y$ Q1 c. v2 v+ E2. 数据预处理与操作
, k& `2 J; W) }0 m+ B3. 制作好数据源3 J, ^; z/ a# r6 x4 b: E
读取标签对应的实际名字1 W0 c% Y$ r* ?- H
4.展示一下数据
  N5 @0 n, R$ F" T. V5. 加载models提供的模型,并直接用训练好的权重做初始化参数2 R( S' B2 `" {$ E% z% h' T
6.初始化模型架构
% E% h- K/ a' H7. 设置需要训练的参数& ?  g# H6 d6 i' c
7. 训练与预测' C; E3 M# J) ^
7.1 优化器设置6 M1 `" w$ T5 _! [- d7 T5 ]
7.2 开始训练模型' p+ m' m6 l" T) u2 l# E# Y
7.3 训练所有层
- }9 Q, S1 _1 {2 \* _( B开始训练$ B. H9 n) P' P7 \# X/ `# |
8. 加载已经训练的模型9 D' J) V# }- S0 t  |9 z  {
9. 推理1 c' v) j3 v# Y- z
9.1 计算得到最大概率" ?, a% t( c- I2 h7 `
9.2 展示预测结果
: |( q) K' w/ e* R7 ^. V3 O写在最后
; b% `0 w) s4 |  _/ Y卷积网络实战 对花进行分类. a' F3 o% j% U- F8 k8 E6 [5 T
本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能
6 i+ u$ V* {! t! s1 W) e
. k9 s5 U: d  h# J+ ~1 p" b" e6 H( w在文件夹中有102种花,我们主要要对这些花进行分类任务; ?/ N# i7 R/ z7 R# K* I
文件夹结构
7 l" i( z- k# T# Z6 H0 E0 }+ E6 I
6 i5 f; j0 Z2 n0 D+ t0 e3 Fflower_data. N" b- y+ H5 X  e* j

3 c' ~$ C" n9 Ktrain+ F. L4 }1 c8 @+ Y. O5 n& A

' J- q7 F, {0 Y1 I1(类别)
6 o! {; Q- ]1 S: o: L! e4 v( a2
& R! c: N/ b3 ]( r. O. rxxx.png / xxx.jpg& s* @/ w$ f# F0 I4 y5 B
valid0 f! C& v/ {; R# T0 M7 C2 u0 p$ K
  L3 r& a$ s" y; l( L  Q
主要分为以下几个大模块
" ~. E( E2 v6 n( @- [- i2 [* z) P. q8 x2 M# V
数据预处理部分( Q5 F6 T1 `: s
数据增强
, A; e0 n: `; h! Z8 r数据预处理
7 H& n# \- |7 @* N3 v" Z网络模块设置
. w. N; K; z( t( m; F加载预训练模型,直接调用torchVision的经典网络架构/ \( J- H+ X8 }
因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务
" X3 w) S; j* s1 U- h* a/ ?/ f网络模型的保存与测试* y9 e, S! f  Y
模型保存可以带有选择性
; k' K' o. s. B: u- b数据下载:: c* I6 N+ ]% w4 u* w
https://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset/ m1 s2 H8 k* O# f
8 l, V, K( h& j$ h9 v* n
改一下文件名,然后将它放到同一根目录就可以了- d* m- E9 h+ I
" {- m7 X$ ]3 H* Q9 F4 C
下面是我的数据根目录
) X7 t5 V! N$ b+ v2 R) t
! e7 J; z: U' e3 n. t) U% A0 m7 V' n  Y& h
1. 导入工具包5 C2 I# @6 c3 o/ I( V
import os
7 l7 ?% F9 R  G* l5 Nimport matplotlib.pyplot as plt
4 o& X/ o/ A( u& T' Y4 S+ d# 内嵌入绘图简去show的句柄$ w6 E* A  d6 i, R: m
%matplotlib inline
5 F9 A, m6 Z7 ^: [( v9 S" j3 nimport numpy as np
$ c  c6 V, o% d1 t6 I& n; f5 G* uimport torch: I" L- S+ _; F3 l
from torch import nn
6 C* j( x0 Q  [9 Y
6 F3 Z* R$ F- T$ [% r& Limport torch.optim as optim  U2 V, h0 b9 X: V+ J) m
import torchvision" H  z- R$ o# n- Z$ k+ b2 k
from torchvision import transforms, models, datasets
0 _: w/ K# `  Z5 w' {1 p* A0 V9 g2 E& y6 i" m+ _
import imageio
7 K0 H9 r3 a  O) ?' U" eimport time1 @- Q- A. T) N
import warnings8 t7 _$ P% ^- w
import random+ P/ ~% A+ A7 M/ m0 U& m$ @
import sys
7 W% j- K0 }# M+ e! Z; kimport copy  V$ X6 O) ]  m( p# x' W0 x2 \
import json
/ P1 [0 H1 `" w7 Z1 v1 ]1 g- E; lfrom PIL import Image
+ f2 Y' r; V( I# J+ b& g+ s9 y3 ?  b3 r* j2 V9 L2 F

7 R  N: l% N4 B' @+ h4 O14 H7 m8 q% {! D' L! z
2
4 q, Y: d5 n. C* e1 g9 Q9 p7 Z3
1 I' b# p& N) H9 Y/ J8 P49 z+ c6 r& B  P- [( J( U
5
3 J7 k: [( t* d# X3 g) h2 O3 @62 F1 F$ J- c2 N$ `4 G: f6 q
7
' f8 G, x% ^3 z; U1 D0 i% Q8
" i2 |5 O: ^* B1 G' z95 V; X/ v' F( v$ Y9 N1 O4 `1 y5 w
10
8 ?2 Q3 R' u, J' i% p! i0 z  T8 f11
/ ^" v3 X2 H# N' X* l/ p12. a+ d! R5 z0 s% i& C" v/ x2 X
13  Z3 O$ h  F. [) R6 r" H
14/ `7 s' Z! L8 S1 A' ?: {2 K
15
* S0 M3 S! G) u16# Y/ c& [- t5 ~: T& E( \
17
1 ^8 V2 K1 [& ^* n+ b% \7 T18
0 {8 @1 N+ k5 d6 N  \19
6 v; o, h/ {6 ~  v) b1 O6 z: d7 E201 D3 g( C. `# p( c$ L3 H% A
21
6 G0 }5 S) v& r( A# c4 X2. 数据预处理与操作+ F2 R- ?  j: l! _, E5 ^
#路径设置
8 a0 I  g! V9 O, f3 F8 ~7 T' y3 ~data_dir = './flower_data/' # 当前文件夹下的flowerdata目录
$ R2 d5 f2 y5 c4 ~  y8 ptrain_dir = data_dir + '/train'( h2 R! H9 g! ?* [
valid_dir = data_dir + '/valid'
) n/ _- I+ x+ G3 s1
( O( O; d8 U  A: `/ v2/ b. j% z: t, ?2 M: d
3
$ u, x3 F( e4 \" H8 U( z( F41 K* I" @+ O3 g% a
python目录点杠的组合与区别
) r3 N0 ^5 K  C, C9 {. w% c5 d注: 里面注明了点杠和斜杠的操作  _) M! @& S. c1 _0 `) Q# J

9 a8 u! A/ j: z2 x1 q# \$ r3. 制作好数据源
( i) X3 V* M5 N# y  R0 `) q/ Bdata_transforms中制定了所有图像预处理的操作
! ~+ c& B) p9 I  b. }5 aImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片
6 Y% O9 S, t$ \" E6 g% cdata_transforms = {9 `& v. {5 c& l: ~% B* [
    # 分成两部分,一部分是训练3 g8 O  f2 z7 B0 q3 x" }3 w
    'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间: m* b  \+ a$ ^- a& L8 s
                                 transforms.CenterCrop(224), # 从中心处开始裁剪
! H$ B. F$ o. i7 X# ~                                 # 以某个随机的概率决定是否翻转 55开! V7 \) r, C& v* R- _# Q4 p! T
                                 transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转% j3 ~4 [  ?1 C& }9 g6 z
                                 transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转. A+ ]9 z; @7 p4 Y& U: m! H
                                 # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相* C' F2 K3 E" \0 T2 o
                                 transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),8 }3 u& g. R8 z7 p
                                 transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB# L! j. x. [# M- M, y' C
                                 # 灰度图转换以后也是三个通道,但是只是RGB是一样的
* f- z( y, F% r                                 transforms.ToTensor(),
: @+ k4 H: _/ T                                 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差
" `' [$ j' ?" |7 K0 u# c5 W6 H                                ]),$ Y& Q5 c% b4 N+ ~$ F
    # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化& C8 O( j* r8 ^# i; ~9 \
    'valid': transforms.Compose([transforms.Resize(256),
, S9 c- Q& |* S/ M: u                                 transforms.CenterCrop(224),; }3 `5 M+ E/ ~% n: v; e
                                 transforms.ToTensor(),
) x. J6 e  M1 r& J9 e                                 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同$ k  \/ g1 [% [: J6 p+ j* ?
                                ]),
* u; N7 [8 M- F6 ?, G}
5 s( s: n6 j2 t( \+ f2 w" ^# W& c5 @9 z  r0 `# M7 f
1& b( o9 r" R( x& z
2
  C6 \/ T' h$ o7 H) [9 y5 h1 R; c3
+ |3 I# j& F" d4. y" d% v6 {9 `: ]- D( ~6 q
5  A" O9 F. ^% t  h& L1 g. a! r
6, ?/ h& {& c/ ]; d0 Y' }2 `
7, \0 q- U; k6 {7 W
8$ @% {  m, _) x
9
$ n$ y! b5 M9 C9 G% ~5 o10' x, k  `6 ~8 K% _. g: b! c
113 f/ H8 Z: u4 B% D0 C) G/ x
12! }# e, m8 O, `3 }
13# r2 |; |5 ^. N! H, A& C) K
148 g6 t  l# o( ~7 U3 p! _% e
15
- S8 X2 J3 u; c1 X16- Q, P% G7 S/ s5 H
17
, n' q# `  A* a# F' X18
3 w& g  k$ A; |" E/ s/ Q19
9 m  m0 _% g' b20
' ?& c8 e- |  J1 k7 k7 F212 A, N$ K. m5 N# K' d  B
batch_size = 8
# Z+ V5 t, j) P0 M  k9 Kimage_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}% p7 h% h  R; p% e
dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}0 A8 i) t9 [0 {* B, P9 [" Q5 S
dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']} , m& H9 q5 r. j$ @
class_names = image_datasets['train'].classes
6 _4 B$ h: A& p# j
5 ~- d- C# W2 v1 s/ U3 p7 x#查看数据集合
/ m! g$ d/ s7 ?1 T4 Kimage_datasets
3 r% l) P5 G+ c$ N
" j+ L) W- W: T$ Q# Y1 G1
2 t4 m9 s( T$ s2 t24 n' q2 G- w' I$ i
3
1 k3 q% Z* X- L& v$ f4
. X4 Q+ Q- E$ A! [0 M5/ B+ n( l$ u1 {2 e
6. N: |, H; s  G) D. ?8 c. N7 v
7
, V) b$ B% k& d7 ?  T& H4 |: B8; f1 g, }! X8 ]. h7 e
94 ~1 X# k9 f% y3 v
{'train': Dataset ImageFolder9 n0 F% |) S' w$ Y! f8 ^
     Number of datapoints: 6552
4 `4 a* q! m, Z) w  u" P     Root location: ./flower_data/train5 F. E) S# G" h# m8 V
     StandardTransform
( n/ x3 D  \4 ~7 e% r Transform: Compose(. x( b- c% |! }% |  T$ c
                RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)1 O0 l- D+ d' P1 F4 h7 g8 P
                CenterCrop(size=(224, 224))7 g' e4 \. v/ |5 l
                RandomHorizontalFlip(p=0.5), h# U( _- o5 H7 S9 C
                RandomVerticalFlip(p=0.5)' m" J; Z+ U' K' E9 a2 D0 \
                ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
; |4 I) y) `; j, W! Z                RandomGrayscale(p=0.025)' J! E. o8 i1 ?* M! u
                ToTensor()7 E! \7 S6 Q! b: `
                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), S2 Q1 d9 T) T6 H4 {7 }0 `
            ),) v4 e. Z1 `& C" q
'valid': Dataset ImageFolder
% X# q( l) O% ~# I     Number of datapoints: 818* c% r1 f2 `5 N# }3 B+ ?
     Root location: ./flower_data/valid
' Q* q* s. E( g( p! _4 J% ~. m     StandardTransform$ N8 _2 i( @) |/ K9 y) s
Transform: Compose(
$ U9 G7 o) `7 m# F                Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)8 h* D2 K$ K  O3 }
                CenterCrop(size=(224, 224))
) w( x0 |# g- ?# i5 F2 R! c0 |                ToTensor()" @3 {/ e$ O  a+ K7 @1 ^
                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
  y2 q  D; _3 x2 {: s6 z4 s: {2 O            )}
* x* J# n, p- X+ |: l! ]/ I- T( U" s
1
2 l& S. r( L: A3 N4 Q9 o2% g' o% {$ d( x0 Q" @  g3 v
37 J. j& J, h6 c) V$ x# y
4
+ s3 _6 U  i+ w2 x, \- `7 M4 m57 N. A3 T9 G6 V. R
6
7 Z$ s3 ]6 _0 A  J6 x1 a; H. v7
2 p5 d% z$ N. p% v8
) ?) |. b( k  w) `9
! z: b' {' x! v# [/ z10. n, A' h$ `- d1 K4 ~# K# Y
11- F! X4 Q% M* G4 b2 `$ P9 I5 P
12
& o: a# u- e4 C  K* I13
5 j' Y: p4 _+ r+ \3 f( ]14$ E1 Z& p" l( G7 w5 U. n* ]" [. Q
15
* {% Y3 R1 L+ y( u# R; h16
3 \4 W! K5 c% f( X# {$ `17
$ m' k* _! ~) h+ s- }18
5 S" k. t8 U+ |198 g5 r* v% u; L9 v8 ^3 v
20' n+ C/ W  Q- n% |1 ^% k5 I
21
" S) R% ^0 Z" Q2 s; l6 J4 ?. m' v221 P2 v" o+ A$ d1 N
23  C/ p# K9 C  x( J) w
24  z& `& }* Y& C. E) i
# 验证一下数据是否已经被处理完毕
  j* M* m, s  R8 f3 o% g- Bdataloaders, I; X- M0 [- H
1
# Q( E7 z$ b2 B. ~( \9 E2
( B% k/ o# c& q) @. Q0 }- X! z{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,/ w5 c1 i, ^; x0 s# N# T- S
'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}
0 K7 a, q, ?% V3 K! I. Z1) ?4 h. i. A) X, n4 c& K: d
22 L" o3 M4 o! \2 Y1 R+ _8 c
dataset_sizes
6 h, K0 _( t7 r1
* l( y* j5 o7 Y' a' f2 o{'train': 6552, 'valid': 818}
: l' \. d5 j; C) A2 Q+ e# [1
4 s2 V# w, F# [. n读取标签对应的实际名字% m; y! i* B1 ?9 N
使用同一目录下的json文件,反向映射出花对应的名字
. F: _" s' @; W+ O1 \- _, v# J" k) \# {  o3 y
with open('./flower_data/cat_to_name.json', 'r') as f:* \2 [- o, x& z4 T4 h1 s6 ?5 Q
    cat_to_name = json.load(f)
% h& y4 u2 c5 }9 a9 k5 X1
7 a! d  [( P* p& ]1 H% C2; R0 }# l- M# b
cat_to_name
9 K, Q1 p# h! L) m. |3 b1; {' w) i, R2 f4 {6 G- i, T
{'21': 'fire lily',
( ?4 s6 y) M& C7 x, G( D '3': 'canterbury bells',3 r+ k/ `$ ?0 X" x, [2 v
'45': 'bolero deep blue',- j7 E: l/ H( ~: x/ @
'1': 'pink primrose',
- M3 p$ w, a6 S4 k '34': 'mexican aster',6 E1 M5 R: B* k
'27': 'prince of wales feathers',$ m. O* g$ P$ ^. p$ S5 e: ?
'7': 'moon orchid',
. @5 @/ F7 t) B+ y '16': 'globe-flower',( v. e! ^; T1 p4 C) o
'25': 'grape hyacinth',
" P4 d. P/ D7 j: y9 |7 J3 Z '26': 'corn poppy',
$ |" I0 R* F  A  @ '79': 'toad lily',
- I, T2 x5 U" p: o" ]5 T0 e1 h '39': 'siam tulip',
8 a9 m; D" }  Z8 a& N; v" i '24': 'red ginger',. @/ ^, J' F& C1 T
'67': 'spring crocus',! C+ O, g- L1 p6 D4 Z9 D
'35': 'alpine sea holly'," S; i  }7 R" \8 b( z' K" ~) ~" P6 T
'32': 'garden phlox',& B$ G. i* N4 B
'10': 'globe thistle',# a6 D! s( t, Z0 G) M4 B2 v
'6': 'tiger lily',5 n# R5 G; T1 v; }
'93': 'ball moss',' W* v9 w# y! I4 Z! R: ]! D
'33': 'love in the mist',
4 H5 a3 a& B6 X" \' H' Y '9': 'monkshood',
& `+ W' C1 B3 [; Q# Q: J '102': 'blackberry lily',
- X! a+ m4 m: w5 i '14': 'spear thistle',8 C& R$ l9 t+ V0 P+ [$ G
'19': 'balloon flower',
. P/ X7 r  r! v# j' A7 X '100': 'blanket flower',
% w( S: P; W+ U0 a '13': 'king protea',- c7 n1 ^$ z) L, t/ F
'49': 'oxeye daisy',9 Q1 T; C( R6 ?) k
'15': 'yellow iris',6 n* R! {4 K( X8 c5 U
'61': 'cautleya spicata',4 v9 H0 g/ i) z( t+ i
'31': 'carnation',
+ o( J5 k: A% { '64': 'silverbush',
0 ^' ?6 D4 `! u, V9 e4 S8 @ '68': 'bearded iris',& m- J2 B3 t- _" {' @7 u
'63': 'black-eyed susan',
5 i# E0 g. z  V4 M '69': 'windflower',4 }: ^' p7 P. q5 |& a9 z6 d
'62': 'japanese anemone',
& [% Y, w8 f3 ^6 A '20': 'giant white arum lily',
7 {' P: ~, P/ V& I '38': 'great masterwort',* @) f4 N/ b; u% M* t0 z2 N  y
'4': 'sweet pea',
  C2 \* N* b: h& J8 i+ [# S2 d  T '86': 'tree mallow',
9 a/ m  ]" p2 r1 I '101': 'trumpet creeper',
2 [7 I, N$ [9 E( z '42': 'daffodil',# `: s7 X. {/ F( f* M* t% m8 ^+ |
'22': 'pincushion flower',
7 g% v2 J1 `0 |( {9 o '2': 'hard-leaved pocket orchid'," F- V4 }$ \2 \( X  N
'54': 'sunflower',
1 o0 u7 o' G! k( S8 E) F '66': 'osteospermum',
/ Z2 x! _7 y( y0 c8 O4 K0 p '70': 'tree poppy',
0 ]" L. h! u' S9 \) _( R '85': 'desert-rose',
, Z: W" s2 O0 Q. I' }( ?3 a '99': 'bromelia',
) I( x- ^4 s! Y3 W1 m '87': 'magnolia',
0 @% y3 C9 e  ?8 w '5': 'english marigold',) n* }3 R; w, L. H8 k0 k
'92': 'bee balm',( c! p: B) G1 P8 q6 C7 i6 D
'28': 'stemless gentian',$ ~3 |3 H8 O. s& {
'97': 'mallow',. ?/ M% j7 H8 o7 W1 R. i
'57': 'gaura',
) [  Y" Z3 d% g; V& s8 H$ r '40': 'lenten rose',
0 R; ?2 k. x% _0 ] '47': 'marigold',
7 p* s( R- R. H4 Y5 y '59': 'orange dahlia',. S, [) N# ^+ K  [3 g8 P, l
'48': 'buttercup',1 f6 ]% |6 Z/ }' W3 \7 N
'55': 'pelargonium',
& g* E5 j4 y6 {) p+ V '36': 'ruby-lipped cattleya',
+ X& o4 e# F0 E) T- V" E '91': 'hippeastrum',
9 c, T6 o1 i) J$ Z/ k  [  Y7 k '29': 'artichoke',* A! g1 L  |" N3 f: ^0 y& l  }/ p
'71': 'gazania',
* e; w/ y4 b  d$ w9 s2 z '90': 'canna lily',
# Z5 j. X, ^( Y4 A! x '18': 'peruvian lily',' ]/ \* V$ S) ]' T% \8 N
'98': 'mexican petunia',
% H/ b4 r* i6 U8 \1 c* S '8': 'bird of paradise',% }8 R2 I" h+ O! _8 J. R- j
'30': 'sweet william',1 W1 X5 @. a4 n1 _( Z/ \) E
'17': 'purple coneflower',, M: C: J# H7 @+ N
'52': 'wild pansy',- U; n% X9 g" }9 M
'84': 'columbine',0 R8 ~9 u4 b7 z# O% o3 }+ I1 a- ^
'12': "colt's foot",* Y. X! l# ?' D
'11': 'snapdragon',
* m: e  _. N9 l. J '96': 'camellia',
3 D0 H! w; Z- e% J. {# ~# g' m '23': 'fritillary',- R/ v' [, p0 E1 e+ r6 H* a
'50': 'common dandelion',
6 C) k5 A2 }. n$ X* N; E$ o7 [ '44': 'poinsettia',! T# X# r  t0 C# \
'53': 'primula',
; g9 d) t3 L1 F( | '72': 'azalea',
9 l/ C  V; [) h '65': 'californian poppy',* |* u" r: w3 n: a4 H8 ]2 @' j8 b; E) z
'80': 'anthurium',+ h% `3 H0 t# n" O
'76': 'morning glory',
5 I0 }7 ]. m7 D: X3 }, Q' ?# S/ Z '37': 'cape flower',
7 o4 ?) T$ o/ |& z6 ]  ? '56': 'bishop of llandaff',1 s0 k# d" h9 ~! p# W
'60': 'pink-yellow dahlia',; W- J* q) W. n
'82': 'clematis',8 K6 l/ }3 j( S' N- C
'58': 'geranium',
# |3 |: T; P+ D '75': 'thorn apple',
4 Z8 L3 v' b: T4 r! l: r( ] '41': 'barbeton daisy',( t2 s8 d% {7 [, {" o/ R! q. u; c
'95': 'bougainvillea',
! c8 s& H: z( s- }* O. w% x+ H '43': 'sword lily',
. {% t3 Z* ?: q. d( t3 P '83': 'hibiscus',
; V* C8 @3 {6 L2 y '78': 'lotus lotus',. g6 ^/ ~2 a# W4 [, W* w
'88': 'cyclamen',2 J2 i$ l. t2 z' z
'94': 'foxglove',- d0 I+ @" q: L/ H# \& K
'81': 'frangipani',8 l4 `8 Y* e( s5 g2 I
'74': 'rose',
0 o8 b; m, R- A" E '89': 'watercress',
3 K& l: d( b6 l" H$ y) Y+ @6 c4 ` '73': 'water lily',
. w" n# J6 q0 N$ J '46': 'wallflower',) r& [0 x3 l6 p1 D7 B) d! Q3 v
'77': 'passion flower',% i9 _( P. ^" p. u6 B$ N! o
'51': 'petunia'}
; Z, c& P6 r. h- i# Q  Y6 _2 O4 n
1
( A& {3 F. r! D% V0 {. D! ]28 W+ s1 [$ Z! k2 v. y! }. v5 `6 l
3
  R2 d# a. f! l( `4
/ y* w" d  p0 P5: ]. N+ w+ ~  i- J
6; }6 ~% N: o6 v1 l3 V( J
7
5 ]$ T6 \' `5 f8- j3 y' z/ w( k0 c- O8 G7 J
9' Q- E1 O  @9 F7 [/ d+ P: g" w  w
10+ c( s1 k" f! p4 x9 i% Y6 \2 U
11" j$ ~/ z4 Y! j$ v' ^
120 W8 x, R9 u2 w- r
13& q( J* ]; N+ R. G
142 r+ i% v" \9 ?5 z! w. e# F
152 k- k- {2 l' n& I! C0 C
16' e' o6 ^- H4 h3 t. q8 M
17
: l9 X2 O; X# h$ @7 z18
1 f4 m  G& L8 Q19
. C3 u% {/ \' j1 v1 _2 A3 X0 ^20% i- y/ i- ]7 d
21
5 N, N4 d, p; t  I' l3 f, X22) H" j$ b% P8 a6 ^: w5 ^
23
. \0 M6 U; X' J/ k245 ]  C8 [( G4 [1 d  ~
251 _# b' \3 ]8 F
26) z- u1 }( }, o3 u, E( \
27
/ d1 T9 S* w4 P9 ^6 b7 A/ ]289 z: X& _! n9 a
29
3 q. g; ]$ S5 r* e30! Z, w  V+ s4 X6 ?+ d# {) _' }2 H
319 e% Z$ W: h- a$ r: C6 R. \
321 N3 K1 T0 I0 l
331 i0 A# B( g" s: p* t6 s! y! O
34
; X; R( ^5 G1 t1 g4 }. b35, @$ R; z# p" Q0 ]7 f1 B
360 H4 w! B! X$ r! A7 e& d
378 ?! U4 ?+ a1 Q1 A. V) c  {
38) ~% M) w+ \0 A) B& a
39: _* e+ {  e% O! e, b4 C: |# |
40
$ z5 j' }) n" I6 x: G1 N- I41
6 g/ A0 u5 N# m42- F% p5 P4 t( s
43" E) B( h+ l6 y) H$ m* t& h
44( k! f4 d# i) J, R% F, C& f
45* W3 S+ d2 y5 B
460 g- K  ^8 ^3 W& R
472 ]# ?) l. ~5 H7 d4 C$ T9 H% `
483 |5 ^( F, S/ R  `, P
49" c, J) y. z+ n& R
50' d6 T: i2 p9 D* L8 ^
51
) M- K$ _) W6 W1 d6 _7 t52
' j8 \3 `4 R; b1 d0 O( I! m) W53
; U$ r4 @0 T1 T' O$ p54
$ H5 |' T6 t$ W- @55
& h3 f( H4 }4 C5 w: _; l8 X56" E  c& C8 T3 h$ w
57
+ U' I- e9 v: ?58) `% n6 i) b$ S$ P+ l6 f9 W5 u
59/ o3 S2 {* N& {7 J1 Q
60
$ h) G) E7 Y9 u0 v! V# o' u3 T61
/ n, d# O4 Z6 P& l$ a62
. i, T9 y! z: p0 B. J- u3 U% P2 _1 b63
! b7 X% F9 R* h0 u: R' }64
' G8 \0 Q7 X/ j0 A! l658 M8 h) O1 j! \6 Y6 d! F: k! [! S
66: p$ [" h& g' S* @  m# C
67
7 i, a/ P$ X% _' ]/ ]5 T& z9 K  a68+ I. l: F+ M* r& C) r& F' [* C
697 p: a2 d0 F) n( C. R
705 Z% Q6 ?& C$ M2 O" p( P
71
+ S  Y( o* _( s+ e( a$ ^# B: B72
, y) ]3 ]! }. p73
" {( ^0 J3 w5 }2 A0 S+ C74
& x; ]2 F. R- `& A75
# z4 T( E& J  N6 ^' t4 s' _76
% Z# L" D+ `. D  H' o77) d, g1 z; n1 j0 ^' ^1 f- Z
788 @8 d: R0 A) X% ]) {* T+ m$ Q
791 {+ r1 g8 L% g) K. C
80
0 y" ~/ ?- c/ S) e7 S3 }+ e5 n9 _. y; t81+ a8 P  r/ U: M5 Y3 H: U& y
82
0 C2 C) C9 W5 Z83
" O' C! M2 l  E5 q" V84" G& g: N8 n9 X- C/ K
85- [' B8 n2 j2 i2 H; A+ {
866 s) h4 |6 V6 A  _' N. v6 I7 n# o
87
4 ?$ {! P0 N& e6 d9 ]88  ^2 ^' @3 f1 W/ J& A! s
89' m( z4 [( J$ Q$ Q  [+ H' V( K7 Z
90
2 I% [5 P4 p1 v1 @, X8 o3 z91
7 @3 }  r! C! w4 c7 ]92
; N1 z( r( L. _1 s* Z; l0 O93  B6 X: }1 H% I6 b# h0 C9 V9 H7 I
94
$ n. ^, N0 b1 k' O" O2 i! b952 e4 f3 i3 x1 ]; j7 r
96
8 q- b: h3 N5 ^1 }' M" _979 _1 U1 }0 R: ]7 I7 ?$ `" V# |
98
" e+ k6 F8 Y3 M7 ?: _+ G% p* u0 V99
1 S: z, `' L  y100
& T; R' {+ m  p" x7 Z3 d; A1018 n; K; x! Y" y2 A  R+ @* C5 H
102
$ D: m0 x2 F/ y, h  u6 Q, Z4.展示一下数据6 L6 ]* x4 d  }# P" \9 a: Z* [. o
def im_convert(tensor):( s3 Z& b" a' k% M  ~
    """数据展示""") V! R; H  @; i& u/ S! g
    image = tensor.to("cpu").clone().detach()
, ], Q6 v9 k1 ]% q    image = image.numpy().squeeze()
9 J; V1 Y4 L, k" f$ J9 m0 g    # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图
+ ~$ |: e6 p5 ^+ O) s3 b    # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)
* }+ J" N) z7 H) T    image = image.transpose(1, 2, 0)' l8 Z" K5 z( R
    # 反正则化(反标准化)
$ W6 z( J8 h, o1 r    image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))
) C0 s/ u" M9 O* ?+ F1 {2 ^; x; v% }) t  F% E: J5 l
    # 将图像中小于0 的都换成0,大于的都变成1
" M' m$ F. l# h+ m" P/ R7 l    image = image.clip(0, 1)
9 u+ r3 l5 z" [( a9 P; W: t6 Z" O( V6 q( H' D! K6 Q) m5 U
    return image
! J  A1 ]4 L- B1 p1
! M6 i2 x( H' f3 e2
4 V1 u% V" C2 \# D3
5 M: z- O( v8 Q# L* B) F4' c8 o6 F+ X$ ?
5+ x7 r# E0 H8 w; i
60 d, _" R. I3 U% @/ ?: c) h2 z6 e
7
. X: S; M: _. E8
. k  k# W( H' A) l) p3 q9
, n' L8 k! m$ ^* Q10' Y; X! f+ Q: z" q4 |* E
11! i# x' G2 J- g
12
# V5 P* F& v0 Y- {, l& j13
5 F; f: N# [8 x14! u3 h! D; n, P5 a2 m1 a% p: |( l
# 使用上面定义好的类进行画图
/ n% @' _( Q: Z7 {6 W5 ~fig = plt.figure(figsize = (20, 12))
+ d0 Q+ l" H& p: Dcolumns = 4( S, B; M0 \, Y, q. G' P3 d: g: t
rows = 2
' q/ S5 l% K/ A! V
# y, k9 }& y1 r, L) |# iter迭代器2 u2 q( k$ N6 E, \& f. K# t
# 随便找一个Batch数据进行展示0 N0 c4 L, F0 [: P+ `/ {4 E' O
dataiter = iter(dataloaders['valid'])
) g' J+ O! p& Y+ Y! E% I/ O( Pinputs, classes = dataiter.next(); d; r5 a' w* p) _' R
  b; N- D% \+ @
for idx in range(columns * rows):
; j: W( X- C# @/ y" N- B' h    ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])
: B7 `( Y) @8 i1 t7 W7 X. e9 Y    # 利用json文件将其对应花的类型打印在图片中5 q% e( _6 V! m; }2 b9 O0 w
    ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])
, a% a4 C( {! e    plt.imshow(im_convert(inputs[idx]))
% g: K3 f  z/ Q8 u( Xplt.show()
) r  P" T) @# [" c" p
: ?* P6 l$ Z: I8 s/ n  x6 q2 B1) R" i" S* {4 F
2
. x1 {) {- B  d, p8 E# p36 L2 `% S. G- h; s) {
4
/ y3 D; M) n, \. ^5
6 {" R0 _$ E& H+ w5 r: ^. ]9 l6
7 @  Y- I5 M6 \8 I( L; s6 O7, {) Q5 u4 \; J
8+ R3 X& C& v( D4 ], ]% p
9
% z: v* R) g8 E1 v% A$ Q5 n; J10
9 d& T; R3 Q: l8 C9 Q( V11
5 @  \% h% _# Y/ F$ F2 C9 F* C; n12
6 ^2 U1 R/ K6 I2 i/ a130 E6 ^1 P% n8 D7 \! r0 R' _
14
/ G. ^0 i- z0 }  f- h+ t15
2 s% x' W3 A5 q+ ~* ~$ W- C16
1 z) ^. ^  n. y/ H3 O# ?6 p, V2 |8 D3 K
0 T: G/ r) Y" ~. u, V# _& J& {
5. 加载models提供的模型,并直接用训练好的权重做初始化参数; ~+ x+ Z! L& @$ S3 v/ [
model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']% }- {! c" A' z5 P6 |& a
# 主要的图像识别用resnet来做
. V- r0 ?) f5 u5 Y4 W  S# 是否用人家训练好的特征
* D9 D% J, y" ^+ Y" Efeature_extract = True6 O7 O/ J# ]% y# v
1$ Z  R2 @2 m3 N4 p8 d4 m+ D
2& o+ i+ ^: M( k0 C+ ^! B! B% T
3
" F# f4 |3 F1 S5 U7 F4" U+ J6 x5 T' W; P0 C
# 是否用GPU进行训练
$ H% m* U% D# F5 `train_on_gpu = torch.cuda.is_available()
  L1 ?. d# k3 L0 L- x( K/ E' T9 y6 K, h9 w& f) f3 D- t% J
if not train_on_gpu:4 w; m' _1 i' I, ]: h# _
    print('CUDA is not available.   Training on CPU ...')6 X8 T- x2 A& j2 j& ^
else:) a  D, d/ H& [2 e6 Y
    print('CUDA is available! Training on GPU ...')
8 Y5 M3 e) |  n) q0 f% V/ ^% P7 c! r( W5 r8 S
device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')
2 @& P2 Y8 l* ]3 d6 e, ~" I" A15 d1 M) ]' |( ~3 d  ?
2/ l; m5 _+ J6 ]3 K( J
3: P, R" C- [# _% I
4
8 T: g+ E6 m% x) @3 R  a0 `$ V  D51 R6 g5 D# l: X9 l5 T% L6 u
6
, j5 x3 ]' p% R# R+ t76 u: q! \' h0 j" O* G% T, h
8
* |5 t7 q' d0 d. S) t2 N9
4 i0 O! \) ^- y( }( k3 j7 ACUDA is not available.   Training on CPU ...
; c9 U4 F5 K! z+ @, S1
  V$ @  D! S# l3 _( l, I9 F# 将一些层定义为false,使其不自动更新
/ ^  g2 D" ?  w& y# M; l% ~" pdef set_parameter_requires_grad(model, feature_extracting):
/ x; p, {3 u" E! c. a0 r    if feature_extracting:6 H5 f$ }5 z' \
        for param in model.parameters():
* i% M. h* V  t. T- P, B            param.requires_grad = False" y$ p1 Y6 g7 z4 c
1$ A( Q! I9 f7 L! R) G0 E9 a4 M
2
. @$ z. U8 P! L) |) Y! k& h8 _0 i3
( ]6 h/ @4 v( F. L" |4 y4
2 r0 }5 m% N2 S* h- K/ k5 c5
  I( \% N6 b8 T) q4 i# d# 打印模型架构告知是怎么一步一步去完成的
0 I" x/ ^) S5 X* e3 `% y& z3 n2 X# 主要是为我们提取特征的
, J! [! h% L2 y9 _" r- D4 U/ R3 j) }/ B% y' ?; X. d( y
model_ft = models.resnet152()
) M3 z' n* y8 ^  Smodel_ft
9 g) g: V# r0 E7 `: B' {1
9 {8 f7 |3 m* H$ q2
* Q  g6 N$ ]1 }: r3
% R! H# E) p( K* l7 g$ @7 A4/ Y% c$ \) v1 Z# e2 p% _
5
3 y& h' D( a. |8 `9 ^ResNet(! g9 M# Z0 `7 k0 H0 Y
  (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
0 k. n. B4 k1 K; H  (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)/ z3 H, L3 t* [8 a4 C, p- I
  (relu): ReLU(inplace=True)& c% f2 P) }$ x0 q# B8 }
  (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
4 v- ?1 i. p) C" T2 M" d, d  (layer1): Sequential(
* E; {  R; r) h% ?    (0): Bottleneck(. P2 U* _7 _) G7 Q& L% _  C' i
      (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)+ a5 Q6 t) J% [' F- J1 G: v/ y  w
      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)+ w) ?8 `# [8 O+ J+ }
      (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)& ^% M) `2 F7 G  z: N+ Q  T
      (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
2 }9 z) |: [* t" ^, Q      (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
; c$ ?7 [3 E& A      (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)# ]! e9 {9 V8 D( k$ \
      (relu): ReLU(inplace=True)
1 N+ [- E- r8 y0 j      (downsample): Sequential(
- U2 D% z2 G6 H3 H3 V' c        (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False): ?( F- C: H- i+ b
        (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
8 S5 @" i! k0 `. j      )8 U" i/ |8 c, N6 X* R' _
    )$ B+ x0 F) `1 R' c4 R' N* f
中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。
1 o& @) v+ B8 [% ^. ?: W    (2): Bottleneck(/ p5 o9 @9 t( D  C; r2 P
      (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)
3 _' _2 t1 g& Y/ p6 T/ Q) S8 ^      (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)7 @: B4 s# f  L0 ^) l+ U* T
      (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)  b# t; z1 A' \- y
      (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)" I! t+ h0 c% D; p- B9 A- o
      (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
! |2 `# r' `* Y+ e) N6 {7 O/ y/ t      (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
$ K2 o" i5 O1 }& }2 g  @      (relu): ReLU(inplace=True)
( Q) _2 E4 }% c" t2 X    )
8 H* A4 q9 t' V% {  )- b0 A/ \8 h' K: `* P' X
  (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))
6 i2 j5 |/ f- d  (fc): Linear(in_features=2048, out_features=1000, bias=True)
. {% z" S) o- s: t). B# l. u4 W* a( ?0 r! A
: z& G; G) s# S2 a
1: Q# D3 \* b' S7 ]! s; q1 C4 A  q
2
2 \. B  l9 a0 ]' \$ S0 I3 ]30 O6 |, _5 w" o4 M8 A
4$ h& a8 m+ `3 E* `  @% b
5
# x, C6 o* J- t0 I5 \; E6
1 R2 ~3 O/ c( J& m6 K; c0 \: Q7( s1 B5 _0 v7 x, h- Q% v3 r
82 X" T; c9 R# s
9
, k& i0 e7 w+ r( `10
( J3 y+ W) \- B& m' M2 X( N: N11
3 c! G, p  M3 A# B- K, h' x; |8 h12" d8 {6 J7 G; Y, G
13
1 M9 e4 O5 ?  j6 w: M5 S2 g7 v14: b9 f* n! j* _. C  C5 S7 _
158 q5 d! z( g6 m/ T' W8 |
16
+ r% @* r9 G& O3 J3 C% e% s17
# m& z' I) x& _: V2 N18
" N) `+ {+ \3 f, a  u19
1 B/ W* {: d7 n4 n/ P0 _20
1 i' I3 k/ d( F! D3 m9 t6 t: S21
, U/ W$ V/ m. [2 }& `22* G+ W* V( P1 F( ?$ c
23
6 v# u) \2 |, u; A$ W24
* O0 K2 A; x8 q( B25
8 _/ m; P4 }1 w; ~+ J- U262 Y& N1 Y9 y- R$ c+ H
27
0 M: [1 t! A- w% Y286 R. H' ~3 O% F: ~# T
29, W! E' Z9 o$ _. b2 g1 Q
30
+ o8 }7 {& G; ^7 ~& E8 V! n31
3 A9 `( L: q+ n) Q, Y' B& l32
5 c* x( M: }$ X6 ]33
  N- L( z5 o6 J. I/ x( t- d- l/ R最后是1000分类,2048输入,分为1000个分类
- R6 U# Q& m) A/ J' G8 @$ ?" G而我们需要将我们的任务进行调整,将1000分类改为102输出  q, j' J, Y! y8 G

# P% ^3 O/ r# E( u( E6.初始化模型架构" r- d- @9 y' l2 }& A6 c
步骤如下:- s5 p5 I: [' y! I% T! i% `
! t$ x) d( E. k# y- W
将训练好的模型拿过来,并pre_train = True 得到他人的权重参数
) z/ m$ @4 y" Q# R可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)' Z' o4 H5 |* ]: ]) u0 i
无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数
' W$ H2 r$ T1 p# u% N官方文档链接9 w% F4 F$ [9 w
https://pytorch.org/vision/stable/models.html) b1 c/ Q7 `. I& B$ t

  w  r9 j$ v. y9 }1 ~( x# 将他人的模型加载进来6 T$ n1 m/ \4 q3 z
def initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):
: b+ Q0 {0 f7 D  y8 q    # 选择适合的模型,不同的模型初始化参数不同# s  g7 w$ u4 k+ ]4 z
    model_ft = None, ?3 p% W' {# s3 e6 S
    input_size = 0) E3 X0 x$ [6 n) R4 r

9 j. y6 K6 ^, E* U: v    if model_name == "resnet":- l: S# C! ^/ S* C% A
        """. o2 L( l# k' b$ e6 J+ P6 N
        Resnet1527 |* B" @% B& O6 Y, h% k1 K
        """8 d9 ^7 A/ o: [  ^9 a8 i
; o; b9 M* e7 p
        # 1. 加载与训练网络1 ]$ z* P7 C/ q/ k% p6 J
        model_ft = models.resnet152(pretrained = use_pretrained)
5 @; _; ^3 j( L" o        # 2. 是否将提取特征的模块冻住,只训练FC层
% X1 j( @9 P3 r( f2 i% z& S        set_parameter_requires_grad(model_ft, feature_extract)
* y7 P$ e! i8 w$ h( z, N, C        # 3. 获得全连接层输入特征2 R% r4 W9 t0 T: W
        num_frts = model_ft.fc.in_features" u7 U5 k# i% p$ V
        # 4. 重新加载全连接层,设置输出102: Y& o! ?/ T: I/ F/ f  f+ w5 q! y
        model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),
9 J* J) h% I6 A) F                                   nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1
. f4 t0 V3 v" g2 M- T        input_size = 224. p, d7 G0 p2 L

0 p# s( f* K- g6 u+ c    elif model_name == "alexnet":) L* C( |' _1 f7 w% b
        """
) P3 R, n2 K3 ]/ j4 ?/ Z        Alexnet  e/ `" ?1 V& i2 A
        """5 m: y" g5 V( F# L$ C$ ^- R$ o
        model_ft = models.alexnet(pretrained = use_pretrained)2 `4 c" c% R  r1 ]5 e1 n; @/ m5 \% C( m
        set_parameter_requires_grad(model_ft, feature_extract)' U0 i, G. K+ M8 `4 x1 V" q

1 q0 _/ o7 n4 _6 I4 X4 D        # 将最后一个特征输出替换 序号为【6】的分类器3 }  F' o) x: @; L9 D
        num_frts = model_ft.classifier[6].in_features # 获得FC层输入  B: w) [9 w6 D+ X
        model_ft.classifier[6] = nn.Linear(num_frts, num_classes)9 e9 H4 A) v9 K9 K* Y
        input_size = 224
, M8 h/ c" H# Z& ^6 U
! f% Q' Z* l5 W    elif model_name == "vgg":# ~3 F2 N+ [/ W8 x& Z
        """* t+ {& Y; w9 y4 p0 [1 X: H6 V
        VGG11_bn) s& R# }# @8 u. f9 I' v+ k
        """( x6 M3 w/ P- h2 Q) K
        model_ft = models.vgg16(pretrained = use_pretrained)" F! b( K; b# M/ a& J% R
        set_parameter_requires_grad(model_ft, feature_extract)9 M3 ]2 `9 Z8 O
        num_frts = model_ft.classifier[6].in_features
4 R5 h9 B8 A: E8 \7 f        model_ft.classifier[6] = nn.Linear(num_frts, num_classes)8 C" K  j6 Z9 U3 h. o8 |# d( d
        input_size = 224
& H8 W% b# Q, Q+ z( y1 I8 s2 Z- O
+ m3 S  T, t5 g! l9 x    elif model_name == "squeezenet":
+ `: v( P  Z1 b4 t; f: `. L6 I        """
5 m' b4 ^! z+ p' ?; ]+ e/ c        Squeezenet
; y1 T" h9 k$ P1 |( r7 n( C$ |        """: t: C$ @* `6 D
        model_ft = models.squeezenet1_0(pretrained = use_pretrained). a9 I% U" t3 n2 N' q+ m5 J. F
        set_parameter_requires_grad(model_ft, feature_extract)
; e3 Z5 ?. I; M& Q7 s0 H; f8 y        model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
) ^1 {' N. a7 d& M! p        model_ft.num_classes = num_classes4 P9 ?1 J) ~% {/ s" B5 S: ~
        input_size = 224& l. U2 G, B0 ?! e
7 u  j; R3 ]9 C% |
    elif model_name == "densenet":
# ]7 U$ v8 j/ n) {        """
) f5 K) `1 p0 C1 p; M+ o4 O        Densenet
6 [$ n# I; t2 q# a/ i3 d' l+ _4 s2 A        """
5 z' p6 u4 Q" W        model_ft = models.desenet121(pretrained = use_pretrained)4 ?  X0 ]0 d  X; T" ~' D
        set_parameter_requires_grad(model_ft, feature_extract)8 [9 d' p, x9 _! W" d5 P
        num_frts = model_ft.classifier.in_features' M, k- {' `1 X( V
        model_ft.classifier = nn.Linear(num_frts, num_classes)
4 n0 o4 q- g1 F9 G& m( S        input_size = 224( G" e/ y3 G2 {8 g8 B, V' ^. K( g( h

: a. Z$ _! X' X, \" K2 r    elif model_name == "inception":) @$ C1 U) M& Z6 K6 V, O# v
        """" y( l/ r* j8 v0 l& A; f
        Inception V3' ~; t7 [6 G' m7 g* e: G
        """6 X0 C& }, m) _4 O9 _7 j5 y  u% ^# A6 E
        model_ft = models.inception_V(pretrained = use_pretrained)
% M! |. q: T7 d% C( A        set_parameter_requires_grad(model_ft, feature_extract)
; F" s) [+ ~3 @% r1 F/ P( C( _  K- C" u5 |
        num_frts = model_ft.AuxLogits.fc.in_features* M1 @0 _1 B8 F% e4 @
        model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)0 H3 \6 s( b0 c4 p; k$ \6 F+ B
) P; _+ F' g/ H6 \$ K- o# j$ n/ s  O2 q
        num_frts = model_ft.fc.in_features5 Q& R( q: B0 o4 g+ a
        model_ft.fc = nn.Linear(num_frts, num_classes)- p/ |& K# v/ S2 [
        input_size = 299* j: Q0 d* I: p9 y: O

+ W, n$ y$ a9 C& f3 x    else:3 g- S- Z' c1 S% u* A0 ?8 S
        print("Invalid model name, exiting...")+ y/ H9 r& [* u
        exit()
. @0 H% F: y1 `, b
/ F/ h& y! T' n! h    return model_ft, input_size: n" p0 G2 e6 C3 D3 T, s2 j0 Y: x+ |
# K( W% A' g" i1 u$ |, p; I
1
- N5 e& ]" e% O9 p4 a2
+ a: k" z0 Y, e5 f3 b7 A32 d  ~& U' G4 F$ R! F$ u
4
) E3 x5 W6 n  [9 Y+ j3 Z51 B8 K% I8 q3 D- u' M4 ?( M
6
: J: U, `: H& D; L7
. V( K" e, i$ O* h$ [: w84 c: w+ G6 X$ x  o
9
9 _2 c+ x& L" P109 R9 {' H: `/ c
11
. n6 g) E; j+ ~12
8 y+ ^; |6 w- M' t% L! q  O13
8 i* K0 D. u0 L4 x+ O& \14
1 S) Q! {1 c9 j( s" N4 L0 c152 L+ W  Y" n4 F  C; h' f
16% Y  f" T+ ?0 l7 F# H4 M4 g
17- r7 t2 v3 n( B) o
18
' j& K: b% r- z! W+ w" [$ D7 Y19
6 j1 h$ t* z( g; K7 {20' H8 {! j7 V) S) {( \/ v0 s
21
+ @) g& G7 K5 A! ^3 {; O' \22
3 x- Q3 K3 z, p2 V: \  [23+ h+ P( L  U- a- |4 d
240 d5 z& v7 B1 Z! }- E
25% |9 l. p; F8 B& B3 M+ O! J
26
8 y: s9 M  Z3 C' L" y. s' x9 y277 k$ x( {4 K$ b6 U: b6 l
28* }% [+ ^: W  Q7 n( S5 H
29) F- h% L0 N  Y1 q2 ~5 E
30) h9 V3 c  P9 i# r( `2 w
31
* H) c3 b5 x" g2 @; F0 s/ l32$ r  {# O1 u" V1 ]
33
& M  J8 T$ C! k) y, Q4 r34  \0 d4 B' D! C0 ^5 M  x; O7 P" P9 t
35- M1 y; D1 U  x8 s7 h" j8 T2 \
361 g4 K% l2 ~; |* C+ F, t0 z$ H
37( S) m$ i' H- t( I
386 f( ^1 p# V: Z: I' n
397 t* {7 L. `4 q" g2 C: Y! j6 e. a
405 e$ |: \% e: S/ y% B8 y
41
) k. F6 S6 I7 t% F8 z3 c3 S42- R5 o2 }+ O  P+ S% r
43
' l6 K: d7 d* S% k, R44: L; Q4 V6 z* V. _+ W( l8 ]! C  A  W
45
9 ^1 j3 x% [# q+ H- E9 y468 Y7 ]$ ?0 i3 U! j
47
; F2 w* W: v3 g$ }: F4 X48& E/ x5 |3 ]* w! g
49" B' V9 a, k2 b. U7 T; S9 s) A/ C0 G
50' m+ [: U1 N3 ^* t$ p$ z+ n
51
0 a5 ]2 `( j# z+ H7 b* i  N4 Q52
2 O8 A7 F$ S& B( w9 s/ a53
: y4 V" k9 e( I4 z5 }548 x. b* z( v; q4 _1 f0 d, u: D. Z
55' P4 ?5 m/ J* n4 \" e
56
! D0 a& j' z# l) S4 W57
& i3 P( f" c9 {8 e; B1 V+ h' O1 F5 t58
; U) T0 b- R2 X6 v# {  j0 \59
' w& j" r/ ^, s$ S. G! h60
) n' L- o: o' i. L: h9 ?8 m61
+ h' G2 P0 h; M9 b3 \6 S- Z62& c! g. c% c1 Q0 B
630 U! ~/ R; ~* Z# D3 K
64' r, c8 _% O' m
65! h5 Q& \9 R; W) Q$ p
66/ P- P) o1 S& o& [8 P& A
67
& f4 o  R( l& p3 W68; O5 q5 n8 q' ?" ^+ A" }3 M6 B+ L, K# [
69+ e* r/ j( I; O( {; ?' U; ^- q
703 T0 p9 U8 e7 u4 @
71
) r# n+ ]2 ]: G" E# _+ s- {4 h72+ b+ l( @$ j9 `, x. X
73
/ N% S1 a4 T5 |# [- x74+ A7 z: s+ T6 U
75
/ {) K) W2 ]& k5 Z/ K76
6 Y& V+ }; M/ [( G77. w& w0 v/ n0 `5 [
78
$ v4 N2 w2 D: i4 P5 E79) o7 d4 c( X6 v7 W, B1 K5 ~' H
80% d1 n( D4 H! |' w0 r
812 n' m3 x/ ?- r: j& e( ^  i5 z
822 O8 C: D% o& J. d3 r4 N2 z; e5 U' @
838 ?2 ~1 N$ W& [
7. 设置需要训练的参数% f( m6 E9 @: \. b# `% ^
# 设置模型名字、输出分类数  V' f) w) S# o; ]  m
model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)
  [$ p: n, r2 j' R$ X6 c# n" A( l- Y, C) o' L* h7 r. d6 n! ~) v9 e6 d
# GPU 计算8 Q  V4 c  F0 X  |% ^; {
model_ft = model_ft.to(device)
8 ]+ \( D) |, F2 _- P7 o+ Z- u( P7 X/ K
# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取
3 h. Q) l" s0 r4 ~filename = 'checkpoint.pth'/ u* J% B. }, n. I. K
$ K1 B; ~2 r* G, Q7 ^
# 是否训练所有层0 V: P; m. p1 B  W3 ?; q
params_to_update = model_ft.parameters()  {, J4 q, N4 H( j
# 打印出需要训练的层. q; B4 d0 p' D8 k
print("Params to learn:")2 g8 @4 h7 y% x2 r1 R
if feature_extract:
% x% A# a0 p. c8 C    params_to_update = []* e$ g$ s: [  b" Z4 ]1 g
    for name, param in model_ft.named_parameters():
8 }) q' w9 U5 c& z9 i        if param.requires_grad == True:
9 `, d$ G% c# ?7 h& p1 K9 U            params_to_update.append(param)! @5 }) X' w( u. Z9 _* m+ f* r
            print("\t", name)& y& ?) d6 ^; [: z% C( t- S
else:: g8 i5 }/ N" j: a5 @8 X
    for name, param in model_ft.named_parameters():
# N# o1 K) c2 s% i2 F$ s        if param.requires_grad ==True:
% j1 v, r" K8 {! L, @' H1 [            print("\t", name)* ^6 `3 X" h& ?: P& n2 S8 W

- Y4 `# t- j' z7 @. N1 d1# m; ]. |8 g% u
2. {' M& W+ Q! G' I6 Z$ N+ z
3
  S1 u) Q) x# Q. D+ G2 @# d4% c$ F/ A# `4 E0 w% F+ ]
5
  B# t" d( V3 y7 `6! J+ G! L* Y& J% E, U$ S
7& N/ n9 H5 m% T' }5 v- L# C, F) P9 [
83 k  ?; k3 l/ E' {4 N% T( d# B* G9 S$ l
9; [$ ]: c; t9 C0 X2 [
10
  o% \6 B! @7 H6 O7 w* G; }2 i0 D3 A11
/ E" w8 \7 a4 Q9 a9 P12# r% W6 T) Y$ R' l! [; q
13% ?8 J- r# ~# V0 ~5 w
14
5 l" b1 T( H( E$ {- J15
' ?. {! T6 B+ _/ i5 _% q, |16
; d9 X4 k3 f3 ]" |$ m' F4 z& l9 ?" E177 b+ t! r) s4 m0 [; X
18
$ p$ j; ]7 b. {, G* N19" h) D9 Y0 ?& Y: W
20, P: p5 m" z4 b7 F) f
21& r& J0 g  ]9 }
22# g8 l5 @0 `7 e/ a
23
, x$ g- H0 j9 ?  l9 q# l! ?0 N& W; QParams to learn:2 z5 e1 m& W# l) v1 E
         fc.0.weight
3 w& d! H+ s& P9 C: q         fc.0.bias
/ s, X/ Y7 t" M% ]# v* ]7 y' p1* V# A1 E5 q1 |4 [: W2 h9 r
2! ~! y9 j1 V' y& ^6 h8 w
3
3 c6 |% x6 W  `2 \" O! @5 K7. 训练与预测7 I9 a. x& h1 x0 K1 c2 p) z
7.1 优化器设置
  T2 G& N6 J  l$ ]' @$ V2 |, I# 优化器设置
2 t- q2 a) E, \optimizer_ft  = optim.Adam(params_to_update, lr = 1e-2)% Y7 h! d3 A6 Y+ R) V8 O
# 学习率衰减策略' F; b/ T! L( h
scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)+ N4 |8 t! a- p
# 学习率每7个epoch衰减为原来的1/10; ^( Y  H! i* `) V5 Q
# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算/ Q' U. M4 B3 h' C. ^: A% p

( ^1 @% Y" J  C8 @criterion = nn.NLLLoss()
8 v% R/ F5 H- P! j5 A" m& E1/ U3 Y5 _+ d! j; z5 ?
2
/ I  l+ F  }; W; |2 m6 N/ V3
* f- T1 T# Y) m' b& y- K8 Z4
* j% L) w( ]7 t  B/ `5. j5 e# q( L- [
61 y! O+ _3 V& m/ [/ n, p  ]& P
7
  [* e; x& |. x# K  ?# Z8
+ B. t' X+ \) w/ H! T3 v# 定义训练函数- _  s+ X# k. v& v
#is_inception:要不要用其他的网络
. X' c3 P% j7 n( W8 o5 Pdef train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):+ @% G1 }1 u! l  a* C" Q3 Y1 b9 z
    since = time.time()
( v: K6 ]% j% G; |5 W  }    #保存最好的准确率# K# H9 I- D' P# W" k
    best_acc = 0* M7 n+ b; p5 U$ w: Q
    """" s* C' ^: ]" k5 `7 q
    checkpoint = torch.load(filename)* `/ x; v) Z. G4 k7 t! f
    best_acc = checkpoint['best_acc']0 x' f; N* \5 s- ~5 r7 n  y- g
    model.load_state_dict(checkpoint['state_dict'])
8 ~0 i% E" A- J' ]& B    optimizer.load_state_dict(checkpoint['optimizer']), n9 @3 ~; q& s' w* @! Q
    model.class_to_idx = checkpoint['mapping']
4 H' r/ D3 ]* r    """3 Y# m6 |$ E7 R( S
    #指定用GPU还是CPU
! h0 {  k5 v/ ?# g1 ~' ?& s/ x    model.to(device)% k3 \5 d- S: W6 Z/ v
    #下面是为展示做的) [1 w9 h: E1 _2 J; i
    val_acc_history = []% \5 @7 J7 F& q
    train_acc_history = []# Y6 A- O# d, K1 K% J1 J
    train_losses = []$ }' s" h0 j6 Y% l" _+ S8 s
    valid_losses = []
* g' i# U3 y' w/ E+ M, m    LRs = [optimizer.param_groups[0]['lr']]# x3 D3 C5 `4 `! G( j  X; Y8 e( x
    #最好的一次存下来" I" F6 h! v/ o' O1 T! J- k7 f
    best_model_wts = copy.deepcopy(model.state_dict())! b) x# I& F4 {5 x$ m0 L+ y
, {1 Z) O- z7 A( U' H0 g- b& `
    for epoch in range(num_epochs):
% {5 u4 p$ }9 S/ e4 M0 J+ \        print('Epoch {}/{}'.format(epoch, num_epochs - 1))
% H, j# K' E/ S        print('-' * 10)2 V' w( {( R1 z9 F: R6 h( Q

+ ^0 Y/ F# M! g3 G' {& J        # 训练和验证5 E% ^' R8 d! @5 H0 X
        for phase in ['train', 'valid']:' e% b2 m" f2 \# Z: x* D; e
            if phase == 'train':/ u1 w# ~' {6 E+ l
                model.train()  # 训练" Y: Y' U; L% W- `  g
            else:2 ~; J& j( m3 v
                model.eval()   # 验证
5 E1 d) s  D6 o) W- G2 m. V* G' p( h! t$ j; X  r, I
            running_loss = 0.0
& W* O$ g& x; D; j            running_corrects = 07 x$ g! {' g6 Y. T+ J9 x

- Y  {3 w5 }4 _+ Y1 e5 z% ~            # 把数据都取个遍
6 S; f& \1 D2 ~& C2 N6 s/ m: w8 g            for inputs, labels in dataloaders[phase]:4 K, ^) Z2 L* W# ?! n* V" D
                #下面是将inputs,labels传到GPU! E1 ~' p0 P2 Q4 D4 ~1 k
                inputs = inputs.to(device)
- J$ K7 M  K& i3 |3 S% U                labels = labels.to(device)
) Z; I- ]$ v4 B2 \5 B7 \2 U/ h# G
                # 清零* j1 [! Y% E$ p/ d2 x# n* d* y: [2 z
                optimizer.zero_grad()3 {2 `: s; I" D
                # 只有训练的时候计算和更新梯度8 {0 S' |$ p7 }
                with torch.set_grad_enabled(phase == 'train'):
5 Z6 N, T8 S8 ?' m                    #if这面不需要计算,可忽略5 f2 f7 b: ]; b7 Y" {% [
                    if is_inception and phase == 'train':
/ ?% O1 g9 y  B( ]                        outputs, aux_outputs = model(inputs)
* J1 }* }/ w2 e( N! H& g                        loss1 = criterion(outputs, labels)1 g( ]  i( K5 V* m5 U% J2 W3 Q
                        loss2 = criterion(aux_outputs, labels)
9 I* w& D2 X+ k( O" T& o# G7 p                        loss = loss1 + 0.4*loss2
2 |' I/ L/ v& c% v1 Q                    else:#resnet执行的是这里& F/ |6 A/ h- e+ i' u( c
                        outputs = model(inputs)
1 |  P: u7 m( A; u( P& @, a; m                        loss = criterion(outputs, labels)/ i+ r3 `7 a, D" s/ N* o! w
; e" ?7 ?) `' p+ m
                        #概率最大的返回preds$ Q' Z6 M% t! i% R' F) n
                    _, preds = torch.max(outputs, 1)
. f% l$ Z& l( N* D2 v
/ u! h0 c( T& K" A  Z                    # 训练阶段更新权重
1 E$ }. M$ G* w( ?+ y3 W                    if phase == 'train':
" |6 c( e! G5 H3 L$ {3 K                        loss.backward(); G# m/ x5 {" N9 G4 M- `
                        optimizer.step()
$ |& Q, q" Z1 l3 I
: S; |7 _. i8 c5 _  H$ L( ?6 z                # 计算损失
' N- Y# c6 a# F7 y9 B                running_loss += loss.item() * inputs.size(0)4 Q  R) L9 k' g
                running_corrects += torch.sum(preds == labels.data)& b% D; m8 A, d

/ @* g% _' N4 [            #打印操作2 J* Z, N( W" l2 v
            epoch_loss = running_loss / len(dataloaders[phase].dataset)
+ I  |7 b) n% t4 h            epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)- V& d6 w# V. b) X& ]% b0 o  W
' c6 ]4 Z$ e6 p3 Q3 M/ N

1 {0 ~/ E& o. H# ~  J            time_elapsed = time.time() - since
, Q- `6 x/ l' W            print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
' \) J3 W3 H/ Q( D            print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))
2 x+ p; _  T2 W' c6 n: X0 q8 H+ v+ L( O
6 {7 R3 e( L! w8 w
            # 得到最好那次的模型& Q$ I, Q  L  x# Y6 A9 N( K/ M
            if phase == 'valid' and epoch_acc > best_acc:% K3 Y6 H2 [7 m
                best_acc = epoch_acc2 l; ?+ V7 I& l! x+ V
                #模型保存
2 R. ], ?" s1 q. z( f                best_model_wts = copy.deepcopy(model.state_dict())
6 I  Q( C8 V0 g' K$ x, }                state = {
8 f2 Y2 b: s; R; j) f- [# Q8 H2 z                    #tate_dict变量存放训练过程中需要学习的权重和偏执系数
% [* C. b/ R1 m# O( T3 g                  'state_dict': model.state_dict(),+ d1 Y1 X; M; T
                  'best_acc': best_acc,$ p  q; h$ z" u
                  'optimizer' : optimizer.state_dict(),' ~: x( u- l+ h& {% Q  ~; }
                }
8 [/ F% L7 W8 v: b                torch.save(state, filename)
5 i, `! A, Z* U; U' j6 f: j            if phase == 'valid':! p! P2 O0 l* L+ T6 k) [
                val_acc_history.append(epoch_acc)& g# r$ E# l0 d: B
                valid_losses.append(epoch_loss)
1 y" s+ H' T# q: E" d& o# e                scheduler.step(epoch_loss)1 ~* L- |$ [$ C* \$ _
            if phase == 'train':
. \  @/ p' n8 {+ @1 X                train_acc_history.append(epoch_acc)
( F% k6 P1 \# W4 L* W4 _/ ^& i, ?6 r                train_losses.append(epoch_loss)$ F- V* q- Z4 N8 M* b2 d
; v( G! k: F2 B+ N- O3 D7 U$ W. S! [2 z
        print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))
( N% e% d7 f% H6 a, X8 ?: g        LRs.append(optimizer.param_groups[0]['lr'])
( }" W8 h7 p* g: r        print(). x1 E  D8 q& A3 C$ l

) U, C2 R5 T! N; M  S  }    time_elapsed = time.time() - since) D, [: Y7 j  L4 n% F' o
    print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
; c" J3 k9 r! }1 Q" e    print('Best val Acc: {:4f}'.format(best_acc))( ~4 S* \3 N. a6 q

; Z9 ~# ^1 \  E    # 保存训练完后用最好的一次当做模型最终的结果
! E# \  S( a6 J# v    model.load_state_dict(best_model_wts)
4 @. H$ A, I* S3 T    return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs * d; P2 t3 o# x

% R5 T& e0 f7 y, w" h/ E% a1 a+ p4 c2 }2 Y
1
5 }9 M- |/ a- G: B+ B2
0 U* U; b2 ?& E3 y8 _2 W3
3 ]) N+ m: M; H6 U3 j9 E" D0 I% d# T45 A- X' [2 T+ C9 O
5, ]3 [% j9 k$ w- R. X
6
- x" c4 T" ^" _( {7
  f" ^+ F% R, }& r& f8$ r! z# H9 F0 P5 Q+ l' E
9
" [, p# J5 e8 g# j104 x+ h2 l4 R) m* _! I0 @* z* U
11' o, C) z3 j8 i4 ~5 x
12: v* L) K+ k- N5 j+ K% K3 v
13
/ ~+ f. o$ }9 h6 H' B+ A! h. [7 I" I4 ]14
" E% t, t$ U9 [* \% g$ R8 k' R15! k8 q: {8 o5 O, ?
16! Z& G9 z8 d& R; X" O
17# ?8 B+ V0 X; c, o5 Z
18
7 ?6 o% N7 j& o; `19
- P7 g7 R0 x" l2 J0 [200 j& g6 g& p1 [6 G  s& G
21
9 c4 p! S+ ]2 t/ R$ d22
4 B3 _! h6 R9 f9 I23# F# ~" O5 G8 B0 h& v6 ]% [* B
24$ @$ Q4 J5 H! A
254 ^. ?' m$ G% L1 N: b4 ^6 Z& |$ U
265 ?6 Y) G9 S1 w( W' u8 F
27
7 `9 G7 n& ~! o( G2 l- _5 ^286 M" D4 ~7 }) G) p! l
29
, _- q9 m1 |+ `& h7 C3 Q30, D0 h) X, a" V; O; j( _
31) Y, V$ N$ |) }0 |- p
32/ B7 V3 m7 a( ^
33, O' h7 u8 \4 D0 U7 j7 w
34) u& ?7 P+ a4 {  m5 \
35
! G! j  }8 Q- w$ W9 D& |36% V5 Z5 q( n, J: w
37
' D) m* q: s7 D38
* W% s7 n- }/ D2 F399 Q9 j; N) F3 \9 J) n, U
40
+ m, l4 `9 k5 L41
6 D1 X: m2 r" ^4 n) _42* X7 z7 n: _  f6 w4 U
43
7 r# s  W- j! B) K; w2 Y44" l$ o$ I( C7 h
454 P9 a7 \+ n2 r: b' ]4 u  u
460 r# ^8 j- y# b$ L2 `
47
. P/ o( N# N4 V48; T9 B* {! B6 H
490 E* U, ?; T4 h, A. A5 q5 _  R  M
504 I  d+ s; n) \. |9 ^- O2 G
51, n: d3 c7 z0 S# S# r/ [3 n! ]' S
52
3 v2 S3 L( v: D7 @0 m& T* t53' v8 h6 u5 d: u+ \" |+ G+ H5 k# u
54
' a7 v5 v/ i8 U* @  Q! K. e; M& @8 u55
' a( [6 H8 V7 [' H  K% Q56
5 M8 Y  c4 h7 r( t57
" Q" g" C% n) e% [' w, A; B/ j# B58& i- l! }0 Y* v6 I7 b
59
$ @& ?. `4 X* |: U* y- i/ r% @60& O+ m# v# e% q6 h8 A
61
; q2 K- u9 K5 U3 R  O6 z# J62' f; n' a/ F: W4 W
635 q  u/ S9 [* c
64  X4 v7 ~! ?! ^
65" N( D2 Z( B2 h* ?2 e
667 y- q4 [" U# Q! w/ N
67
3 a3 g1 D+ N" ?68
( U) d% B; E% ^% `: R69
/ E; g7 g4 t! ?* ~5 V3 o) ~& L70
, v* T: \5 B% x71
# ^% {8 G  O% J. ]72
8 a+ s+ }/ Q$ `! I$ n73
; Y1 g( @# |+ b  z. y745 @( G* t9 A: [# |
758 q( Q8 ~9 z+ `6 {3 l
76' Q5 B% E  j0 p' O& g
77
; A: a" k9 e/ Z" G: V' H78
4 ?) k+ W* r4 M% Z& R: T79
- `( @1 J: c& U3 c0 g" K0 v80
* c* g- T3 h5 H81
% R) S8 A' D) [+ m82# m% W2 r8 C5 J7 L7 f) H
83
$ B* F/ U: K( x84$ L6 z: R3 O6 Q8 R
85
( X8 T. S$ m8 Q86
- m% O/ [2 o) Q& n+ W& c  W% }( j87
5 D% E; [3 r7 X* f. y0 C0 Z. {88
- ?# K, Z" h/ G3 G89
) O' X( Y  |/ K5 z( k904 b! X% L. B; C/ g2 {: [
91
: u/ E' ~' o& h8 G" d92
; x# b6 p+ q2 y+ V1 F0 q93
% h) L0 J9 W7 s& N1 D! h94
, V2 T3 Y' n( o95& t/ D; A* K: n' s1 ^& f
961 r( U) k1 m. |; j" D2 @% `
97
7 i( Q$ M$ Q9 {  z6 \% \* a& g98
5 h) v, b4 P4 D99
; l9 r) v1 Y6 _/ d100
. |  d3 T# {2 V5 e101/ z5 N( A, I( r! C6 u- F/ _8 v
102
/ W& L0 K. N4 p$ C8 s3 Y103
+ b0 D5 v& X: F( i- z104
9 B; n( ?% x* t, R+ f105. m, d+ B9 y1 Z1 t+ ?) _
106# a; ~' i& f$ k* S! e
107  }/ O& O# ?3 R# X/ B7 f
108. I- ~! ]5 n% J, Y
1095 A3 N& P' j1 e* R
110
3 W. r6 G" Z1 o8 R111
* H( A+ L6 T% c. L2 k, `112
' m7 k, q' Y; K7.2 开始训练模型, B* l, n7 n* S, ?) l. w; b
我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次" U3 f) Q" e! j6 T/ V
1 g7 V" {' [4 k$ S% o. z
#若太慢,把epoch调低,迭代50次可能好些
! ?* X) x( l+ V+ T, h% D$ U#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合, _- {! A8 v: W+ i- \/ m
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"))
; P$ e6 F2 |- ~' ?
* q2 v* y9 M9 D" ~$ |1; E' b8 B- X# u; e
2
4 w7 f7 W8 b) G3 B3' l, c. b* G4 G5 K. ~+ e. m7 E
48 r0 v2 [+ q$ w
Epoch 0/4
- ^$ Z# B6 X' u6 K/ a----------
( R1 |& W1 v  UTime elapsed 29m 41s
9 E* w4 w6 N, U' ~train Loss: 10.4774 Acc: 0.3147
" ~8 [' t& ]4 LTime elapsed 32m 54s0 n) F% v- j2 E- J5 f3 @
valid Loss: 8.2902 Acc: 0.4719
, M' l1 |( ~7 aOptimizer learning rate : 0.0010000
* M1 Q4 m, S" H
) f9 _+ F  |8 Q% O/ z) h( x3 [" |Epoch 1/4
/ i/ V8 C. Z, L! L6 Z' M* X$ k0 h5 S----------
4 H7 E# K5 Z6 g) b( y8 jTime elapsed 60m 11s
, \7 H' C9 p9 L& z  c2 w+ n+ Q+ Vtrain Loss: 2.3126 Acc: 0.7053
7 n, C& R& N; j; E0 Z7 m9 MTime elapsed 63m 16s- T' k, z. u; G9 E3 ]& u
valid Loss: 3.2325 Acc: 0.6626
! S7 M+ Z& x* x" Q2 k+ `Optimizer learning rate : 0.0100000
4 {; {& q6 _$ ~* L0 Q( W$ D
" |8 v/ B4 P( S8 a5 J& i/ p- B6 IEpoch 2/4
3 C# ?/ y4 |5 _+ H5 |2 l1 B----------0 C! `0 q6 s% ?6 z% ?
Time elapsed 90m 58s7 C" w. i" V6 L7 b' s: f% b& R. s1 A
train Loss: 9.9720 Acc: 0.4734
, X+ Q" }. {: R8 yTime elapsed 94m 4s
! ?/ K5 {7 l' H! w5 l+ Vvalid Loss: 14.0426 Acc: 0.44139 Q; x1 n! U3 s4 I; ?
Optimizer learning rate : 0.0001000) ]( @0 ^" V) G/ ~2 a6 y0 g7 Z2 u
; B9 z" `9 O1 Z5 S  [1 P: I# _2 {9 |
Epoch 3/4
! }1 `0 s& z. O  T: _----------
2 Z$ P3 M/ C0 j. b, ^Time elapsed 132m 49s
0 L7 o7 h, i; ~7 M$ Itrain Loss: 5.4290 Acc: 0.65488 F& I3 {. F0 C- M& k/ p5 b
Time elapsed 138m 49s
7 u+ F, G8 y! b: d. N+ Avalid Loss: 6.4208 Acc: 0.60270 z6 k+ l0 ^. h1 L* x! h4 S
Optimizer learning rate : 0.01000007 C; d9 Y7 w2 |( E* ^5 \
6 Z; M; Q/ S! V2 M5 D7 q
Epoch 4/4
9 l. A/ }' e* h" V+ {----------
5 A+ H; _* j: fTime elapsed 195m 56s
+ R6 `* X# e# n6 N# htrain Loss: 8.8911 Acc: 0.55198 [% ^3 S, P, u5 ~. S5 e: G9 |" O
Time elapsed 199m 16s$ I/ I; x4 Y0 t3 X' {1 S
valid Loss: 13.2221 Acc: 0.4914
& L0 R0 m7 j* Y+ [7 }' hOptimizer learning rate : 0.0010000) f1 X. C1 R( e. y( {

& D+ x& @! U8 v& i5 @Training complete in 199m 16s
: [( U. A7 d* \, H; pBest val Acc: 0.662592
4 h& m3 \# q' n" i* m+ m# o% v4 u. V# a' \1 t6 n
1# N6 w8 E* K( ^. T( Y: H
2! L8 Q3 O9 Q( q: G# l$ l) B
3
/ D; M6 V7 J$ x' l" I- i4. P( B( I0 O' g& o$ I
5$ g; Z. }" w8 z- [- ^
6
) S- [3 r- c8 w7 ]+ y2 Z9 m, ^9 v7
: p& J: S$ b. a% z7 D8 b9 Q% X' P: r8( w; E1 Z" W1 m. d4 z; P  s, t  S* M
9
; q- E  I- r! _5 }) q* G  o; a10" f0 M! p1 ^# ^4 w) P& L, R( K, B2 n/ z
111 W2 ?  U* v9 l: f4 @7 ]/ C
12  X% K: k1 ?1 M! y
13
; |. U* \, T: X! g' j. G141 i3 N& o  s, B3 \
15% |& D9 D# o# u5 x  f" j
16  _# Q+ Q9 h( \  U7 e4 e! c
17
9 ?$ S; I0 s9 d6 y  |+ n6 u185 F& S( A9 n* J
19
- k1 ]& k, {: k8 k+ N20, [8 K4 X/ h2 B( B
21& q( q2 }0 Y" ^" R# W: ]
22
; E7 |: r* u% [6 C23- N. B5 }3 Q0 g( p/ ]5 S
24! r( W1 ~7 g$ u  I& J
25
& u/ X# r* i; ~7 x3 d& n26
5 P9 C: C3 m- Y8 c- c8 w+ Z& Q27
; D# J$ C2 ~, o7 E7 K0 m3 F  x* t28
2 g8 x9 M4 `3 m' D# o$ w29
) R: H: o9 g6 R5 `; z$ `3 s5 Z! j30
/ t' y5 }7 c# _( M31
/ w+ _( N5 x% |: I7 e) {32
0 f; Y: D" e2 v33
3 @* m: Y% [& b8 N0 W34
) N4 X. n* T) P: k2 S35+ ?0 U' T5 S5 ], M' f6 ]' Z. ^1 W: `
361 {- }- M( }. J) w2 M' y3 p
37, T4 s3 P" y" D4 X* f% \. Q7 {- S
38: h* p: A4 X2 @
397 w4 i: ~' h5 A6 _  C
40
2 o. G: r" G* `0 w( e6 ^419 W5 _' Q# O5 |3 J) T3 S3 z$ L
42: F6 C7 g. G7 I; g2 F8 [
7.3 训练所有层
7 ?9 I" e' {) p: [# F! t4 D1 v# 将全部网络解锁进行训练
% P$ i7 I1 c  {* m+ z) Q7 v# ?for param in model_ft.parameters():
6 [$ y+ }. }: f; [6 E( r    param.requires_grad = True1 W. f6 m4 ^/ X7 J' c/ j* p
) f) J: q: t% H1 u1 |) b! ]1 |
# 再继续训练所有的参数,学习率调小一点\8 B4 J/ ]# J9 n! h' _
optimizer = optim.Adam(params_to_update, lr = 1e-4). y1 w  A9 X4 q! H! g' j! C
scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)
6 I1 a$ \* u- U% V2 s: l4 v; ~; P, ]. w4 V: C, \
# 损失函数- j; N. s9 ~& Y0 B
criterion = nn.NLLLoss(); A7 u* j" ?& c" x% i
1
0 T' v: u3 [  A, ^5 A- n) H  a2( A7 |9 ?, j" [& V( {& A. U
3" F4 h" s& U: N3 L9 m  C' z
4
# ~) ?1 a0 D* s! N9 ]5
( r+ [1 t* ]2 Y3 L6. _5 I, g; ]! p1 b* ^* ?6 Z
7
8 R+ @; j1 Z$ E9 f3 f9 o8
! }! X, T( A: x8 O3 |3 N9
% N, k  b) n+ k2 D% m10
0 i/ Z! h. r6 E. G# 加载保存的参数
2 k' v+ {' ?" P/ f, q# g4 p& G/ D# 并在原有的模型基础上继续训练2 Q8 \$ j! I9 C% m
# 下面保存的是刚刚训练效果较好的路径
4 ]/ C9 Z" K$ `, d7 bcheckpoint = torch.load(filename)
) v+ T! Z* p9 b7 I6 `" u; {& a. pbest_acc = checkpoint['best_acc']+ c* S' u# m) i6 P7 `
model_ft.load_state_dict(checkpoint['state_dict'])
% N! D' O- t( `' I# [optimizer.load_state_dict(checkpoint['optimizer'])
9 f; D" ~0 l% r1 z1
5 s' T3 Z# c, f5 l! b% H2& t) G, }% J( V6 U9 Z  j6 a
3- x- ^: Y& A% `8 ?
4/ v- f; a7 z8 O9 `1 P% C& X: a- r
5
2 \# B- X( J/ O% Z6
7 R/ w+ A4 t: h* p2 }7/ a: G; p  U' G0 r) M: o* _
开始训练
; ]* i  R' L# H$ A1 N. R7 z  W注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考- j; V+ U0 ?) f/ p

) j- ~: y3 _7 _6 m4 a& qmodel_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"))
! I: V; V0 A2 \9 q( S. R1) H5 Y& M+ \* e5 p; v9 q
Epoch 0/1& w8 A- S3 D- d; X
----------
% m! y; o& Q7 R  xTime elapsed 35m 22s
7 R/ c1 p8 U( Mtrain Loss: 1.7636 Acc: 0.7346
$ O& z' a) [2 M" z! u: zTime elapsed 38m 42s
3 ?  q- M. O1 J/ G# H/ Pvalid Loss: 3.6377 Acc: 0.6455
, b* n. l: `; S3 ~% t) G+ IOptimizer learning rate : 0.0010000" @; z- e" D8 z. [

5 l* E* C* O5 y3 u+ K6 J- y% YEpoch 1/1% p2 l6 \- j; ?+ {' N- i
----------
- z1 S' A* v" @* gTime elapsed 82m 59s
3 y$ f9 ^3 @+ A9 z' E, T$ Q$ wtrain Loss: 1.7543 Acc: 0.7340
0 l% {' u" E  J2 g) ATime elapsed 86m 11s
$ z. r* s7 s4 b! B& wvalid Loss: 3.8275 Acc: 0.6137
; X7 P9 j( }8 N3 h7 XOptimizer learning rate : 0.0010000
8 H' E7 v, V5 ^
  q- m+ m) \4 t4 G+ X& P; iTraining complete in 86m 11s$ J  ~2 I- X% ^9 x* n5 T, [
Best val Acc: 0.645477
9 v2 n# j) D8 c7 ]) Y% E, S8 Z2 I! ^6 w
1- w1 x) k) o/ G( o8 r/ R' p
2+ f; d) j2 D. o2 o% I, R
3( n5 h' t* N! r1 T& k9 C# M0 J
4
5 x0 T1 U6 j/ B0 ?! ~  n/ j5
& L1 \. @5 I6 \' Q. P: a65 I( `6 x7 Y. `& b9 J
7' x* l; [2 |( a: u' @% o- Q
8
& E: F; s4 R9 u  x; t9 B4 d% q+ |9
$ `% d; X. J! ^# o9 @3 z' I9 s10
- \  a/ ^% H% ~3 E11
/ d3 N7 e* Y6 L; |7 P$ a12+ K" I7 O9 O9 A6 A
13
9 d. i  o) O( U5 [4 Z; O) {+ @14  D) [. U; t& w* [% {- M" ^
15. J( Y& G- O/ @! a* \& x/ J. A" y& w
165 }1 e4 T5 P5 R8 Z
17  c7 d( M9 r( W6 V4 b, s( j# A. c# u
18$ K  v3 t& W7 h% ^( W
8. 加载已经训练的模型  a- v' m: M0 b+ P! Y, q
相当于做一次简单的前向传播(逻辑推理),不用更新参数
% z2 D4 F# }9 ~  U' }5 E3 z% w; H( d+ v' I
model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
; n& }0 g' U- o# N' T4 C6 z% o
+ Z. H8 T  a4 D7 ]3 X# GPU 模式2 V; m" ^5 }9 f" Y' b7 ^
model_ft = model_ft.to(device) # 扔到GPU中( R7 G5 l/ k+ i! X, i& g

. A6 g2 i2 y/ g! ~4 a: H' h# 保存文件的名字
/ g8 r9 A, o, z' @# S8 Jfilename='checkpoint.pth'
* ]& W. y' e# n4 |
' I% Y( m8 K( ]# 加载模型
( Q% {) D+ r' |/ P$ h- ~checkpoint = torch.load(filename)9 P; n1 s5 @: z* d* _' F0 ~$ X% z
best_acc = checkpoint['best_acc']( O& x; \5 K  _) q
model_ft.load_state_dict(checkpoint['state_dict'])
, Z' U) f8 T+ F  e9 i* h1: j. l3 h; q( ~
2
% |7 E5 c% y# J$ B' `3
, ~. z. i5 A1 C9 j$ p42 ?" {* ~3 o& p( |' L& O5 P: g
59 D; M1 [! u0 i1 i' X3 S8 V+ I
6! J+ \1 r. l% T! L/ _
77 \) `; W+ X4 K/ x/ i3 a. r: c
8
4 u6 {' w6 u3 ^' C$ b" i: M7 N" s9' i# C* D4 }/ Q6 A2 |0 p# m  m4 M
107 b" _6 W, A; O0 Y! f, |8 h8 z
11
+ \5 V6 T5 l- J+ P% v12
- x) k* T2 H, C" o1 _8 `/ j9 N<All keys matched successfully>
% U" L" ?# `. X; ~. P) {1
' n9 v( _! @5 \( a0 y' r& A6 ?def process_image(image_path):
% v4 s  J0 g- M( p    # 读取测试集数据$ {' D) Q: W6 c6 |* B
    img = Image.open(image_path)' I2 j& y$ x* H: n5 H3 E+ D9 V& m! K
    # Resize, thumbnail方法只能进行比例缩小,所以进行判断
+ x; |$ p; _3 N# q: w    # 与Resize不同
% J/ e5 M) ]% {8 Y, e' J    # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小, W( d) G) s) \! w
    # 而且对象调用方法会直接改变其大小,返回None1 U2 D( ]- J4 E3 ]9 g) w) K& f; k
    if img.size[0] > img.size[1]:$ q0 S: ]+ I5 j9 S1 S
        img.thumbnail((10000, 256))
/ J0 r; ^! {, _5 K: d5 d  f    else:
* T1 [" o5 ^& ~- [! ?. ]        img.thumbnail((256, 10000))/ `9 E( S% d( q- E" U% U

% Y7 Y. R1 I: t* s9 H( X    # crop操作, 将图像再次裁剪为 224 * 224+ s. F' s0 J' p4 u5 G
    left_margin = (img.width - 224) / 2 # 取中间的部分
6 ]; M' b% T) ^$ X! O! ?    bottom_margin = (img.height - 224) / 2 # a5 }. g1 p5 I. h" @% h/ L# O
    right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度
" g$ l7 Y7 K# {9 h  B) y$ q    top_margin = bottom_margin + 224( M9 p# b# y9 [: B* I- }1 y! A
# D0 e* \2 K& F: R3 F
    img = img.crop((left_margin, bottom_margin, right_margin, top_margin))
/ t& u1 a. t8 x& ~2 s3 @8 {  _' T# p' v* N* z
    # 相同预处理的方法$ V8 a, G" e9 j3 p* B
    # 归一化& M5 t# @. J3 T$ {* P1 Y$ n6 [- K
    img = np.array(img) / 255& A2 J5 u1 w: Q2 c) l
    mean = np.array([0.485, 0.456, 0.406])' }) N9 v3 ]; M
    std = np.array([0.229, 0.224, 0.225])
, q+ q4 U6 D; _    img = (img - mean) / std9 ^( u* J# a& y' V! f0 \2 g
  N9 U3 |2 [( h
    # 注意颜色通道和位置
) {& |  M2 ~# G  q: |# M    img = img.transpose((2, 0, 1))& ?( r& _, Z- J; v

9 e! s0 \* C0 n8 @$ L) l& o) c2 f    return img
& [. i& Y9 Y* c$ \  B8 {( f4 e/ ]) w1 U0 V, n
def imshow(image, ax = None, title = None):/ j8 @$ {8 Y! t4 v; B3 d
    """展示数据""": x' N) k0 O( m2 ^. `  J! R$ Y7 j# g
    if ax is None:1 O7 N6 W% `, ]0 }
        fig, ax = plt.subplots()
( o. p6 x5 L& {$ y# S) O( o# C! _2 {6 |- `! c# D: R9 n
    # 颜色通道进行还原
0 G7 w) |6 F9 j1 L* n% b    image = np.array(image).transpose((1, 2, 0))
# N% [' d4 u* @% C8 c$ {/ I; `6 L, k2 q$ {/ m
    # 预处理还原
, S0 q  u. N& P" R    mean = np.array([0.485, 0.456, 0.406])
2 \" {% X7 L5 I0 b; u, ~    std = np.array([0.229, 0.224, 0.225])
1 X- L- \) P) f    image = std * image + mean
5 u% d  x) P& V& ]! u    image = np.clip(image, 0, 1)8 Q2 d# p$ Q. X# L
7 f. F" f. M- C2 y
    ax.imshow(image)
0 L) M: c, t, Z* L    ax.set_title(title)
3 j) b* F. M4 J. d  c
8 X& \0 n! k) q  r4 |. ~- T    return ax
( l% j, D' [" w6 p( X6 T% M7 p7 P2 U
image_path = r'./flower_data/valid/3/image_06621.jpg'
+ }/ F* |2 f8 z3 b$ `! l# fimg = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理
  o( m4 w6 E  M5 S0 oimshow(img)
6 _! U4 R: ]5 G* N3 \% S' K, ?" t: Y# z1 e; l% o& y& J
14 V' |# b) T, t5 B5 F/ e: S
2% x; _! u# P% V9 B0 v
3
$ B/ a8 [3 C6 t/ b" S0 S1 k. w* J47 ?; F: t3 `+ h( J
59 k. ^3 v0 w& v8 R4 z5 @. S
6
4 A+ x/ ^  _2 J5 J. \1 s7$ h, Q3 c3 R" d: A: K) e
85 {$ m' Q. e1 v( l2 }* L: Z0 ^
9
3 `# g  `& S! K/ e10# @8 H8 G! O9 e. l9 z: a* B
118 K2 c3 M+ ^, h$ L
12
9 k+ U* E, v- D4 Q! I2 F130 b; Y) u( g; V; m
14$ p5 k+ b5 }+ V6 k5 |
15
, T& I% {) b$ X16; q, T2 a$ F8 u3 a6 I! F
17
9 |) D; b7 f. j% ?: I9 T9 \181 E6 q# X. U. b7 ~3 ?
19
# f6 [) `- x5 k6 \4 M2 l1 ]& `* O# t20
9 T6 @# }+ M, q) F( i) q7 \; g! z21  a& r5 ^, p1 v7 O, f6 U
22
2 u2 p: Q3 Q, e% s23
9 S4 v7 A9 x5 g2 m9 Q3 F) Y24& X( D: L+ Y6 p  _) `2 u
256 U' z( y2 y7 k( |% v$ I
26
8 C% }; v+ {: P0 V27) K* ^- h1 M# y8 @5 }
287 g8 \8 t# [) x( z
29
) p7 q' B& O: y6 N307 m7 d3 d0 O: x. f
31
5 {) O0 ~: y5 }' g4 v: @% h325 N! U* [4 l7 z- t/ ~& u
33( J* x/ s- W# M9 p5 c  A7 J
34
! b+ ]" i: p+ p) q& @9 |6 x35/ e8 j/ ^# _! j- r- A8 _, e6 d
36
& l" h8 ?1 ~& S( v3 K37
- w/ p" f6 l) C38" q9 d/ B- `+ X6 [# C
39+ Q8 y! Q+ C% Z+ {
40
$ [5 i* L6 i% Z0 z) O41
1 b, N, B+ c2 u) g' ]7 C42
3 ~2 V" ^1 S! a0 I$ m: n- |! h+ z43  [8 J! t  @$ y* g! w" |( ^
44; u' S- p- x% k5 a8 E
45  n/ _8 p6 F. _) b3 `' Q) ^
46# k5 S* T: i+ Y
47
9 }/ E& r; s6 @, [! N- }! ]48
8 p, B% x8 P# o- b49
; o4 l8 r6 ^4 }8 U" y- g505 j2 `& N8 o) @  Q# f2 [
51
% w2 V, k; J# Y- P4 _521 d9 R0 d$ A) A0 o3 n0 W, t% C+ ]
53+ ?5 M3 h# ~1 L+ Y" y1 D/ M
54/ C. o5 j3 w# P# q9 S
<AxesSubplot:>+ _9 K! S( l4 @% C9 V" i
1
! {+ z- W+ I! L7 Q+ d) b; _- }$ x
8 @# k2 i! \, c9 T上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确
+ z. W$ H5 X' a- X; ~: y/ h  z( G7 v
img.shape
) `; s8 O) V0 v- s. P+ q  [16 D0 ~# q. a; X  H+ ~9 Z" v$ q. x
(3, 224, 224)2 ?2 b* S$ z8 |$ g3 m5 R, q* E
1- ~* d9 M; r% E7 e; M4 ?
证明了通道提前了,而且大小没改变! I, _5 U4 l: _

9 z: s+ r# R: W5 m1 K9. 推理: x) o( Y7 n. |4 V. [. k
img.shape* q$ J5 y8 e/ D
/ N4 r& f$ `+ d
# 得到一个batch的测试数据
8 }- h( P# F  o: y! o" hdataiter = iter(dataloaders['valid'])& B7 U3 S1 T0 a( }
images, labels = dataiter.next()
* f, h! `- b  y% c0 _  E; a; O
model_ft.eval()- X9 y/ z  d/ f$ Q& P, X7 y
8 `! k8 W' e0 _" `9 h
if train_on_gpu:
( }+ m: f- ?, {) M$ \    # 前向传播跑一次会得到output+ k; z4 m' q" l- Y
    output = model_ft(images.cuda()), ?$ P7 {- m7 K& q: h
else:8 }6 C, S" H1 y5 q
    output = model_ft(images)
1 D4 g. E+ q; [! `9 m/ N
6 Z6 t1 j1 L- V3 x1 g# batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值% E0 }+ x" w2 f
output.shape
: o8 M5 U- {! j! I# m& e/ C
9 m% f( f) i; L: m0 e0 D/ w  G18 T3 _9 E7 `8 U' O
2
# Q# m" K+ b% ]" f3. R4 G* x) X  {# x' ?6 B) O2 E
4: K+ F9 d0 S4 f' v) R8 W
5& y" S+ ]% C) Z- \3 |8 N. e9 B4 x5 u
6' `. w2 L& ~- k/ n9 o3 Z0 V$ N9 y
7
) J6 d4 F& W5 ~8 q0 K8
, ?7 S7 q7 B# B- e  B9
$ q6 o% _: M: Y10
9 J4 u! }) n9 A/ |: x* w, h0 G11
$ f- f# y& v: w- d8 m- r  l2 u" P12& s- _2 Y7 N( C# K8 r$ y; A
132 U( y9 w0 ?/ |
14
" r& a8 P2 }: w- s: e2 g15! l8 E8 l) \+ H8 @$ E; E# `1 t; M& I
16
2 b1 G5 H4 O) itorch.Size([8, 102])
* b. l$ N& S# U9 ~% w; B1  g1 u2 a7 b: J' X2 e- G
9.1 计算得到最大概率
' j/ h8 g+ N2 D* o8 V$ x_, preds_tensor = torch.max(output, 1)
( u5 b0 E- c. d' x* i! @7 \' S
: f+ q; W9 P, A3 w% T( g2 ?preds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量
, [( h, k* E! ^8 g- @1
1 L' a4 U8 n- C4 k. h27 O, d0 r, s( A( h6 Z! R+ t& b
3+ T4 u$ t$ R6 M. `' M
9.2 展示预测结果: Q, L. Y: X% }+ [
fig = plt.figure(figsize = (20, 20))" ?4 s3 x! W1 {' X' P8 }5 x! X& M3 d
columns = 4/ b; x- m* V- v: T+ m2 \$ `
rows = 2
- C5 A1 v# n( \- E3 j/ l! ^& N7 o3 x( S) G
for idx in range(columns * rows):
7 k( t$ M" q/ T' j" u8 B    ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])
" l6 w; L# R6 {6 v    plt.imshow(im_convert(images[idx]))6 \" y5 b2 p& ]- h
    ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), ' g( o& ]  L" i
                color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red")). ^& F. L+ E0 C0 e1 l4 E! G9 Y7 L
plt.show()# s2 ^" [! n5 o! {
# 绿色的表示预测是对的,红色表示预测错了
2 \4 X5 Q- [8 I8 g1& C# V0 J: [0 q; e4 i
2
/ g" G3 S/ J1 h4 X3 u2 i3
  u$ @4 C$ u' z1 f) t4* ^. {( ^- {( r
5. C6 v! b* I. ~% |9 O' O( B
6
  X8 _9 @' T$ e& |/ g# w72 P! O) L5 e/ Q$ v# l
8
" G. ?& l+ s7 O4 R& b. G! G9
1 n( J) p2 \1 p8 p4 h10+ \% O# h/ ]; J0 k' Q9 _; m
119 x9 F0 [. l5 k0 m. s, F
  ?  U$ m: g. @: _4 ?, F

4 i  W4 M# ~6 F: L# P
1 q5 t4 c) i& d————————————————0 z! C, r+ }2 ~6 I- F, Z9 K- g
版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。/ n- X, C4 i: O4 S
原文链接:https://blog.csdn.net/LeungSr/article/details/126747940. x: {; c' C  a- ?) w/ |* x* o

% @. p# L/ O0 Q6 [9 a! q. S6 s" a4 K$ e. \( I- q: g





欢迎光临 数学建模社区-数学中国 (http://www.madio.net/) Powered by Discuz! X2.5