数学建模社区-数学中国

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

作者: 杨利霞    时间: 2022-9-8 10:41
标题: 【深度学习】 图像识别实战 102鲜花分类(flower 102)实战案例
【深度学习】 图像识别实战 102鲜花分类(flower 102)实战案例" i# M& L7 n+ v# M$ h, C; m
1 {2 q/ M# W* k+ o' a8 a
文章目录' Y) c/ k  y% Q! e) {# h
卷积网络实战 对花进行分类
( v/ g* R& m8 \0 n数据预处理部分
+ o0 C* _8 B  h6 V3 @$ d# R; Q网络模块设置
2 w0 D% t$ A$ A' \网络模型的保存与测试
5 {3 s0 P" t! U/ w8 O+ s! w/ A数据下载:
4 {$ ]( C6 d0 Y: a' g1. 导入工具包' l( e+ n: J5 q+ B$ Q7 L% |
2. 数据预处理与操作
  e) W( U) A$ y4 s* z3. 制作好数据源7 Z+ r) q# b7 M, i: {6 |
读取标签对应的实际名字
' b; K+ o! t* N: F+ H4.展示一下数据
' f1 ]6 ~' ^& H- O: v) D- r5. 加载models提供的模型,并直接用训练好的权重做初始化参数
& A( c2 l! s, q, `% k. O6.初始化模型架构1 d+ m2 V$ E" P% R# K6 c
7. 设置需要训练的参数
9 I7 l4 B+ ]% v3 d/ i# K7. 训练与预测
" w" P5 y4 ^  @3 S7.1 优化器设置" `$ d0 {9 u: f- U/ j
7.2 开始训练模型
" o+ |. E9 ~% A; t7.3 训练所有层
; ?5 u/ c- U5 j4 \1 S! A开始训练
, h. D- o: X# c. o8. 加载已经训练的模型! t9 r% [' U* s$ u4 h) T
9. 推理2 Z7 Q6 e- y; E( M6 e4 ^
9.1 计算得到最大概率7 E. N+ z. G3 Q+ O) X3 m. ~
9.2 展示预测结果
& k7 f% R& Z8 I& K# p8 a& V% q# g写在最后$ t: Q4 p$ N9 G: @$ D
卷积网络实战 对花进行分类1 n9 i& @2 V( r) ?, H7 F* w
本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能
/ O8 _. l* r7 M: k7 K, [* U
/ ^/ X: u) Z  c+ g% _) Y4 u: i3 ?在文件夹中有102种花,我们主要要对这些花进行分类任务) `5 X$ `" ~$ E  ?
文件夹结构
1 u5 d# a/ R" m1 ~8 d" Y* S4 V1 D( H# z
flower_data. [% g% O0 n9 B& Q

, n3 o  A1 q! z! E( Utrain5 o5 z7 v+ W& A

- E4 v! w6 P3 @4 i8 c1(类别)/ l  J8 ^6 d# T7 c/ ~5 w0 v3 k2 L+ u
2
9 O) P' Z. k  `% X# K) x/ y4 s# hxxx.png / xxx.jpg1 L5 ~: f8 E+ r0 K8 X8 u
valid
4 n/ T* `3 P% P9 n0 d, s& n
# u, p# R% Z; X1 p! I- K# p+ v7 F% M主要分为以下几个大模块
$ E' P2 [7 F6 ?' Q- N2 i9 J0 \* j  U  N6 I# k9 t
数据预处理部分
6 \1 V4 O6 q6 s; g4 G/ G3 M数据增强& O2 H+ l; O: _2 L& k* C# B7 q! Y
数据预处理
, u& W. s& }; y- Q) ~5 Q网络模块设置
+ q  l# n* E5 M  r# Y; v; `: ?加载预训练模型,直接调用torchVision的经典网络架构5 n# t2 D% n6 M8 ~; s3 `' l0 V
因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务
4 z& o( b+ y. N$ j* t. {  o$ o网络模型的保存与测试% k" Y* X  P: j# X# d% z
模型保存可以带有选择性
' \, q% g! D/ g% G/ K8 E2 c$ ^/ W数据下载:$ @3 y" k+ K% T9 `7 F% r+ a8 q
https://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset1 `0 D" X9 k$ ]+ J# Q

, x' g3 X. H0 {% ~9 K8 L改一下文件名,然后将它放到同一根目录就可以了
7 L2 Q5 ?8 w+ N2 }
+ Z3 @. I8 c; ~下面是我的数据根目录
1 {" N  y. ^/ n+ }7 m$ a; e# r' j& v0 S4 H
+ {9 M, ]: ^: H$ [: D5 V% C. ?
1. 导入工具包5 a. U  A% J" }% g) C! ]% I
import os1 c, L3 g) T" {' s. Y3 d
import matplotlib.pyplot as plt' U. t, A1 q3 S! n3 O
# 内嵌入绘图简去show的句柄  _( u* K( r2 y+ k
%matplotlib inline
, P1 Q# m- \# Jimport numpy as np
6 y: \  D3 M( w; }( f0 x, Qimport torch
9 v# ~( F$ r/ c, hfrom torch import nn
& H+ Y( X: {& A; o, Q% e( e, p; b
' p. v/ m( b2 ?import torch.optim as optim
4 N, Q! d7 C2 K9 @* l/ `& w" Z! Limport torchvision
! X4 T& m5 j/ a1 Wfrom torchvision import transforms, models, datasets9 W0 s" L  E  s9 s2 x

, _" R3 V7 h% q, t1 h; E: ^' }5 gimport imageio
2 l& O9 U8 W3 ~& A/ Mimport time) G) e: ?" ], V' X; a% ~
import warnings" I* F- h- T" c
import random/ q' q' E2 O7 b3 I5 M" d
import sys; b' t' D' n4 I! Z: C8 N# S3 T* O2 X
import copy
7 Z7 F' O: X0 I  Gimport json* S$ b9 v+ i/ f( t/ T
from PIL import Image/ M8 ^# h3 w6 _, \& \" h
0 O# i% C; o1 D2 h

# X" F5 d, \5 t6 i; m2 ?% B1
$ M; H/ K3 e( {20 ]2 P/ }5 L& I  |! N
3
  o2 `: U9 B; @9 Z$ E49 `3 p  t. M  |9 {
54 M6 _) C" j8 U% x5 \
6
& H6 \2 [6 I7 O5 F7
2 B; {& ]& _, U0 w* p" P8
5 o, f6 x, i$ @. R; |$ g7 L9/ \7 _  k& s4 q' T" i/ j
106 k! X/ B# H8 @" V6 y' F
11# b: q1 _7 c5 ?9 O- E7 p$ B5 h
12
- |( g6 H7 h' ?4 o5 A13
# u% M7 [: x& N. L14$ j3 n4 ?; G3 v3 W, Y
15( B. p7 O& E& a0 q3 f
16
5 Y' e6 [( s9 _7 {3 Z9 n$ Q17
& x3 J/ E, x- p0 G4 {* h18: p) C9 R" @0 D; }4 J
19
3 q2 q7 B5 \9 {; f1 [) C' y20
# z+ G" c6 N. m3 e; J& m4 X21
( b  w6 [, T/ @# c- h2. 数据预处理与操作/ C% x) k) D7 s6 U0 A0 F$ {! }) _
#路径设置5 ?4 s; ]( b" A, v- K7 D
data_dir = './flower_data/' # 当前文件夹下的flowerdata目录
4 K6 U9 o- q; c# ttrain_dir = data_dir + '/train'
  @& c+ D" O/ n% ]! Jvalid_dir = data_dir + '/valid'
# K, m( I5 O5 A& k% @1! F7 V6 a- X3 A1 h7 j/ l2 r
2) |+ A; V# ^1 q* D$ O% Z6 i1 [+ k
3, j0 v7 g+ G5 Z  J
4/ K4 P& T' y7 F7 ^5 m- ]* [1 |
python目录点杠的组合与区别
: P+ \% S& V5 I9 z5 {7 `注: 里面注明了点杠和斜杠的操作
7 f6 a% ^* w* w! _) }4 Z" j) N8 D- }% F: D8 D$ w8 v2 K" K
3. 制作好数据源, P7 y$ m) n5 |) V1 ]( P
data_transforms中制定了所有图像预处理的操作
9 `  z; v& V- s6 g- oImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片, y, O! W# _$ x8 U  J& f# q
data_transforms = {3 ^& w# c3 P+ o" t1 I, J
    # 分成两部分,一部分是训练4 J0 C8 E' ]: c- R
    'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间
% }2 L  }, g! s4 X! C2 c                                 transforms.CenterCrop(224), # 从中心处开始裁剪' k0 F. L8 n5 e; M. i, L2 a
                                 # 以某个随机的概率决定是否翻转 55开
; J" A* D7 ~2 Q- S, R                                 transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转% j3 V# _8 N1 B; t
                                 transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转
8 E1 J9 @- |2 ]- a/ M) r- q                                 # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相
5 K3 c- m7 H, \( O# c) _                                 transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),
, i8 l1 G* K7 V8 w; _6 U                                 transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB
0 P& \* z1 V. e* X) R                                 # 灰度图转换以后也是三个通道,但是只是RGB是一样的
* [; S" n5 _, L  w% `# @                                 transforms.ToTensor(),2 o" ]6 ~$ Z( y! I" `6 \$ M* h
                                 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差
: s# ]3 }. w' e9 n3 Q                                ]),  k/ |/ r( j5 D3 V9 c9 f
    # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
7 ?: E' x5 v( g2 C    'valid': transforms.Compose([transforms.Resize(256),
* y3 ^( d% p& B6 i4 X                                 transforms.CenterCrop(224)," p  H9 z5 F. _5 k' [( ?% \
                                 transforms.ToTensor()," F, L2 ^  c- m1 U% X) s
                                 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同$ O8 ?+ G% y3 _: G+ H, E& E5 J6 Y9 }
                                ]),: |: l- B* k7 x* J6 X" N) f
}& g7 U- j, U" W, m! ~: W; A

8 N0 I+ w! h. ~/ O5 s" @6 J! `1# U2 L2 G+ r6 R/ u% ~. ~6 \
20 r- Z% |# Y( [, X0 z
36 `5 A8 @; y) l/ a" V- [
46 {) P6 N7 A9 H2 X) _
5
7 v: ?1 Q# w% h: o% a6
1 f" _: |2 G# e8 d9 B9 K73 @2 F" B& S" w
8+ ~% b' ?, c( @. b
9$ o. k, q2 |! k, @
10, i3 c8 h/ n. H  h8 L
11; A) ?) ~. G$ j& E. L) R
12% c0 @& @& ~9 d7 }1 W/ J7 ~
136 I& V- [0 g% \
14
7 D  d2 b2 M3 t7 P2 ]. S  P15; D( v- s+ n, K# O2 e% O2 \
16
6 R% B/ {& d- `17
! A6 K2 g/ K0 U+ E$ `5 H7 ~18
1 L4 S/ p% l3 |19
8 u. q: g; V) C( j$ a* ?2 Q20
3 ]5 c. j- q! @2 a21
% Q8 d5 c& ^/ Y- ~9 f+ }& lbatch_size = 89 R6 ^3 ]; r. _
image_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}
9 j2 }: Y7 e1 R/ \dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}
. V& J# N* E4 ~6 r; A, kdataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']} # N4 y! f9 ?' f# a+ l8 R' w8 \
class_names = image_datasets['train'].classes) V0 B1 Y* f1 I7 H8 K
, z( k: i) D% o1 A% G
#查看数据集合  R  o! C# d% {/ `) p1 ~5 r
image_datasets
+ @5 w3 X! |- {8 V6 ~
) l' J* \( a& a& ?+ s1
6 Q0 _8 ?, z4 J& G20 \" K! W- M% {1 g
30 ~) L" A1 b; V& q/ s
4# }: w$ D' n1 S9 J
5& r$ Q$ H! h+ t! ]
62 m8 B  P7 f' ?2 O
7
/ D3 ?/ Q% I3 A6 c$ U8
' o% |' V" m% A" k" z9* J- I& T! L: X% [* I& m
{'train': Dataset ImageFolder2 Z2 M( R0 M& z5 o" ^9 p
     Number of datapoints: 6552# E8 b7 F- }- c& {% ~
     Root location: ./flower_data/train
& }1 P) B- n7 t) C/ h     StandardTransform
' U, U" v* \9 L0 M* F Transform: Compose(
: g+ O+ ^- x# L$ r5 r% T                RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)
! }& j( s* w, R! i( u) a0 ^; k                CenterCrop(size=(224, 224))4 _0 E1 X1 K& D" ?5 j$ y
                RandomHorizontalFlip(p=0.5)9 Y2 f" l: V* h) f  W# U
                RandomVerticalFlip(p=0.5); }# z4 g( u+ R
                ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])- e  f- w! b% x2 }
                RandomGrayscale(p=0.025)8 g% C- V3 c( M! ^7 m% K* o
                ToTensor()
" u) W( ^' k' s8 c" _1 L, n: u3 n# C                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])7 _4 I! z0 z- B$ h" q
            ),) G# B$ Y8 e- c- @5 P5 n. B
