数学建模社区-数学中国
标题:
【深度学习】 图像识别实战 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' g
1. 导入工具包
' l( e+ n: J5 q+ B$ Q7 L% |
2. 数据预处理与操作
e) W( U) A$ y4 s* z
3. 制作好数据源
7 Z+ r) q# b7 M, i: {6 |
读取标签对应的实际名字
' b; K+ o! t* N: F+ H
4.展示一下数据
' f1 ]6 ~' ^& H- O: v) D- r
5. 加载models提供的模型,并直接用训练好的权重做初始化参数
& A( c2 l! s, q, `% k. O
6.初始化模型架构
1 d+ m2 V$ E" P% R# K6 c
7. 设置需要训练的参数
9 I7 l4 B+ ]% v3 d/ i# K
7. 训练与预测
" w" P5 y4 ^ @3 S
7.1 优化器设置
" `$ d0 {9 u: f- U/ j
7.2 开始训练模型
" o+ |. E9 ~% A; t
7.3 训练所有层
; ?5 u/ c- U5 j4 \1 S! A
开始训练
, h. D- o: X# c. o
8. 加载已经训练的模型
! 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( U
train
5 o5 z7 v+ W& A
- E4 v! w6 P3 @4 i8 c
1(类别)
/ l J8 ^6 d# T7 c/ ~5 w0 v3 k2 L+ u
2
9 O) P' Z. k `% X# K) x/ y4 s# h
xxx.png / xxx.jpg
1 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 J
0 \* 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-dataset
1 `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 os
1 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- \# J
import numpy as np
6 y: \ D3 M( w; }( f0 x, Q
import torch
9 v# ~( F$ r/ c, h
from 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! L
import torchvision
! X4 T& m5 j/ a1 W
from torchvision import transforms, models, datasets
9 W0 s" L E s9 s2 x
, _" R3 V7 h% q, t1 h; E: ^' }5 g
import imageio
2 l& O9 U8 W3 ~& A/ M
import 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 G
import 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 ?% B
1
$ M; H/ K3 e( {
2
0 ]2 P/ }5 L& I |! N
3
o2 `: U9 B; @9 Z$ E
4
9 `3 p t. M |9 {
5
4 M6 _) C" j8 U% x5 \
6
& H6 \2 [6 I7 O5 F
7
2 B; {& ]& _, U0 w* p" P
8
5 o, f6 x, i$ @. R; |$ g7 L
9
/ \7 _ k& s4 q' T" i/ j
10
6 k! X/ B# H8 @" V6 y' F
11
# b: q1 _7 c5 ?9 O- E7 p$ B5 h
12
- |( g6 H7 h' ?4 o5 A
13
# u% M7 [: x& N. L
14
$ j3 n4 ?; G3 v3 W, Y
15
( B. p7 O& E& a0 q3 f
16
5 Y' e6 [( s9 _7 {3 Z9 n$ Q
17
& x3 J/ E, x- p0 G4 {* h
18
: p) C9 R" @0 D; }4 J
19
3 q2 q7 B5 \9 {; f1 [) C' y
20
# z+ G" c6 N. m3 e; J& m4 X
21
( b w6 [, T/ @# c- h
2. 数据预处理与操作
/ 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# t
train_dir = data_dir + '/train'
@& c+ D" O/ n% ]! J
valid_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- o
ImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片
, 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 \
2
0 r- Z% |# Y( [, X0 z
3
6 `5 A8 @; y) l/ a" V- [
4
6 {) P6 N7 A9 H2 X) _
5
7 v: ?1 Q# w% h: o% a
6
1 f" _: |2 G# e8 d9 B9 K
7
3 @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 ~
13
6 I& V- [0 g% \
14
7 D d2 b2 M3 t7 P2 ]. S P
15
; 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 Q
20
3 ]5 c. j- q! @2 a
21
% Q8 d5 c& ^/ Y- ~9 f+ }& l
batch_size = 8
9 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, k
dataset_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& ?+ s
1
6 Q0 _8 ?, z4 J& G
2
0 \" K! W- M% {1 g
3
0 ~) L" A1 b; V& q/ s
4
# }: w$ D' n1 S9 J
5
& r$ Q$ H! h+ t! ]
6
2 m8 B P7 f' ?2 O
7
/ D3 ?/ Q% I3 A6 c$ U
8
' o% |' V" m% A" k" z
9
* J- I& T! L: X% [* I& m
{'train': Dataset ImageFolder
2 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
1
1 f1 K3 i- d- L9 I7 R* ]
2
& C5 D5 k6 k4 ~7 E+ x
3
! S9 P, Q: ~/ t% _
4
+ ]' K% g+ A0 v8 e& _. q
5
. W: A4 s3 C' ~# v* N% E
6
+ S! o# J& Z! D% Q! z, _& X
7
6 S, g* ?/ l* U! C# F5 F
8
$ 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 y
12
6 j& x# p& j3 \0 m7 V. f/ H
13
/ u4 C$ ]/ N5 N% d8 h2 i
14
7 F) e( E; ^: v4 ?) o# G; J3 ?
15
; d' |/ ?7 p3 V8 R" v
16
! 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' |
19
8 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% H
23
& W- \/ k5 L# T( W6 ^* p1 F8 m
24
3 e) Y; m& i4 w+ @1 a8 M2 t8 K' c% h
# 验证一下数据是否已经被处理完毕
, P# g) m9 k: M# x( V# h, F9 q
dataloaders
5 N/ L A4 i9 H- [. F9 r
1
* M% G7 S" z; v p' b% X+ u$ S
2
1 }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% d
dataset_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! t
1
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; `. k
with 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
2
3 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! K
1
1 o0 q5 J) W( ^
2
$ y+ C& d2 s. j; q4 U5 v" k
3
" R6 V$ R3 |# Q8 _
4
5 I. a$ l0 b Z7 Z3 ^" T3 C" L+ a
5
+ 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( f
8
5 K; f d9 K5 `. ]- I' \: Z
9
( ?% h- Z8 h4 Y
10
0 m$ z" |% G% C2 B
11
) l$ p/ \9 P5 e
12
" V: I: Y3 ]* H% g6 j
13
' C$ n) T3 F6 y. B
14
$ }2 ]: A$ S/ J
15
" c- ]; f* J: K& t! ~
16
9 S* K1 N- |( r* ]. Y( U
17
/ ~6 T- W1 f% [
18
! J* l: J6 J+ o, A+ U
19
* F# I% H0 @: K- o
20
, ]7 w$ d+ H7 ]/ L( p0 n( H7 N
21
6 _" |* V! g3 ^- S; |% C9 A
22
8 u3 r9 r' M7 n* ?% P
23
5 j' B% p: Y0 L7 J- q6 w' R; H; B- _
24
- L# _* R) Y; }; Z9 z3 M, ]2 Q
25
4 p* T4 x! R" N0 U1 [5 R, T
26
$ ^* U- T8 D) G c" C4 M+ u# s! n
27
) S8 w4 C+ p. b9 `
28
$ a3 D8 c* z8 p9 y b+ @7 V5 Q
29
2 L) [/ i( g6 M& J
30
: P2 R' q4 {* o7 _. ^5 X* R; W: B, [
31
U E6 q+ }+ }8 G) m) L
32
: d7 b4 ~. E- c2 r0 R4 J
33
4 R& Z" o% i- |, C) [, C
34
) O( i- B- Y0 I" J" k$ n( u X
35
5 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$ T
38
$ ^: \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. ]" H
42
9 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* u
46
: d* H$ z0 }* H- W
47
5 }' Q0 P; |& c* t
48
S2 u2 L& j+ G3 p
49
; d, |* w( B% s v/ T3 U5 J. S
50
: P' B0 m4 [7 K
51
9 z- U* E6 B }7 X
52
0 e* e( a# c5 M7 t# h7 r0 t
53
! I( Y1 |" x6 p4 V& B2 H' Z$ |
54
' a+ ]1 v7 [" h" [! H
55
, O/ H; _7 j: Q9 C$ O* a% K: a
56
( b* j8 w; N9 l+ T4 z
57
4 r W$ i" p- {$ t
58
+ c" P9 i3 _2 J- i+ ]) x
59
- S; j0 {7 h! z. p6 a
60
: ?: 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) N
64
5 ~/ v4 v# P% Q- z' \$ u
65
% 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$ q
70
' W* t5 ?* z5 f5 i( X
71
, X$ ?- O/ t) p- e+ U
72
3 d0 n5 z. a# D- ]# C& o t
73
y" q% Z6 h% a% ~( i
74
# ^( ?" V7 p# C# y; I5 P
75
, @: H! q& z0 ?2 |- Q& ~
76
' {$ X% ^* M# \
77
) G/ l u. x6 s- l
78
5 ^5 x# l) |% `, |2 J
79
. i: Q5 H! D. c+ L0 ~3 ` R
80
9 b) ]* @3 M1 D, y; U; }
81
8 `, v0 M/ `* U* {: z( U( A
82
" |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- e
87
, A L# e" O: z7 R4 h/ f- J+ _# Y
88
5 _3 Z6 s. n0 k+ q
89
3 n. D# m2 `3 k. x/ M
90
+ P7 n6 p2 {8 {" N0 e9 z/ k
91
. w) E: n% z% g$ y
92
" D" ]" b8 \& C1 a, E- D
93
* r: C: m9 j4 M
94
9 N, {% h! T4 Y# ~: K
95
( M1 x. x' ?; V$ x D, J$ Y( \/ V
96
, R4 a% c, \; d6 W; V: ^6 V
97
0 n8 u5 B+ S, K/ N4 D' o
98
; h9 p$ S w$ I: S
99
+ n4 V; f( `- p
100
7 _2 ?5 n7 k% F& }
101
% J9 E- ^) v. b; Y9 V
102
8 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/ R
4
1 R3 A9 \+ g5 c/ n
5
l1 M& ~4 q0 e0 A( ] D$ X$ e
6
7 ?& A6 W) r7 U, K! `
7
P7 e& b( [" Y$ V
8
- U. @6 n1 h4 z
9
% Y+ X% Y) \, Q( Z0 X
10
' f) B7 H* b4 W1 m* ^% o$ O% M
11
2 Q, i! l! V! p
12
9 b0 _/ ?0 J7 @6 T+ n# D
13
) [ r; W5 g/ m. y
14
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% y
columns = 4
6 \- 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 W
dataiter = iter(dataloaders['valid'])
; |* u( D1 i& k: D0 e7 y
inputs, 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 p
1
/ X! G% I. H% a2 {" O- Q, X2 K. }, _
2
7 d3 w& [: T5 M
3
1 I- A) Q& \& D/ V
4
: i/ P+ X+ a, V, D! h
5
$ O$ K8 w3 _) \% Z, R% b0 `5 {
6
1 j% E) b: n% G
7
6 ~5 Z; f& V- U+ Y: U
8
. o) @9 ^' X) w
9
, g/ T0 S" c; b0 K4 k. |3 i' C
10
- W! \! z1 M' p: Y. L
11
2 z4 K& h, ~/ A) N( D
12
+ z* w5 q% q0 x S6 B: s4 j, w
13
& R3 u- { w3 s1 M- a9 p+ k6 F1 |% J
14
) 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$ O
model_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 N
feature_extract = True
" v. g: X! u2 e5 ]# J/ B" p5 d! G1 Z
1
: O: m% t( W/ x1 \
2
# d8 F5 o% e1 i
3
5 H4 M: v0 b! k4 j( j) ?3 ~
4
3 ~0 {2 K; i: k# n8 h: u
# 是否用GPU进行训练
3 k7 k* q" l y
train_on_gpu = torch.cuda.is_available()
0 t- }1 u; y; m; P. Z
5 d8 b: v" _- ]. a+ W5 l5 u
if not train_on_gpu:
0 h3 x7 ^4 z& ?4 E( k
print('CUDA is not available. Training on CPU ...')
/ k% z+ {$ Q4 a: P
else:
" 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& C
device = 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 c
4
% A* F' F2 S, g7 O. o) D) B# B
5
$ h/ U" y: c! P$ G& s. m; U
6
- I/ I9 ?% a9 u: E
7
+ `7 f; [4 j+ p: Z' `! n" } b
8
6 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+ ?
2
9 [. 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. l
model_ft = models.resnet152()
; w% |! ]5 h6 U- \- U% D
model_ft
1 C: T. Z/ W9 ^* s: N5 ^
1
( ?( R; L) O- h: q4 u/ Z
2
- R# t- T/ |1 O$ E
3
7 @9 H6 \ S( b$ K- U
4
: f6 g1 ^1 @6 b. L
5
; [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 }( T
5
! R; B5 u4 G1 `. {. j% c
6
9 `3 X- |* G8 ]* C
7
9 O9 I5 X$ }9 q! ?; G
8
6 E# u# O1 v- e" P# c# J5 H6 S7 H
9
5 _+ s4 i6 ]) H+ o$ p
10
8 A3 r7 }$ Q0 g$ _9 T. n
11
0 k) z- B0 _! q% ?
12
4 F8 l+ a" k# ]7 j/ K& ^8 _
13
3 L/ l1 q; I1 V/ |2 b3 M- p
14
9 T. _3 x# N& d" l1 r9 Y% t# A
15
( 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" U
19
. s* K6 R" E2 ]
20
3 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: l
25
, 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; h
28
0 H; o, E1 ^* D; q& l. e
29
0 h4 x* {9 N5 W7 m% c
30
5 u2 `6 P: h1 V( V
31
- s4 c! p( `& }- C6 t4 t
32
% 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_features
6 @: 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 = 224
9 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 V3
4 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 l
6 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 T
1
6 I9 t2 R! B3 P9 ]
2
' I1 L+ h; A* Y( d2 ?
3
, V* F4 P! H# H2 O
4
2 N, D+ Q4 _( i
5
+ ]: m2 ?1 t5 V9 W7 ?' u
6
/ Q* \, c: B) R5 Y9 q. v
7
R5 H, M7 g: n1 r& |& V5 p
8
' w. \$ ~! g5 e2 ~, H5 M: A0 ~2 r5 ~
9
P& u( {9 D7 P3 |$ y
10
! ~9 `: m! _1 ^: p& G9 Z
11
: b- j" O, `9 F" @' d# R
12
6 t3 o3 n5 ~; F2 O
13
. n7 M+ p, P$ d$ K
14
! v7 H% S( m7 o0 D& a/ L- j
15
: j, ]8 L2 [, N
16
4 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 J
20
n5 P4 L/ u: C- G+ l
21
! Y, d4 L" ]( M W' K
22
4 g; N+ N, ^# F# N1 K
23
6 F: ]9 i: u5 A) A0 t
24
( `3 _: A e T6 O( v
25
$ U3 T/ P8 i9 n8 \1 S
26
) a' `+ W* P# j- {# [. h# u
27
6 `$ E6 y! T" b0 }/ g$ C8 |
28
# J5 E& L- ?9 L4 i! m
29
0 W& \; _: X6 T; n6 ^: |8 g
30
/ Y( E2 }0 {8 T6 b1 Q1 l! v
31
! [2 m, Y! \0 `
32
0 s( n: z: X# F
33
& H) [+ q" _1 B" Y
34
9 l2 l2 v2 g4 D$ ^
35
& J8 s/ a8 K/ C( A E7 K
36
. u; p' M% y$ S1 W
37
u' A: _# N; c5 T7 }2 Z2 ^
38
$ Q) @- Z4 d" {7 F6 i' {
39
" l- j. m4 X# ?% g
40
5 H# @3 T. ]3 ^& ]8 L
41
1 w/ P% b3 B# j
42
( T$ }+ s& b4 g+ w) H
43
/ o2 f2 Q" P7 b) h5 R* }
44
7 f+ N5 |8 I7 A% P; K* J G
45
& B) T1 ]2 j7 f1 o, S8 C
46
6 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) f
49
/ P5 P/ _) a* c% x
50
# t+ Z2 v( |4 j, a5 T: ^
51
& X" C8 B1 U- D. b: l
52
4 m3 _9 x8 G! O; G, S [" B
53
( {$ D0 ~7 w. j, W: j% \: a6 L2 A
54
9 m! X: T/ j- T& `) n" K, K
55
: \9 \% Q- P/ B. ^' U; m
56
: C: A% F0 v6 S! g9 K3 i; p; G% M
57
( a# R+ A. |/ O- D
58
, j- j& \8 c# U; |/ C
59
, N: o8 x' W7 o
60
6 ^9 m1 q; |2 J. _
61
. \! l; @: {- a# N1 k
62
G8 K2 |9 P: O; n( z& V2 @, D
63
) ~. A4 {7 x( v! u! t
64
z1 }: ?2 K3 w# P! `6 p
65
) q2 Z" g$ g ~$ z/ p& v
66
, \. y* J/ b8 P7 Q) ]8 W
67
6 \' ?$ O3 ?5 j$ E" k
68
7 V. d& L! a+ @ ~3 b. Y9 a
69
$ G& l7 w/ s% q2 ~
70
& u8 |3 I% h6 n! g' x
71
: h8 v v0 L7 |' ~8 s
72
6 p7 a4 s8 q4 A+ k5 g: F- v: a
73
: k1 r1 b; _2 I
74
8 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( A
79
! |! S9 _. _! E! [6 e
80
9 x. m" B9 i3 ?- M/ W. k; \6 e1 R2 m) R
81
2 M/ I4 J; k c% e, A" j
82
# W1 a$ T5 ~# a5 ]2 [1 d& w
83
0 e: Q4 G4 r. g: x7 C% q. t" N) k. g
7. 设置需要训练的参数
% 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 z
model_ft = model_ft.to(device)
8 N( ~+ d3 r* d! h
& F( @( m. n2 {
# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取
- Q$ Z6 G, J4 k
filename = '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 m
if 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 a
2
: V& r4 w7 ^1 z. x K: ?* |
3
, J* m$ Q7 E# q% J
4
9 @5 A" w3 J* o, g) X
5
5 p- h2 r9 U1 J4 E9 N
6
2 l5 W9 p8 H; l1 }
7
& O5 c# I8 y9 o8 [
8
2 b3 U4 q& Q0 t2 G
9
! p. \8 N: }. S& c7 x# d
10
% Z* d/ R+ P5 T; d& v1 P
11
; T# `8 h! G# w* d0 ? x# v
12
7 \' B9 p$ K: {) G4 C
13
) C: \' t, Z5 ~6 |# D
14
# U( D% C% s q- a/ y
15
: j- K! [( @/ [- Z; W
16
. ]4 l4 W5 E0 d5 y* Q
17
' ~; Y" V( ~* S, f$ h0 E$ ?
18
1 ^* w3 B2 L O$ Y
19
) U$ `5 a/ h: V% [, A# {
20
: o0 ?, A, w& d$ D2 b) k6 ]. d
21
+ d z$ p1 I& Y: i8 c+ A
22
. 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.bias
1 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 R
7.1 优化器设置
1 ~" f9 k2 m& t) a$ s
# 优化器设置
$ |) G$ ^( N0 u! M2 o
optimizer_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 S
scheduler = 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 W
1 E8 \ g7 B. e7 S2 Z
criterion = nn.NLLLoss()
5 Y7 _2 t) n# q* h! G
1
9 S! y B* B0 [5 G% h" Q
2
" o* v4 S1 A" W% S- S4 f v$ t
3
. T) C6 @, C) N
4
8 H8 \4 v9 m& T5 K$ w1 R. R0 i' p
5
; n( _- o: z& O8 i) [, ~! W
6
3 E3 d$ L& A$ ^
7
7 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还是CPU
0 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 = 0
4 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*loss2
3 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$ q
1
2 c1 ~0 J F) C: a/ U$ G
2
! z4 Z% n# B" M% f: J# `
3
2 o% P% K2 y, Z u6 U# L9 @) E
4
* y. P$ ?5 \7 B( E5 t7 m( f
5
6 j, p6 L- V$ ?
6
- I$ Y' u, |" c; }# x" {4 h
7
0 u' V |, U* c% a2 D
8
/ w! W7 x- s. t e& c; j* f
9
! P2 ?" t0 x# Y& n" s
10
9 a. A j# C, q
11
6 o. `* d% \' i4 r8 O
12
" [7 p) Q3 l0 q' E" l/ ^
13
/ V: l% G* K6 t. x: O" k
14
6 _7 P" t; J) D' D3 O4 c
15
; c1 L! ?" O9 u! a( r' K! x
16
/ U. x( X7 _6 z0 Z! A
17
0 ]: \8 K9 u( a% |0 P$ f/ N
18
" w0 Q( v5 g5 v
19
* z5 }$ q1 V* b' _* e9 n
20
# x0 M2 [5 D0 e1 f* r% g# P
21
% E6 L1 M) p/ X# M
22
5 G v2 R* _) |8 f- X* v
23
* H3 [5 r7 ^% ~+ }8 r, m% D
24
& o6 q( @1 ?. U7 K
25
9 \3 v9 }) e- Z$ x
26
3 O: ~9 l+ P0 u, c
27
# w D( _1 H. }# W) I! W
28
6 Y! r5 R q9 l% C; V4 C
29
B0 }) c7 `2 V2 r; W$ v* y" G0 U
30
/ ]. f7 Q* B2 g' ^" `* D+ U7 h, W
31
0 H. C5 U, \( W7 ^9 c
32
O8 h3 M% S1 o7 d
33
4 T$ H& R# J- S/ k, _
34
% j) b5 [+ w* ? F1 Z& L1 T J
35
- M T8 w- @- O. p ~& b
36
3 ~' _: u( r! w6 a( {- u
37
' i3 C" p5 o5 @9 ?* Y4 u
38
; c' a% v9 e. g# V x1 B
39
' ] I. z9 `# S5 v9 @+ e) |
40
. _+ L: g9 Z% q [+ i
41
6 m0 p, v8 @3 r' p
42
; ~6 q' @/ c0 l
43
0 k- C) ~: Q5 `) e1 x. `
44
* `$ [# t$ u1 z7 y- r
45
7 W' g3 G* a( Q8 E8 u
46
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 d
53
, |" `2 U$ [* H
54
1 U- x+ {7 X( Q) V5 m) A
55
: q* J, {) t3 [3 d) A6 B
56
- `7 N% ]$ F. w; H, W# h
57
. t' D) q) k! s/ s. D1 V
58
. C- \, _% R# [ p6 a, Y n4 q
59
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$ T
63
+ 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) M
68
b/ G4 ]# e, Q+ }" }+ T
69
) s, c( d# F' p
70
& M- |3 o! N; ~9 m$ L
71
' A- t9 I3 x9 E' F
72
! C! i p( K' F7 S; E8 c
73
( _# {: r! R! I. D
74
( 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 a
78
$ K" {) u. U; c6 E" i# ]
79
6 D ]# p* z. M$ K/ Y; ], ]" u
80
0 s- H: M8 o) G
81
6 X" d( Q1 ^2 A4 f
82
& Q5 A" `; n2 }9 _9 s, m( ]
83
5 f0 ]3 b P/ m% B) O! `
84
' I/ d3 N: y9 P8 S0 b* v+ U. R' t7 }
85
5 T) h8 z8 M, o q4 s' }- i( q4 W1 a
86
( ]. k* |( S: ^* f
87
: Z; `( d6 Q$ `1 V* A$ F! R3 S
88
4 @0 ?& ^( R) b) g+ f0 R. s4 m I
89
. y! e y5 K$ {
90
' Y: r% ~2 {& U9 H ?" t+ Y$ r2 E7 D
91
" N9 N% N7 J( _' P; T1 R
92
3 ]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 Z
95
9 R/ k7 S/ s A2 O' J" S
96
3 j) t/ \/ R0 u; {$ X8 ?
97
2 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 A
100
4 R' T {: G8 u+ H$ w3 [& ~
101
G4 S! v0 T" W, G9 ~% a( a
102
, T: X4 ` K) F3 a2 {; R
103
4 f7 o. \2 `0 O" e) K4 X' e# y
104
4 X% c* y6 I, D+ \) Y; l# n
105
! q9 l4 T( o! X8 `8 k- P
106
* n7 M) c6 G4 J
107
% |) 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 c
112
' H6 F5 Q, z' ?+ l
7.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+ r
model_ft, val_acc_history, train_acc_history, valid_losses, train_losses, LRs = train_model(model_ft, dataloaders, criterion, optimizer_ft, num_epochs=5, is_inception=(model_name=="inception"))
1 |1 N$ S4 m8 l, y
9 b0 ?6 O$ W# }! T! @7 E [6 B
1
- N3 z' g- u# [6 f; \! I
2
! `% s' ]% G9 N3 a C/ m
3
$ G; d) z5 o$ h9 O- j+ Q/ t
4
. 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; p
valid Loss: 8.2902 Acc: 0.4719
6 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+ L
Epoch 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.7053
8 @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 L
Optimizer learning rate : 0.0100000
a& d7 w1 s4 G" n1 j3 ?1 `2 H, c
$ e) {! B& ^$ ]) a: l, h3 S- J( g
Epoch 2/4
2 n9 ]9 j/ M$ a& K, ?: A
----------
* m% h) z( m" B- N- p/ k
Time elapsed 90m 58s
; z5 a" b8 j" ?. w
train Loss: 9.9720 Acc: 0.4734
9 M, V: s2 M7 D, {
Time elapsed 94m 4s
0 X; z4 o- h2 V) g# d
valid Loss: 14.0426 Acc: 0.4413
9 p+ R/ `3 j# s$ t' r' N4 V1 A
Optimizer 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# j
Time elapsed 132m 49s
' }9 d4 H: ^3 l: r' F' K
train Loss: 5.4290 Acc: 0.6548
) \: c1 h1 z( e8 j: B
Time elapsed 138m 49s
, r6 S: t7 M, |" y. }
valid Loss: 6.4208 Acc: 0.6027
6 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/4
8 s: t6 L3 { k
----------
* [* U5 a4 o" d7 t7 J, S! t, G
Time elapsed 195m 56s
6 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.4914
9 X, G: q) P9 q; i) T2 t
Optimizer learning rate : 0.0010000
1 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, w
Best val Acc: 0.662592
3 s4 M% r: r* v& D/ s& m
: Z% R6 K1 n2 V3 g
1
9 k. u" S& J+ J0 l/ l
2
) j9 k* b O4 T" j
3
. G, Z, [9 N- N- W8 z8 p% @: a
4
; z! C& X$ f0 }6 J# l4 i. } q
5
, l; |5 Y) U" ?' P5 Y' A
6
# V3 L0 z w1 u3 c. `; k3 {: P
7
6 ?( ^* 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 k
11
0 x) T; G1 p* J' n
12
+ a5 _ {) V9 ]* Q, E. o% T& n
13
5 m+ u/ V6 A( K$ n, i$ h5 {
14
7 ^9 U' F+ O* [/ @
15
; e" b6 ^ d+ \1 ]6 y+ S
16
2 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% U
21
* z. D( O; M) A- W
22
) Z+ F( b0 H( L7 N
23
1 C4 B" I$ ^; r. _. J6 I8 |" N
24
3 o2 a" d3 \; H+ L+ t
25
0 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- \
29
4 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$ \
33
1 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 G
39
$ 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/ V
for 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 r
optimizer = 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 N
criterion = nn.NLLLoss()
6 v/ U2 q: P; V& Y6 V- q
1
+ i7 `- M- I8 j* f- g0 B |: B X
2
Q8 w- a+ K$ v5 X% l- e
3
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& n
8
7 k- O/ x. \5 L- u+ [& r y# S
9
2 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
1
2 a0 v0 w& G/ j
2
8 S4 e/ e$ X' P
3
+ p7 f/ U9 x- {- D s% J+ s6 e
4
' V1 T2 J: h- p. w1 d/ Y
5
" 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/ I
Epoch 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 }# l
Time 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/ K
7 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 k
train 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 {
2
5 X2 s6 l8 P6 O. p# P
3
- I* ?7 ?5 I) _% L: E& W- n8 J3 b
4
. u' ^- K; N. d, t3 Z
5
6 |: a! ~0 _6 k" D B2 \' J
6
: {! {+ C& M) Q: F# v( U
7
( X3 s' d$ k) j, f7 {0 c
8
; `- @+ q" B- ]. a" E; o
9
( 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 C
12
1 W7 M, y E, I' T
13
5 S- K8 Z# D4 J9 F
14
o9 ]* Q* @0 Z8 x7 L
15
1 o, N% @& G7 ?$ @4 o
16
% 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 `! c
model_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# _+ O
model_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! H
best_acc = checkpoint['best_acc']
& D5 B, v; o9 j& r
model_ft.load_state_dict(checkpoint['state_dict'])
5 y% o! {+ |$ W
1
1 ~+ ~: V7 U) o" a) o5 z/ G
2
1 A! ^7 [& e' c& }
3
$ ~" j& }% c3 b& ~
4
3 ~" 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 v
8
( V- K e* k& Z: t1 Z% t
9
* j9 k! Y% ~+ \* ~8 S
10
P9 _, f: X" q- N! `% m
11
6 g5 W$ y3 s6 m, O
12
9 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 * 224
8 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 p
def 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' ], V
imshow(img)
+ v- s% z, W1 F# G# b4 B
; ]1 d! u# A! v- \2 _. \
1
( j% | n+ Y2 Z. a
2
) T) F& I# v5 F
3
( Z5 L1 E! }: U8 a) d9 s7 b
4
& I6 I8 F, H! G
5
! O Z) p: C$ ^1 h: H* J0 o) g
6
/ K' f: h7 G6 n/ I
7
3 b4 }5 b$ ^' d# I w
8
% P1 C! @, e3 D% u/ W
9
& t3 O' h: `' T* Y" P! N
10
: X0 H/ J1 N3 P- ?9 k
11
, N2 `- l: L' w
12
5 ~: D. F& C5 U
13
4 L" v3 }) G' p. b
14
* w7 v$ ~0 M1 D, G
15
7 S4 G/ I- Z6 \9 D; p
16
- e1 t5 K8 ^5 E' c5 Y% K' X
17
1 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 b
22
4 [8 D5 F7 d& s- y% Z
23
) ^. G% B3 X( T# W
24
# _) Z. Y7 \$ G( m4 F8 A
25
5 r L7 B q+ D- {2 X
26
' V/ O# A$ c& |, v6 f2 U
27
' n) q& N1 `1 p/ ?5 s
28
" O$ {8 x5 F( b- Q( H Q2 N
29
; B% x7 E( F4 C# q
30
8 c' P# t7 T) |. L& L- w/ O
31
* G. n% m0 m3 |- G( y
32
; z8 Z+ r1 o; k) g' j' u
33
" z9 C! f U X k$ n4 @
34
, a; A. u) x8 Y& q5 N B2 E/ R, B7 t9 {
35
4 b( y0 B( e8 T
36
( v# O% y i1 t: q( s" X
37
3 k4 e, e' x* x$ G8 V9 ]
38
5 K( N; t% L. c: k7 ^# t0 }
39
7 e) V0 U) R" g
40
9 w2 C. ~& y+ X/ D8 Z) S
41
0 F- |0 p# b2 z& \# ]) p
42
0 ?& W% c" d0 ]7 \
43
& C4 G+ ^$ D' e" K7 E$ p- x
44
$ Q2 z ]- }8 @4 B
45
$ @" n* r/ Z2 o0 b; t" z$ @" {9 b
46
* q' M! u/ X4 k' P2 U6 S, U6 r
47
1 n( G0 n9 c' U3 k
48
* `# h$ N, S9 G w8 C
49
0 b$ o! X( `$ f/ g
50
" B1 _4 t) I/ A9 R$ ]
51
7 g8 }6 P' U$ u4 Q" j
52
. _( v0 g9 ~, L& U' U6 Z
53
0 O8 p! G2 w' k6 f
54
$ U/ L0 |! o7 M( l! W/ C
<AxesSubplot:>
0 ~( I& n- [2 u; L/ Y, s
1
' |' _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 d
6 M( e& l4 `* c4 m: ` I
img.shape
) @* v S+ U( B* r
1
! 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' s
images, labels = dataiter.next()
: C& w1 Q3 F/ d+ X! v- P& Z& W
9 L- ^, a7 N9 l
model_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 z
output.shape
2 N4 j! o$ E( w7 j3 |+ v
- W* e8 l0 X N0 s G
1
( b1 N6 G7 _8 F* u2 b) K
2
+ v' Z% ?+ d5 V7 X
3
. q3 n! ^! |" k8 Y. D. {
4
* t( @; ^; `# s0 X4 O _
5
3 V" F' W/ }9 p& R# p
6
9 ~3 d- ]' k& o# S8 Q* I7 r
7
4 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) n
10
- p& G( E9 ]+ o3 v1 ~+ l& z! Y
11
" ?) v% }$ j7 B" F
12
' z2 A' i% m. E$ W* B0 U/ Q: f
13
/ @5 J4 P+ Z g- q' I
14
, G" J/ r$ _& o- X c8 C0 I
15
& F" F! p! f' E3 L8 H7 q9 E1 U
16
5 m! S. y/ y* Q) _/ d8 V
torch.Size([8, 102])
5 P/ q) a6 v6 [' Z5 g
1
/ 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( L
2
9 U! T5 J( a8 k6 b+ E
3
* 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- G
columns = 4
% A) N% \1 O: u) e" P4 Y. E% C
rows = 2
" ~" P9 b6 s8 v; l0 C+ H
0 ]0 c1 |# [5 O9 U. `$ L3 `; G
for 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 }, r
1
7 t" D( A w! [: K
2
, n% D! r2 u/ H3 B
3
. C5 @1 G# p2 v0 i0 l
4
( h8 Z i' l* F( |- {
5
8 U8 c% {2 {, q% E
6
2 J$ S! N _' x
7
* {9 j7 k- n7 @+ M: t
8
9 c$ v& }! o+ E r
9
9 ?. U( W7 i* K% D# d+ \
10
4 d# e5 q8 H) c! i
11
1 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