数学建模社区-数学中国
标题:
【深度学习】 图像识别实战 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+ E
2. 数据预处理与操作
, k& `2 J; W) }0 m+ B
3. 制作好数据源
3 J, ^; z/ a# r6 x4 b: E
读取标签对应的实际名字
1 W0 c% Y$ r* ?- H
4.展示一下数据
N5 @0 n, R$ F" T. V
5. 加载models提供的模型,并直接用训练好的权重做初始化参数
2 R( S' B2 `" {$ E% z% h' T
6.初始化模型架构
% E% h- K/ a' H
7. 设置需要训练的参数
& ? 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 F
flower_data
. N" b- y+ H5 X e* j
3 c' ~$ C" n9 K
train
+ F. L4 }1 c8 @+ Y. O5 n& A
' J- q7 F, {0 Y1 I
1(类别)
6 o! {; Q- ]1 S: o: L! e4 v( a
2
& R! c: N/ b3 ]( r. O. r
xxx.png / xxx.jpg
& s* @/ w$ f# F0 I4 y5 B
valid
0 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 N
import 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 n
import numpy as np
$ c c6 V, o% d1 t6 I& n; f5 G* u
import torch
: I" L- S+ _; F3 l
from torch import nn
6 C* j( x0 Q [9 Y
6 F3 Z* R$ F- T$ [% r& L
import 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" e
import time
1 @- Q- A. T) N
import warnings
8 t7 _$ P% ^- w
import random
+ P/ ~% A+ A7 M/ m0 U& m$ @
import sys
7 W% j- K0 }# M+ e! Z; k
import copy
V$ X6 O) ] m( p# x' W0 x2 \
import json
/ P1 [0 H1 `" w7 Z1 v1 ]1 g- E; l
from PIL import Image
+ f2 Y' r; V( I# J+ b& g+ s
9 y3 ? b3 r* j2 V9 L2 F
7 R N: l% N4 B' @+ h4 O
1
4 H7 m8 q% {! D' L! z
2
4 q, Y: d5 n. C* e1 g9 Q9 p7 Z
3
1 I' b# p& N) H9 Y/ J8 P
4
9 z+ c6 r& B P- [( J( U
5
3 J7 k: [( t* d# X3 g) h2 O3 @
6
2 F1 F$ J- c2 N$ `4 G: f6 q
7
' f8 G, x% ^3 z; U1 D0 i% Q
8
" i2 |5 O: ^* B1 G' z
9
5 V; X/ v' F( v$ Y9 N1 O4 `1 y5 w
10
8 ?2 Q3 R' u, J' i% p! i0 z T8 f
11
/ ^" v3 X2 H# N' X* l/ p
12
. 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) u
16
# Y/ c& [- t5 ~: T& E( \
17
1 ^8 V2 K1 [& ^* n+ b% \7 T
18
0 {8 @1 N+ k5 d6 N \
19
6 v; o, h/ {6 ~ v) b1 O6 z: d7 E
20
1 D3 g( C. `# p( c$ L3 H% A
21
6 G0 }5 S) v& r( A# c4 X
2. 数据预处理与操作
+ 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 p
train_dir = data_dir + '/train'
( h2 R! H9 g! ?* [
valid_dir = data_dir + '/valid'
) n/ _- I+ x+ G3 s
1
( O( O; d8 U A: `/ v
2
/ b. j% z: t, ?2 M: d
3
$ u, x3 F( e4 \" H8 U( z( F
4
1 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# \$ r
3. 制作好数据源
( i) X3 V* M5 N# y R0 `) q/ B
data_transforms中制定了所有图像预处理的操作
! ~+ c& B) p9 I b. }5 a
ImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片
6 Y% O9 S, t$ \" E6 g% c
data_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; c
3
+ |3 I# j& F" d
4
. 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 o
10
' x, k `6 ~8 K% _. g: b! c
11
3 f/ H8 Z: u4 B% D0 C) G/ x
12
! }# e, m8 O, `3 }
13
# r2 |; |5 ^. N! H, A& C) K
14
8 g6 t l# o( ~7 U3 p! _% e
15
- S8 X2 J3 u; c1 X
16
- Q, P% G7 S/ s5 H
17
, n' q# ` A* a# F' X
18
3 w& g k$ A; |" E/ s/ Q
19
9 m m0 _% g' b
20
' ?& c8 e- | J1 k7 k7 F
21
2 A, N$ K. m5 N# K' d B
batch_size = 8
# Z+ V5 t, j) P0 M k9 K
image_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 K
image_datasets
3 r% l) P5 G+ c$ N
" j+ L) W- W: T$ Q# Y1 G
1
2 t4 m9 s( T$ s2 t
2
4 n' q2 G- w' I$ i
3
1 k3 q% Z* X- L& v$ f
4
. X4 Q+ Q- E$ A! [0 M
5
/ B+ n( l$ u1 {2 e
6
. N: |, H; s G) D. ?8 c. N7 v
7
, V) b$ B% k& d7 ? T& H4 |: B
8
; f1 g, }! X8 ]. h7 e
9
4 ~1 X# k9 f% y3 v
{'train': Dataset ImageFolder
9 n0 F% |) S' w$ Y! f8 ^
Number of datapoints: 6552
4 `4 a* q! m, Z) w u" P
Root location: ./flower_data/train
5 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 o
2
% g' o% {$ d( x0 Q" @ g3 v
3
7 J. j& J, h6 c) V$ x# y
4
+ s3 _6 U i+ w2 x, \- `7 M4 m
5
7 N. A3 T9 G6 V. R
6
7 Z$ s3 ]6 _0 A J6 x1 a; H. v
7
2 p5 d% z$ N. p% v
8
) ?) |. b( k w) `
9
! z: b' {' x! v# [/ z
10
. n, A' h$ `- d1 K4 ~# K# Y
11
- F! X4 Q% M* G4 b2 `$ P9 I5 P
12
& o: a# u- e4 C K* I
13
5 j' Y: p4 _+ r+ \3 f( ]
14
$ E1 Z& p" l( G7 w5 U. n* ]" [. Q
15
* {% Y3 R1 L+ y( u# R; h
16
3 \4 W! K5 c% f( X# {$ `
17
$ m' k* _! ~) h+ s- }
18
5 S" k. t8 U+ |
19
8 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' v
22
1 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- B
dataloaders
, I; X- M0 [- H
1
# Q( E7 z$ b2 B. ~( \9 E
2
( 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. Z
1
) ?4 h. i. A) X, n4 c& K: d
2
2 L" o3 M4 o! \2 Y1 R+ _8 c
dataset_sizes
6 h, K0 _( t7 r
1
* 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+ O
1 \- _, 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 X
1
7 a! d [( P* p& ]1 H% C
2
; R0 }# l- M# b
cat_to_name
9 K, Q1 p# h! L) m. |3 b
1
; {' 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! ]
2
8 W+ s1 [$ Z! k2 v. y! }. v5 `6 l
3
R2 d# a. f! l( `
4
/ y* w" d p0 P
5
: ]. N+ w+ ~ i- J
6
; }6 ~% N: o6 v1 l3 V( J
7
5 ]$ T6 \' `5 f
8
- 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' ^
12
0 W8 x, R9 u2 w- r
13
& q( J* ]; N+ R. G
14
2 r+ i% v" \9 ?5 z! w. e# F
15
2 k- k- {2 l' n& I! C0 C
16
' e' o6 ^- H4 h3 t. q8 M
17
: l9 X2 O; X# h$ @7 z
18
1 f4 m G& L8 Q
19
. C3 u% {/ \' j1 v1 _2 A3 X0 ^
20
% i- y/ i- ]7 d
21
5 N, N4 d, p; t I' l3 f, X
22
) H" j$ b% P8 a6 ^: w5 ^
23
. \0 M6 U; X' J/ k
24
5 ] C8 [( G4 [1 d ~
25
1 _# b' \3 ]8 F
26
) z- u1 }( }, o3 u, E( \
27
/ d1 T9 S* w4 P9 ^6 b7 A/ ]
28
9 z: X& _! n9 a
29
3 q. g; ]$ S5 r* e
30
! Z, w V+ s4 X6 ?+ d# {) _' }2 H
31
9 e% Z$ W: h- a$ r: C6 R. \
32
1 N3 K1 T0 I0 l
33
1 i0 A# B( g" s: p* t6 s! y! O
34
; X; R( ^5 G1 t1 g4 }. b
35
, @$ R; z# p" Q0 ]7 f1 B
36
0 H4 w! B! X$ r! A7 e& d
37
8 ?! 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- I
41
6 g/ A0 u5 N# m
42
- 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
46
0 g- K ^8 ^3 W& R
47
2 ]# ?) l. ~5 H7 d4 C$ T9 H% `
48
3 |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 t
52
' j8 \3 `4 R; b1 d0 O( I! m) W
53
; U$ r4 @0 T1 T' O$ p
54
$ H5 |' T6 t$ W- @
55
& h3 f( H4 }4 C5 w: _; l8 X
56
" 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 T
61
/ n, d# O4 Z6 P& l$ a
62
. i, T9 y! z: p0 B. J- u3 U% P2 _1 b
63
! b7 X% F9 R* h0 u: R' }
64
' G8 \0 Q7 X/ j0 A! l
65
8 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 a
68
+ I. l: F+ M* r& C) r& F' [* C
69
7 p: a2 d0 F) n( C. R
70
5 Z% Q6 ?& C$ M2 O" p( P
71
+ S Y( o* _( s+ e( a$ ^# B: B
72
, y) ]3 ]! }. p
73
" {( ^0 J3 w5 }2 A0 S+ C
74
& x; ]2 F. R- `& A
75
# z4 T( E& J N6 ^' t4 s' _
76
% Z# L" D+ `. D H' o
77
) d, g1 z; n1 j0 ^' ^1 f- Z
78
8 @8 d: R0 A) X% ]) {* T+ m$ Q
79
1 {+ r1 g8 L% g) K. C
80
0 y" ~/ ?- c/ S) e7 S3 }+ e5 n9 _. y; t
81
+ a8 P r/ U: M5 Y3 H: U& y
82
0 C2 C) C9 W5 Z
83
" O' C! M2 l E5 q" V
84
" G& g: N8 n9 X- C/ K
85
- [' B8 n2 j2 i2 H; A+ {
86
6 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 z
91
7 @3 } r! C! w4 c7 ]
92
; N1 z( r( L. _1 s* Z; l0 O
93
B6 X: }1 H% I6 b# h0 C9 V9 H7 I
94
$ n. ^, N0 b1 k' O" O2 i! b
95
2 e4 f3 i3 x1 ]; j7 r
96
8 q- b: h3 N5 ^1 }' M" _
97
9 _1 U1 }0 R: ]7 I7 ?$ `" V# |
98
" e+ k6 F8 Y3 M7 ?: _+ G% p* u0 V
99
1 S: z, `' L y
100
& T; R' {+ m p" x7 Z3 d; A
101
8 n; K; x! Y" y2 A R+ @* C5 H
102
$ D: m0 x2 F/ y, h u6 Q, Z
4.展示一下数据
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: t
6 Z" O( V6 q( H' D! K6 Q) m5 U
return image
! J A1 ]4 L- B1 p
1
! M6 i2 x( H' f3 e
2
4 V1 u% V" C2 \# D
3
5 M: z- O( v8 Q# L* B) F
4
' c8 o6 F+ X$ ?
5
+ x7 r# E0 H8 w; i
6
0 d, _" R. I3 U% @/ ?: c) h2 z6 e
7
. X: S; M: _. E
8
. k k# W( H' A) l) p3 q
9
, n' L8 k! m$ ^* Q
10
' Y; X! f+ Q: z" q4 |* E
11
! i# x' G2 J- g
12
# V5 P* F& v0 Y- {, l& j
13
5 F; f: N# [8 x
14
! 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: D
columns = 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( P
inputs, 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( X
plt.show()
) r P" T) @# [" c" p
: ?* P6 l$ Z: I8 s/ n x6 q2 B
1
) R" i" S* {4 F
2
. x1 {) {- B d, p8 E# p
3
6 L2 `% S. G- h; s) {
4
/ y3 D; M) n, \. ^
5
6 {" R0 _$ E& H+ w5 r: ^. ]9 l
6
7 @ Y- I5 M6 \8 I( L; s6 O
7
, {) Q5 u4 \; J
8
+ R3 X& C& v( D4 ], ]% p
9
% z: v* R) g8 E1 v% A$ Q5 n; J
10
9 d& T; R3 Q: l8 C9 Q( V
11
5 @ \% h% _# Y/ F$ F2 C9 F* C; n
12
6 ^2 U1 R/ K6 I2 i/ a
13
0 E6 ^1 P% n8 D7 \! r0 R' _
14
/ G. ^0 i- z0 } f- h+ t
15
2 s% x' W3 A5 q+ ~* ~$ W- C
16
1 z) ^. ^ n. y/ H
3 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" E
feature_extract = True
6 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 F
4
" 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" A
1
5 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 D
5
1 R6 g5 D# l: X9 l5 T% L6 u
6
, j5 x3 ]' p% R# R+ t
7
6 u: q! \' h0 j" O* G% T, h
8
* |5 t7 q' d0 d. S) t2 N
9
4 i0 O! \) ^- y( }( k3 j7 A
CUDA is not available. Training on CPU ...
; c9 U4 F5 K! z+ @, S
1
V$ @ D! S# l3 _( l, I9 F
# 将一些层定义为false,使其不自动更新
/ ^ g2 D" ? w& y# M; l% ~" p
def 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 i
3
( ]6 h/ @4 v( F. L" |4 y
4
2 r0 }5 m% N2 S* h- K/ k5 c
5
I( \% N6 b8 T) q4 i# d
# 打印模型架构告知是怎么一步一步去完成的
0 I" x/ ^) S5 X* e3 `% y& z3 n2 X
# 主要是为我们提取特征的
, J! [! h% L2 y9 _" r- D4 U/ R
3 j) }/ B% y' ?; X. d( y
model_ft = models.resnet152()
) M3 z' n* y8 ^ S
model_ft
9 g) g: V# r0 E7 `: B' {
1
9 {8 f7 |3 m* H$ q
2
* Q g6 N$ ]1 }: r
3
% R! H# E) p( K* l7 g$ @7 A
4
/ 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 ]
3
0 O6 |, _5 w" o4 M8 A
4
$ h& a8 m+ `3 E* ` @% b
5
# x, C6 o* J- t0 I5 \; E
6
1 R2 ~3 O/ c( J& m6 K; c0 \: Q
7
( s1 B5 _0 v7 x, h- Q% v3 r
8
2 X" T; c9 R# s
9
, k& i0 e7 w+ r( `
10
( J3 y+ W) \- B& m' M2 X( N: N
11
3 c! G, p M3 A# B- K, h' x; |8 h
12
" d8 {6 J7 G; Y, G
13
1 M9 e4 O5 ? j6 w: M5 S2 g7 v
14
: b9 f* n! j* _. C C5 S7 _
15
8 q5 d! z( g6 m/ T' W8 |
16
+ r% @* r9 G& O3 J3 C% e% s
17
# m& z' I) x& _: V2 N
18
" N) `+ {+ \3 f, a u
19
1 B/ W* {: d7 n4 n/ P0 _
20
1 i' I3 k/ d( F! D3 m9 t6 t: S
21
, U/ W$ V/ m. [2 }& `
22
* G+ W* V( P1 F( ?$ c
23
6 v# u) \2 |, u; A$ W
24
* O0 K2 A; x8 q( B
25
8 _/ m; P4 }1 w; ~+ J- U
26
2 Y& N1 Y9 y- R$ c+ H
27
0 M: [1 t! A- w% Y
28
6 R. H' ~3 O% F: ~# T
29
, W! E' Z9 o$ _. b2 g1 Q
30
+ o8 }7 {& G; ^7 ~& E8 V! n
31
3 A9 `( L: q+ n) Q, Y' B& l
32
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( E
6.初始化模型架构
" 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
Resnet152
7 |* 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_classes
4 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_features
5 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 a
2
+ a: k" z0 Y, e5 f3 b7 A
3
2 d ~& U' G4 F$ R! F$ u
4
) E3 x5 W6 n [9 Y+ j3 Z
5
1 B8 K% I8 q3 D- u' M4 ?( M
6
: J: U, `: H& D; L
7
. V( K" e, i$ O* h$ [: w
8
4 c: w+ G6 X$ x o
9
9 _2 c+ x& L" P
10
9 R9 {' H: `/ c
11
. n6 g) E; j+ ~
12
8 y+ ^; |6 w- M' t% L! q O
13
8 i* K0 D. u0 L4 x+ O& \
14
1 S) Q! {1 c9 j( s" N4 L0 c
15
2 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 Y
19
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
24
0 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 y
27
7 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/ l
32
$ r {# O1 u" V1 ]
33
& M J8 T$ C! k) y, Q4 r
34
\0 d4 B' D! C0 ^5 M x; O7 P" P9 t
35
- M1 y; D1 U x8 s7 h" j8 T2 \
36
1 g4 K% l2 ~; |* C+ F, t0 z$ H
37
( S) m$ i' H- t( I
38
6 f( ^1 p# V: Z: I' n
39
7 t* {7 L. `4 q" g2 C: Y! j6 e. a
40
5 e$ |: \% e: S/ y% B8 y
41
) k. F6 S6 I7 t% F8 z3 c3 S
42
- R5 o2 }+ O P+ S% r
43
' l6 K: d7 d* S% k, R
44
: L; Q4 V6 z* V. _+ W( l8 ]! C A W
45
9 ^1 j3 x% [# q+ H- E9 y
46
8 Y7 ]$ ?0 i3 U! j
47
; F2 w* W: v3 g$ }: F4 X
48
& 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 Q
52
2 O8 A7 F$ S& B( w9 s/ a
53
: y4 V" k9 e( I4 z5 }
54
8 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 W
57
& i3 P( f" c9 {8 e; B1 V+ h' O1 F5 t
58
; U) T0 b- R2 X6 v# { j0 \
59
' w& j" r/ ^, s$ S. G! h
60
) n' L- o: o' i. L: h9 ?8 m
61
+ h' G2 P0 h; M9 b3 \6 S- Z
62
& c! g. c% c1 Q0 B
63
0 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 W
68
; O5 q5 n8 q' ?" ^+ A" }3 M6 B+ L, K# [
69
+ e* r/ j( I; O( {; ?' U; ^- q
70
3 T0 p9 U8 e7 u4 @
71
) r# n+ ]2 ]: G" E# _+ s- {4 h
72
+ b+ l( @$ j9 `, x. X
73
/ N% S1 a4 T5 |# [- x
74
+ A7 z: s+ T6 U
75
/ {) K) W2 ]& k5 Z/ K
76
6 Y& V+ }; M/ [( G
77
. w& w0 v/ n0 `5 [
78
$ v4 N2 w2 D: i4 P5 E
79
) o7 d4 c( X6 v7 W, B1 K5 ~' H
80
% d1 n( D4 H! |' w0 r
81
2 n' m3 x/ ?- r: j& e( ^ i5 z
82
2 O8 C: D% o& J. d3 r4 N2 z; e5 U' @
83
8 ?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 _- P
7 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 d
1
# m; ]. |8 g% u
2
. {' M& W+ Q! G' I6 Z$ N+ z
3
S1 u) Q) x# Q. D+ G2 @# d
4
% 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 [
8
3 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 A
11
/ E" w8 \7 a4 Q9 a9 P
12
# r% W6 T) Y$ R' l! [; q
13
% ?8 J- r# ~# V0 ~5 w
14
5 l" b1 T( H( E$ {- J
15
' ?. {! T6 B+ _/ i5 _% q, |
16
; d9 X4 k3 f3 ]" |$ m' F4 z& l9 ?" E
17
7 b+ t! r) s4 m0 [; X
18
$ p$ j; ]7 b. {, G* N
19
" 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; Q
Params 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' p
1
* V# A1 E5 q1 |4 [: W2 h9 r
2
! ~! y9 j1 V' y& ^6 h8 w
3
3 c6 |% x6 W `2 \" O! @5 K
7. 训练与预测
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& E
1
/ U3 Y5 _+ d! j; z5 ?
2
/ I l+ F }; W; |2 m6 N/ V
3
* f- T1 T# Y) m' b& y- K8 Z
4
* j% L) w( ]7 t B/ `
5
. j5 e# q( L- [
6
1 y! O+ _3 V& m/ [/ n, p ]& P
7
[* e; x& |. x# K ?# Z
8
+ B. t' X+ \) w/ H! T3 v
# 定义训练函数
- _ s+ X# k. v& v
#is_inception:要不要用其他的网络
. X' c3 P% j7 n( W8 o5 P
def 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 = 0
7 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 B
2 \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' c
6 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_acc
2 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+ B
2
0 U* U; b2 ?& E3 y8 _2 W
3
3 ]) N+ m: M; H6 U3 j9 E" D0 I% d# T
4
5 A- X' [2 T+ C9 O
5
, ]3 [% j9 k$ w- R. X
6
- x" c4 T" ^" _( {
7
f" ^+ F% R, }& r& f
8
$ r! z# H9 F0 P5 Q+ l' E
9
" [, p# J5 e8 g# j
10
4 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' R
15
! 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 [
20
0 j& g6 g& p1 [6 G s& G
21
9 c4 p! S+ ]2 t/ R$ d
22
4 B3 _! h6 R9 f9 I
23
# F# ~" O5 G8 B0 h& v6 ]% [* B
24
$ @$ Q4 J5 H! A
25
4 ^. ?' m$ G% L1 N: b4 ^6 Z& |$ U
26
5 ?6 Y) G9 S1 w( W' u8 F
27
7 `9 G7 n& ~! o( G2 l- _5 ^
28
6 M" D4 ~7 }) G) p! l
29
, _- q9 m1 |+ `& h7 C3 Q
30
, 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 D
38
* W% s7 n- }/ D2 F
39
9 Q9 j; N) F3 \9 J) n, U
40
+ m, l4 `9 k5 L
41
6 D1 X: m2 r" ^4 n) _
42
* X7 z7 n: _ f6 w4 U
43
7 r# s W- j! B) K; w2 Y
44
" l$ o$ I( C7 h
45
4 P9 a7 \+ n2 r: b' ]4 u u
46
0 r# ^8 j- y# b$ L2 `
47
. P/ o( N# N4 V
48
; T9 B* {! B6 H
49
0 E* U, ?; T4 h, A. A5 q5 _ R M
50
4 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* t
53
' v8 h6 u5 d: u+ \" |+ G+ H5 k# u
54
' a7 v5 v/ i8 U* @ Q! K. e; M& @8 u
55
' a( [6 H8 V7 [' H K% Q
56
5 M8 Y c4 h7 r( t
57
" Q" g" C% n) e% [' w, A; B/ j# B
58
& 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# J
62
' f; n' a/ F: W4 W
63
5 q u/ S9 [* c
64
X4 v7 ~! ?! ^
65
" N( D2 Z( B2 h* ?2 e
66
7 y- q4 [" U# Q! w/ N
67
3 a3 g1 D+ N" ?
68
( U) d% B; E% ^% `: R
69
/ E; g7 g4 t! ?* ~5 V3 o) ~& L
70
, v* T: \5 B% x
71
# ^% {8 G O% J. ]
72
8 a+ s+ }/ Q$ `! I$ n
73
; Y1 g( @# |+ b z. y
74
5 @( G* t9 A: [# |
75
8 q( Q8 ~9 z+ `6 {3 l
76
' Q5 B% E j0 p' O& g
77
; A: a" k9 e/ Z" G: V' H
78
4 ?) k+ W* r4 M% Z& R: T
79
- `( @1 J: c& U3 c0 g" K0 v
80
* c* g- T3 h5 H
81
% R) S8 A' D) [+ m
82
# m% W2 r8 C5 J7 L7 f) H
83
$ B* F/ U: K( x
84
$ L6 z: R3 O6 Q8 R
85
( X8 T. S$ m8 Q
86
- m% O/ [2 o) Q& n+ W& c W% }( j
87
5 D% E; [3 r7 X* f. y0 C0 Z. {
88
- ?# K, Z" h/ G3 G
89
) O' X( Y |/ K5 z( k
90
4 b! X% L. B; C/ g2 {: [
91
: u/ E' ~' o& h8 G" d
92
; x# b6 p+ q2 y+ V1 F0 q
93
% h) L0 J9 W7 s& N1 D! h
94
, V2 T3 Y' n( o
95
& t/ D; A* K: n' s1 ^& f
96
1 r( U) k1 m. |; j" D2 @% `
97
7 i( Q$ M$ Q9 { z6 \% \* a& g
98
5 h) v, b4 P4 D
99
; l9 r) v1 Y6 _/ d
100
. | d3 T# {2 V5 e
101
/ z5 N( A, I( r! C6 u- F/ _8 v
102
/ W& L0 K. N4 p$ C8 s3 Y
103
+ b0 D5 v& X: F( i- z
104
9 B; n( ?% x* t, R+ f
105
. 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
109
5 A3 N& P' j1 e* R
110
3 W. r6 G" Z1 o8 R
111
* H( A+ L6 T% c. L2 k, `
112
' m7 k, q' Y; K
7.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 B
3
' l, c. b* G4 G5 K. ~+ e. m7 E
4
8 r0 v2 [+ q$ w
Epoch 0/4
- ^$ Z# B6 X' u6 K/ a
----------
( R1 |& W1 v U
Time elapsed 29m 41s
9 E* w4 w6 N, U' ~
train Loss: 10.4774 Acc: 0.3147
" ~8 [' t& ]4 L
Time elapsed 32m 54s
0 n) F% v- j2 E- J5 f3 @
valid Loss: 8.2902 Acc: 0.4719
, M' l1 |( ~7 a
Optimizer 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 j
Time elapsed 60m 11s
, \7 H' C9 p9 L& z c2 w+ n+ Q+ V
train Loss: 2.3126 Acc: 0.7053
7 n, C& R& N; j; E0 Z7 m9 M
Time 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 I
Epoch 2/4
3 C# ?/ y4 |5 _+ H5 |2 l1 B
----------
0 C! `0 q6 s% ?6 z% ?
Time elapsed 90m 58s
7 C" w. i" V6 L7 b' s: f% b& R. s1 A
train Loss: 9.9720 Acc: 0.4734
, X+ Q" }. {: R8 y
Time elapsed 94m 4s
! ?/ K5 {7 l' H! w5 l+ V
valid Loss: 14.0426 Acc: 0.4413
9 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$ I
train Loss: 5.4290 Acc: 0.6548
8 F& I3 {. F0 C- M& k/ p5 b
Time elapsed 138m 49s
7 u+ F, G8 y! b: d. N+ A
valid Loss: 6.4208 Acc: 0.6027
0 z6 k+ l0 ^. h1 L* x! h4 S
Optimizer learning rate : 0.0100000
7 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: f
Time elapsed 195m 56s
+ R6 `* X# e# n6 N# h
train Loss: 8.8911 Acc: 0.5519
8 [% ^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 }' h
Optimizer learning rate : 0.0010000
) f1 X. C1 R( e. y( {
& D+ x& @! U8 v& i5 @
Training complete in 199m 16s
: [( U. A7 d* \, H; p
Best 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- i
4
. P( B( I0 O' g& o$ I
5
$ g; Z. }" w8 z- [- ^
6
) S- [3 r- c8 w7 ]+ y2 Z9 m, ^9 v
7
: p& J: S$ b. a% z7 D8 b9 Q% X' P: r
8
( w; E1 Z" W1 m. d4 z; P s, t S* M
9
; q- E I- r! _5 }) q* G o; a
10
" f0 M! p1 ^# ^4 w) P& L, R( K, B2 n/ z
11
1 W2 ? U* v9 l: f4 @7 ]/ C
12
X% K: k1 ?1 M! y
13
; |. U* \, T: X! g' j. G
14
1 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 u
18
5 F& S( A9 n* J
19
- k1 ]& k, {: k8 k+ N
20
, [8 K4 X/ h2 B( B
21
& q( q2 }0 Y" ^" R# W: ]
22
; E7 |: r* u% [6 C
23
- N. B5 }3 Q0 g( p/ ]5 S
24
! r( W1 ~7 g$ u I& J
25
& u/ X# r* i; ~7 x3 d& n
26
5 P9 C: C3 m- Y8 c- c8 w+ Z& Q
27
; D# J$ C2 ~, o7 E7 K0 m3 F x* t
28
2 g8 x9 M4 `3 m' D# o$ w
29
) R: H: o9 g6 R5 `; z$ `3 s5 Z! j
30
/ t' y5 }7 c# _( M
31
/ w+ _( N5 x% |: I7 e) {
32
0 f; Y: D" e2 v
33
3 @* m: Y% [& b8 N0 W
34
) N4 X. n* T) P: k2 S
35
+ ?0 U' T5 S5 ], M' f6 ]' Z. ^1 W: `
36
1 {- }- M( }. J) w2 M' y3 p
37
, T4 s3 P" y" D4 X* f% \. Q7 {- S
38
: h* p: A4 X2 @
39
7 w4 i: ~' h5 A6 _ C
40
2 o. G: r" G* `0 w( e6 ^
41
9 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 = True
1 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 a
2
( 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 L
6
. _5 I, g; ]! p1 b* ^* ?6 Z
7
8 R+ @; j1 Z$ E9 f3 f9 o
8
! }! X, T( A: x8 O3 |3 N
9
% N, k b) n+ k2 D% m
10
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 b
checkpoint = torch.load(filename)
) v+ T! Z* p9 b7 I6 `" u; {& a. p
best_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 z
1
5 s' T3 Z# c, f5 l! b% H
2
& 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% Z
6
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& q
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"))
! I: V; V0 A2 \9 q( S. R
1
) H5 Y& M+ \* e5 p; v9 q
Epoch 0/1
& w8 A- S3 D- d; X
----------
% m! y; o& Q7 R x
Time elapsed 35m 22s
7 R/ c1 p8 U( M
train Loss: 1.7636 Acc: 0.7346
$ O& z' a) [2 M" z! u: z
Time elapsed 38m 42s
3 ? q- M. O1 J/ G# H/ P
valid Loss: 3.6377 Acc: 0.6455
, b* n. l: `; S3 ~% t) G+ I
Optimizer learning rate : 0.0010000
" @; z- e" D8 z. [
5 l* E* C* O5 y3 u+ K6 J- y% Y
Epoch 1/1
% p2 l6 \- j; ?+ {' N- i
----------
- z1 S' A* v" @* g
Time elapsed 82m 59s
3 y$ f9 ^3 @+ A9 z' E, T$ Q$ w
train Loss: 1.7543 Acc: 0.7340
0 l% {' u" E J2 g) A
Time elapsed 86m 11s
$ z. r* s7 s4 b! B& w
valid Loss: 3.8275 Acc: 0.6137
; X7 P9 j( }8 N3 h7 X
Optimizer learning rate : 0.0010000
8 H' E7 v, V5 ^
q- m+ m) \4 t4 G+ X& P; i
Training 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/ j
5
& L1 \. @5 I6 \' Q. P: a
6
5 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 s
10
- \ a/ ^% H% ~3 E
11
/ d3 N7 e* Y6 L; |7 P$ a
12
+ 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
16
5 }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 J
filename='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* h
1
: j. l3 h; q( ~
2
% |7 E5 c% y# J$ B' `
3
, ~. z. i5 A1 C9 j$ p
4
2 ?" {* ~3 o& p( |' L& O5 P: g
5
9 D; M1 [! u0 i1 i' X3 S8 V+ I
6
! J+ \1 r. l% T! L/ _
7
7 \) `; W+ X4 K/ x/ i3 a. r: c
8
4 u6 {' w6 u3 ^' C$ b" i: M7 N" s
9
' i# C* D4 }/ Q6 A2 |0 p# m m4 M
10
7 b" _6 W, A; O0 Y! f, |8 h8 z
11
+ \5 V6 T5 l- J+ P% v
12
- 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
# 而且对象调用方法会直接改变其大小,返回None
1 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 s
3 @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) / std
9 ^( 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# f
img = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理
o( m4 w6 E M5 S0 o
imshow(img)
6 _! U4 R: ]5 G* N3 \% S
' K, ?" t: Y# z1 e; l% o& y& J
1
4 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* J
4
7 ?; F: t3 `+ h( J
5
9 k. ^3 v0 w& v8 R4 z5 @. S
6
4 A+ x/ ^ _2 J5 J. \1 s
7
$ h, Q3 c3 R" d: A: K) e
8
5 {$ m' Q. e1 v( l2 }* L: Z0 ^
9
3 `# g `& S! K/ e
10
# @8 H8 G! O9 e. l9 z: a* B
11
8 K2 c3 M+ ^, h$ L
12
9 k+ U* E, v- D4 Q! I2 F
13
0 b; Y) u( g; V; m
14
$ p5 k+ b5 }+ V6 k5 |
15
, T& I% {) b$ X
16
; q, T2 a$ F8 u3 a6 I! F
17
9 |) D; b7 f. j% ?: I9 T9 \
18
1 E6 q# X. U. b7 ~3 ?
19
# f6 [) `- x5 k6 \4 M2 l1 ]& `* O# t
20
9 T6 @# }+ M, q) F( i) q7 \; g! z
21
a& r5 ^, p1 v7 O, f6 U
22
2 u2 p: Q3 Q, e% s
23
9 S4 v7 A9 x5 g2 m9 Q3 F) Y
24
& X( D: L+ Y6 p _) `2 u
25
6 U' z( y2 y7 k( |% v$ I
26
8 C% }; v+ {: P0 V
27
) K* ^- h1 M# y8 @5 }
28
7 g8 \8 t# [) x( z
29
) p7 q' B& O: y6 N
30
7 m7 d3 d0 O: x. f
31
5 {) O0 ~: y5 }' g4 v: @% h
32
5 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 x
35
/ e8 j/ ^# _! j- r- A8 _, e6 d
36
& l" h8 ?1 ~& S( v3 K
37
- w/ p" f6 l) C
38
" q9 d/ B- `+ X6 [# C
39
+ Q8 y! Q+ C% Z+ {
40
$ [5 i* L6 i% Z0 z) O
41
1 b, N, B+ c2 u) g' ]7 C
42
3 ~2 V" ^1 S! a0 I$ m: n- |! h+ z
43
[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- b
49
; o4 l8 r6 ^4 }8 U" y- g
50
5 j2 `& N8 o) @ Q# f2 [
51
% w2 V, k; J# Y- P4 _
52
1 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 [
1
6 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 K
9. 推理
: 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" h
dataiter = 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 G
1
8 T3 _9 E7 `8 U' O
2
# Q# m" K+ b% ]" f
3
. 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 K
8
, ?7 S7 q7 B# B- e B
9
$ q6 o% _: M: Y
10
9 J4 u! }) n9 A/ |: x* w, h0 G
11
$ f- f# y& v: w- d8 m- r l2 u" P
12
& s- _2 Y7 N( C# K8 r$ y; A
13
2 U( y9 w0 ?/ |
14
" r& a8 P2 }: w- s: e2 g
15
! l8 E8 l) \+ H8 @$ E; E# `1 t; M& I
16
2 b1 G5 H4 O) i
torch.Size([8, 102])
* b. l$ N& S# U9 ~% w; B
1
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. h
2
7 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 g
1
& C# V0 J: [0 q; e4 i
2
/ g" G3 S/ J1 h4 X3 u2 i
3
u$ @4 C$ u' z1 f) t
4
* ^. {( ^- {( r
5
. C6 v! b* I. ~% |9 O' O( B
6
X8 _9 @' T$ e& |/ g# w
7
2 P! O) L5 e/ Q$ v# l
8
" G. ?& l+ s7 O4 R& b. G! G
9
1 n( J) p2 \1 p8 p4 h
10
+ \% O# h/ ]; J0 k' Q9 _; m
11
9 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. S
6 s" a4 K$ e. \( I- q: g
欢迎光临 数学建模社区-数学中国 (http://www.madio.net/)
Powered by Discuz! X2.5