'valid': Dataset ImageFolder
) u( e. F/ H& `" H! l3 a     Number of datapoints: 818
) j" p' D) U. _0 ]1 x     Root location: ./flower_data/valid
5 L- {2 U+ o# T6 \     StandardTransform
) p9 W+ l3 g: g) C5 k. ^ Transform: Compose(
! B4 V8 j9 m& i6 U; x9 M8 _                Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)- X0 M" u( R! |, y5 ]0 o
                CenterCrop(size=(224, 224))/ i  q( S' i7 }4 _; o5 O+ C
                ToTensor()
+ X* i( I; [% ~7 R                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]); T5 V5 C; ?/ z$ U! Y' n/ s; M
            )}
" F) n/ Y. m# j  _0 {" K& n) {* a1 v2 ^! Q$ Z
11 f1 K3 i- d- L9 I7 R* ]
2
& C5 D5 k6 k4 ~7 E+ x3! S9 P, Q: ~/ t% _
4
+ ]' K% g+ A0 v8 e& _. q5
. W: A4 s3 C' ~# v* N% E6+ S! o# J& Z! D% Q! z, _& X
7
6 S, g* ?/ l* U! C# F5 F8$ d9 u, w) ?" K! U
9
, c: u1 j' M# u4 K) m" Y+ ~10
( b5 D( V: D- a4 d6 ~11
6 w! _! y) u8 `: a5 y12
6 j& x# p& j3 \0 m7 V. f/ H13/ u4 C$ ]/ N5 N% d8 h2 i
14
7 F) e( E; ^: v4 ?) o# G; J3 ?15
; d' |/ ?7 p3 V8 R" v16! h* |. U' ]8 |, y- E5 r3 k' T! H3 J- C
17  O, M7 G9 D- {: b, p( ]3 n. P$ k4 h" Z
18( \2 m8 t8 e' |
198 p' T9 e' z! B! b1 V
20; `, {3 U% ?, O% F
21" u/ _- g. _  \9 a0 ]# ~( v# J! D$ T$ {
22
0 m( s2 j7 g1 A% H23
& W- \/ k5 L# T( W6 ^* p1 F8 m24
3 e) Y; m& i4 w+ @1 a8 M2 t8 K' c% h# 验证一下数据是否已经被处理完毕, P# g) m9 k: M# x( V# h, F9 q
dataloaders5 N/ L  A4 i9 H- [. F9 r
1
* M% G7 S" z; v  p' b% X+ u$ S21 }6 }" \) m: z2 H" C% z7 v
{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,
4 H) D; I4 ~  k3 A6 n 'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}2 g2 P( _8 `, a& i
1- F5 k3 d) S+ c, V7 _' Y
2
+ A" z' J# V$ u1 c% ddataset_sizes: _4 q9 W3 V# b2 X" E2 x) r
1* n7 ^2 x8 z! ^" Q
{'train': 6552, 'valid': 818}
7 Q0 I/ [; X% h! a5 s% l: l! t1
5 j9 D& H, }* x) W: G0 I读取标签对应的实际名字7 l9 L# z" O5 R7 g, C3 U( c4 @  W
使用同一目录下的json文件,反向映射出花对应的名字, i1 e- ?5 t+ P4 o

