- 在线时间
- 1630 小时
- 最后登录
- 2024-1-29
- 注册时间
- 2017-5-16
- 听众数
- 82
- 收听数
- 1
- 能力
- 120 分
- 体力
- 565688 点
- 威望
- 12 点
- 阅读权限
- 255
- 积分
- 174929
- 相册
- 1
- 日志
- 0
- 记录
- 0
- 帖子
- 5313
- 主题
- 5273
- 精华
- 3
- 分享
- 0
- 好友
- 163
TA的每日心情 | 开心 2021-8-11 17:59 |
|---|
签到天数: 17 天 [LV.4]偶尔看看III 网络挑战赛参赛者 网络挑战赛参赛者 - 自我介绍
- 本人女,毕业于内蒙古科技大学,担任文职专业,毕业专业英语。
 群组: 2018美赛大象算法课程 群组: 2018美赛护航培训课程 群组: 2019年 数学中国站长建 群组: 2019年数据分析师课程 群组: 2018年大象老师国赛优 |
【深度学习】 图像识别实战 102鲜花分类(flower 102)实战案例9 A& z4 W5 v4 ?2 u& ~' F6 u
" q4 f& p( w9 L! b
文章目录
- c% j; d7 s8 ^8 g4 d8 H0 K+ b卷积网络实战 对花进行分类4 L2 h3 s& o' E# V& t
数据预处理部分
6 f: E3 _; f0 u) R: j# u网络模块设置( _) A- ^4 ?7 y6 W) W
网络模型的保存与测试
1 s! l4 s9 r+ @4 \数据下载:$ T) r8 U/ l' R _+ I: `" |- P
1. 导入工具包
5 K0 R: ~+ v$ o$ E9 S! H: _4 e2. 数据预处理与操作
. d$ N* E' Z# Z! ~4 d* l3. 制作好数据源% [! ]; g4 `7 f/ d+ U) r
读取标签对应的实际名字9 F9 G6 F# J1 u, V! Y% W
4.展示一下数据
' s- \5 X0 w) @5. 加载models提供的模型,并直接用训练好的权重做初始化参数; `1 Q6 @; F6 y* d0 x2 ?
6.初始化模型架构
; K4 a( p+ z- F: g5 ^: V3 _7. 设置需要训练的参数8 W# j6 {4 Q2 l" U) N
7. 训练与预测* q! }1 D3 m: f0 {# Q
7.1 优化器设置
% M" ?" D, C% l2 l) E: n7.2 开始训练模型+ J/ O' n! z; m: n* U: F0 a6 R1 T
7.3 训练所有层2 W: F) X( z; G: n O' O+ I& S
开始训练
6 _; ]& Q$ O% Q8. 加载已经训练的模型
7 j! m" Z# z: l' w* j, S/ z5 R9. 推理
1 D+ n+ n0 H# B- v* e4 f( i2 H$ \' J9.1 计算得到最大概率# u/ f9 P" x/ J
9.2 展示预测结果
; k0 z/ l" k' b: L: g! x写在最后
+ E! F: h# t0 Z卷积网络实战 对花进行分类: j+ v) T V% b1 X7 Z# R
本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能
0 s, r; y) _9 G4 @- Y+ J. W; q+ p4 _: w6 q. `
在文件夹中有102种花,我们主要要对这些花进行分类任务
) Q, E3 s( |3 m' U9 }: S% E文件夹结构' l% I5 p( N/ h$ q1 X$ a
/ P% N; V7 m+ s0 X! {4 ^
flower_data& m4 @( r9 Z/ z7 C* q% H* z/ r
+ V1 ^" K0 K" V" b2 ltrain
/ K$ H" I9 ?- T6 g
9 M. Z% f6 L- u. `: d% z5 ~' {& N1(类别)
1 H6 u/ E- u5 A* Y2 }; t2
Q5 [# r- ~2 }6 I5 J+ C7 Dxxx.png / xxx.jpg
/ E' J* i4 p. T" g$ d, u. q' ^valid
& i) j7 w0 s9 h2 T/ J
7 E9 \- r! f5 U主要分为以下几个大模块( q3 R' H, G% ]. W6 S A
8 d# y9 k# w# F0 a, E7 y数据预处理部分+ k/ q2 y" f# @, T9 m
数据增强
; g$ d/ s$ L8 h( z数据预处理! d6 k0 o% `* q' s! [
网络模块设置# y4 m# r& m& S4 t
加载预训练模型,直接调用torchVision的经典网络架构
# b, w) c) X( ]因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务
! o' d; L' i( a' O, g8 O网络模型的保存与测试
4 _: A. {3 g# N: t( W模型保存可以带有选择性# o1 o' |& `/ T8 Q. o% k5 L
数据下载:' V# r# F& V% q. W$ ]$ p
https://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset/ s7 r) {% u/ T1 ~* ~- e
+ }1 e: @+ z# H9 x6 x/ ]5 e
改一下文件名,然后将它放到同一根目录就可以了" l* }' Z" [. A5 M/ i4 D
. F Q1 Z+ U9 {; r% G
下面是我的数据根目录4 z" J8 |) U) E. v) |
3 U' A' l" f7 T% x$ t' i4 ]% W
8 n. P) ~. _: Y- S# B1. 导入工具包! `7 C3 v# |' u' ?& o. M* f
import os
) n8 J) k+ P8 a! I; N% Uimport matplotlib.pyplot as plt
: B+ u4 [( l* s9 T6 s4 F4 f7 {# 内嵌入绘图简去show的句柄; b4 ], r7 D j! X9 N" q6 M
%matplotlib inline 3 J) f7 _6 J5 w: u4 O! ]
import numpy as np& y. S2 @) t) H( Z2 I
import torch2 R2 E+ f# g" j( W7 K
from torch import nn
' u. }8 c7 F7 i* A- m8 ^
8 I+ I. t9 I$ j' O5 h$ N" rimport torch.optim as optim
. {0 m1 C7 v1 X2 m9 K2 |7 H! limport torchvision) k1 O! v& p/ g
from torchvision import transforms, models, datasets
+ k4 F+ Z! A# @8 V1 g! A
) G' ]/ |2 Z5 S* i3 b% W$ simport imageio) H- ?. J2 K, F4 m& D
import time
' g3 Y7 y& W' |$ B# J, P# F& Qimport warnings
* _2 Z3 r8 c9 z$ J4 D% a9 Timport random
?# q, {# q; o1 T% x Jimport sys) f0 ]. P& c; h U; i; p/ ^7 S
import copy
4 S' a9 g( \& z* Z B2 S) g% Limport json
5 _7 j* w+ b5 Q4 Z# Ofrom PIL import Image
0 L; R) ~0 W! F& n. J( X5 X2 i+ p& x0 G4 Q8 q) ^
' O: ^1 E8 c7 {0 t1
5 n( s% K! t, E' G' ]8 a# v2 l- i2/ P7 C& O, z O" \+ b% j4 F( P
3
$ G( t! b9 N+ Q! e+ v7 k4 M$ X# n" G: J
58 z0 V3 M. W) H2 ^2 c$ [' j2 ~7 {
6
3 {: {+ [/ u2 H8 A7$ z( L: ]* J: }9 E2 q
8' ^- m3 _) K+ F7 v8 P9 L6 M+ w: d
9
/ z( S5 ]4 y9 g( f2 n7 T6 y) ^10! {% b+ l- R5 M& Q% @6 O5 g5 i
112 Z1 P* r$ u9 `$ x/ G( h6 V
12
1 R& M! \; u; d# X2 Z! O6 o138 H# c1 O2 k H R0 ?- I* _
14" E! g0 _0 {0 Q: w) I- R! Q
15
$ G% x3 r8 ]4 b( `. p& L8 @6 g16
/ ?/ S( r M+ t7 V2 i6 V+ S17( ~3 v, g! K: ?4 E5 ^' p2 f
18( y4 w% o9 c, Y4 x6 a2 w2 V
196 w- l3 l1 O/ @0 E( J
20
, W! ?; X3 b4 E% g5 Q216 y9 c L; O$ C$ J; k. q' i; d
2. 数据预处理与操作
0 y/ e2 H9 a' T0 h7 D: v#路径设置2 F8 m! Q* j: L% m4 w% A
data_dir = './flower_data/' # 当前文件夹下的flowerdata目录
y$ b% t+ H3 R6 `3 {train_dir = data_dir + '/train'
0 f2 L- U4 R; |1 a) H: ?( Xvalid_dir = data_dir + '/valid'% z- O5 ^1 W8 L3 i8 c4 e
1$ ?1 F& K$ W! ^) I% w2 `8 t
2
! a {8 Y7 j6 e6 l" F3
0 t- l6 _7 e8 z- ~5 f; t4+ \; o1 e$ T! C, k: C0 v* _- N2 k
python目录点杠的组合与区别
! C3 w/ R0 y& ~! f注: 里面注明了点杠和斜杠的操作; h# h5 d, t$ E
( P3 _! p: T8 g5 Z; |/ p& ~9 ]
3. 制作好数据源
% @6 P1 T' e5 J' mdata_transforms中制定了所有图像预处理的操作
! D3 Q9 z4 r9 Q1 e( HImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片
% N7 Z& Y$ y- L; \3 l: ddata_transforms = {3 L9 l$ ] ~3 J* B
# 分成两部分,一部分是训练9 ~* X- _6 Z- P+ _' t; C( P
'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间
9 W0 Q% y w# ^" U4 m& M* S, p transforms.CenterCrop(224), # 从中心处开始裁剪
+ H' ?# P: V# L # 以某个随机的概率决定是否翻转 55开$ C4 I, B9 m- j+ A$ I7 f
transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转2 E! K5 Y9 N8 B. f" ^% j5 ]! Q: l3 T
transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转 e5 k, F7 z9 z9 N6 z" H2 G
# 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相: S: B" u3 V, a+ _/ N; T
transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),. r7 G8 Q- W# n
transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB4 i" j j- o$ w& x
# 灰度图转换以后也是三个通道,但是只是RGB是一样的; \$ |' [9 ?" V2 \! Y
transforms.ToTensor(),
% Z. P! h4 v4 G0 [; o; a/ S transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差, t' V5 W8 |/ g* e. q
]),+ ~* J, w) o, s/ G6 d
# resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
9 C2 H; Z1 f6 S& @" I- @ 'valid': transforms.Compose([transforms.Resize(256),1 J5 L# g( J: {: w2 I) N
transforms.CenterCrop(224),
2 |6 y+ T/ S( N$ d( ]; t1 K+ d/ u transforms.ToTensor(),1 k9 p+ v6 n) l# s2 z/ A; R
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同+ `) ~& Y) N, |" Z2 P1 }6 g) s6 T; O
]),9 {. o1 K! c- }% ]( o7 Z& G0 X8 H
}
9 k( B L2 W2 b+ X0 w+ K/ I5 f( j5 P& |/ c, T% [ j6 D" t
1
. S1 v. p3 U! d# h2
) _/ l5 v, i5 m% g3/ X" B0 f$ t3 Q# z- ?
4
* }/ y: r- b. S2 `# Y; c/ q7 c5
, H n: G; J `2 |) {6
* H7 B+ y4 g3 I. s- V76 b2 I, V$ f+ U; d/ y1 G; J
8
. W. c2 p1 M2 o5 }3 C9
# ?! I8 ^6 t2 q10
! v/ c$ S& K" \11
* b% \: s+ i: Z& K128 z% L* P2 M9 o5 j% o5 h+ M
13
# R: l" s- h; p: l3 E4 N! p# o. h14
9 [7 F3 J; S5 y( x. @150 [; \0 N, j3 Z* w" k
169 K% U& F( `0 {9 a& L( S Y
17& Y% w7 |0 F5 I# t( P. c
18/ H, Y1 b4 f: { o8 h# Z* R
19$ v) H a& b, @% G1 P9 w
20
) _2 {& ^. Z$ I. d9 U7 X) ?" d( P1 y213 x h# D5 [1 E/ h+ g
batch_size = 8
3 b4 V) I$ _/ Z" P! ]3 Uimage_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}* F: X! O) U9 P8 ?6 `1 s$ m* L# G
dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}% V, a9 J* h' ~3 T8 E% m$ D& y
dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']} * j8 g; O+ T5 b
class_names = image_datasets['train'].classes6 q$ g7 q( B9 W4 e# A3 M' O- s
& V( A& U- s Q$ E#查看数据集合
0 {1 w* w0 D! e8 g7 R9 a% Zimage_datasets
8 v3 V4 r {( `! [7 Z
0 P1 T1 o8 |9 R0 [+ u1
* k1 B& s' f! Z. i" w1 j27 m6 {9 Y, k2 k8 V
3
( H$ x8 {4 }. ]7 S, A2 ]( |4
: }; u$ Z. G# P a* C' e5
) t# l% d' D& i9 `- x6$ Q2 U, I. }1 W* q
7! I4 J+ T4 o" \
8% c3 M7 y6 T# C% k: c d3 a
9 ~7 R9 R" d' C+ [: t- W$ q
{'train': Dataset ImageFolder; U H5 v9 c3 A; N% G( |
Number of datapoints: 6552: n/ h; H B3 x% W6 T& p, I5 q
Root location: ./flower_data/train
8 x. M4 Z, a' y7 z StandardTransform
: j" y' p/ F% `; ]3 B Transform: Compose(
" F# R" W% \" w RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)
: q2 n; \7 J ?' O CenterCrop(size=(224, 224)): c, X* N ^" D9 w
RandomHorizontalFlip(p=0.5)
& m3 E; G6 A* m' g, h/ R. D9 D RandomVerticalFlip(p=0.5)) W/ \& @9 W5 J$ [0 A" p: Y
ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
- L+ M; S# z2 \. Z e( i* Z RandomGrayscale(p=0.025)1 Q, ^# u2 f! _
ToTensor() K; M6 z* f2 T
Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])- Z$ C; ? ^0 v" H' X; R" \
),
4 ?$ t( U# r D 'valid': Dataset ImageFolder" L2 W O6 V8 f0 [" Q% K
Number of datapoints: 818; \: y; T, l- L8 a
Root location: ./flower_data/valid) L. X* U* h& t$ Y
StandardTransform
! G0 c/ _2 u# h5 p' P4 ` Transform: Compose(. a0 \5 \ F4 d @6 s8 D7 s
Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)
$ I) {7 a* R% m7 v) v' F CenterCrop(size=(224, 224))
! _' u" j3 s* `2 }- T) u9 h ToTensor()
# r! S2 }6 j, t# F: n Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])0 g: p3 v) ^: k3 Q/ ]2 g
)}
/ ]2 y0 D& f- U: P7 c! I
* {& e4 V! i8 A3 D1
& q! M- P) e- c2
$ i: n2 r. ]! B/ @3; n8 a8 C, {/ ]8 g3 H f
4
( \8 X4 [2 H: }- K P) X4 M5
1 \# ?' w% ]- d) z/ @: b* P/ {6% L% i( n5 V+ }, C
7. r! }0 e r" R4 J8 N# ~
8* V% n" G# O; k2 @9 A
9- w- c: h8 Y# o5 o; w# n# z* h
10
6 s! l! n' m. r112 e! [ p" F R, B# e
120 N% _; }7 j. h1 D2 |# ^6 f! r4 h
13
; P- F8 i7 t) t142 R6 y$ u% h7 J7 T( ^
15" |+ F! A0 n. R- T; c7 |
16
$ ?9 ^- r3 x9 o) a/ m9 z178 ~; s. A+ s6 ]" n3 n- Q) Q" Z
18
% |0 n# _5 a A19/ h' j, a& z+ A( g7 Y9 M: n
20
3 W6 g+ g& _ f21& D. q4 B' }' N, ?) y& @& S! k0 W- R5 ^
22
$ |6 X0 P6 [5 Y0 s7 x23$ o" o7 w6 i3 T/ n0 F3 q% ^; ~
24
2 O& h. ]) e# W. f K# N0 L# 验证一下数据是否已经被处理完毕3 p7 z# _2 j$ i" p
dataloaders
6 ]5 b% W. z9 l- ?( R' ]1
0 u+ S' H( C6 g; H- s; W9 ^4 n8 p26 E3 z- @6 c" v
{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,
5 ?0 ]6 P, ?1 U, `6 m$ d 'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}
1 `: \+ _* ~! ]( w, Y$ a17 N+ c& @8 b' j1 M' X
2) {! V6 _3 H* y& o& }: n
dataset_sizes; R8 Z1 {- m9 R- Y1 V
1
4 [5 P* w# V+ s/ w+ X{'train': 6552, 'valid': 818}
7 ~9 N5 D5 Q) v% r$ ~4 t. }1
: @% Q; q# `3 J; |5 H& c读取标签对应的实际名字
+ V( U( G" h7 \. p& l使用同一目录下的json文件,反向映射出花对应的名字
6 [3 o: X8 q1 F- E0 S3 f4 L0 Y, K/ m G
with open('./flower_data/cat_to_name.json', 'r') as f:' `# P3 d' J; v/ P) q+ V7 S9 n
cat_to_name = json.load(f)4 w- ` r% D7 _. |' `6 s: y4 W
1 d! G( T) t8 ]( N6 n6 p/ P( I
2
]2 ^1 p. s- e6 Ycat_to_name2 N% Q( g6 e8 B0 L
1
: {* {& {6 b3 A4 o; h' t2 V1 a* `( {{'21': 'fire lily',- u" ^5 q9 |. U0 Y e
'3': 'canterbury bells',' W! n, o `; V7 f- S2 K, O
'45': 'bolero deep blue',! E+ E( Q" H2 g
'1': 'pink primrose',0 g- L2 T8 c, k6 G- n
'34': 'mexican aster',
' f7 |) N. q* g3 u6 _ '27': 'prince of wales feathers'," K$ g% Y/ W5 y2 ~! b9 y
'7': 'moon orchid',2 g$ u" s! k, S# Q% f( l# J
'16': 'globe-flower',
# J( W f' j) g# _ '25': 'grape hyacinth',4 [1 e% ^0 t) M% K6 ^/ T* l# Q6 r4 N9 J
'26': 'corn poppy',
) N( B; a$ Z" L/ Q '79': 'toad lily',
$ ]6 d6 _6 A; \' W( N1 ~ '39': 'siam tulip',) v" m- R/ [# F; d) O3 z. F4 L8 e: f
'24': 'red ginger',
$ R! I0 _0 Y) H% R4 S7 R '67': 'spring crocus',
* {6 E) r- j9 h '35': 'alpine sea holly',
% W1 ^& W: W- J! d '32': 'garden phlox',: S* ]! d% g6 X1 u2 c. ?6 ?, r
'10': 'globe thistle',
- k$ ~( P& O5 g# f '6': 'tiger lily',
; c9 _$ ^0 U7 Z/ z1 S4 ? '93': 'ball moss',
- o! ]9 j; c" ~ '33': 'love in the mist',
/ b7 n* w# v( G '9': 'monkshood',
. A3 r$ c8 \7 d. K1 c( n0 U '102': 'blackberry lily',
@, p4 x+ ?/ {- a' f( k6 m, O: x '14': 'spear thistle',
1 V. S* Z. T f6 c '19': 'balloon flower',
* Z. [* }: W! ]5 s4 H. @$ a '100': 'blanket flower',
* n% P3 n7 e; O. U6 T '13': 'king protea',
0 _: P" t+ n1 f J" C '49': 'oxeye daisy',% \5 j2 T3 b$ ^9 c/ t) J
'15': 'yellow iris',
1 p& ^ D( V/ O; G' r '61': 'cautleya spicata',
" f" [9 {2 D1 t7 R6 f; Z0 Q '31': 'carnation',, S1 F; B: G, R
'64': 'silverbush',
3 O1 ]1 w. y* }: X( u1 k, Q8 ? c( L '68': 'bearded iris',: i! H6 ?$ J' v* a* W- J2 T! {! y5 d
'63': 'black-eyed susan',
! N$ m, t5 m/ E0 l/ k# d9 R2 | '69': 'windflower', |6 k( M+ `6 y! B4 j
'62': 'japanese anemone',
/ Z& @/ V. A- k- G: O '20': 'giant white arum lily',
K2 R- D& v! P( [ '38': 'great masterwort',9 J, L, } f; n+ f0 \ m
'4': 'sweet pea',
5 [! Z. }7 d1 W ~ '86': 'tree mallow',
( m9 ^. l& Q- P. K '101': 'trumpet creeper',! q+ i7 a. }$ ]
'42': 'daffodil',
) r9 \8 W1 L$ \1 o '22': 'pincushion flower',+ Q) J* x$ {1 c( G% F) H8 e" w
'2': 'hard-leaved pocket orchid',' }: h, a5 |' X% t; C
'54': 'sunflower',
# Z8 H6 @0 g8 H0 ^1 O7 [ '66': 'osteospermum',
5 m% V4 G- S/ L '70': 'tree poppy',8 a% j& N* L. g8 F9 G
'85': 'desert-rose',
' A4 r) p5 y0 ?& y& U '99': 'bromelia',
0 M, a5 r; ^2 q' i5 G! _" u '87': 'magnolia',
7 H. L5 Z( A0 @' a '5': 'english marigold',
t Y3 N( o, _; y( j '92': 'bee balm',
# V& A$ k( R c' u" V% o1 J '28': 'stemless gentian',+ u) d; w# c6 L2 A$ ^5 n
'97': 'mallow',3 g- ^* d1 h9 |9 @7 e3 S
'57': 'gaura',
0 |$ c* c) j$ u& J+ y, z, Y '40': 'lenten rose',
8 H' s/ x/ T9 C* { d4 k: f '47': 'marigold',: s, A! E; T1 E& b/ ~4 `
'59': 'orange dahlia',
8 O5 W+ D6 @0 F" R& [ '48': 'buttercup',8 U# u$ H* @9 e4 R6 Y# n# J
'55': 'pelargonium',7 A, `6 k1 W Y2 s5 ]2 t! }# }
'36': 'ruby-lipped cattleya',
7 {% @: x/ e: j9 x '91': 'hippeastrum',
: a/ e0 K% W ~2 W Q$ s% g; z) X/ D '29': 'artichoke',# Z, h4 f" b# h2 |* @. w3 d% d
'71': 'gazania',
" q+ ?& Y# r: F) c1 |* Z" q '90': 'canna lily',
5 h8 q! K; R7 e, @: p! W. J '18': 'peruvian lily',; d4 s1 w- a1 O: k7 u: ]- @
'98': 'mexican petunia',% V& R+ k& f: S# l9 {, E7 h
'8': 'bird of paradise',
; S8 V% i4 F8 m) b; A '30': 'sweet william',
6 V/ I \ G. p- [" }2 d5 d8 G '17': 'purple coneflower',5 P7 Y% F, L G+ K8 F p
'52': 'wild pansy',& u8 t* q" r, W% ^+ j) n' M4 D
'84': 'columbine',4 D" T( O+ D) C M1 g- b9 S* ?
'12': "colt's foot",2 M" M4 }8 m- ~9 L
'11': 'snapdragon',2 V0 _9 U6 U: Q
'96': 'camellia',
! _/ w" L3 M! s. X5 _6 ~5 T6 V '23': 'fritillary',5 ~) r K2 Y x" I' s' G$ J0 N
'50': 'common dandelion',& f, ?; b) X9 q# h: l
'44': 'poinsettia',/ @. o5 n O' u
'53': 'primula',
% q2 |/ T0 h" z4 j" y '72': 'azalea',; ]7 E+ t; R+ |
'65': 'californian poppy',
. O3 k+ q# a. | q! J; Q' X '80': 'anthurium',
/ }/ c( C: Z: v9 E '76': 'morning glory',
% `( N1 c$ _9 z8 s$ \2 J' @7 j '37': 'cape flower',
, M$ m5 [) J& D, o '56': 'bishop of llandaff',' E8 P. W0 X8 L+ w
'60': 'pink-yellow dahlia'," Y; ~0 X: y" W+ \) ^
'82': 'clematis',' j2 o% u1 d4 c) }5 |+ Y B
'58': 'geranium',
) s" E) G V" @# b& s8 f) } '75': 'thorn apple',/ I% g3 k, k1 l, u$ S+ }+ O$ R
'41': 'barbeton daisy',
0 W1 u2 f: E0 T, r. e, h '95': 'bougainvillea',
: S5 c6 |9 |4 I7 _$ a '43': 'sword lily',; i: C9 H& E7 U& Y5 r+ ?9 Z
'83': 'hibiscus',
1 K) q' d8 G& q, l) J5 T2 c- a '78': 'lotus lotus',2 r$ \" m g/ L+ Q
'88': 'cyclamen',; S! J! Z, P. a) ]5 r- C8 Y2 X4 Z
'94': 'foxglove',/ E* m% y |0 e7 o9 o
'81': 'frangipani',! N) x' ?7 G# x- O. Q/ ]( d
'74': 'rose',( Z4 R6 [; [1 s" @5 |
'89': 'watercress',7 D/ M; G2 N" t4 Q& y7 y$ J: B
'73': 'water lily',
+ Z+ E! c" x! E6 n3 {. m '46': 'wallflower',: C N% E& U# O/ }7 l
'77': 'passion flower',
9 t$ m( n5 A% E( E '51': 'petunia'}# r( K$ A G2 S8 R$ g# [) ]
+ W+ A# n3 h4 ~ K
15 Q* I) N1 C8 r9 p+ i
23 X% F" J* L* U! y1 i5 I. |, [
3* K' R" P5 [4 ?. k* E
4
" r7 T9 o% E: T8 i" X5# I% }- J; f* D8 @1 z* z$ f; A3 b
6* [. e6 Z& o0 O9 M$ Y
7
- t8 K4 W- [: u% l$ m8- g" l' [( d4 l X. _) }7 o
99 a' i5 ~" n: E6 f6 b @& X
102 A+ \% c5 R, k* G9 |, R% e! a
11& K7 q$ P: B+ X$ w
122 g; W* d6 K2 g8 x
13% a6 Z$ D3 a6 P5 R B
145 }4 u+ W2 @1 b7 g
15
/ e+ V/ }( e0 }16) x! L' `6 h0 A2 t/ |
175 G/ g! C* B! A" g' Z1 Z* \& o4 D
18, Q H% l, t( T9 ]1 W, b
191 x5 f4 B' N7 v7 N( s" t
203 c5 d% O p4 l+ c+ _2 H
21% a# W1 _, a* a3 W! K: x6 S
22* v2 x8 @( |' C* k8 {5 q+ D X/ U) C
23& ~& B9 ?4 P# p- ?
24
! h& N% ^8 _+ H) H, e25% i, j2 r/ L8 _% S8 Y6 U g
26* b6 x% |8 o2 D( t9 P, Y
279 E; E% m; i* X. C, `1 }* U/ s
28: @( Q; Y" A; T. }; n1 L! u
29
H) e( [& d0 e+ V8 _5 p3 h306 Z9 ?* |2 }. r, y' d7 \1 J0 D. e
31' t3 s9 _, p1 K5 m# q) v& ^7 F* g
32
( z0 C4 [' _" U) [4 k33/ s' d" B3 c0 ~0 p
34- w. n1 v2 X; [2 ~/ \9 h/ H4 ^
350 l0 t- R3 @4 p# Z
36' V- ~: A- ?) K$ `: f+ Z4 C
375 z2 J5 Q" N; u0 U4 u$ {- U' P
38
& t9 ^( R9 q3 m* O39
' ~ d; M& x/ R1 I40
) }3 R( G' w6 _2 J5 @" k' R: U41
. U2 `& d2 t- `4 r0 a1 c420 L( v( w( ?% Y) G; W
43# p# }2 n0 o7 i4 N! W% j& y6 l9 k0 Q' x
44" g$ x6 W, D! O" t7 E) V6 w$ a
45+ N" g8 D3 B9 H+ s6 U" f {6 ~+ ]/ Q0 @
468 D B- H( ^: _& Q
47( C* t' T! Z- a
48
/ ?9 ]9 }5 O! o: g49
+ n0 a3 j8 V0 D$ A. }; q# I1 b" ~3 j50
9 V3 L$ u. \6 D, M/ N. Y& Y% ?" Y51
. z% Y$ x! G8 Y7 L/ E52! B) C/ u9 p$ `* K
53# j- A' [9 p" b ]# k4 Y! j& ~: {
548 @9 _" p- ] _- n
555 @2 {) k( ?& l* S
56' `- ~$ d+ y! d
57
6 K* t' F7 G8 P* \& i! Z% Q: V( ?58 |$ a0 L2 }, g: e0 V3 s! g2 l* g
59! _0 ?4 E/ [. Z% }* S
60/ T ]+ b3 F- u
610 Y( v/ u/ A. M1 }9 H
62
/ C& J2 s8 f9 ?2 _- q$ r4 \63
* x. d6 J' O1 C7 N64
3 u6 A/ Z& e# H7 \2 Z1 G8 E65
8 U0 ?# _3 X L G: H J. m66
+ U. y- T0 |: B+ l8 \. r67+ t* y J2 `( g+ m
680 c. k; a5 `' J( j& A1 C* s( e! h
69
7 |; m7 h6 z6 O& b70
2 g" b9 O) `! ?4 h; T71' ?- b5 A* L* {2 W$ B* T6 K
721 K4 J/ b1 \6 B' g9 B
73# G5 y# e$ d9 D- D& s% u
74
+ |/ A; Y& k- ]; {# p+ C75; m: d0 h# ? S' ]2 D2 [0 W
764 R, l3 }$ Y+ m+ @. ~! p
77# O2 L% a8 \# r5 i9 a `
78
4 Z O4 `, n' }; }* f79
* f7 K2 w- c0 u6 n q' i) K80
" _- b' w7 I7 T81
* K2 G( Z7 M% d ?( P! L82- f3 D) t9 x, J- O
83
) ]" k, T6 |: ?) k! i3 V, x. L# o84+ n3 a. }( A# M! b. C* U" @
85
?. D$ o7 M8 y0 p- p8 W$ n" t9 Y86
" y0 i0 N1 m4 {4 w8 j87
^! r; w) K# V5 z( p88
9 s b5 z ] t. C0 A9 G$ o89
# k; B0 O% l1 W5 I( d2 ^90
7 Z9 p( f% X& D* U2 G5 m4 w3 ^91, x$ J8 u! v+ b
92
- G" e h) w) i" a' g93
3 e8 U1 K- ^; G; a; r94! V, X3 ^2 b6 u' P# Y$ b G0 I6 r( E
95& B: b8 B+ ^! L. Z& J T$ Y# X
960 n6 ^2 [+ U/ N) s) K0 z
975 ^$ J, a, i4 O% d) m7 G6 r2 ^9 W
98
1 V2 j- U+ W% b, ?992 Q1 _3 Z! B! z) `. n: u, q
1009 F' r4 v, r# o% N$ o: ?& f& w+ X7 B
101
/ s) Z* Z3 }( Z- M) L102; P& |) P7 ]% {$ K& c% J8 p
4.展示一下数据
U) d& T' P1 K' q( e6 Cdef im_convert(tensor):1 v1 W+ [& z; b8 ]' M( u
"""数据展示"""
. ?: {. b, k. V6 T+ {, B image = tensor.to("cpu").clone().detach()
4 B' v2 G5 a% A( I- \ image = image.numpy().squeeze()2 J( m% D, J" ?% O5 J
# 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图; _' [8 R @! Z9 q8 a2 F! @# Q8 O1 y
# transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)6 k3 ?! o3 Q( O6 k2 F* \, N
image = image.transpose(1, 2, 0)
2 ?" @ k4 ^2 I* w! | # 反正则化(反标准化). g( k( i* k8 {5 }4 `1 u% Y" b
image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406)). e' N. ?3 w: C @) n) B
( h9 C7 U* r1 i% N2 I; I # 将图像中小于0 的都换成0,大于的都变成1
/ A1 d. X7 _5 {* y image = image.clip(0, 1)" M0 h7 Z1 ~. Y& R+ W1 e4 s
5 x/ M4 T4 b0 O7 R; m3 u
return image6 b' i/ O( M9 l. e9 L
1
1 L& j/ {+ c/ H' F4 R2( E8 J* A4 x8 X
3) o5 H2 ]' d, J
4' ~% h- |, F2 g. D$ i
5* [4 B# G3 J$ L6 r( R
6
' u( X# {2 [& R7 K* e% \; f7( g3 v9 K9 F7 ]5 c5 r
84 j7 _+ ^/ E" t8 B( T: |$ }
9" S! x5 k0 @6 h
10
) D7 h; J g* G. u11
5 D) D- M( d& w7 G* Z7 V6 c12
/ \/ K1 w' C6 O" c* Z! J8 a2 C# S13- {/ Z0 @* Q; E9 t
14
0 `7 V' Q0 f* `! L# 使用上面定义好的类进行画图7 q# p+ X3 X8 F. z
fig = plt.figure(figsize = (20, 12))
& H6 N( u. j' @5 F) Jcolumns = 4
/ d; s! a, U8 ?) h, W- T4 Nrows = 2) j! h2 S! \* v; f' h. z
$ U" @2 b# q5 {3 x! M3 c0 [5 L# iter迭代器$ z3 v1 h* \! \! h3 B. @9 O3 B
# 随便找一个Batch数据进行展示 e& K5 w/ V' \; k5 [
dataiter = iter(dataloaders['valid'])
- v8 f1 u/ }3 B# I/ a: b& L. _inputs, classes = dataiter.next()
- P3 q, U T) T& z$ }0 o, M% Q3 d5 h
for idx in range(columns * rows):& Z" _# I c3 Y0 {, T. E/ }+ H* H
ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])% h) g1 e# @. [7 f
# 利用json文件将其对应花的类型打印在图片中
; y& d& g5 {; |4 P ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])8 I- F p0 K% u/ ?
plt.imshow(im_convert(inputs[idx]))3 [( Z! X' W" i( a: m, B) ? s
plt.show()! z6 }5 b; S( \- x& a0 C
+ Z# n+ d+ O# f1
( z# d9 F' i0 f, n3 b0 r2
' \- p( t( k8 f8 x( \33 f7 X6 \+ X" \% i+ E5 C/ e9 l
4; H: F3 [3 p* k
5
2 y* @# R6 w" F8 A" g; Y$ G6 q68 Q3 n1 |2 j+ l
7" N0 r+ N, F8 z# v/ y
8
7 ?! I5 f* [4 S$ Y/ Y% {, x( t, M. z' M9 `9
) f4 E; E% ?; Q& T* D10
" ]+ `3 m& h* d6 C# T" P; @! E11' q" c3 x/ f! d- J& S! H
12" ^# d; L0 H( N9 q# b9 V2 ?
13; E S; c* ?! ]
145 E. K4 f G& H! @, T. @; P
15
2 z% M' i# u0 s16
/ N( @2 p. J" P5 g( i7 ]2 @2 m) z) K/ ]7 [, p( P7 X l }
- B' Z0 C1 l2 y+ l
5. 加载models提供的模型,并直接用训练好的权重做初始化参数* g0 F4 D4 v- H& o3 T: M$ [1 m
model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
4 ~: j$ p3 M! H4 b& G# 主要的图像识别用resnet来做
& h3 [% t! K/ q5 b# 是否用人家训练好的特征
. q: _/ s! N2 o1 N- _# |) Ofeature_extract = True( U6 ]# j5 Q( H; C
15 i* q+ E) O% N) A$ c2 p& T
23 M/ y1 ?" { f- I$ I: o1 w
3, K' U: v# {0 v
4$ }9 ~) b& S: ]2 @% @( t; O/ ^
# 是否用GPU进行训练
3 a, b8 r m8 p+ n& C+ ^train_on_gpu = torch.cuda.is_available()2 ^- j" c2 y' A6 R* [6 d
4 Z6 h; O. V* [) D) Y8 l0 R' s
if not train_on_gpu:& x, k+ k6 x0 [9 {% l( i, g2 h$ B
print('CUDA is not available. Training on CPU ...')2 d$ \2 N; l* h5 C( N2 c
else:8 v# k4 q! T' r0 q7 |$ ]% n( Q) Z
print('CUDA is available! Training on GPU ...')
/ W7 T; ^7 C: b" o3 s- t; |2 X" @; }0 e0 ]: i3 T4 R3 @
device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')
% g7 }5 k8 n6 O" r# d m, H1
; W; p) G6 d r2
) N% j% J H3 \0 Y- u6 d( c3
& g5 H6 P+ c5 N5 e% X# }45 [" k4 C* f6 C
5 ^/ k/ d' f0 \/ r4 W+ N5 n9 q
65 t( _5 d; l9 Z8 Z* |9 Z" [' A
7: Z. {0 A$ K- z! Y0 K
8
1 Z3 o, H8 p9 { o- B9
& h" o, a: s/ W( k; U5 O, C+ m* q. XCUDA is not available. Training on CPU ...- ^5 I- k R( b) B( Q% f
1+ y/ [% M; J. T7 ^- {
# 将一些层定义为false,使其不自动更新9 u1 C C0 u! p& D
def set_parameter_requires_grad(model, feature_extracting):& o3 c9 b4 x6 y* h K! v( ?
if feature_extracting:
; I& ?8 u B' s5 y' J9 ^: Y# H for param in model.parameters():
8 {' x: `% A. g' Q: f: _ param.requires_grad = False& f k* J( F5 @8 n; F
1! O& t' o2 ]6 W# [8 N# i
2
% i! d6 q/ i. q30 E! e/ B; h* e' R5 Z2 |
4
9 S. [; J9 {; m" W" Z4 \% Q, ?5
O' v7 B2 i4 R9 m# 打印模型架构告知是怎么一步一步去完成的
# _5 [. X% ]+ b+ V# 主要是为我们提取特征的
/ T8 @$ j( a e7 r8 J% H! S2 b+ N" K8 b4 Q
model_ft = models.resnet152()
( s" P. R" j8 |/ K6 [! O; _model_ft
& T/ z2 T/ v- Q6 C6 g1- n: }6 M' T. z- ~2 p+ H# {
2# U. F! H+ E( n
3
4 b a* Z* s1 c- O4; Q! I W/ w+ X' @0 S1 p. E1 N
5* `/ g& A/ L. a) a# i {' Z" V
ResNet(2 Y; K5 @( w7 ?) d$ v% G# e
(conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
7 q2 k5 W& F4 D+ | (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)$ H9 |- ^: [( i- I
(relu): ReLU(inplace=True), P+ Z2 U1 d& C7 t( I" H8 g8 k
(maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
& V' F+ T/ z! Q u, @+ t (layer1): Sequential(5 F* C: s" h* U9 X* T
(0): Bottleneck(
u! c7 M- H, ?) a* ]( c6 Y0 ? (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)4 P1 x3 \) O6 X% z2 P9 r) I5 E5 J) l
(bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
& k, m- R" ^1 b5 \ (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)+ E' x* [1 {& M
(bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)1 k. K; X5 C# f! z) V2 w
(conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
( |: s7 l$ [( U# Q (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True): `) o. z, }1 q& m9 R3 U" c: v2 p
(relu): ReLU(inplace=True)
3 m2 v- a' |- z3 x) E* G+ F7 o (downsample): Sequential(. A+ W* ~6 S: S( g2 `4 G
(0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)3 f6 o% o1 h9 p1 g, Q
(1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
- ^' |4 S h& W )7 G% {4 l) d! x
); P0 L+ e# X- q& v
中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。, { X5 T# f E6 g. i: l; K3 U: }
(2): Bottleneck(
0 w$ u: K6 m# ^% c, e% x& I, j (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)2 p8 d! N* z1 z6 d2 P q+ v' Y
(bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
4 n; ^9 q) i% V Q3 J (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
( I% ~7 v: A2 r* t/ O/ Z3 @7 H (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True), p/ z9 ]3 o1 P9 |6 V) @! N
(conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
& s1 p' H& `# V B( W0 B! z4 B: C (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)$ C+ T+ n4 M" ?
(relu): ReLU(inplace=True)) d5 C: a1 ?) M E# I! T( ^
)
+ }+ t8 c0 `4 a( A) o9 s )
6 I; T+ f( H2 j5 H, R5 e; C (avgpool): AdaptiveAvgPool2d(output_size=(1, 1)), Z1 i" X3 v# i# i
(fc): Linear(in_features=2048, out_features=1000, bias=True)7 ]( c8 w( E7 K: d! b7 ]2 F7 C: \
)
# y) C, u% @# r) p% c9 _7 B4 b; S( i4 ^9 D5 ?1 P5 L4 P
1
: Y: z* p2 N: u/ w: t2
e3 q( ^; N2 P3 ]; _, \31 t; v6 b- a+ F+ c7 @
4
7 I3 P, Q V- C$ f9 R5" _# y7 o5 I1 T' V. X, X; M
6" W! q+ X2 I* p0 I- N V x' U
7: E E& g' g) g( r8 S, X$ p( l
8& g& t6 w1 p/ M v: i5 d( h
9
8 Y1 w( [5 |5 h+ \. E10% L# _( s5 V, N% J. Z P5 y; E
11! j6 y/ ^7 S" M* Q3 E! ~1 p J
12
( c' D$ I" J1 q. I8 P1 r- Y139 n3 o9 _- ?$ i5 q1 k) O
14
! ?/ s9 ~ Y3 o4 ]15) n+ g/ h; [ `, H
16
$ G& d, L: b+ a2 T7 _2 W17
# B1 m7 `& @5 N( C' v \. D! g18; S4 ^" B! _0 N
19
) E' M `0 X& [0 M9 Z20
2 q l9 L, l6 `# c6 ?21 ^9 ?# e8 d1 g3 G9 d Q
22# U( }9 I) P! M$ t; i
235 R: u4 J* m3 L& _) ^% i
247 n4 L( n+ s8 B& P# q
25( j8 r. _/ `% Z' _/ D" A; a
263 D& G% z' j1 h: ~, p/ W. E& U+ z, l( E
27! v. O, O2 j0 H
28
2 E2 u4 _0 a& f4 [9 h291 e5 [- Z1 K- t' p! R* g
30
! H4 _. [9 r( `$ k31$ b/ V- K# j9 Q( d. ?7 b% D
32 `* Q5 @- o0 w* m! m8 y
33; r# J& J' f: J) \0 H v( l
最后是1000分类,2048输入,分为1000个分类
% `3 k( g6 w1 N' y而我们需要将我们的任务进行调整,将1000分类改为102输出
+ ~/ X: k' \" v: j/ y. | y6 v
* J& L6 ]8 v2 [5 G& _6.初始化模型架构
+ }- \8 N: [+ h9 I5 P步骤如下:6 _2 J: F; c4 K( O* a" P
" }) j% K) f7 I$ T9 L- ~
将训练好的模型拿过来,并pre_train = True 得到他人的权重参数
8 q1 B1 n5 J8 U0 B& {( a5 h可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)
& N& f& @ a3 ^无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数% d( Q, C H3 W! v1 C" G
官方文档链接
" J, k( }0 R) j2 lhttps://pytorch.org/vision/stable/models.html |9 D1 Z# C7 `! N* k P
' p, f3 A6 o+ {1 O" v# 将他人的模型加载进来
2 j. Q1 I1 b H1 t- sdef initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):. X! m% l1 ^9 K0 _
# 选择适合的模型,不同的模型初始化参数不同2 l5 h' D+ G7 p$ P3 g: l4 g
model_ft = None; P4 R% E3 ?, E, V% [6 U
input_size = 0
4 H3 o; {/ t/ l& D: E5 _& w2 C" {" w! E/ k4 i4 |
if model_name == "resnet":/ \1 Q* I& L( W4 E
"""
- t2 D9 _2 a2 Q# W/ G4 W3 C, C Resnet152
' E- N/ G6 G7 T* u: i; _# U """* }: u8 k# ?! r1 n3 m; M
) O4 ~3 j$ e. @ C3 f
# 1. 加载与训练网络
- e9 ~0 s4 G7 D) ] model_ft = models.resnet152(pretrained = use_pretrained)
( _* i4 B$ Z) i9 [% g # 2. 是否将提取特征的模块冻住,只训练FC层2 }( ?+ i# S* }3 _' v1 i4 }
set_parameter_requires_grad(model_ft, feature_extract)
& C- K, m0 G% m: I" I # 3. 获得全连接层输入特征; v! n7 @3 B& ]
num_frts = model_ft.fc.in_features3 f: j; Z( B' b/ t+ v
# 4. 重新加载全连接层,设置输出102: U# ~4 F" w/ m
model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),, m3 w: E' E4 Q. V+ Y( {
nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为11 z ?4 M5 `( r L2 A" j' r
input_size = 2243 B8 {+ Z- i; k; r+ a1 k- k) A+ u3 F
8 I9 T4 P* u7 c/ R* |# O elif model_name == "alexnet":
* X9 v$ H3 _ p* N7 I7 q/ j# F! Z0 s- M """
3 i5 F, X R7 m/ ?" t; w Alexnet
0 x# ^" r7 w' } """4 F3 z! K0 f5 Q5 T+ [
model_ft = models.alexnet(pretrained = use_pretrained)! s J; L& j' i6 A/ y+ w' }& ^
set_parameter_requires_grad(model_ft, feature_extract)
( o' N# O4 |# Z
; X) A8 B( E. H1 q # 将最后一个特征输出替换 序号为【6】的分类器
, P: u9 D! R) C1 O) t) h W# F num_frts = model_ft.classifier[6].in_features # 获得FC层输入
: B; I% q; S* |1 M1 O5 _ model_ft.classifier[6] = nn.Linear(num_frts, num_classes)# _& G5 a- W. D- f- m( _( K
input_size = 224
{1 c: s* O& ^1 W; n5 }* b
. v( r7 O' E8 {1 S2 U elif model_name == "vgg":
6 f5 y: K% y4 i4 ]$ p """% b* x+ X2 D7 v1 C
VGG11_bn
: p/ u, A3 _9 q9 X' N """! t" |) t+ N' n0 q
model_ft = models.vgg16(pretrained = use_pretrained)
3 u6 L4 A! `0 V4 H1 {0 ] v set_parameter_requires_grad(model_ft, feature_extract)3 s O8 A W& b6 D
num_frts = model_ft.classifier[6].in_features$ H5 D* l/ ?1 }( u
model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
! R/ V$ K9 U, ~! }. @ input_size = 224) H$ [& ?; E- F1 V r. c
' x; a3 o7 ^; [3 X
elif model_name == "squeezenet":- l3 w6 {% m4 ]2 w- h8 ^
"""
/ h% |- Q8 |$ F' }" w) E4 ~ Squeezenet
+ V) n% U# F' w$ O4 r3 H+ V9 u! J """( _" T. `" B4 b" ]& q% f& C$ ]
model_ft = models.squeezenet1_0(pretrained = use_pretrained)7 T6 p0 ~9 U: I' S6 i
set_parameter_requires_grad(model_ft, feature_extract)$ D2 p' S! R$ j6 f' ^$ q N; d
model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
5 x0 ^3 O9 X8 Y model_ft.num_classes = num_classes& N; |. B& Z+ T' Z
input_size = 224: z, }0 {( m6 W- s( W/ I) U" I
4 {( E9 I# T4 M5 G+ y/ b elif model_name == "densenet":
" u8 q1 U' v- U+ }+ r """
" w4 b. c7 z4 x Densenet" G9 `) m* ^/ V, \# d! R; {
"""" T. U% Y" c2 H2 R, c( O
model_ft = models.desenet121(pretrained = use_pretrained)# Z) G1 x t8 B# x' x
set_parameter_requires_grad(model_ft, feature_extract)
/ q, y! y" O& t" z num_frts = model_ft.classifier.in_features
8 Z8 ^$ B5 m. h2 S model_ft.classifier = nn.Linear(num_frts, num_classes)" O8 J- T& M( t0 t7 g2 r
input_size = 224) C& x# K- l4 I2 ^+ b, Q
. `& z' K$ F" S elif model_name == "inception":8 [5 {0 t6 o1 u, j1 f; h& {6 Z
"""/ @- L' ~7 |- B' ]" \; h: f9 k
Inception V3
" P+ ~" u+ n7 S1 v) |1 |' O4 q! P3 I6 H. | """
+ J+ `7 ?, l$ y+ A5 `$ Q9 E model_ft = models.inception_V(pretrained = use_pretrained)4 i: {5 Y4 c1 i/ G
set_parameter_requires_grad(model_ft, feature_extract), D1 W; { @ W5 k
: B; _: h6 u7 H, h$ z num_frts = model_ft.AuxLogits.fc.in_features
( ^; M8 j$ b ~) J9 e model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)
4 N9 y1 I, ~" X [5 p3 P
0 p& ~. O$ B; h( F0 U num_frts = model_ft.fc.in_features8 f4 Z S4 Q; Q- K
model_ft.fc = nn.Linear(num_frts, num_classes)
6 A& F# P6 O1 k) Z' [ input_size = 299
+ \7 V; x2 W! T0 R. N% Q, R& `# P/ k( @ X
else:
5 T$ l& h! X- r# O0 n print("Invalid model name, exiting...")% s# l$ ]& ^, a( h. l
exit()
6 _. C. z$ g' }) ]6 o& U+ D0 S U) C" r5 G$ R' l
return model_ft, input_size
8 _% c" X$ e L: S9 K& E+ R. M7 X" B8 [6 d' W! ?1 r/ n
1& M0 \- m2 f& O+ U. x, |" f
2) a. ~ Z ~( O% j6 V' _
3
/ D9 b* `3 n+ d$ i& d) h4 S4, W8 d9 q2 e/ G: v X2 B8 r
5! v3 X/ J, k' f0 Y1 j: O: e4 ?; P
62 ^' H4 a( X# Y" }' N& C$ I
7
) ^- z! w6 U. g' f; i1 J! S; F8
5 P; h8 X$ }4 Z/ ]1 u1 l- a, a. \9, z/ A& j+ A3 w1 @5 |" Z+ d
10
% B6 _4 Z( _# `) f112 {/ _4 A6 N7 J9 L
12
. v* x- ?. f6 r+ v* J' F& U13
8 f3 V6 F* Z, n% d; n2 n14
4 V! V$ v- }% D v15% I2 J2 h' o/ k- U
16
3 I" Z+ J7 C4 w. H17+ X& E; Y/ W% \5 x' z
18$ l% x$ }/ s) U* e7 j$ ~& b9 U9 w/ d
197 w( z- J+ Y& H
20
& V' W4 e" f, _5 w# V21
( L! C' s, P" e9 L" M227 |( Z7 ^3 C4 W; H1 r2 v0 `
23
, Q5 u% X' n. k( @* Q24" |1 S9 @8 V0 g; H$ V! D0 } F
25
+ q$ C2 K, w' ]2 \# W9 w26) o* ?/ V. v1 |4 C, j
273 \) k9 Y0 |* @3 f9 d
28
! A( d) o$ g$ c* U! u3 C' {295 B) U, D8 `! [' m5 [& h
30! o0 N! z+ z) j5 R0 p- x) X
31
4 w; \9 `" C; ?32. S) p( {) f" O! `+ B, ~
33( ^1 t w9 P( _% y3 M4 f
34! Q+ a+ p, g5 A4 ~+ m
35% I Q4 x) K) M( t+ _0 M4 E! l
361 o: y+ W- E8 a1 K6 ^
37
4 P$ M* w3 @+ f/ k& b0 B38
4 M0 N' L) n6 x, M39* ]* L) h0 \# a$ @: }, |
40
5 G# R) ^- y/ M41
9 P; |: E p8 @4 m: H42. N0 K/ n9 l( b) c, L3 s6 b
437 z8 t5 k9 U! w! F" o' M2 R* @
44
8 j; F; n0 e% k: F6 F455 s( P9 N/ C: u* Y
46
# @: r, _- x- Z# |47
0 k- ?+ R- L; A! D48/ w( z7 {. K: e: q
490 r6 m( k3 q- ?' }# p2 f
50$ F9 ~" z% g' O( s7 Z. W, e8 M' `* G
51
) d3 E: y2 s& p, f1 @, {9 d0 x52# R: u8 G" b l# b/ q, N
537 @8 y6 M# p8 G3 c4 B
54& E5 R. r& @6 e
55
2 G$ |& d; P6 l) {6 w567 E; {2 E$ x* i* `1 U
57
0 ] m& X% F- y3 c* r6 v58
( \- \- \' q& E3 z0 x* O59
2 A: I( z& x, x% Y5 v) a60) b3 X7 K$ {! m, q/ v9 v
61' n6 Q: u; w0 F6 W d, T2 Z
62
3 P, h7 i! J/ h. F* \632 j% i9 z, L( \# ~
64
; Z% @, v( J- m/ b6 i8 S% ]65, c7 j9 {! m, n7 B* ~# O: e3 v9 B9 }
66
3 o- u! n4 ?' U: x, _' m* z( k67% o% `0 l4 _! ~
68
+ |& R/ }% a& ?7 P* ~! Q69! J" S$ y1 ]& p( c2 L% c
70* q; ?* o% a. U2 ?
71
% x, c5 _. v: a, b2 Z72
1 b$ {0 m6 F3 m/ B73
$ i$ V8 ]; C; B% H74
6 H. a( w2 f, t( {0 s75
/ o1 A( w/ K2 G3 I1 L2 t% {76; g3 m8 x; B' v" r/ Y
77: n! @2 q0 ~ ?) w' _) U, }& i5 S) j
78# g5 Y# X x/ \( }* V$ ~
79; v9 ^. u& t9 m: \# x3 F- n
80$ _+ }6 x& t9 M3 _& N2 b$ @
81
) T6 ~ @! t7 K' l8 z$ F V82# `+ p( b; x$ s
83
1 Y1 e( J6 R" @" [' C7. 设置需要训练的参数) E" }/ H0 m- V: ?1 s
# 设置模型名字、输出分类数
5 Z8 l# B9 x5 n; ymodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)' t8 V" {. F/ K' V4 W& {! E
0 z; T# l$ J9 ]" y' z3 R4 Y: I# GPU 计算2 f# p& j" h8 G( D. b/ y- X
model_ft = model_ft.to(device)* b/ r0 q# j; V9 `" s9 |2 c
, C2 k2 _' I7 j. L3 R# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取; e6 j6 {3 S4 J3 ~2 I5 B- q
filename = 'checkpoint.pth'
$ j+ l; _" h0 |4 |( l' K- _* o
* u: C8 `8 J. d$ D% P' w3 i# 是否训练所有层
% C. @, ^8 r' `7 v- A8 I8 T4 _( tparams_to_update = model_ft.parameters()
# q& _* l1 C- K0 U' H# 打印出需要训练的层6 n' _) b& k8 }. \6 [
print("Params to learn:")
) p) U3 e4 I" ?& j* jif feature_extract:
5 E, U& x& j& p6 u- ^( y/ O1 L k params_to_update = []
$ S; s1 ^7 D: m9 `2 ~! Z6 ] l for name, param in model_ft.named_parameters():
8 Z$ n* S* G2 k! d2 N/ }9 Q if param.requires_grad == True:
5 N/ N4 F8 i$ c* P; _' o) H6 p params_to_update.append(param)
7 g, F8 ]; `( U: s/ E+ w+ v print("\t", name)
0 L7 u" |6 x6 belse:
+ [5 n; w# \* \" {; ]) i for name, param in model_ft.named_parameters():
a4 S- W" f# }6 U. I+ d if param.requires_grad ==True:( E$ B3 W5 q8 Y3 `$ o# V
print("\t", name)
, p+ g* F" ^7 z5 H
. [4 T/ z6 }. `& W( W& k10 @/ @ C4 e' k. m$ W
2
# U" s( k' o, t7 d3: c' z ]5 h# d- z- h
4
Z" O7 V: a& j! h* L) X5# D: C( J( ~9 k9 I# F% ]$ \7 x
6' t+ ^- }$ j5 S! f
7
& }% e; W) z t1 ]; _! u- P7 l8 W: H2 r/ m* t) g# M
95 i) e/ }3 D( r3 i
10& b. A/ z9 r7 h7 C- ]6 u
11& C2 W5 @* ^, {9 u$ b/ _
12" {4 ?) s) i. ^, n# v. Z
13
* `' N0 ^9 d8 U' X) l1 L14
' }! U% b( ]' h6 C$ n* H15" u5 w4 j/ O5 P
161 p+ H5 |0 F2 s6 F' z6 j& G+ R
17* c8 A1 G' M0 O& Z' X, d' m8 x. K b
189 }0 e4 y# [ I) X$ ]. v! |% C) [$ l
19
! r2 r6 K! @' F! Y2 ^: t q20, f# c2 f6 P1 m$ p7 x8 @
21
1 [# V. D3 W& N O22
- H6 T# r L9 H i23- V% Z0 M4 w" R$ a& z$ [
Params to learn:
1 E9 E8 f a; y1 @+ P( T1 w1 x1 ~+ D fc.0.weight( a+ g. `7 X. E6 `# z
fc.0.bias) C7 G5 ?, N6 Q
19 Q/ b* E3 @2 A' ^
2
0 T7 j1 \# Y# b; P3$ ?/ z4 i: {: n$ @/ v4 ]/ R
7. 训练与预测
: @. ^2 n) j$ ^" {) w5 t, w7.1 优化器设置. g& f) `. `3 [3 X- @) K7 [
# 优化器设置/ I5 S7 O6 p0 h4 O) u, _
optimizer_ft = optim.Adam(params_to_update, lr = 1e-2)
. g9 [; M5 J% C) z* i1 ]$ m. b5 j+ a) j# 学习率衰减策略
4 `0 U/ h+ b0 ` K1 Z/ dscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)
! h1 s; {: c) ^6 |9 t# h7 _1 {- ~# 学习率每7个epoch衰减为原来的1/10
" G; T3 k% b5 q7 s$ ] i' \. O# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算! B) M' G* A+ q( v% m
( S# c4 Z ]* e3 {- y
criterion = nn.NLLLoss()% N% ~3 v7 k! `+ l) d h
1& Q7 m( q7 \" B' n! \& g. F
2, d' t; } k& m: ?4 N
39 ]/ O$ w# T: c# A U% n
4
% F% ?( b. K" R$ _5
# E6 K6 @4 s5 M# s I2 e61 Z' |+ E: a5 a1 S! W/ s
7; m0 `! H, F$ p7 v( \2 a: E
8( @9 H% |7 F2 r: K
# 定义训练函数9 t; y' v" B, G, ]$ n. R5 L+ G% c7 ^1 c
#is_inception:要不要用其他的网络
. K- C/ J- p/ C6 F2 n" ndef train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):% W/ q6 s9 i1 G0 i# f) W
since = time.time()
; `& y, t* X" u$ c" a #保存最好的准确率
9 T' F+ \+ L. N |* b& k5 h best_acc = 0
# A: ^, Q4 `3 m4 L """1 W$ G, H1 w8 j( r4 X* H
checkpoint = torch.load(filename)3 ?" ^; w8 ]$ x8 C6 ?( Z# c
best_acc = checkpoint['best_acc']0 b) p+ G, o! ]- G ]# V! ?* L+ a$ v
model.load_state_dict(checkpoint['state_dict'])
" F! j/ }, _% n/ X; U+ Z optimizer.load_state_dict(checkpoint['optimizer'])8 f: X. k" t) q, A$ b- G: Q
model.class_to_idx = checkpoint['mapping']* T/ y) I! u* q4 ^0 u
"""" D+ w% l3 L9 R" R
#指定用GPU还是CPU8 N9 `" R% z5 r2 q0 E3 P8 j. i
model.to(device) G. J z4 i7 K% c7 U* ^
#下面是为展示做的# b" B: t8 u. G: O/ X2 Y; i# ]
val_acc_history = []
( b0 N; Y6 |0 G! v6 i1 x train_acc_history = []2 u! B& N' t% {9 c( q
train_losses = []
9 f- o9 [0 _9 d valid_losses = []
- \2 u' u o! ]( k1 ^- C3 A LRs = [optimizer.param_groups[0]['lr']]* e8 C) Z: `3 o- H& e1 t
#最好的一次存下来
$ U& y! ^/ M1 w A8 D% G. B5 f best_model_wts = copy.deepcopy(model.state_dict())
2 Q4 X$ \+ C# x( _, Q ]
' p# m9 J( e6 Q/ f for epoch in range(num_epochs):/ G/ y/ N2 m2 I( S3 A' C
print('Epoch {}/{}'.format(epoch, num_epochs - 1))) K% ] R$ m5 K C4 Y; M
print('-' * 10)- C! C3 q; A% R2 T
6 b9 _9 a& G9 f$ b; C8 }# M# G8 s # 训练和验证
8 l$ h; m6 s* ~0 N for phase in ['train', 'valid']:
3 ~* q' G, ~& ?3 h# M- i! o/ ? if phase == 'train':- t# Z- j5 Q) t" t1 j
model.train() # 训练3 |2 s+ \' z$ T6 a) j6 K' e! y2 T
else:2 D9 I# E% }7 v& L" A5 t
model.eval() # 验证! [2 {( i6 g8 U9 C
0 B' e9 c. W/ h
running_loss = 0.0 j; S& O5 g3 I: y+ r. @
running_corrects = 0
/ o7 b2 v7 L) J/ w2 R
$ J! F. s. \, B# ?6 U& `5 u7 I # 把数据都取个遍
' [ d+ D% i5 w2 b, F0 B: k9 S2 R for inputs, labels in dataloaders[phase]:9 c* Q) E- |3 D- ~/ m
#下面是将inputs,labels传到GPU
1 r) Q, H6 ~3 S6 u" [' }0 U. H3 v' e inputs = inputs.to(device)
( A4 W1 e" q3 g/ h5 E labels = labels.to(device)/ j: ]2 H4 `8 I& A
: u* F- |4 U! Z9 S2 ?/ W& M
# 清零
+ ^7 ?" H& l; a+ { optimizer.zero_grad()
, Y9 }7 L4 F a* R+ o # 只有训练的时候计算和更新梯度
# J! B, {' Y. H with torch.set_grad_enabled(phase == 'train'):
. T$ e Q1 L% f: k5 I$ J# n6 h #if这面不需要计算,可忽略
: Z' x6 [' N. o1 O, B if is_inception and phase == 'train':$ W# `% \. K+ n2 G4 k# w
outputs, aux_outputs = model(inputs)
& H5 X% i' u8 D loss1 = criterion(outputs, labels)
* ~" e" I) P7 M0 D loss2 = criterion(aux_outputs, labels)
0 B0 R3 i) {$ g/ d6 J" [ loss = loss1 + 0.4*loss2, W2 H# u; h$ T
else:#resnet执行的是这里
" k; P! d3 f' a9 {+ l outputs = model(inputs)7 W: A3 l Y F# q' C
loss = criterion(outputs, labels)
3 D( G, F0 |( h, N# \
& v: A0 f* i8 E* b #概率最大的返回preds4 W6 Y+ S0 z% j
_, preds = torch.max(outputs, 1); y* ] z& d. X% p7 C0 E4 K
6 Q( c2 \3 d9 e' J) U$ C0 e # 训练阶段更新权重
) ^4 D- O& Z1 i/ G/ L% k) N if phase == 'train':
8 m- L$ Z6 z. V; y, }$ c5 Q2 [4 e loss.backward()
4 p) P- u3 O, u0 Z! A optimizer.step()
) T* U/ d& [/ f s7 _0 f0 _2 f
& t5 I/ I: B9 x4 \ R* Q # 计算损失+ ~, p h7 F* H
running_loss += loss.item() * inputs.size(0)/ f9 `! D( |+ o2 d
running_corrects += torch.sum(preds == labels.data)
) a9 x3 f! Q& o, R' j2 K8 m; k y& q
1 s) Z! c( {: N6 J; H9 @7 V( P- D #打印操作) B" X) y0 I! M& B: i
epoch_loss = running_loss / len(dataloaders[phase].dataset)
% Z3 C" d/ ?3 k$ ]' Z$ y T5 e# Z9 [ epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)' P: G6 F* N! H
! L% ]: n! e% Z2 d& w' n2 n, y' r. V$ ^" q! D0 K1 a9 v
time_elapsed = time.time() - since' M. t4 O" _5 A6 k7 K+ k
print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))1 h" `; N& J' b1 E5 ?( B
print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))# Z5 S$ o) C g
2 H! b0 n3 u' T. j) a r o; H7 J1 Y5 a7 @3 ?$ e
# 得到最好那次的模型
; q& K) J9 \, |+ }+ B6 M if phase == 'valid' and epoch_acc > best_acc: P, T$ _4 f4 C# X4 y
best_acc = epoch_acc m: i- H3 P5 h R- H. j f+ L
#模型保存( n2 q2 w0 P- k7 b! b3 ?* s
best_model_wts = copy.deepcopy(model.state_dict())
2 y) S) X7 B* _4 @) i3 z5 n state = {8 g0 K! P) p$ L; C
#tate_dict变量存放训练过程中需要学习的权重和偏执系数
+ n4 Z, e. L! n8 R+ O7 I9 Y 'state_dict': model.state_dict(),
1 O' w6 w3 e" [$ \3 Y% [ ` 'best_acc': best_acc,% {6 W4 h1 f N" O( |
'optimizer' : optimizer.state_dict(),* C9 Q& X5 ^5 G7 l% p6 d' I
}
( c* y9 a& U: }% L torch.save(state, filename)
: S+ ^0 k1 D( p if phase == 'valid':; X; ^; z$ m( g) H% F( B1 d
val_acc_history.append(epoch_acc)6 {+ ]4 C& v1 `/ m* n& q4 ^4 k
valid_losses.append(epoch_loss)* f8 q! @% T8 x. x: H
scheduler.step(epoch_loss)1 w) K# f% A! f" v$ W! ~6 w5 \
if phase == 'train':0 V# M- M0 U3 u. t5 g1 D$ Q
train_acc_history.append(epoch_acc)
# c4 t C) Z: N3 [ train_losses.append(epoch_loss)9 b$ m4 |' p+ Q X* S* i
1 D) L- h1 |# x2 c% \
print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))6 W8 c; B2 I% [( f# _" @
LRs.append(optimizer.param_groups[0]['lr'])
6 L" @) u8 I \* @4 D print()# @3 B: Y3 y; M, v* \5 l) l5 H( @
, q& n, `' z- M1 L+ @1 L- M time_elapsed = time.time() - since4 W+ V; @ u2 l8 \+ h4 e1 L
print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))3 y; A+ L" t; `6 M9 X2 t8 g h
print('Best val Acc: {:4f}'.format(best_acc))
; H7 p ^ A, P) [9 U+ ^
- f& ?& D& W- c! m. l+ u # 保存训练完后用最好的一次当做模型最终的结果' q: d: I+ u/ U% S/ y# ~
model.load_state_dict(best_model_wts)8 r' v, \5 C+ @$ M$ u: ^/ g' {: r
return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs
: |) D# c! `. L7 O" k
0 R6 p4 \% |% H: ]8 t
7 N* Q! ]2 b& S. y" W/ u1
a2 ^) Y4 [( K, X1 j7 G7 \8 t: ~) @2% @* q% j( D2 r( e
3
' R- l! A, }8 ]: a1 Y. @4
9 Y. q, ?' e+ j; k. P53 }/ X2 [* D [: h: y
6& Q" C8 y& T: U! G
7/ |9 U. Z# |) v
8
4 ?' V& g, K' E' k$ o$ ~: \" r9+ ]2 } ^5 }" S: O5 v- g
10* |* M. c; Q4 U" ~2 c+ D
11
: w* ] v5 o4 k4 c, k5 {" s120 i9 \5 T9 B& ^6 I
133 c, d( a$ X5 e5 P8 o/ O2 x. g
14
. b2 z) h3 |) y% d9 }% a15
* ~: m5 ] {- m; M/ X ]1 Q4 U16 x/ p" L R' q8 r- e1 C% ]9 b
17
( d" @) l9 Z& U! A9 d! o18% _4 U$ S. ^! b5 K
19# x& r" v: _) Y
20% V# c' x0 `$ a1 f
21( }; i. u0 O3 B& y8 N* o
22; B# x6 @+ G+ j7 L+ w
23) C' e F2 `8 F4 y, e
24
+ r0 ^. i; t' A1 H25- E" ^, g x( R; J2 Y
265 z* k" o. R" L1 c4 q2 H
27$ r9 m+ S/ w2 O+ t% r+ r8 A# F
28
3 G- J* H, e0 L29) {" L- T) z9 U$ G, p Z$ k; x
30% m ?6 e c' ^* c# Y5 `
31
9 l1 x& G3 G* @4 C* {0 _32
4 D# `- C- e/ p U33" m' d2 k3 g( X+ X6 E& e% G
34- d9 Q4 \9 j0 K/ d; D$ C) Z$ w1 P
35* _4 ?6 |- u8 m# Z! R! ?
362 B; T6 S- @* U
37: d. Y2 B( j# Q h
38; O6 M& I! O* L+ j" Y9 s
392 n8 Y" R. B8 i$ z. T+ ?' f4 u
401 T6 N9 B9 w8 @3 P' W$ X D
41
& q) V& J" H# B% J1 ^4 A42
# b2 B: ~/ k& k- m: p& J! E, b) V435 x8 p' U# t2 @* e
44
" j* Y/ _; u; L& T1 ]% q; A45! }7 \ }$ Y7 z+ E, B% }. x' j) n5 i& g
46" A! }: o5 x0 q) @2 Z9 c
47 @4 Y2 g9 o) ?1 i2 o/ @
48
2 j2 ?8 C! v: O( b R# j7 |49# \! c3 u# ]) L0 K0 Y2 \
50' z; W$ @( \0 M, {$ y
51
7 M/ N P$ Y+ V6 T7 a52
: y: ?6 v* m$ G; o53
. G3 c: v/ R3 q9 } J54
5 p4 J- r, x9 b* o( H! B55
' B# e1 x$ w' l8 C+ L56
% y$ f. @! X( w: r57
1 _2 r' E2 D; H' R58 s* l3 a0 e8 X1 I. y' Y
59
- _% F2 ~9 F: Q3 p- U60
& Q3 ]+ K, g: b' i5 l4 U* w61
/ H: P" C4 Q: T62
6 d5 ^7 W! c% M4 Y2 t63
% T( E1 n! g P- A64/ h& e6 [+ Z8 ^& G: I W! w
65
% Z* P- |$ i$ [! L, M# X& ^66
p. i- x* _" a* n674 X+ w; y3 u5 \4 {; ?7 H9 M& K* l
68
; Z W4 K+ M0 h0 @6 e" P( [69
& C1 b; t5 `0 }1 X, X! C705 A/ E8 Q2 L S+ D7 s3 o
71% w, V! x' U( q) G- R, |
72
! F$ V8 |! g# ^- |9 u737 b4 A+ d8 k9 ~
749 S3 o+ b: P- V- H- Z- c9 q
75
$ V8 [' d0 Y, {" d2 [6 K( u76
& W7 H+ g+ P. H4 B1 G776 g0 \" c0 s, w9 c
78
/ N H# ^. P( q) B" J79
! S* e' D+ \, e- a. h; }* X80
6 _5 e% I( o) Q" ]: ?: N& x81. c4 [3 j7 A% }, h% h; G
82/ f# a7 j4 N- X& C0 R, i! n. Y# o
83
4 W4 K$ j: W( d: E844 j1 Q' P' Q, x# W/ j+ o# }
85, S/ q. w+ ]6 D' O2 R: t1 h" f
86. \. H L3 d1 I. l
87- s. ?$ F3 R! Q9 m) w7 n4 C- T
88
2 G8 R; X( |4 O8 `890 \* {8 ^1 z& C9 Y8 O
90
2 u3 k( Q1 J6 z5 S7 ^( R91
1 y7 \5 [9 C5 U Z92
6 F! G; }2 e4 b( D/ F93
1 u) {- S# c6 C& K( V' S/ {94- w3 j& ^# n. y/ Y( N5 J
95
, @" m+ L- O7 K, W3 o7 S4 t96
2 T: g! w+ I/ Z, C2 B7 l5 q97& R" V; O. }1 V* U
98$ [* C1 F% [$ [! h: h4 Z7 }
99
4 O1 w: @' [! d: c- ~- l! R8 n+ k& P( X100
& I2 N, s7 l- l2 f& ]; |101
0 n0 O$ k9 A" d( ?: I102
" B/ B7 l& {+ P* C" T103# d/ I; H0 D% g: s2 Q! P, X8 q
104
0 C9 A# ^: X# k105. l- p! l2 @: w8 `3 M0 @
106/ n. I) q" c6 V! _# j
107. n4 ?0 b3 a) s2 s
108
/ P5 u+ w. X7 f6 ~109, s+ b" }/ ?. p( w, J+ M
110
2 V: U$ b" P4 ~: c$ I. v1113 a& h ^7 C/ ]$ c! Q
112
% o7 \, }2 L: q" C7.2 开始训练模型2 Q5 e0 ?1 u& |$ f8 b
我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次
, y3 u, ^ V) { w
$ u5 u# k' s9 \/ f/ c b3 t4 F#若太慢,把epoch调低,迭代50次可能好些
) ]5 i& k/ w/ o$ H) I) L: C! p#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合9 A& O$ w+ \7 D- ]
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"))
9 v% s+ t" O: }. p
: \4 }$ h# r$ I& |" q1
( B; _- w2 y+ ]( A4 a7 s- V# b [2
& f/ U) X8 Z. h7 U# Q5 [) @32 r( e1 x7 @# L8 W
4
% O# X3 u) G2 H+ NEpoch 0/4( t* i7 b L8 k% f) E
----------! m4 ^- I" O9 d
Time elapsed 29m 41s& u8 |4 Y5 P, s2 R: c) E' i2 D
train Loss: 10.4774 Acc: 0.3147
]4 _& r j4 S! a0 ?2 g7 O. \Time elapsed 32m 54s
+ T4 ~" G) H9 q+ Wvalid Loss: 8.2902 Acc: 0.4719 G' \' r: f% g3 X% i8 A% z
Optimizer learning rate : 0.0010000
0 Z8 ?! g- R ~' o/ T( Q8 v4 f5 [ _; q. F1 m0 d; u' O
Epoch 1/4
& L/ N- c3 t" f----------% L. L& z# Y! L9 L3 W3 H
Time elapsed 60m 11s
# b4 h" C' \6 L8 q% \( R9 h+ Ztrain Loss: 2.3126 Acc: 0.7053
4 s! Z U! B4 s4 J, O4 eTime elapsed 63m 16s) O; |& @5 j0 f3 y/ Y
valid Loss: 3.2325 Acc: 0.6626
4 E" N& g3 W2 eOptimizer learning rate : 0.0100000
- q( v! w$ B _0 H2 z7 a4 w, M7 R( G+ \' |% Q$ Q
Epoch 2/4 I/ R# L0 j+ J- s$ v: O
----------
1 c2 E' ^$ H' m0 H" u9 ETime elapsed 90m 58s: M4 S- W# C; @/ T
train Loss: 9.9720 Acc: 0.4734
, W$ R# u5 s# DTime elapsed 94m 4s' i0 y# f# M0 \' M$ m' |
valid Loss: 14.0426 Acc: 0.4413
* \* }3 ^2 ]6 A: i; o mOptimizer learning rate : 0.00010007 G: Y' U. q- P. X5 m+ J
- O+ \, p9 `# f( \* B9 p$ a1 R8 R d# ^
Epoch 3/40 q, a% q4 o( Q. `
----------
/ V2 H+ i+ p* C+ Q {% O% M7 ~: iTime elapsed 132m 49s
( q* _8 ]( e: }& E/ ~8 d# Utrain Loss: 5.4290 Acc: 0.6548
4 c' B/ w: a) VTime elapsed 138m 49s
+ Z7 O. w, R- O8 D5 j' Vvalid Loss: 6.4208 Acc: 0.6027
# o. r& l5 }# ~; R1 v0 T! DOptimizer learning rate : 0.0100000
4 }+ k5 A8 W/ V0 t# J& U
1 m1 i4 C. z/ [" Y! {1 jEpoch 4/4. l" G9 I- |! L+ i. K: H1 u( N1 ^
----------+ t) `7 o1 s m5 |1 W& T. Z! v
Time elapsed 195m 56s( P, L+ r0 C5 q0 f& e1 o# I
train Loss: 8.8911 Acc: 0.5519' |$ D/ I9 ~4 Q5 g/ J* U; H1 [) I
Time elapsed 199m 16s S0 F6 f! }5 I( U, u0 W% H( k b
valid Loss: 13.2221 Acc: 0.49148 p5 h% N/ B, G5 N% d9 |
Optimizer learning rate : 0.0010000
2 h$ u: H- T: n! ^) m# o3 Q2 G/ `# W. Z3 m( X, B' L0 n& D
Training complete in 199m 16s
; c* d$ K; P( i- |) |$ UBest val Acc: 0.6625922 t$ A/ [2 a6 G) y& h
- Z- @- k* T" o H/ @1 v0 l3 Q( d
1
( s6 u. U/ ]; l2 x% f% a2! B$ s \' w2 W# s
3
( b3 ^4 X X, K" B1 {% d2 @4
: Y/ H: g# N2 ?+ Q* {3 g5
: } r9 {3 t$ |! b6
6 k3 R; I# O, v6 Z: L7* z) a, _5 N4 ?5 h4 }
8
' V' i+ f8 s3 ]1 y9
! U4 [5 q: X, w# o0 c8 x10/ R+ v* x e1 ^: i
11
& }9 ?8 @/ b# r+ \2 r" Q5 H. J$ U12
7 W4 N1 M: Y* `13
4 x' i- |7 @) P3 D; C4 j14
7 W; }" x v8 W* X: z" d15/ l8 C# v0 e$ {1 o7 U+ A( m
167 [5 ]8 r6 X# _' H0 d' T3 S
17
! d- C* T1 h0 u18
+ r+ l3 K5 U9 i7 q19
1 i0 ?. G5 h' f$ O. |20+ t2 [+ e/ S1 o6 k' U! P# ?
21
! ~& C2 c; N7 y. Z: K/ n% C# z22
) F& v$ R' k) F; a23
5 `9 V9 [* w, c1 C24
8 G s# D) \( E u2 a4 |1 t25
5 D8 y. J4 w8 D, c7 O0 T) y268 i' m& N" _' ^) K" ?
27
# y% K( W/ h; E: }1 D& [5 O9 I) t280 q0 F8 N" F6 P2 L* x( O1 ]
29
5 P/ J' ]! J1 Z/ M30& m$ ]- }* V, u# g0 H# a: b A
316 T2 P% J- R$ _' u2 f/ v
321 S3 q. f: R0 O% i, `
33+ F. ?3 }8 W" m5 s
34
& Y/ a" k: a$ a1 L, T5 D35& ]6 V9 ]9 h8 ~% y3 }9 P
36: S. X9 B2 i2 x$ J
37
* k% G: W- d( d/ n% R2 ^( j5 j38
: w5 F& ?9 U7 T! d0 G( Q$ ^5 q396 W" G6 u/ z/ q' r
40
|( p( f+ G: N4 o8 R6 u41
" }6 d' O: ]/ S3 H" |; K42' u @, o1 C6 [8 x+ j
7.3 训练所有层: H& D0 Y. Z2 I5 \; R5 d6 v
# 将全部网络解锁进行训练
( a; ]6 t' n) Z, h' K+ Zfor param in model_ft.parameters():
$ ~+ m& k$ V& s4 Q param.requires_grad = True
, A9 @* A& R3 c2 {- h, `2 }2 o3 I& M- ~& Q# h4 E! B) h
# 再继续训练所有的参数,学习率调小一点\+ o2 O; m) s- U1 `
optimizer = optim.Adam(params_to_update, lr = 1e-4)
, s3 H! z7 R: f( o- c. U, j! jscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)
0 W0 n' d& T/ B7 M% x7 v' i
, l' ]* k5 j1 q7 O8 k( [# 损失函数' }4 d1 [6 _5 L, Z: O: @3 x
criterion = nn.NLLLoss()0 P, o0 z7 X) L6 \
1
8 Y8 ]0 u) b0 X- u" ^. s1 G2
3 w$ M \$ h: K' T3
, _2 E2 r3 W3 O# W3 W \- X7 D4
7 p! |3 A" R7 \: G3 @5
) h. k( c( M8 l! t6
0 B$ M. M* k; x77 b( @! X1 s& H" U! J/ |
8
2 f! x4 }8 e. J y4 u) g4 m9" S% d* l; B5 ^. K6 s% G9 `, G9 K% S
10
* q! R3 A% E0 e+ H [2 w5 I# 加载保存的参数* o! }( h0 b- R1 U' a6 Z
# 并在原有的模型基础上继续训练. p" v$ v3 d- ?) k0 s5 Z6 a$ d* v
# 下面保存的是刚刚训练效果较好的路径, T$ }4 |* }# S
checkpoint = torch.load(filename)
* u9 L5 d. y+ X# v9 Y t) sbest_acc = checkpoint['best_acc']2 O E1 e# ~+ _
model_ft.load_state_dict(checkpoint['state_dict'])
% G0 H% b4 v1 K* a9 }% Xoptimizer.load_state_dict(checkpoint['optimizer'])
, P3 @' f, T# x* L, C9 h0 f1
- S; q! E$ X. T2) n2 L+ ]( o9 H' u7 k6 J5 U
3
( z6 R" O3 M" y8 Q' u# r4
( ]8 L4 W, T! i! w/ c; P/ i* r7 E5
: y. G" p) c! o @# t0 O& K6% W3 ^. `1 @# D
7
9 A6 ?0 G. { [$ K8 d) R开始训练# |2 Q% X: h6 d1 F8 k* V. `* B& H) f
注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考
2 q* `4 y1 b( o4 k) M* L; _5 `6 A6 i; O
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"))
% }/ @( [6 E, Q7 {9 c$ A# k/ R4 K7 n1
4 o: j# w2 Q- ZEpoch 0/1
& d4 @! F2 s* e& m$ w$ N; s----------" P d% g6 ^! A& E
Time elapsed 35m 22s% R! J# L& o9 o: ?: Q4 S4 S
train Loss: 1.7636 Acc: 0.7346& Q# n1 Y. ^( f/ v7 r
Time elapsed 38m 42s
- c8 _* R c# P9 z+ K. G0 Q2 Bvalid Loss: 3.6377 Acc: 0.6455
5 r* P0 T! V/ S. |+ Z0 vOptimizer learning rate : 0.0010000
$ D3 P2 M8 ?6 w! }8 D: _2 M+ ?0 H6 P, s4 r! M4 g
Epoch 1/1. u* b3 }# Y( f. F; r
----------% T _ q/ U4 N
Time elapsed 82m 59s
% }$ o& [, g0 ~& {& Z; U5 f! ytrain Loss: 1.7543 Acc: 0.7340( m5 w3 a; I- K. K/ O7 Q
Time elapsed 86m 11s- m" F+ g/ r" e: O* [* A- U y
valid Loss: 3.8275 Acc: 0.6137
# x7 q7 W. p! d: ]8 {Optimizer learning rate : 0.0010000
$ ~9 \9 ^2 p2 _
5 K K8 f; ~; Q' o6 B1 g# DTraining complete in 86m 11s
* G) E8 P8 {: o# P' B7 o2 z" m& e NBest val Acc: 0.645477/ X! ^+ A1 N# O- D; F2 m5 J
. n4 f5 m r: I) \4 A9 v n19 r$ w! L( t8 x+ Z* k
2
, q* `. E" D7 K0 L/ i0 d3
$ _) _* N( n$ c! c) R' s4
D2 Z' y# F* t3 S5! M, {# U1 U" l0 `# t
63 P, B, N+ H$ b5 o9 t, Y
7; ^, ?+ x8 ]! H" P
87 ? M# P; { c% |
9
! d8 A2 [4 O: j4 @, c10
( u2 ]# r) M, p7 H4 S114 O1 P* G9 |' ^ F8 ]) w z1 t
12+ S, I9 o' t$ P: o; K$ e$ [$ _
13
0 A! t( I; v2 W* ]6 R8 X9 ~14
4 t! t8 @( e3 E* f/ m! ?, {3 T15. @; e1 T" ]" j
16
# Q0 u' j3 c x17
: A( N7 n( ]7 R( J7 n, I18
) E0 s5 e T. d* c9 q8. 加载已经训练的模型3 R! h6 X2 H/ l
相当于做一次简单的前向传播(逻辑推理),不用更新参数 q3 o5 M) \9 ^
) h/ c3 E3 q8 C' ] Omodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)( B9 y9 g) g* O
; }9 K' A: |1 Y/ |! H& _( N* \+ T' j# GPU 模式+ x2 i, k2 N1 K8 c
model_ft = model_ft.to(device) # 扔到GPU中/ U# {2 Z" y) q- t1 Q
! a5 W2 f6 l. j' L u8 ^# 保存文件的名字3 S: X5 C3 M) _( s% C
filename='checkpoint.pth'$ G* }" `; d) f) d) _# x. ?- W
* H3 H4 l6 ]1 A5 J6 e: x
# 加载模型
# T9 F, `( i, O- l1 l3 P d0 [checkpoint = torch.load(filename)
x# F4 ]# S- W {5 D/ u# h. B Ibest_acc = checkpoint['best_acc'], R, i% V6 K$ A
model_ft.load_state_dict(checkpoint['state_dict'])8 o/ y9 x- N( T4 j8 b/ h4 }: n% }
10 E* I( \2 t* t
2
: Q3 N$ D$ w- @! R/ s5 {3
7 V# j/ |3 ~0 q3 c3 y: ~& t4
6 M0 u0 Q$ P5 p$ r5
, B: a0 G. v: W, r62 r: d1 {5 l. t. l2 _
7
+ h1 a+ j" m. Q& \& I: p3 N! j4 M8
! [4 ^' Z- ?: H- Z+ E, e2 c5 Z& O9
; |- k) Y0 i! `- K1 K% ~10
' x v( F, g: f; D0 f: V& P, h113 c/ i1 ?' m' T* O" ?
129 w( ^( ^8 |, U( B' \( V3 j
<All keys matched successfully>( F, L4 `% D b I$ y; P. H
1
5 @2 w, V) c& p* O2 `. i$ Bdef process_image(image_path):# s' A* @8 U% D4 g
# 读取测试集数据
- s, b9 D: r; E" q img = Image.open(image_path)
0 B; v' u" u: v # Resize, thumbnail方法只能进行比例缩小,所以进行判断5 g/ l7 g* e5 X' @1 Z5 X9 y! _
# 与Resize不同* z- j9 S. n- P2 Y
# resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小1 h* Q4 J. {# n1 h- B
# 而且对象调用方法会直接改变其大小,返回None5 ?& B: }1 |1 a- w& @2 [' J% h
if img.size[0] > img.size[1]:" O- n5 k: `" v2 Q
img.thumbnail((10000, 256))
1 ^& A& k: ?' ], A- `9 s! c else:4 n; u, q, {% U A5 d( A% [( x
img.thumbnail((256, 10000)), R/ N0 I& m' Z- ^
) Y! M0 ?* ]% M' d, I2 X- j
# crop操作, 将图像再次裁剪为 224 * 224
! l9 h. o1 C9 ?; f: `! b left_margin = (img.width - 224) / 2 # 取中间的部分& | |1 }5 |- F/ o/ f4 S. M' `
bottom_margin = (img.height - 224) / 2
4 D5 n" [, p, c% B" T right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度+ R6 H1 I G+ ~: Z& k: B! z6 s$ g* F# I
top_margin = bottom_margin + 224+ i* Q2 r4 h! ]. L% v
/ x$ a- V) e/ L8 n4 k img = img.crop((left_margin, bottom_margin, right_margin, top_margin))
, P5 e9 V: B0 ~7 Y% `: z: i2 ?. G7 R8 a1 \
# 相同预处理的方法
% x4 {% x @6 q0 R7 [$ R& L # 归一化
; ?7 o" I; J; p) y* B" h% l3 f img = np.array(img) / 255* U% X0 v( @6 d1 U! C6 ^
mean = np.array([0.485, 0.456, 0.406])
+ M) B( U B7 `$ g6 E& x& k# m! ^ std = np.array([0.229, 0.224, 0.225])) p9 `! f! m: a, X) W# g
img = (img - mean) / std
5 u* m7 c4 q& P+ A. L( F% p: n: c- t: S, j
# 注意颜色通道和位置
. Y2 {0 |+ {5 s' A7 }! K img = img.transpose((2, 0, 1))( E9 h3 R0 \9 g: o' e
' H" e, w% [4 k return img
( S7 i7 |3 d1 B, H5 O' z
$ L0 p+ a& X, j1 x0 D; | @( G0 H' Ndef imshow(image, ax = None, title = None):& f$ k+ |' e' E: h6 b! O7 l0 o
"""展示数据"""
5 i) Y L4 y d1 H4 K0 @ if ax is None:3 ?6 K( u9 L k8 y
fig, ax = plt.subplots(), P& x: t$ G. N/ q- x( T N
- `( K1 L/ v4 X' C: L+ g2 s
# 颜色通道进行还原$ V |1 Z1 a* g# t
image = np.array(image).transpose((1, 2, 0))+ ?9 x- L. s9 ^
' S* W, N2 |* G' } # 预处理还原8 h! R' j; Q( R) A* h6 i+ B! F4 o
mean = np.array([0.485, 0.456, 0.406])) n h7 T3 c8 c, m% R# Q+ t
std = np.array([0.229, 0.224, 0.225])
' b( |& w# `; I- q+ _" E9 `0 n image = std * image + mean
* {% l# s! B& A' p: k image = np.clip(image, 0, 1)# p8 e$ L* o3 B( j0 q" @* h/ ~
! K% J& k3 a! x5 F+ B1 N ax.imshow(image)
. I H5 m' e, k1 I" q+ O9 l ax.set_title(title)
. t9 Z5 _* f8 [) i9 M8 g0 S( R7 O/ T6 x# o
return ax) U4 t( k& m+ v. b( T+ s! y
1 Y5 a# g3 \; T4 V3 Q+ M( o
image_path = r'./flower_data/valid/3/image_06621.jpg'
0 y, t2 i9 S: Bimg = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理3 M( ~1 d b( j- {
imshow(img)( A; n. d' ?& f8 G. @3 G0 x" ^
( Z( j1 E8 ?4 I9 Z$ u h
1
% [2 r# u3 j% s! P X2 z) k" [9 H: J2* l1 E) t2 X: F ~' o8 Y; s
3% C; ]" F/ N& h3 M+ r
4; W: S1 q1 g7 i0 ^* e) Y% F
5
+ Z: z, f/ E5 b8 V+ W$ h2 {6! Z7 W4 M! R/ [ ?! \# d
7
- `* I' d9 H1 R# ?9 b/ B1 @4 x8 E8
3 h, I% u% x0 W r7 a9+ J4 m0 h$ P* T+ O3 O8 k6 X+ w
10) [8 w, @ J* F$ [6 y% v' ]
11% _0 ]0 z' D8 O' N5 H$ l1 i9 Y
12
$ s0 \0 a9 f) }/ U13
- `; i8 y" Q8 C C, r. ^ f* o3 c147 y2 q/ Q3 r6 |$ F' y9 Q$ \9 {
15
7 T. L: b1 N. Q( g' j16( N6 c8 }" A9 H
17
M, ^. e' m D$ o3 W) R& B* @5 f18' I9 u1 S2 O7 W/ Z/ k
19
( g: O6 Q5 M# B3 V. s" A# U' R200 F7 [3 ~, |7 E- w+ {
21
9 e# a2 Q* t# u5 h22
) @( _8 f; s B6 C/ p' H23( `, }, l; z7 @6 O
243 N* J. n8 p6 T/ m" u5 [ ~
25
& t0 l& p7 S5 i# ?. r26/ e) M9 l6 f, u h* @
27
; E% [! D% s6 E28
2 K1 d' y0 w2 Y7 R, Q; v2 H3 n& Z295 ]2 Q# _4 N5 j0 k/ n
300 W$ R5 I/ b* r: H
31( h8 g& c" ]% H% x( l
32
# M, ~2 e' s5 Z33
5 B4 u7 |( R- b$ y34
6 s- ]4 C, ~5 L3 ]0 k35
9 N1 z3 ^( T ~5 ~36) g' ?( C8 w- f, |5 E( i3 c2 _
371 h k7 F1 a+ @! c
38* G. G. W5 A+ Q [3 d
39
- W$ h8 {. m3 ^2 N40! k8 D* b/ a% Y g" q& W
41
3 |- |* P4 P C/ E. w& D1 i42
2 E5 L9 ~9 a: e' w0 F7 J431 O4 u$ S: ? I
44) ?% Y! ?5 v# o& K% ]6 z4 q
45# f3 u) a$ L) m0 f
46 a% i6 I$ E" G9 L* |& ~+ X: x
47
4 W; Z$ \6 @+ [' |- c48' E) i' h/ F( i7 Q8 R# A5 d+ ^. k
49
7 X, B* T# p/ B) ^% @& q% M, O- r50' K5 I4 L! X2 i( g" h( q
51
8 n3 G7 V0 o0 G+ h52
& |, U6 B# D& ^( g53
4 k8 d K7 b( [, L* ]54 K/ t$ P# {# z' \& U2 n0 }2 H
<AxesSubplot:>
# {4 M) L7 H1 K. _7 K1
5 H9 k# ?+ C, u7 ?) q+ A( z6 V2 s, N& r
上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确5 }9 N5 o4 Y! b1 d3 m
- z0 m" ~1 t2 D6 x- a
img.shape
% k) n0 n+ M) Y, z6 x; u' F1& E3 \. l* f! |+ M% G
(3, 224, 224)9 d) b( b0 e1 v& u8 h( `, D2 |
1 _ |: e; j0 X0 C3 |* A9 E4 F B! O
证明了通道提前了,而且大小没改变
* l7 ^/ w0 k3 }# W1 Z1 ]2 o( n6 K, L1 g7 j0 M
9. 推理" h5 T- I# u& e5 E
img.shape
0 y! {) _' p" D2 f2 I( ]/ ^4 }7 |4 T B" l. j* g0 {
# 得到一个batch的测试数据, u) ~ m& n6 c9 C; [
dataiter = iter(dataloaders['valid'])
0 }2 u/ J; {; D( Bimages, labels = dataiter.next()% _4 p _; U6 h
+ o7 n, @" p+ omodel_ft.eval()6 K' ^ A4 ?' P5 @1 N$ v) Y: M
: b0 b& Y& |1 B0 Z, i' e
if train_on_gpu:
5 S# V. U4 S/ j/ Q, o1 k # 前向传播跑一次会得到output
2 Z3 [6 T& o7 }7 R/ K output = model_ft(images.cuda())# ?7 ~3 q, y3 D/ ~& J3 r* c
else:; |7 ^" p1 f5 W4 `
output = model_ft(images)! g* @7 D8 b+ e" R, W1 z* t
8 d! m E; T3 C( j' r
# batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值
- @8 Y4 s! }. b( K2 J% @0 U9 N6 y: eoutput.shape3 e( U% p) a* v
) s9 ^0 q$ T) j5 S6 ~8 j1# D/ A; w* ~) X$ J! m
25 d' @% e5 K/ x* t s4 u5 j
3
$ E( N6 X1 O0 w, _" j+ n4
% [; h; f9 d1 K# e& t5
+ }6 A1 E2 s) ~3 g( }/ K- s60 w. B, ^% D9 q
71 O/ O9 P( T" J0 F. U0 p
8
% T6 F. E, m) x/ B$ Z) k9- ^; y, w! T, c- J; J% a2 s6 N# T
10
9 j- Y2 Y! S @8 K. y. h) X& O111 K, Z* Z; b8 x: y0 x! _4 X0 j
12 }* l5 W, |! C* J3 K6 R
13
- j- ^- z7 j3 e) _141 q4 j# ]4 x+ E7 [ l# t
15
- ~) C% G9 I# l! t166 B4 M; t$ t. W! r' X1 ]
torch.Size([8, 102])( `. g- A7 s7 z! x) s# t
13 V, S7 E: ]2 o
9.1 计算得到最大概率0 \( y( k; p3 }# g' u- [
_, preds_tensor = torch.max(output, 1)# B5 ~% C9 i9 L' G _
) p. p1 b, c! D+ a9 _+ b# [
preds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量
1 Z% n1 y4 l7 T4 x) A1
0 L; Q. X& z* K. A; V2 [* }, ?* w$ G( d! d
3
0 y. s; E |, }$ D9.2 展示预测结果
$ {) }0 r a- s/ D( I, \: |fig = plt.figure(figsize = (20, 20))
4 X ]+ B1 w6 f# L. R7 V& Rcolumns = 4
" S% T9 ^9 g! A" e$ v4 L, Grows = 2- A0 ^' l/ K% [: X
- f. b1 w7 o4 pfor idx in range(columns * rows):
3 K! Y2 n8 c0 z+ d# x- h ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])
+ m1 q/ f3 l* w plt.imshow(im_convert(images[idx]))( Y2 C5 N) s1 b- M
ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), . p ~* q5 j5 b r- O. s0 O
color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))
% i) n6 [! v, `* _; I/ y0 gplt.show()" |/ M3 W+ N& e* h( d/ M
# 绿色的表示预测是对的,红色表示预测错了
" ? `+ s, ]# G( u: v0 f x17 \. D8 x2 F2 }3 O2 v1 ?
2
4 s7 v% Y+ ?) V3 w1 A; D' R3
5 j. y. a3 D; F8 e" O; k3 y% h/ n- t4
, y0 V3 R. I7 q0 |5( q3 ~, D$ `$ |& t+ v+ {7 t+ F
6/ o \% A, q7 v4 a, D" k6 C
7
* V( N4 `/ X' d3 e8
8 N0 |% ]4 `$ C9! k3 i/ d) G% k+ n& m
10
6 w% K6 m. M; \11
- D9 w8 K, e: i1 e! E& r3 u( U2 `( H" p: n& J. c5 P# o
, b$ t0 p& O; ^; z4 p+ C' U7 p3 f. K2 V. Z! N$ V' \$ a- b5 B0 s7 i
————————————————9 w" r& W$ H: i0 a1 W' ?
版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。( c5 v/ o: t% G4 g8 u2 K
原文链接:https://blog.csdn.net/LeungSr/article/details/126747940
# I' f9 A) p' e! k% }
' W8 C: ~0 P7 y0 f0 J' x ^5 i* K* s0 L) W. D3 g$ z
|
zan
|