$ o! y% _) [# k0 o; `. kwith open('./flower_data/cat_to_name.json', 'r') as f:
. [: D: b9 K7 k% V8 V    cat_to_name = json.load(f)0 }' _! e+ s5 l5 O: w. p: d
1; f/ U/ x' o9 O' T* J
23 p  H8 E( r+ j( }* V
cat_to_name; W% S; D1 W/ Y5 B1 \# x. c
1
2 v; {% G& a% d, Z8 \8 X# A{'21': 'fire lily',; S$ o& [9 T0 n! I+ v/ w  E! D
'3': 'canterbury bells',' e% c$ j) J' b9 Q- N; k
'45': 'bolero deep blue',; ~, A# C# I: f) e
'1': 'pink primrose',
. V& L8 c, z7 }4 h, ~7 U '34': 'mexican aster',% A/ z: J1 [: |6 _; S+ P# h
'27': 'prince of wales feathers',- a2 G' q$ e" a" R2 U
'7': 'moon orchid',
- s. p# g+ M) k; s1 d( n '16': 'globe-flower',/ J" P0 E/ _2 a- b9 G( O
'25': 'grape hyacinth',
! e" y& e6 s3 |& I' S: c1 Q '26': 'corn poppy',
) w& `( _* \& y7 D- L+ Y '79': 'toad lily',, v8 o2 V/ T7 ?8 B. K9 \
'39': 'siam tulip',
( R: ^. o  \7 l0 k8 @ '24': 'red ginger',
1 H  F. B4 \. w1 g7 [) e# q '67': 'spring crocus',
/ l5 I% a" s& ]/ e '35': 'alpine sea holly',& H; i( _( b( I* g! X1 ^
'32': 'garden phlox',* a) A+ F0 s" j  Y7 ~$ F! X5 w
'10': 'globe thistle',
: C& G- H# Z2 M+ t% O+ e '6': 'tiger lily',
# [! X4 f7 y& _6 t& } '93': 'ball moss',; G& c9 A3 p& @( Z2 n& `2 Q
'33': 'love in the mist',
1 U) W3 k0 I+ C2 E- s$ J- V/ _ '9': 'monkshood',: G% o1 r. C5 I3 B2 X" d6 u. }
'102': 'blackberry lily',
( o( x7 j. }; X$ \. \: S '14': 'spear thistle',' a* |9 n/ g# T1 A
'19': 'balloon flower',5 n: {! X7 @# d/ s  D; x
'100': 'blanket flower',
7 {! H2 t2 W; r3 Y '13': 'king protea',
4 [9 ?8 k2 W5 s; n4 o '49': 'oxeye daisy',
1 V5 {' l' o! L5 r' ]! E '15': 'yellow iris'," ]5 m) n: A8 e# ]4 M
'61': 'cautleya spicata',3 `. W0 o: ^' e* B, @5 ^
'31': 'carnation',$ U' g# d! Q! z* G% Q; \. i0 S
'64': 'silverbush',
/ J2 F: c! k2 d  |) J0 l1 \" }8 m '68': 'bearded iris',
$ o0 s$ a( y8 K. R '63': 'black-eyed susan',. ^/ l+ [$ o( O# ~% r- M
'69': 'windflower',! D4 w/ h, h8 H9 D9 h
'62': 'japanese anemone',
5 b% M( {. d; p '20': 'giant white arum lily',# T  Q. B: A9 R9 m, Q) R
'38': 'great masterwort',
" d2 J- o( z2 n8 \: w* R '4': 'sweet pea',: A! u$ {* e( U/ [$ k" R
'86': 'tree mallow',
" ?  m  `, F  g+ m& U( v '101': 'trumpet creeper',
; V; M  B+ f" I8 E '42': 'daffodil',
" D3 i+ s8 y& N '22': 'pincushion flower',1 D  c9 [4 ^3 Q! f* [3 n1 S
'2': 'hard-leaved pocket orchid',
2 G: E6 t3 }; u+ D( x '54': 'sunflower',* F$ H! u2 y6 `! E- R) L
'66': 'osteospermum',
# |. H+ v6 o8 R0 a/ b  m$ u& o '70': 'tree poppy',7 z2 |. b5 e9 s4 M( o; G, V$ a
'85': 'desert-rose',
/ P3 f" G# F3 _8 t8 \$ U '99': 'bromelia',3 ?$ |. p4 x! v/ l' \/ }$ X; o
'87': 'magnolia',8 e. f; V' a8 L
'5': 'english marigold',
+ ?" n; q4 A4 b5 H1 L '92': 'bee balm',8 }& C" v: ]: q8 P) _+ `
'28': 'stemless gentian',
4 y! X( q- K2 R+ @ '97': 'mallow',& d  v7 z6 t& V% V" _3 ~/ B1 ^8 @
'57': 'gaura',
" q+ [7 P3 d( \3 L! O '40': 'lenten rose',8 X4 g( b: s/ R# ]
'47': 'marigold',% v' U  J) @+ j  A5 D4 w/ L
'59': 'orange dahlia',
- x& q; U% e$ Q7 J6 B7 Y4 T. [ '48': 'buttercup',
/ T' m; w- n- s# ~/ P% Z3 E% ] '55': 'pelargonium',
9 e) h) o0 K# ~" @ '36': 'ruby-lipped cattleya',3 B' W9 p) S& \
'91': 'hippeastrum',% R# W% m: F1 G6 ^) t
'29': 'artichoke',( }* Y- n* ~4 }/ t2 B; J$ ~7 A
'71': 'gazania',
5 H* z1 C. I) j/ K! w; k '90': 'canna lily',0 ~2 p+ G) M3 F% z- v
'18': 'peruvian lily',% u* f" j. {( {" \
'98': 'mexican petunia',
8 k$ t  h: R5 o1 Y, X/ M '8': 'bird of paradise',
5 v. H( d( Y  |2 {8 R '30': 'sweet william',* h  ^2 \$ a- L5 Z+ h1 w
'17': 'purple coneflower',: `* ?& o5 h! k9 {/ N1 H1 _% w: C  l1 \
'52': 'wild pansy',
8 s( T1 R0 p/ a) m2 F- A6 [4 L! D '84': 'columbine',
7 \* F8 V2 ?$ b, ]* \6 |2 B* ~3 T '12': "colt's foot",) R" E9 H. H% Z; r3 n1 l
'11': 'snapdragon',1 h( k6 l- v, Q: V
'96': 'camellia',
2 T3 ^6 d& f3 @5 ~" }2 a '23': 'fritillary',
) p7 g6 I' O) d( f& S '50': 'common dandelion',& d) X  _8 r1 D& d: ^) x
'44': 'poinsettia',
! g$ [9 @4 h1 g& `( ?, C '53': 'primula',
) L2 ^6 I5 V0 S0 i. d '72': 'azalea',. b% B& `& E* {( j5 O* |
'65': 'californian poppy',7 I" q+ o2 R. v  L
'80': 'anthurium',
& o# ^' d. _5 u! [  e6 q '76': 'morning glory',
* J0 U+ L" U( d '37': 'cape flower',
% z: j$ `5 C5 H/ B8 z) z '56': 'bishop of llandaff',
# P/ Y- z* H" T  T7 } '60': 'pink-yellow dahlia',
  v; P/ o/ L4 \! Z! j4 \" B3 w '82': 'clematis',1 w' a. x" j5 `0 o
'58': 'geranium',
  s8 Q- R% k/ ^5 V1 {1 o '75': 'thorn apple',, f0 W! \$ {( c
'41': 'barbeton daisy',7 z) G% D% ~! x) J  U) F' \4 \
'95': 'bougainvillea',8 B6 k3 V& w0 D& E" q* @
'43': 'sword lily',8 u1 V9 h  k4 k2 @8 V  Q& M, F; i5 e6 E
'83': 'hibiscus',
0 \* N3 L$ Q- Y7 S '78': 'lotus lotus',
  h2 [0 X/ D, Y$ M" X '88': 'cyclamen',
- d$ r/ q+ I: ^$ z/ b0 C9 g) O '94': 'foxglove',, O$ V8 ^' P: @2 e, G' P, r' a; r
'81': 'frangipani',
& A' S* n) w% I '74': 'rose',1 Z8 l% H  Y- M' N; b! {
'89': 'watercress',8 y4 _+ U$ }9 p7 P* x/ t) D5 }
'73': 'water lily',& _+ I) s( H+ R1 q
'46': 'wallflower',
8 N& S  h  Q) N% e '77': 'passion flower',, v1 p' P7 M. S* Y$ ~$ v4 r
'51': 'petunia'}: _5 Y5 ^4 O2 F" m' o0 p& s

- K) |! `2 B9 T! K11 o0 q5 J) W( ^
2
$ y+ C& d2 s. j; q4 U5 v" k3" R6 V$ R3 |# Q8 _
4
5 I. a$ l0 b  Z7 Z3 ^" T3 C" L+ a5
+ o( t9 W8 E( k/ \% d6 W8 \6
- e3 O& r' F8 M, A7 u7 s) q1 ~7
5 I7 H/ M2 o/ |7 f( f8
5 K; f  d9 K5 `. ]- I' \: Z9
( ?% h- Z8 h4 Y100 m$ z" |% G% C2 B
11
) l$ p/ \9 P5 e12" V: I: Y3 ]* H% g6 j
13' C$ n) T3 F6 y. B
14
$ }2 ]: A$ S/ J15" c- ]; f* J: K& t! ~
16
9 S* K1 N- |( r* ]. Y( U17/ ~6 T- W1 f% [
18! J* l: J6 J+ o, A+ U
19
* F# I% H0 @: K- o20
, ]7 w$ d+ H7 ]/ L( p0 n( H7 N216 _" |* V! g3 ^- S; |% C9 A
22
8 u3 r9 r' M7 n* ?% P23
5 j' B% p: Y0 L7 J- q6 w' R; H; B- _24- L# _* R) Y; }; Z9 z3 M, ]2 Q
254 p* T4 x! R" N0 U1 [5 R, T
26
$ ^* U- T8 D) G  c" C4 M+ u# s! n27
) S8 w4 C+ p. b9 `28
$ a3 D8 c* z8 p9 y  b+ @7 V5 Q29
2 L) [/ i( g6 M& J30: P2 R' q4 {* o7 _. ^5 X* R; W: B, [
31
  U  E6 q+ }+ }8 G) m) L32: d7 b4 ~. E- c2 r0 R4 J
33
4 R& Z" o% i- |, C) [, C34) O( i- B- Y0 I" J" k$ n( u  X
355 e, u, S) o- T3 z1 @, L6 i
36. p1 D0 c& j. N' x( E
37
; _8 e6 M0 ]/ H8 e6 v9 A+ A* I$ T38$ ^: \1 U6 f7 H* X
39' k1 m3 A. e" E+ X
40
0 @4 k, W: U  [# w  |41
, F. \6 V7 O# w. ]" H429 f/ A& h( L: d/ g/ w
43
% U3 }/ y& n0 `1 _44, Y) T* M6 O$ q8 {
45
4 w8 D# p+ e: w' S/ `! r* u46: d* H$ z0 }* H- W
475 }' Q0 P; |& c* t
48  S2 u2 L& j+ G3 p
49
; d, |* w( B% s  v/ T3 U5 J. S50
: P' B0 m4 [7 K519 z- U* E6 B  }7 X
52
0 e* e( a# c5 M7 t# h7 r0 t53! I( Y1 |" x6 p4 V& B2 H' Z$ |
54' a+ ]1 v7 [" h" [! H
55
, O/ H; _7 j: Q9 C$ O* a% K: a56( b* j8 w; N9 l+ T4 z
57
4 r  W$ i" p- {$ t58+ c" P9 i3 _2 J- i+ ]) x
59
- S; j0 {7 h! z. p6 a60: ?: Q/ k6 B8 l3 a
61
* ^3 s& Q: d* f, }' Y5 ~62
( C$ H" N: Y( f) k+ C$ r/ \63
& Z( {" c% F& w5 Y- p) N64
5 ~/ v4 v# P% Q- z' \$ u65% N! \; c5 s3 r( o8 u, `, i
66" M* |. N, o5 E4 S4 c
67- C& |3 R1 ]) n. d7 ]) B$ G1 q4 d$ F1 W
68
7 \/ N7 [+ D4 G7 ~( C8 p. {69
( \2 O& v0 y& x% A2 a$ q70' W* t5 ?* z5 f5 i( X
71
, X$ ?- O/ t) p- e+ U72
3 d0 n5 z. a# D- ]# C& o  t73
  y" q% Z6 h% a% ~( i74
# ^( ?" V7 p# C# y; I5 P75
, @: H! q& z0 ?2 |- Q& ~76
' {$ X% ^* M# \77) G/ l  u. x6 s- l
78
5 ^5 x# l) |% `, |2 J79. i: Q5 H! D. c+ L0 ~3 `  R
809 b) ]* @3 M1 D, y; U; }
81
8 `, v0 M/ `* U* {: z( U( A82" |2 K7 F; Q4 |, Z- R5 A8 Y7 y; W
83% e* K% x* Q) F; G; H' Y
84; X" Y- x) |0 X; a" b9 t
85
% J+ U# [- }, U4 p* J8 J% {86
! W2 n$ Z' \6 L, e4 W' V- e87, A  L# e" O: z7 R4 h/ f- J+ _# Y
885 _3 Z6 s. n0 k+ q
89
3 n. D# m2 `3 k. x/ M90
+ P7 n6 p2 {8 {" N0 e9 z/ k91. w) E: n% z% g$ y
92
" D" ]" b8 \& C1 a, E- D93* r: C: m9 j4 M
949 N, {% h! T4 Y# ~: K
95
( M1 x. x' ?; V$ x  D, J$ Y( \/ V96
, R4 a% c, \; d6 W; V: ^6 V97
0 n8 u5 B+ S, K/ N4 D' o98; h9 p$ S  w$ I: S
99+ n4 V; f( `- p
1007 _2 ?5 n7 k% F& }
101
% J9 E- ^) v. b; Y9 V1028 m' t+ H* I1 f- A# a
4.展示一下数据, a8 P3 p4 i! k# T$ v
def im_convert(tensor):9 ?3 v! J. g) v, ~/ l& y, w0 T2 m
    """数据展示"""
6 J; o/ o: }' U8 y" R3 d    image = tensor.to("cpu").clone().detach()8 I0 n5 ^6 F2 H1 l8 Q
    image = image.numpy().squeeze()
: P8 d& G9 o, h" L3 O    # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图
* X* f2 A% F0 H, K6 c7 c' F7 @    # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)! ^$ w3 z# H% x* a* S
    image = image.transpose(1, 2, 0)1 ^, d8 p; O, ~4 b3 j: O
    # 反正则化(反标准化)
! h; Z. u6 J! T6 B7 [2 C, N    image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))
& ?" T. c. Q8 E5 x; Z6 v1 u# C
" `% Z! i7 ]. A3 b* m    # 将图像中小于0 的都换成0,大于的都变成1) J$ m0 [8 t  j
    image = image.clip(0, 1)0 m2 Q5 ?7 x2 f5 W# N3 i
$ t' g3 X# x  R6 _
    return image% n3 i, V4 m7 I0 J+ ]+ `! r- u; @' n4 B$ C
1
. `# K6 P, X+ ]+ E$ k! I5 {2' [0 L5 x5 w3 w- p5 ^
3
! B/ F9 m+ Y/ R41 R3 A9 \+ g5 c/ n
5
  l1 M& ~4 q0 e0 A( ]  D$ X$ e67 ?& A6 W) r7 U, K! `
7  P7 e& b( [" Y$ V
8- U. @6 n1 h4 z
9
% Y+ X% Y) \, Q( Z0 X10' f) B7 H* b4 W1 m* ^% o$ O% M
11
2 Q, i! l! V! p129 b0 _/ ?0 J7 @6 T+ n# D
13
) [  r; W5 g/ m. y14
9 t  o, W( X" b4 }0 I# 使用上面定义好的类进行画图5 x! M) b( I, ]+ t0 a8 x: l( c
fig = plt.figure(figsize = (20, 12))
& t- L8 c% Q9 f2 S: z. r% ycolumns = 46 \- s* l. Y; o
rows = 2
5 y/ Z& B; s2 w/ X' ~8 E$ Z3 p$ L: A# b( w- C
# iter迭代器8 V" k1 e) M: @, g. C& i( v0 _
# 随便找一个Batch数据进行展示
5 W8 z* ?; R( w7 Wdataiter = iter(dataloaders['valid'])
; |* u( D1 i& k: D0 e7 yinputs, classes = dataiter.next()
4 `  N! w6 l# y; K8 w+ C  c# z7 p; L' p  O
for idx in range(columns * rows):
* h/ y7 i; Z  L    ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])
- }* y$ O" l& N' X" p. y2 R" n    # 利用json文件将其对应花的类型打印在图片中& Y8 G. M6 h5 }4 T
    ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])) v; p" k$ O5 q3 p" N7 Q7 X, R
    plt.imshow(im_convert(inputs[idx]))& U: l+ A( n7 J0 ~( F! M/ k" r2 m
plt.show()
: w& O+ b+ q- y; J2 p$ S8 p
$ s& Q5 v6 S' V) E4 p1
/ X! G% I. H% a2 {" O- Q, X2 K. }, _27 d3 w& [: T5 M
3
1 I- A) Q& \& D/ V4: i/ P+ X+ a, V, D! h
5
$ O$ K8 w3 _) \% Z, R% b0 `5 {61 j% E) b: n% G
76 ~5 Z; f& V- U+ Y: U
8. o) @9 ^' X) w
9
, g/ T0 S" c; b0 K4 k. |3 i' C10
- W! \! z1 M' p: Y. L11
2 z4 K& h, ~/ A) N( D12+ z* w5 q% q0 x  S6 B: s4 j, w
13
& R3 u- {  w3 s1 M- a9 p+ k6 F1 |% J14) M+ k- c, j3 n) Q8 M, j8 v
15
% n% g' k; \& J- g) ?* ~16* H: P- H' Q; h6 q. |9 e3 E% L8 H

* z! |4 X" q/ w* w
+ e2 l4 Z. Z8 }4 Y' O- ?5. 加载models提供的模型,并直接用训练好的权重做初始化参数
2 A: J: J  @2 n% k$ Omodel_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']) {0 f; t! T, I% h% t  e
# 主要的图像识别用resnet来做3 v& n. `# R0 r6 B7 K- y' s
# 是否用人家训练好的特征
0 J0 ]+ `/ v0 S7 Nfeature_extract = True" v. g: X! u2 e5 ]# J/ B" p5 d! G1 Z
1
: O: m% t( W/ x1 \2
# d8 F5 o% e1 i35 H4 M: v0 b! k4 j( j) ?3 ~
4
3 ~0 {2 K; i: k# n8 h: u# 是否用GPU进行训练
3 k7 k* q" l  ytrain_on_gpu = torch.cuda.is_available()
0 t- }1 u; y; m; P. Z
5 d8 b: v" _- ]. a+ W5 l5 uif not train_on_gpu:0 h3 x7 ^4 z& ?4 E( k
    print('CUDA is not available.   Training on CPU ...')
/ k% z+ {$ Q4 a: Pelse:" e* P: l& K# T2 g4 C& W
    print('CUDA is available! Training on GPU ...')
, l3 F: j# k# ]1 L/ V& I
- w; f4 z' D/ f2 x' g& Cdevice = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')7 Y2 H2 {+ y% ]. y, |
1& x+ U! V# l7 G" H9 p: T
2  {  M  D: C; a
3
# B! l7 l- h$ s) C1 c4
% A* F' F2 S, g7 O. o) D) B# B5
$ h/ U" y: c! P$ G& s. m; U6
- I/ I9 ?% a9 u: E7+ `7 f; [4 j+ p: Z' `! n" }  b
86 T* D1 A# S" ]3 N* D) S: V
9
: f6 S, N6 X8 o+ A' G9 _8 |CUDA is not available.   Training on CPU ...5 q; [9 y+ {$ {0 |  v1 O' |$ |  z, c
1  O- k( B) \4 j$ B3 i" r
# 将一些层定义为false,使其不自动更新" S3 G  _8 E6 ~. A. B! t0 @
def set_parameter_requires_grad(model, feature_extracting):5 U! v/ `0 E: K5 |' v" m) F% M  P
    if feature_extracting:
. q, f4 D, K+ e* X- p2 Q        for param in model.parameters():/ m  q1 X9 W4 W) _% B. x& J
            param.requires_grad = False" y) [8 V3 ]) i; U6 Y& V9 `
1! N9 d5 U, Z: h4 x7 Z* W# \2 W9 Q+ ?
29 [. O5 ^- k8 z6 ]! n
3
  k/ a: E" `8 [3 n5 X  {! }4; B" ^4 B0 P9 k* h+ O; _
5* O' ~6 U. J5 P# {' v+ ^( I
# 打印模型架构告知是怎么一步一步去完成的
, f5 M' s  U; n% d" t& N# 主要是为我们提取特征的
  U  w7 D: z6 r: ]3 {  I& M
: p+ i+ i( ?# [' o" Q# f. lmodel_ft = models.resnet152()
; w% |! ]5 h6 U- \- U% Dmodel_ft
1 C: T. Z/ W9 ^* s: N5 ^1( ?( R; L) O- h: q4 u/ Z
2
- R# t- T/ |1 O$ E3
7 @9 H6 \  S( b$ K- U4
: f6 g1 ^1 @6 b. L5; [0 @9 E6 v; l* T! l$ ~8 l7 S
ResNet(
6 E" C4 t- P' k6 {* L9 I5 X  (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
8 ~: }$ f  q. @. [. }/ n  (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True), s6 G/ T) A3 ^
  (relu): ReLU(inplace=True); B% |  p1 K$ ^+ H
  (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)5 |7 B$ x8 D: Y1 h) O
  (layer1): Sequential(: V( g; {; V0 ^  M" a* i
    (0): Bottleneck(' `" H  o( J* k% [3 P3 y
      (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)
4 z2 v* m: i- u. c4 J, J7 n" ~( P      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)/ E; Q& V) P% X  y
      (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False): F* v1 U$ t! ~" M+ [! Q
      (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
3 i; E* W$ W( O( g2 ^  D. \3 q      (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)% r8 z% F1 s$ y  e( k
      (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)# c' L3 B3 W9 U6 j2 t
      (relu): ReLU(inplace=True)6 w  T$ F$ z7 s3 T
      (downsample): Sequential(3 J& b$ u% |1 G# r1 e- K+ ]8 [
        (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
; W; L! E3 M2 q  e& m" R- P" ?        (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True), P. }: x+ u3 |5 V) R" X# I
      )
% Q% s# i8 A$ T' E. L    )9 }, h" v* z, [. `& F7 T
中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。
/ U2 {3 C4 y! y; K' ?  @    (2): Bottleneck() F+ _! A4 F8 K6 M0 Y; r! Q8 q
      (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)
* A+ A8 d& |* N( w" _/ y. {: C4 A      (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)3 g. c6 x! v  Y& Q' ^" W/ S
      (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)2 N; K9 n+ x3 o8 G/ Q2 k' Q( p5 c  F1 K
      (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)0 f; z1 a4 X- [0 @
      (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False), ~/ N/ E1 [4 q/ s* Q1 S, p/ j
      (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
) _9 B( f0 @) Y1 U* N, K6 r/ L& l      (relu): ReLU(inplace=True)
. w8 p) B  T& R# J# y    )8 Y  Y: X4 k0 W& N$ h
  )
! y( r+ {5 ]) L$ P  (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))) @! k2 Q% O5 l- v- S& Y: M' [& E2 F
  (fc): Linear(in_features=2048, out_features=1000, bias=True)
0 C5 h) r! i6 N0 a6 ?+ _)* \9 B' ?) i7 b$ @% X
" `8 ?& G+ W$ G6 v4 _& i/ m
1+ b- s) W/ t3 K- K8 {' U) I' P8 R
2( P. I& S0 P' Y0 j) o, B
3
, {. }7 S7 F5 W5 ?# z$ h" A0 O( b9 }4
. v# |9 \4 }( T5! R; B5 u4 G1 `. {. j% c
69 `3 X- |* G8 ]* C
79 O9 I5 X$ }9 q! ?; G
86 E# u# O1 v- e" P# c# J5 H6 S7 H
9
5 _+ s4 i6 ]) H+ o$ p10
8 A3 r7 }$ Q0 g$ _9 T. n110 k) z- B0 _! q% ?
124 F8 l+ a" k# ]7 j/ K& ^8 _
133 L/ l1 q; I1 V/ |2 b3 M- p
14
9 T. _3 x# N& d" l1 r9 Y% t# A15( t6 P; n6 k  J; A8 |! q
16
3 h5 z7 x' Y0 |1 ?17. c1 s0 n1 T& U- c
18
6 p/ ~# A9 X7 y" U19. s* K6 R" E2 ]
203 R, G4 d* p: R4 b
21
+ s: j6 }3 W5 \! s$ |22( S: ]% N$ ~( [. j% i0 g$ N
23* C& t9 O# T  P5 r- b* a
24
* h5 y$ A! K9 h8 w! q7 L: l25, s2 h% }8 e8 y8 U& F8 w0 `  @
26& M, U  ]. U" c/ J, [7 c0 H
27
  Z2 v% t* L8 T% _6 |  Z# m; h280 H; o, E1 ^* D; q& l. e
290 h4 x* {9 N5 W7 m% c
30
5 u2 `6 P: h1 V( V31
- s4 c! p( `& }- C6 t4 t32% k! H* o3 b; c% ?
33
' U% ~6 |4 k4 G0 e7 x# D& [; p最后是1000分类,2048输入,分为1000个分类
0 Y, j9 W" G& Y5 x7 c7 L& J! w而我们需要将我们的任务进行调整,将1000分类改为102输出5 A4 U# H0 T, H, o. r
0 S; _9 w; h/ w$ Z' x, X
6.初始化模型架构
. |, L9 ~, G% f1 R- x步骤如下:* U, W& G& R9 f+ f+ _
9 M- k6 v! m# ]8 D4 G9 p# D  c
将训练好的模型拿过来,并pre_train = True 得到他人的权重参数9 V9 w8 j3 {  }# x0 t. t! j
可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)) p4 k# P5 ]5 f- O+ W* |3 ^
无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数
# a0 g% k" ]( A( T官方文档链接" B: B  h* n0 m
https://pytorch.org/vision/stable/models.html
' V4 b+ c( Y, l9 U! l/ a
+ n' e) B4 |  @) N) ?; U3 g, [% g# 将他人的模型加载进来$ E- f- K7 U% O2 \% E; c5 J
def initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):# Z. k5 h) u* {1 B; W
    # 选择适合的模型,不同的模型初始化参数不同
" r  u  F, N' J    model_ft = None
5 l( r. t; T9 Z6 A* _    input_size = 0) D0 j, U0 ?5 Y7 M: `0 l
: G" D( g7 \, o! w) @: K
    if model_name == "resnet":: S$ t4 m8 f6 i" R8 ~" F
        """
& p9 ]' m3 ?! r# [4 q        Resnet152+ o( Z1 ]( v; s6 M. K6 e' }
        """% B2 Q8 t' {' R8 a% I

5 X, R5 G8 H# C6 f        # 1. 加载与训练网络% S: W6 e, z0 l" T& R# m; F" n# ~
        model_ft = models.resnet152(pretrained = use_pretrained)
6 v4 K* p8 o1 L: Y: ^& n        # 2. 是否将提取特征的模块冻住,只训练FC层. l) ~* ?; Q4 F$ o. l" X3 Q, n$ X$ A
        set_parameter_requires_grad(model_ft, feature_extract)0 x! @+ G4 q) C' O
        # 3. 获得全连接层输入特征
: I* {  i# G' p* `$ N6 g  C        num_frts = model_ft.fc.in_features
4 c$ x  H: s7 ]3 K% Q        # 4. 重新加载全连接层,设置输出102
6 ]+ J4 i( E. ]7 c2 I6 }        model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),% M  @$ q- E( b: h& o9 H- q
                                   nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1
' H7 |3 I7 O' }* B        input_size = 224  P% ^! g3 `5 V' l& B1 k6 u
" X5 y% @, `* ~3 c4 C) G. \0 H
    elif model_name == "alexnet":
" |, U# F0 j7 w) C/ ~        """
5 H: q2 ^9 V  f        Alexnet
0 O6 L& X% W# ]        """
- E, S/ ?. {  g# S& o8 p        model_ft = models.alexnet(pretrained = use_pretrained)$ t' F$ s: T) X" d/ k8 J
        set_parameter_requires_grad(model_ft, feature_extract)
  a/ L8 \& K( m1 G. X. l4 B) M
2 e0 ^+ @8 f1 J& H        # 将最后一个特征输出替换 序号为【6】的分类器
9 K; g0 a: |) y; Y: o2 J        num_frts = model_ft.classifier[6].in_features # 获得FC层输入1 B" h% q! Z" o- J! q
        model_ft.classifier[6] = nn.Linear(num_frts, num_classes); E8 E( ~% c  E% f7 f0 T2 M" y
        input_size = 224
8 a0 Q6 h/ I+ B2 q. {( O' I
4 R6 E6 d: L3 H: h) f    elif model_name == "vgg":
3 V: m- R6 I. a. j- m% R! Z% ]' @& c        """
! \5 a+ @* S, ~( \9 E; J" L- q        VGG11_bn
% v% a9 m  J9 Y9 F& Y4 Z; G        """/ c- C6 d* j" C+ t" d
        model_ft = models.vgg16(pretrained = use_pretrained)6 P  O& k3 B' U. E
        set_parameter_requires_grad(model_ft, feature_extract)$ {0 c5 T9 e. o  D
        num_frts = model_ft.classifier[6].in_features6 @: B; U  ^( y& |& i* b4 V/ h: J
        model_ft.classifier[6] = nn.Linear(num_frts, num_classes)- b3 c$ t* `7 o  g- p  a. A
        input_size = 2249 j0 L. h$ B- @+ G2 z; z5 c
+ O0 d/ a  ^$ `8 O+ c: _
    elif model_name == "squeezenet":3 i% A  U3 @% f) {% V: _
        """
6 {9 \) w/ u: G  q" X, w/ N# ]: S5 G        Squeezenet
9 r# R+ w) v5 z3 D        """7 Q7 F) H1 c4 j8 s
        model_ft = models.squeezenet1_0(pretrained = use_pretrained)5 s0 q& O+ r9 |. Q% Y$ E
        set_parameter_requires_grad(model_ft, feature_extract)
. j' W6 h8 ?; g( O: e# v* u        model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))& N! A$ N2 g; s4 Q3 b
        model_ft.num_classes = num_classes
5 x$ q" ?5 B) X7 B+ I        input_size = 224& |# p! Y6 m4 q3 P7 x$ b  p2 n+ d+ n

  M! G  t* V' I7 m4 z* a: q    elif model_name == "densenet":2 w' B1 V: g0 y6 m8 b+ A) |
        """
3 C, w. C) \( p* |( L        Densenet
3 e: B. z" A2 |& m, M        """0 K  f# u0 ?! H' N* z- b  |7 Q
        model_ft = models.desenet121(pretrained = use_pretrained)
9 n3 [( k2 P# Z' Q6 n& k# T/ M, {        set_parameter_requires_grad(model_ft, feature_extract)/ ~" |# i1 E- v1 A. e
        num_frts = model_ft.classifier.in_features- `4 O5 o, U% h' K/ q) W
        model_ft.classifier = nn.Linear(num_frts, num_classes)2 c) Y# c% S! q% l( b8 d
        input_size = 224# |$ u# R1 I+ O6 ?8 m

. r, i, s9 f: G  f' ~: Z* k" r    elif model_name == "inception":+ R" t+ m% G  r
        """' y; E5 X, O5 r0 F
        Inception V34 y$ g) ?& F4 T' T/ ?' p7 C
        """$ E" e, A  E8 m$ G3 D/ d" a
        model_ft = models.inception_V(pretrained = use_pretrained)8 s% v+ Y$ Q3 \7 i7 g* ^! f
        set_parameter_requires_grad(model_ft, feature_extract)' y' A3 u& X/ j8 e& ~
* ~' v9 {5 s; {1 {% @
        num_frts = model_ft.AuxLogits.fc.in_features, L4 a) B" P% D: c- L; {' Q
        model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)
4 `5 L% g6 R, h" e! s: v% u% W5 |, e( p; G$ l% I
        num_frts = model_ft.fc.in_features+ Q, m! N% i; r" P7 ^
        model_ft.fc = nn.Linear(num_frts, num_classes)
1 O: M% ^$ A! g2 r        input_size = 299
: B" U  P5 [8 i! m9 l6 c1 O6 |- J) ^: _: @" [. X* h
    else:! b9 s8 O; i" l# [
        print("Invalid model name, exiting...")
+ j0 ^7 |- H) W/ n        exit()+ i0 }/ {0 }8 `& ^1 ^& p

) l; t$ I& S% g8 m5 p- u' S: j3 U: }    return model_ft, input_size- u# L# R* s& e

7 X3 g( p1 ?( h: {, w9 D1 T16 I9 t2 R! B3 P9 ]
2
' I1 L+ h; A* Y( d2 ?3
, V* F4 P! H# H2 O42 N, D+ Q4 _( i
5
+ ]: m2 ?1 t5 V9 W7 ?' u6
/ Q* \, c: B) R5 Y9 q. v7  R5 H, M7 g: n1 r& |& V5 p
8
' w. \$ ~! g5 e2 ~, H5 M: A0 ~2 r5 ~9
  P& u( {9 D7 P3 |$ y10
! ~9 `: m! _1 ^: p& G9 Z11: b- j" O, `9 F" @' d# R
126 t3 o3 n5 ~; F2 O
13
. n7 M+ p, P$ d$ K14
! v7 H% S( m7 o0 D& a/ L- j15
: j, ]8 L2 [, N164 a6 U2 v+ A7 b  Y. |) s* v4 b
17
9 k2 ~% v  q6 ?3 ~" ~8 |18. E  j) E" b" o- E7 O
19
+ z, @  B. G: u4 d& j, m3 u8 `0 J20
  n5 P4 L/ u: C- G+ l21! Y, d4 L" ]( M  W' K
224 g; N+ N, ^# F# N1 K
23
6 F: ]9 i: u5 A) A0 t24( `3 _: A  e  T6 O( v
25
$ U3 T/ P8 i9 n8 \1 S26) a' `+ W* P# j- {# [. h# u
27
6 `$ E6 y! T" b0 }/ g$ C8 |28# J5 E& L- ?9 L4 i! m
290 W& \; _: X6 T; n6 ^: |8 g
30
/ Y( E2 }0 {8 T6 b1 Q1 l! v31! [2 m, Y! \0 `
320 s( n: z: X# F
33& H) [+ q" _1 B" Y
349 l2 l2 v2 g4 D$ ^
35& J8 s/ a8 K/ C( A  E7 K
36
. u; p' M% y$ S1 W37
  u' A: _# N; c5 T7 }2 Z2 ^38
$ Q) @- Z4 d" {7 F6 i' {39
" l- j. m4 X# ?% g40
5 H# @3 T. ]3 ^& ]8 L411 w/ P% b3 B# j
42
( T$ }+ s& b4 g+ w) H43
/ o2 f2 Q" P7 b) h5 R* }447 f+ N5 |8 I7 A% P; K* J  G
45& B) T1 ]2 j7 f1 o, S8 C
466 v+ b  |  ?8 @: S3 E
47" Y6 U$ H9 j, m0 b" U1 J% D8 g" O3 |
48
5 e  d2 l2 m$ u) f) f49
/ P5 P/ _) a* c% x50# t+ Z2 v( |4 j, a5 T: ^
51& X" C8 B1 U- D. b: l
524 m3 _9 x8 G! O; G, S  [" B
53( {$ D0 ~7 w. j, W: j% \: a6 L2 A
549 m! X: T/ j- T& `) n" K, K
55
: \9 \% Q- P/ B. ^' U; m56
: C: A% F0 v6 S! g9 K3 i; p; G% M57
( a# R+ A. |/ O- D58
, j- j& \8 c# U; |/ C59
, N: o8 x' W7 o606 ^9 m1 q; |2 J. _
61. \! l; @: {- a# N1 k
62
  G8 K2 |9 P: O; n( z& V2 @, D63
) ~. A4 {7 x( v! u! t64  z1 }: ?2 K3 w# P! `6 p
65
) q2 Z" g$ g  ~$ z/ p& v66, \. y* J/ b8 P7 Q) ]8 W
67
6 \' ?$ O3 ?5 j$ E" k68
7 V. d& L! a+ @  ~3 b. Y9 a69$ G& l7 w/ s% q2 ~
70& u8 |3 I% h6 n! g' x
71
: h8 v  v0 L7 |' ~8 s72
6 p7 a4 s8 q4 A+ k5 g: F- v: a73: k1 r1 b; _2 I
748 b& s. O; J/ y5 ?
75' {! I, j' f3 o/ b* u
76' d4 h$ a4 T4 k4 |/ u
77: r/ h6 C7 E& s% T
78
5 A# [6 R. v: A4 m( A79
! |! S9 _. _! E! [6 e80
9 x. m" B9 i3 ?- M/ W. k; \6 e1 R2 m) R81
2 M/ I4 J; k  c% e, A" j82
# W1 a$ T5 ~# a5 ]2 [1 d& w83
0 e: Q4 G4 r. g: x7 C% q. t" N) k. g7. 设置需要训练的参数% X# |7 h  [6 ~  ?
# 设置模型名字、输出分类数! {# L: |( L) D5 G
model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)
& w: Q4 ]6 q# h; ~
4 W0 s4 r6 o) k4 ]# GPU 计算
5 F9 j" O( J7 zmodel_ft = model_ft.to(device)
8 N( ~+ d3 r* d! h& F( @( m. n2 {
# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取
- Q$ Z6 G, J4 kfilename = 'checkpoint.pth') l$ l4 @4 h( s  L

* S( m0 r8 ?% ^6 G0 i# 是否训练所有层
# X. D2 H. Z2 V$ Z8 Q% ^params_to_update = model_ft.parameters()$ n3 t+ t/ M& u& U& c
# 打印出需要训练的层& r9 z5 \% l/ ]9 `
print("Params to learn:")
% c4 _- J* ~6 p" y& n" i  mif feature_extract:
- x; [/ ~" N3 L4 q7 J  x8 m- b    params_to_update = []
9 m5 `% \9 |( [# w7 S8 q4 g6 [    for name, param in model_ft.named_parameters():
# d8 h8 g  R$ j; ~" r; i        if param.requires_grad == True:+ F, d3 |2 T1 b% p* n1 ~
            params_to_update.append(param)
. u) f; [  V- l# }7 A            print("\t", name)) ~, ]" V* C- w
else:3 a7 y! o0 v  M1 N9 v: i- e9 a: d
    for name, param in model_ft.named_parameters():. |! |) p6 O; F5 p* D# E7 A
        if param.requires_grad ==True:# m" H9 p8 a* l% d
            print("\t", name)
/ T7 U: l. g$ B! W/ t  B7 S, N8 _2 E, T, Y5 y
1
. f0 o0 b& H! ^+ S$ z2 a2: V& r4 w7 ^1 z. x  K: ?* |
3, J* m$ Q7 E# q% J
4
9 @5 A" w3 J* o, g) X5
5 p- h2 r9 U1 J4 E9 N62 l5 W9 p8 H; l1 }
7& O5 c# I8 y9 o8 [
82 b3 U4 q& Q0 t2 G
9
! p. \8 N: }. S& c7 x# d10
% Z* d/ R+ P5 T; d& v1 P11; T# `8 h! G# w* d0 ?  x# v
12
7 \' B9 p$ K: {) G4 C13
) C: \' t, Z5 ~6 |# D14# U( D% C% s  q- a/ y
15
: j- K! [( @/ [- Z; W16
. ]4 l4 W5 E0 d5 y* Q17' ~; Y" V( ~* S, f$ h0 E$ ?
18
1 ^* w3 B2 L  O$ Y19
) U$ `5 a/ h: V% [, A# {20: o0 ?, A, w& d$ D2 b) k6 ]. d
21
+ d  z$ p1 I& Y: i8 c+ A22. g* [" S/ U% `) {4 X3 K5 h8 O
23! w% X. s8 K1 M: H) v  h
Params to learn:2 Q6 Y2 S. {+ o1 j4 W
         fc.0.weight
6 y6 G& M5 g+ u         fc.0.bias1 N9 k1 Y0 g# {4 g: u. r2 D0 h  P
1+ p+ u3 T- o3 l/ @7 F
2" K* a: D& p, G% k0 y* M
3; L, k0 S3 q) C" p5 O' I) g
7. 训练与预测
7 b8 J" X' C; n6 \1 J  R7.1 优化器设置1 ~" f9 k2 m& t) a$ s
# 优化器设置
$ |) G$ ^( N0 u! M2 ooptimizer_ft  = optim.Adam(params_to_update, lr = 1e-2)* S) d/ q+ s+ L( @- w
# 学习率衰减策略
2 r/ p  E! s' i  v2 K* p* e0 Sscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)
" `  l& `+ b. o4 D2 }$ Z  J- K# 学习率每7个epoch衰减为原来的1/10  O4 N) l' ^8 b$ b1 C
# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算
7 h7 \( y1 `6 O8 U3 W1 E8 \  g7 B. e7 S2 Z
criterion = nn.NLLLoss()
5 Y7 _2 t) n# q* h! G19 S! y  B* B0 [5 G% h" Q
2
" o* v4 S1 A" W% S- S4 f  v$ t3
. T) C6 @, C) N48 H8 \4 v9 m& T5 K$ w1 R. R0 i' p
5
; n( _- o: z& O8 i) [, ~! W6
3 E3 d$ L& A$ ^77 J9 A7 a: l% o
8
# S- Z; x; q2 W$ B# 定义训练函数( e6 [4 Q3 t+ ]* S
#is_inception:要不要用其他的网络3 r' G2 j4 Z; f$ Z
def train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):
7 ^6 y) w7 H# k    since = time.time()) h1 W2 a, c* g, o0 j) R
    #保存最好的准确率
) ^0 O) U/ F# N    best_acc = 0: ]8 u" m3 u4 r+ U' {
    """6 Q- T5 ~2 O, J% y
    checkpoint = torch.load(filename)" M- i' I$ [$ ]% J5 I
    best_acc = checkpoint['best_acc']
3 d2 ]& S  h0 J- @' p1 f    model.load_state_dict(checkpoint['state_dict'])
& I& D" @" Z$ \, @    optimizer.load_state_dict(checkpoint['optimizer'])
+ k- b% R2 c  c" U4 Y1 @/ I, C- @    model.class_to_idx = checkpoint['mapping']6 h) m  J9 T! g( C+ D+ w
    """
! H: T" U0 l2 ^, s4 z, l+ u* Q    #指定用GPU还是CPU0 S4 k. ^' R: B. O
    model.to(device); _. p$ r  w& \
    #下面是为展示做的
: Z2 a: m/ o$ ?& K3 S) w" E( S    val_acc_history = []. r4 }# p$ W* z1 b& i+ y
    train_acc_history = []
/ ~+ Q! p! d! x  k2 p' O8 M    train_losses = []* H* L( M9 Y6 j# P( G5 W
    valid_losses = []
: E' i) j+ S  @* X2 b! t    LRs = [optimizer.param_groups[0]['lr']]& M7 p: w- c) z9 A8 W8 K
    #最好的一次存下来
& S6 z; k: T# C* Q) k+ ^( z& K    best_model_wts = copy.deepcopy(model.state_dict())' l9 e8 v5 v1 _) }. O3 x9 \: ^0 u
$ ~* \) v7 `: X) F, c# Y
    for epoch in range(num_epochs):% F9 t( A, f. ~0 g/ K  X: s
        print('Epoch {}/{}'.format(epoch, num_epochs - 1))  H; K1 C5 ^  }$ h4 T3 V% i
        print('-' * 10): u, L+ |/ q5 c; ~9 F* d9 O9 @" l

+ e+ k* F+ v; G2 x1 B" a9 j0 a        # 训练和验证( @* m# V3 ]* t% ^. R$ `
        for phase in ['train', 'valid']:
2 o$ i" ?3 m: A; J- ]            if phase == 'train':% _, f7 J' B2 N5 T5 \4 O
                model.train()  # 训练3 t/ b) l4 p- W, a% U2 K
            else:
) A1 v/ B1 ?( |4 {: w                model.eval()   # 验证4 _( ~4 u8 y% m9 Y

) ~: `6 z7 W! i# `            running_loss = 0.0) P) I. S7 W: i7 B. i
            running_corrects = 04 T! O7 }! j0 O4 O

5 _0 [/ o8 l  k6 l6 c- T0 s. B; m! F            # 把数据都取个遍
- p& `$ |0 S! }9 M/ O% w            for inputs, labels in dataloaders[phase]:* `/ K! p8 w2 `, H6 l: l$ A  A8 u
                #下面是将inputs,labels传到GPU# M0 ~% {3 I" |* h. r5 x
                inputs = inputs.to(device)& s  M" k* j7 Q( o7 h1 ~5 [7 r
                labels = labels.to(device)7 w" c& l9 {8 S. k4 `
. Z0 D- h, ~# q/ t: y# R
                # 清零; b0 {( ?- G* }& ~) s4 T2 O, f. A
                optimizer.zero_grad()- x! j: z  _$ b5 D
                # 只有训练的时候计算和更新梯度
( R* {4 N  x1 b- ~4 {% i                with torch.set_grad_enabled(phase == 'train'):
& W, W7 Y  H* o1 A" L* K* w                    #if这面不需要计算,可忽略
& t6 J( s' T3 I' g" y5 f, I; W                    if is_inception and phase == 'train':& ~! t2 v2 i$ @9 {
                        outputs, aux_outputs = model(inputs)
! \6 k, _+ Z: B. E                        loss1 = criterion(outputs, labels)4 M2 q! U) B) y4 y6 w1 H
                        loss2 = criterion(aux_outputs, labels)
( _6 s  f  m- p; r4 U                        loss = loss1 + 0.4*loss23 k9 f6 ]8 Y5 n" r; M- ?, I( o% d, \
                    else:#resnet执行的是这里
' d% [$ z$ H% R" W; h                        outputs = model(inputs)
, [& S! ~1 L8 P# a  |4 C* \- ]) R1 f                        loss = criterion(outputs, labels)% M8 m! D( l: q7 R
- k- e  Q* y+ F4 Q
                        #概率最大的返回preds
! ]0 r) Z4 x3 g: T; T5 C                    _, preds = torch.max(outputs, 1): x9 A' j. e# {. P
4 q& I# I$ e( v/ s" k: Z2 ?* ~1 J# T
                    # 训练阶段更新权重: g  k" K7 [3 v" o6 z" ?
                    if phase == 'train':1 V. @* y3 H; y
                        loss.backward()
- `6 f- U! y/ N& e1 m                        optimizer.step()( W: |# u3 _( {

& B" w, q9 J3 L" u) v                # 计算损失
% F- Z6 Q4 V- A, B  S0 z% m. s                running_loss += loss.item() * inputs.size(0), Y4 q6 l- |6 [  M& s" V
                running_corrects += torch.sum(preds == labels.data)7 Z( ?  _9 j+ p, R7 X( {" W
6 X; q" D% _9 `2 A; V  U& l) F
            #打印操作! U; [  E) Z$ U8 S2 W# M# L
            epoch_loss = running_loss / len(dataloaders[phase].dataset)7 n9 W* h6 U8 Z0 g
            epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)  p0 T  `- }5 y, ^" d

$ W! x4 E! b- f* b0 b" R6 E2 l8 R) G
            time_elapsed = time.time() - since, Z* \; T$ f- C
            print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
, Q6 v: N8 b" N8 C            print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))! T/ U2 `# @8 ]8 X% g4 G5 C1 x! k

. @$ F% l) D4 J$ R- P3 m/ h- `% g+ v
            # 得到最好那次的模型6 `$ s5 P  i4 u; }2 U  W
            if phase == 'valid' and epoch_acc > best_acc:
" U. R6 ?$ x: ]  C- E# }1 y                best_acc = epoch_acc
" W$ v) U! R; X0 _                #模型保存" K2 ]' S' a- ~! z$ h9 U4 `
                best_model_wts = copy.deepcopy(model.state_dict())
: h; y& I9 k( y0 S$ ^                state = {
0 ?* i8 M5 Z# j& H$ l2 ^                    #tate_dict变量存放训练过程中需要学习的权重和偏执系数
9 N4 d; y. T4 y2 b! J7 t                  'state_dict': model.state_dict(),
. M% i2 ~$ j5 Y, X8 G" k3 N                  'best_acc': best_acc,
, }1 z1 d/ W1 k* Q                  'optimizer' : optimizer.state_dict(),2 B# K3 C* F5 W9 c: t  o7 i
                }+ l4 H) c( w9 e9 g; m# d
                torch.save(state, filename)
/ |- S- n) d  i- U  S. S% _' Z* I            if phase == 'valid':
$ z: D9 ]  G% J; g/ \% E9 }                val_acc_history.append(epoch_acc)9 l$ t8 U3 F7 M# L. u- P, S
                valid_losses.append(epoch_loss)1 ~/ r4 O9 @3 y, ~' a$ J
                scheduler.step(epoch_loss)! u' ?1 j, R% I! `. f; e
            if phase == 'train':( p$ p+ M# x0 D" Q, @
                train_acc_history.append(epoch_acc)
$ Y  t  ]  J; L  t+ q7 {                train_losses.append(epoch_loss)
  t, m* x4 H6 C/ R4 J4 @+ T. J& u
        print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))! }$ _6 D: w9 D) g7 _1 J& R
        LRs.append(optimizer.param_groups[0]['lr'])
- m$ O8 d/ A( }$ Y        print()
; i- q+ S, I0 L1 D5 k1 Y$ f* F2 y
2 P; A/ E* D. ]0 g8 m) K) c& G    time_elapsed = time.time() - since
4 g9 F. v1 _% c4 Z: U! l6 E, `/ _7 ~+ y    print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
7 q, \% S. a. j% [    print('Best val Acc: {:4f}'.format(best_acc))
" Q! M, \; T2 a4 }
9 R/ d  v. I4 n. s% ^+ s, z    # 保存训练完后用最好的一次当做模型最终的结果
' C& b4 u# M# a$ L% q4 l1 E+ Q    model.load_state_dict(best_model_wts)5 h! q. E# x" q, v+ R3 {8 ^7 ?, X& I
    return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs ' d5 [: d& j& G
" a1 j9 }7 V3 L

6 |; V# i9 \* T0 Q# F$ q1
2 c1 ~0 J  F) C: a/ U$ G2
! z4 Z% n# B" M% f: J# `3
2 o% P% K2 y, Z  u6 U# L9 @) E4* y. P$ ?5 \7 B( E5 t7 m( f
5
6 j, p6 L- V$ ?6
- I$ Y' u, |" c; }# x" {4 h7
0 u' V  |, U* c% a2 D8/ w! W7 x- s. t  e& c; j* f
9! P2 ?" t0 x# Y& n" s
109 a. A  j# C, q
116 o. `* d% \' i4 r8 O
12" [7 p) Q3 l0 q' E" l/ ^
13/ V: l% G* K6 t. x: O" k
146 _7 P" t; J) D' D3 O4 c
15
; c1 L! ?" O9 u! a( r' K! x16
/ U. x( X7 _6 z0 Z! A17
0 ]: \8 K9 u( a% |0 P$ f/ N18
" w0 Q( v5 g5 v19* z5 }$ q1 V* b' _* e9 n
20# x0 M2 [5 D0 e1 f* r% g# P
21% E6 L1 M) p/ X# M
225 G  v2 R* _) |8 f- X* v
23* H3 [5 r7 ^% ~+ }8 r, m% D
24& o6 q( @1 ?. U7 K
259 \3 v9 }) e- Z$ x
263 O: ~9 l+ P0 u, c
27# w  D( _1 H. }# W) I! W
28
6 Y! r5 R  q9 l% C; V4 C29  B0 }) c7 `2 V2 r; W$ v* y" G0 U
30
/ ]. f7 Q* B2 g' ^" `* D+ U7 h, W310 H. C5 U, \( W7 ^9 c
32
  O8 h3 M% S1 o7 d33
4 T$ H& R# J- S/ k, _34% j) b5 [+ w* ?  F1 Z& L1 T  J
35
- M  T8 w- @- O. p  ~& b363 ~' _: u( r! w6 a( {- u
37' i3 C" p5 o5 @9 ?* Y4 u
38
; c' a% v9 e. g# V  x1 B39
' ]  I. z9 `# S5 v9 @+ e) |40. _+ L: g9 Z% q  [+ i
41
6 m0 p, v8 @3 r' p42
; ~6 q' @/ c0 l430 k- C) ~: Q5 `) e1 x. `
44* `$ [# t$ u1 z7 y- r
45
7 W' g3 G* a( Q8 E8 u46
3 ?# a* G* m  [47
& b1 ?3 G2 z( \* _48$ c: n! [% p2 ^' e" g
49( A  N  }! C8 i/ n- J* }. k
50# _6 [& G% K& L! F
51: \1 j2 C# n5 m: E& p$ v
52
2 N8 C& [; _2 v2 H8 d53, |" `2 U$ [* H
541 U- x+ {7 X( Q) V5 m) A
55
: q* J, {) t3 [3 d) A6 B56- `7 N% ]$ F. w; H, W# h
57. t' D) q) k! s/ s. D1 V
58
. C- \, _% R# [  p6 a, Y  n4 q59
0 f& M9 o  x1 c# `60& Z, I' s5 c  ~# k7 S
61
) M" v& ?' X" k: `# L# \62
& X) D0 g0 O$ Y# v7 r$ T63
+ q1 @& q0 V( T7 ]$ `64# ^1 M( _: J6 I( q% T
65. }1 _; d7 P3 h2 Z% }5 A
66
( x" @( z6 R* r! @67
6 o1 b  Y. J1 U) M68  b/ G4 ]# e, Q+ }" }+ T
69) s, c( d# F' p
70
& M- |3 o! N; ~9 m$ L71' A- t9 I3 x9 E' F
72! C! i  p( K' F7 S; E8 c
73
( _# {: r! R! I. D74( q' J# q# f7 |8 R( R0 f
75# _: Y/ `! t1 P' W. T
76! b: y7 e8 R+ D* T( \+ R. g9 ^
77
, s" d, U0 y$ I& H6 a78$ K" {) u. U; c6 E" i# ]
79
6 D  ]# p* z. M$ K/ Y; ], ]" u800 s- H: M8 o) G
81
6 X" d( Q1 ^2 A4 f82
& Q5 A" `; n2 }9 _9 s, m( ]835 f0 ]3 b  P/ m% B) O! `
84' I/ d3 N: y9 P8 S0 b* v+ U. R' t7 }
855 T) h8 z8 M, o  q4 s' }- i( q4 W1 a
86( ]. k* |( S: ^* f
87
: Z; `( d6 Q$ `1 V* A$ F! R3 S88
4 @0 ?& ^( R) b) g+ f0 R. s4 m  I89
. y! e  y5 K$ {90' Y: r% ~2 {& U9 H  ?" t+ Y$ r2 E7 D
91
" N9 N% N7 J( _' P; T1 R923 ]4 {3 g3 H& d2 j9 y" ^
93& t# O7 ], v( E& D. Q# Y4 |2 C& n
94
6 H6 a7 V9 _% l% j  Z95
9 R/ k7 S/ s  A2 O' J" S96
3 j) t/ \/ R0 u; {$ X8 ?972 W) [  o4 Z5 l+ X
98' z, [4 e. L3 s6 d+ L( g# ]/ ^: l4 H
99
7 P  D: z& q# k. |4 T2 A100
4 R' T  {: G8 u+ H$ w3 [& ~101  G4 S! v0 T" W, G9 ~% a( a
102
, T: X4 `  K) F3 a2 {; R1034 f7 o. \2 `0 O" e) K4 X' e# y
104
4 X% c* y6 I, D+ \) Y; l# n105
! q9 l4 T( o! X8 `8 k- P106
* n7 M) c6 G4 J107% |) S7 X5 b2 G! x9 B- i
108' Z% B+ F: i# }( S3 `+ T- ~# z& j
109$ o+ N0 I, j7 ]& n) ]
110" ~: _4 I# j+ ^6 \$ w: p
111
  D0 W! A, @5 r3 n3 ?3 c112
' H6 F5 Q, z' ?+ l7.2 开始训练模型
! Y& g- u; |. \; T+ j4 u: P' B我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次3 W5 I( a( i0 Q# ^% i5 H: v

9 E* M. b3 g  z9 Q5 z7 P; H#若太慢,把epoch调低,迭代50次可能好些
$ @% L3 {* E6 D+ k" u#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合
; R4 ^" s# T: p+ rmodel_ft, val_acc_history, train_acc_history, valid_losses, train_losses, LRs  = train_model(model_ft, dataloaders, criterion, optimizer_ft, num_epochs=5, is_inception=(model_name=="inception"))1 |1 N$ S4 m8 l, y

9 b0 ?6 O$ W# }! T! @7 E  [6 B1- N3 z' g- u# [6 f; \! I
2
! `% s' ]% G9 N3 a  C/ m3
$ G; d) z5 o$ h9 O- j+ Q/ t4. H1 u6 `& B) ^1 l  z2 u# A
Epoch 0/4. ~  s! D! A$ i
----------6 j" F4 Z+ |& R. Z2 p
Time elapsed 29m 41s& A0 ]2 o- M5 [) b! i
train Loss: 10.4774 Acc: 0.3147) p, I! \- e, w) M& G4 R, U9 G
Time elapsed 32m 54s
$ J" C# Q) a3 T; pvalid Loss: 8.2902 Acc: 0.47196 K: i% t# g: Y5 T: C  d9 D
Optimizer learning rate : 0.0010000
5 P0 q0 C4 s4 z8 I) k
4 Z8 e+ J% [4 B3 R; z1 r  O" Z+ LEpoch 1/4. ^: a% O/ `- t- v+ q7 W+ _$ y
----------; f0 i. L6 r& m) U) F* |
Time elapsed 60m 11s
" ]) a5 f6 u* C& J/ f8 z6 M7 ^train Loss: 2.3126 Acc: 0.70538 @4 L" z/ P5 E# K; b
Time elapsed 63m 16s  s" v# s2 x) s0 j0 R4 b) K
valid Loss: 3.2325 Acc: 0.6626
3 e6 Z4 ^$ W0 ]: o  D" m4 c: {5 LOptimizer learning rate : 0.0100000  a& d7 w1 s4 G" n1 j3 ?1 `2 H, c

$ e) {! B& ^$ ]) a: l, h3 S- J( gEpoch 2/42 n9 ]9 j/ M$ a& K, ?: A
----------
* m% h) z( m" B- N- p/ kTime elapsed 90m 58s
; z5 a" b8 j" ?. wtrain Loss: 9.9720 Acc: 0.4734
9 M, V: s2 M7 D, {Time elapsed 94m 4s0 X; z4 o- h2 V) g# d
valid Loss: 14.0426 Acc: 0.4413
9 p+ R/ `3 j# s$ t' r' N4 V1 AOptimizer learning rate : 0.0001000
4 Y! t7 a1 a& A/ ?6 c  U2 x9 F/ z! L+ `( ^' ^
Epoch 3/4- d( u& o" m9 S
----------
# X- R: S/ T4 r3 y% J# jTime elapsed 132m 49s
' }9 d4 H: ^3 l: r' F' Ktrain Loss: 5.4290 Acc: 0.6548
) \: c1 h1 z( e8 j: BTime elapsed 138m 49s, r6 S: t7 M, |" y. }
valid Loss: 6.4208 Acc: 0.60276 b. V, X3 B& l: z& O0 a" {- \( Z
Optimizer learning rate : 0.0100000' {7 ]" l! T$ O
: v4 Q" E4 u+ I6 U. ?9 _! j
Epoch 4/48 s: t6 L3 {  k
----------* [* U5 a4 o" d7 t7 J, S! t, G
Time elapsed 195m 56s6 Q+ ]9 b: D1 n' E2 ]
train Loss: 8.8911 Acc: 0.5519( ?% Y; `4 q" u7 k/ S$ Y6 _0 F' i
Time elapsed 199m 16s# D6 I% z9 U2 w% ~% q8 F
valid Loss: 13.2221 Acc: 0.49149 X, G: q) P9 q; i) T2 t
Optimizer learning rate : 0.00100001 A  O3 v4 D- R2 j7 Z4 |
' `. T( v% P/ O' [
Training complete in 199m 16s
1 o% s# r2 O+ n  e6 L& G0 v  J; S" O, wBest val Acc: 0.6625923 s4 M% r: r* v& D/ s& m
: Z% R6 K1 n2 V3 g
19 k. u" S& J+ J0 l/ l
2
) j9 k* b  O4 T" j3. G, Z, [9 N- N- W8 z8 p% @: a
4
; z! C& X$ f0 }6 J# l4 i. }  q5, l; |5 Y) U" ?' P5 Y' A
6# V3 L0 z  w1 u3 c. `; k3 {: P
76 ?( ^* s, R& y% K3 D2 m+ K2 x2 V6 E8 Q
8  a5 w, x* ~( O* Y3 K- s$ }
9" Q- k: f4 `2 W4 p& G
10
; h/ z" ^) G# O9 u( L6 k11
0 x) T; G1 p* J' n12
+ a5 _  {) V9 ]* Q, E. o% T& n13
5 m+ u/ V6 A( K$ n, i$ h5 {147 ^9 U' F+ O* [/ @
15; e" b6 ^  d+ \1 ]6 y+ S
162 I! l5 b" }) d" w
17, T0 t; K; D% u6 {- I( T, `0 W
18' M& k3 W& w; @7 V8 U: o6 W
19: \1 D! q- z$ N: x
20
" U( k- x) E5 P$ j$ u& M# f% U21* z. D( O; M) A- W
22
) Z+ F( b0 H( L7 N23
1 C4 B" I$ ^; r. _. J6 I8 |" N24
3 o2 a" d3 \; H+ L+ t250 t# v/ B& G) Q4 h3 j
26
! A& y7 b0 b9 b8 b  {27; a/ Y; S8 \- N! }' A1 D% F7 S$ O
28
2 S/ ~2 e. ]+ ^8 E- \294 B1 X) ?$ T2 t0 ?, S, Q
30. f4 b# I- R2 s% H/ \8 T( L& I- P
31
' b# @# V4 @5 N8 F- j1 t$ ?32
7 @+ O) u+ c" x8 L; J( Q2 [) p( n( y$ \331 f9 W2 H6 F3 A* j, \# L4 U
34* I: S2 Z( V+ l: P% z, q- l
35+ `1 N9 q# q) D
36
& `7 J: Y- K+ ^* w% L+ O& {37
! z- }6 G4 D% g3 m5 a9 {1 ?38
7 e' t6 l0 ^  f+ ]" \* W# m) i4 G39$ P' t# u6 W% y4 h& a( i: ?4 W
40, g8 A5 c9 f/ T8 w3 E, x
41
- \! i" A# M  C& h, h2 v& w* Q! ?42  Y" c+ R, f  q$ o
7.3 训练所有层0 V3 X. c* `8 l
# 将全部网络解锁进行训练
4 ~3 v) m3 S' p/ Vfor param in model_ft.parameters():! V) e% [4 F6 r  Q! @* R
    param.requires_grad = True- ^4 M% P) p9 ~& n

5 O( u9 R& o% Y# X8 Q3 }7 n# 再继续训练所有的参数,学习率调小一点\
, \  M: h1 s6 M3 roptimizer = optim.Adam(params_to_update, lr = 1e-4)" r: T5 l, Z8 m( l) x3 W
scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1); B0 C/ Y. ^2 [- h) b

( J- y3 u5 f: A+ {' ]# 损失函数
: Y- P0 f" T; w* v+ e. N7 j: V5 Ncriterion = nn.NLLLoss()6 v/ U2 q: P; V& Y6 V- q
1
+ i7 `- M- I8 j* f- g0 B  |: B  X2
  Q8 w- a+ K$ v5 X% l- e3
5 j1 p! M: B0 l" D; N2 Y* ^4
) K  c- z- ~5 c$ ~5- X3 c- e, v2 K
6/ P0 }* B  u3 t9 Y, Z/ [$ g, F2 O/ ^+ B
7
+ z, Q% o' P; b* {* i) q" Y& n87 k- O/ x. \5 L- u+ [& r  y# S
92 R1 r0 [+ p2 a& n$ B( G$ x
10
+ `% k4 ^  H& l' E* d: R: Y$ F# 加载保存的参数
0 ^9 N8 o' j- }$ I3 z9 h- P# 并在原有的模型基础上继续训练. O# o, ^4 Y9 }; [
# 下面保存的是刚刚训练效果较好的路径
3 B2 n9 a3 j" j- {checkpoint = torch.load(filename), z3 q2 z7 _1 N" m* m$ ]
best_acc = checkpoint['best_acc']0 i. h, f% D1 \& ?) h( \
model_ft.load_state_dict(checkpoint['state_dict'])  N( ~0 }! b/ i7 D
optimizer.load_state_dict(checkpoint['optimizer'])$ e' _& w9 Q/ R. L8 K
12 a0 v0 w& G/ j
28 S4 e/ e$ X' P
3+ p7 f/ U9 x- {- D  s% J+ s6 e
4
' V1 T2 J: h- p. w1 d/ Y5" F6 G0 G2 e9 w0 d7 c4 L
6/ _4 a- o0 e) r. r3 p
7
1 J  k: X0 |& Y5 K+ Q开始训练
: v+ R" m; V/ P5 S* X7 n/ o2 j* u9 P) q注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考: L. X% m+ P5 \/ O7 q% d$ b

( Q4 K% U. m- g9 ]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"))1 t; W2 k. [" b8 c8 ]3 ?
1
7 D3 K) m3 Z4 A0 J3 O- }2 U/ IEpoch 0/1
! @5 D, n8 l* F/ g7 `9 M/ ?----------* D9 y, {" I6 w" ?
Time elapsed 35m 22s
: F- E- R; @1 }train Loss: 1.7636 Acc: 0.7346
/ `  ~5 @; c+ E% i; r) i7 }# lTime elapsed 38m 42s. U7 n5 d# R2 \9 M2 v
valid Loss: 3.6377 Acc: 0.6455! O  u4 v* N0 |( _
Optimizer learning rate : 0.0010000
7 M5 j7 L1 _0 d/ K7 w& P2 g) Q# v7 n
Epoch 1/1
% x. g" W2 D  o& v5 L----------: w+ r4 d% U0 ^, v/ _/ G
Time elapsed 82m 59s
& Y, s& \& E' T8 ktrain Loss: 1.7543 Acc: 0.7340* e% E0 Q0 |9 b$ z6 A. o3 n
Time elapsed 86m 11s, T6 N- {: c' K: H
valid Loss: 3.8275 Acc: 0.6137
# c* }/ G7 f9 {" i! J4 f2 @Optimizer learning rate : 0.0010000
: g1 L; m7 C  I1 J; }5 c+ }5 R, J9 ]
Training complete in 86m 11s) m$ W$ L0 D$ ~6 H6 I% K& p% J* R
Best val Acc: 0.645477
& E) u3 ?9 V2 z! c% _5 R% ?. p) p$ ?$ \
1
6 Q' q9 k. b9 \& M5 e; D4 {25 X2 s6 l8 P6 O. p# P
3
- I* ?7 ?5 I) _% L: E& W- n8 J3 b4. u' ^- K; N. d, t3 Z
5
6 |: a! ~0 _6 k" D  B2 \' J6: {! {+ C& M) Q: F# v( U
7( X3 s' d$ k) j, f7 {0 c
8
; `- @+ q" B- ]. a" E; o9( K; Z  o# @& `. q8 ?6 U' A, s
10
$ w9 v% b; Y/ p7 @11
" j+ m# D7 k  e- R& U: j5 U! ?3 C12
1 W7 M, y  E, I' T135 S- K8 Z# D4 J9 F
14
  o9 ]* Q* @0 Z8 x7 L15
1 o, N% @& G7 ?$ @4 o16% b6 y; F- d9 f4 w) ^
17+ c7 t2 y/ H* j/ {0 d* n! L
18% }' B. p8 i- g  Y
8. 加载已经训练的模型% R( V( x1 [" J2 k, ~' r1 i4 v# L
相当于做一次简单的前向传播(逻辑推理),不用更新参数* s: K0 k& O* k+ l8 S$ o8 `6 m

" {5 s& n4 j5 h% L0 `! cmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
7 p- _6 }0 ]8 t+ [4 v1 U! n/ z+ ~" |9 i. R; d& `3 g
# GPU 模式
6 e$ t& y$ t' W# _+ Omodel_ft = model_ft.to(device) # 扔到GPU中' z) e3 Q9 R! P
% D4 \( _- b; {2 y
# 保存文件的名字3 X( t/ \/ y% Q9 l1 V( h: {
filename='checkpoint.pth'; z( y  I3 S# ?9 Z2 y+ t
' f1 t7 T- ^6 m" ]  ]
# 加载模型( o! T. ^, g1 _: E- q- {
checkpoint = torch.load(filename)
4 D1 I1 x- C$ m, }+ a! Hbest_acc = checkpoint['best_acc']& D5 B, v; o9 j& r
model_ft.load_state_dict(checkpoint['state_dict'])
5 y% o! {+ |$ W1
1 ~+ ~: V7 U) o" a) o5 z/ G21 A! ^7 [& e' c& }
3$ ~" j& }% c3 b& ~
43 ~" N/ m: y. V
5) s: ?2 b2 D) B) d7 f
6$ F  O2 d! D1 P$ U( C
7
! i. f% _! Q. h2 v8
( V- K  e* k& Z: t1 Z% t9* j9 k! Y% ~+ \* ~8 S
10
  P9 _, f: X" q- N! `% m116 g5 W$ y3 s6 m, O
129 d0 Q0 d7 A8 B% m# n' u+ ^9 {
<All keys matched successfully>  k; |' p7 ]% M6 b+ {
1% C' b& {- l; f% W
def process_image(image_path):
' l4 Z% J/ h: K3 o    # 读取测试集数据
9 S2 L2 N+ w$ s, n5 A% y9 D    img = Image.open(image_path)( A; b" ]+ i; l
    # Resize, thumbnail方法只能进行比例缩小,所以进行判断- N- J. p# i% g4 u8 M* H3 V( R& o2 o
    # 与Resize不同, G9 B, M$ b! l) ~! o
    # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小
. P% s* m- T9 C. P. K! d    # 而且对象调用方法会直接改变其大小,返回None
0 {" x7 [3 |- q" e+ _    if img.size[0] > img.size[1]:
0 Q2 A* @- T) V        img.thumbnail((10000, 256))5 C8 z( W, T. i) [# X/ {
    else:2 S4 t- j0 z  R$ m/ q
        img.thumbnail((256, 10000))$ v* w2 t3 @1 k" [

0 `8 _, {% x) e) |. t3 M    # crop操作, 将图像再次裁剪为 224 * 2248 O/ L1 }' N+ C
    left_margin = (img.width - 224) / 2 # 取中间的部分; n/ _4 U, t8 d- y
    bottom_margin = (img.height - 224) / 2 ! C( k' m) t, _5 w2 @
    right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度% ?  p! N3 n; V+ e( W1 @
    top_margin = bottom_margin + 224  U7 e' j# ^; W# c# x; a

1 K) b! i+ p7 S4 p. h9 L    img = img.crop((left_margin, bottom_margin, right_margin, top_margin)); ~" Y1 I; T- M. e
  n: f) C/ {+ }4 u+ l
    # 相同预处理的方法3 Y$ a7 M' z4 ~) y4 ~: D
    # 归一化
7 |9 Q, M- B; @. R    img = np.array(img) / 255
) |" Y# ~' O6 L: D7 t0 g' ]* ~9 B2 n    mean = np.array([0.485, 0.456, 0.406])+ T) U  n4 M8 L) ^
    std = np.array([0.229, 0.224, 0.225])
7 T: K4 U6 k7 H# F    img = (img - mean) / std& b0 t* y/ H9 C% }% _8 g9 |& r' ?

0 m! C1 S! n" j+ `" b    # 注意颜色通道和位置  l- m2 n# W' {1 n; p8 `
    img = img.transpose((2, 0, 1))* H$ I; E6 P+ u  B; k3 p( C# W- t
& c3 _! G  N2 Z* O
    return img
( l# K1 Z7 n: @% x8 C* k
0 h! w) F1 s2 e/ q, j3 D; ^5 M7 pdef imshow(image, ax = None, title = None):
. w* ^: ~, m, \' s* Q' I7 j! }    """展示数据"""
  Y3 ?" X; P8 z% i6 k$ p    if ax is None:
% C* w! Q0 T  I3 \1 |: M$ \4 e        fig, ax = plt.subplots()
3 n; E* v% K0 U: N8 W) }/ H' _* e! }+ F7 U8 d: d0 q" {0 M
    # 颜色通道进行还原
, S8 W. Z! v! O    image = np.array(image).transpose((1, 2, 0))
0 b& @4 S- w  D. @$ F
5 X5 ]4 j; v/ z; C    # 预处理还原: |$ J: K( q9 `1 H% T9 i
    mean = np.array([0.485, 0.456, 0.406])- w, z+ F( M- H5 o  u5 t, ^- S0 W! q
    std = np.array([0.229, 0.224, 0.225])
. J: T0 t9 l0 W8 [" [) S$ U4 H2 B    image = std * image + mean# @! X9 s6 L. U7 U3 a; \
    image = np.clip(image, 0, 1)
+ G8 B, I) e! w7 ~! Y  Y; {
3 b! B; g' C- p    ax.imshow(image)" p' s/ b4 |6 {% L5 I+ m' `
    ax.set_title(title)
5 n; r/ k. d: N
6 M" P. _" r6 _* o& i    return ax* S4 B8 `& ?) |# v
- _: o3 j$ S5 f6 X' Q
image_path = r'./flower_data/valid/3/image_06621.jpg'& P; u7 i( E8 N* J
img = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理
- ^- X* K8 ]4 h: i1 l" n' ], Vimshow(img)+ v- s% z, W1 F# G# b4 B

; ]1 d! u# A! v- \2 _. \1
( j% |  n+ Y2 Z. a2
) T) F& I# v5 F3
( Z5 L1 E! }: U8 a) d9 s7 b4& I6 I8 F, H! G
5
! O  Z) p: C$ ^1 h: H* J0 o) g6/ K' f: h7 G6 n/ I
73 b4 }5 b$ ^' d# I  w
8% P1 C! @, e3 D% u/ W
9
& t3 O' h: `' T* Y" P! N10
: X0 H/ J1 N3 P- ?9 k11, N2 `- l: L' w
12
5 ~: D. F& C5 U134 L" v3 }) G' p. b
14
* w7 v$ ~0 M1 D, G157 S4 G/ I- Z6 \9 D; p
16
- e1 t5 K8 ^5 E' c5 Y% K' X171 E5 E3 j0 A8 T+ g0 {
18. k# e0 i8 C/ ?
19- h7 o4 Q! M4 i7 M: W  T8 |
20( J) G8 ^3 ^# C+ P; a& Y+ w
21
) d7 r! n+ c0 K& v$ H6 b224 [8 D5 F7 d& s- y% Z
23
) ^. G% B3 X( T# W24
# _) Z. Y7 \$ G( m4 F8 A255 r  L7 B  q+ D- {2 X
26
' V/ O# A$ c& |, v6 f2 U27' n) q& N1 `1 p/ ?5 s
28
" O$ {8 x5 F( b- Q( H  Q2 N29
; B% x7 E( F4 C# q30
8 c' P# t7 T) |. L& L- w/ O31* G. n% m0 m3 |- G( y
32
; z8 Z+ r1 o; k) g' j' u33" z9 C! f  U  X  k$ n4 @
34, a; A. u) x8 Y& q5 N  B2 E/ R, B7 t9 {
354 b( y0 B( e8 T
36
( v# O% y  i1 t: q( s" X37
3 k4 e, e' x* x$ G8 V9 ]385 K( N; t% L. c: k7 ^# t0 }
39
7 e) V0 U) R" g40
9 w2 C. ~& y+ X/ D8 Z) S410 F- |0 p# b2 z& \# ]) p
420 ?& W% c" d0 ]7 \
43
& C4 G+ ^$ D' e" K7 E$ p- x44$ Q2 z  ]- }8 @4 B
45
$ @" n* r/ Z2 o0 b; t" z$ @" {9 b46* q' M! u/ X4 k' P2 U6 S, U6 r
471 n( G0 n9 c' U3 k
48
* `# h$ N, S9 G  w8 C490 b$ o! X( `$ f/ g
50" B1 _4 t) I/ A9 R$ ]
51
7 g8 }6 P' U$ u4 Q" j52
. _( v0 g9 ~, L& U' U6 Z53
0 O8 p! G2 w' k6 f54
$ U/ L0 |! o7 M( l! W/ C<AxesSubplot:>
0 ~( I& n- [2 u; L/ Y, s1' |' _9 j/ S1 U+ \( R1 M; Y  F/ w

7 ?% l. u# H" \0 `上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确
8 e4 M$ P8 K" `2 i- z" a& k6 d6 M( e& l4 `* c4 m: `  I
img.shape
) @* v  S+ U( B* r1
! O! Y& z. n6 z7 P" A/ h(3, 224, 224)9 z( G' Q5 l8 g8 T
1
; V+ f+ O& ~8 m( p& z证明了通道提前了,而且大小没改变$ @5 X/ N9 M- M+ ?
6 D1 ?. x* `% y# S
9. 推理  G' }# }& M% O+ ~) _: x* G1 @
img.shape; H" ~; w$ N' L3 G* @& s' k
. r9 ^: A" H5 S, ~: C, l" S: x4 D
# 得到一个batch的测试数据6 O4 S3 r6 b6 g" b' l$ K
dataiter = iter(dataloaders['valid'])
6 X* ~( d' p+ A) |% g' simages, labels = dataiter.next(): C& w1 Q3 F/ d+ X! v- P& Z& W

9 L- ^, a7 N9 lmodel_ft.eval()0 I1 N2 n1 A" a, Z9 @2 c
  X! q9 y7 A! }. K6 i
if train_on_gpu:4 \+ ^; ^  _- u7 l& S% ^
    # 前向传播跑一次会得到output
( n  Y! h% M' C9 R/ N% r    output = model_ft(images.cuda()), D, y$ t& u, U( G& J, _  _
else:& L3 K6 Z: S; ]1 T  _
    output = model_ft(images)
6 S7 H5 [% ]& ?3 }
& I; R* i% o6 X* b" H4 W+ `2 @# batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值
- S/ H( a5 G4 K, }0 x) D1 |7 zoutput.shape2 N4 j! o$ E( w7 j3 |+ v
- W* e8 l0 X  N0 s  G
1
( b1 N6 G7 _8 F* u2 b) K2
+ v' Z% ?+ d5 V7 X3. q3 n! ^! |" k8 Y. D. {
4* t( @; ^; `# s0 X4 O  _
5
3 V" F' W/ }9 p& R# p69 ~3 d- ]' k& o# S8 Q* I7 r
74 Z/ k6 o' M2 f' d0 I" V
8' Q3 \1 c9 p! D3 O% ^0 l# P
9
5 X# F% D" ~4 G4 Z, j) n10
- p& G( E9 ]+ o3 v1 ~+ l& z! Y11" ?) v% }$ j7 B" F
12' z2 A' i% m. E$ W* B0 U/ Q: f
13
/ @5 J4 P+ Z  g- q' I14
, G" J/ r$ _& o- X  c8 C0 I15& F" F! p! f' E3 L8 H7 q9 E1 U
16
5 m! S. y/ y* Q) _/ d8 Vtorch.Size([8, 102])
5 P/ q) a6 v6 [' Z5 g1/ X/ V2 T. L& @4 ?( e
9.1 计算得到最大概率
- _+ v: ^# O) q5 n0 Q_, preds_tensor = torch.max(output, 1)
; V4 L+ ?1 b- F; O. h: g* B( l% U/ C3 d/ ]5 C
preds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量4 a( i) b# s' |6 Z/ i. ^5 r
1
' f; f7 I3 q9 P$ g$ L( L2
9 U! T5 J( a8 k6 b+ E3* P# I2 u1 m$ R
9.2 展示预测结果; p' ~' U4 f: _
fig = plt.figure(figsize = (20, 20))
1 P" [( f0 o# \: j; B4 ]3 B' e& M- Gcolumns = 4
% A) N% \1 O: u) e" P4 Y. E% Crows = 2
" ~" P9 b6 s8 v; l0 C+ H
0 ]0 c1 |# [5 O9 U. `$ L3 `; Gfor idx in range(columns * rows):
, ^# I& H/ I' @    ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])
/ ]0 m, `5 X8 h; ?    plt.imshow(im_convert(images[idx]))
4 b! k7 E0 T$ G4 x, |4 [6 w    ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), ) b6 c* ]# w1 ?2 {: h; m. Z6 q
                color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))/ O# g" [0 |2 W9 K4 l# Y' x, e4 S/ W
plt.show()
2 R& V$ X& N* n. f. j( ]# 绿色的表示预测是对的,红色表示预测错了
" l" q  ?  X  w7 z# x  }, r17 t" D( A  w! [: K
2
, n% D! r2 u/ H3 B3. C5 @1 G# p2 v0 i0 l
4( h8 Z  i' l* F( |- {
5
8 U8 c% {2 {, q% E62 J$ S! N  _' x
7* {9 j7 k- n7 @+ M: t
89 c$ v& }! o+ E  r
9
9 ?. U( W7 i* K% D# d+ \10
4 d# e5 q8 H) c! i111 S- f$ N" a$ F/ u
: [( d; [; {1 Y- O

" n2 I9 h: ?, I: S$ a3 ]7 o3 p1 N( C6 |; m
————————————————) G, ~/ O* X9 t0 d
版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。
" ^7 E  d, w2 f( r# P- n原文链接:https://blog.csdn.net/LeungSr/article/details/126747940: c* h( v. r. C( p
" y7 r2 S! M7 ~, g% \" m9 A3 e! A% b# u
; ^, g. A% W! J$ A





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