在线时间 1630 小时 最后登录 2024-1-29 注册时间 2017-5-16 听众数 82 收听数 1 能力 120 分 体力 565683 点 威望 12 点 阅读权限 255 积分 174928 相册 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)实战案例 3 V7 [6 {$ d# g% b/ }# F% \
M b8 a! Z* _$ I. s6 k5 b& Q0 O1 P) F
文章目录
/ W" ~) r& @& v9 S* u- s8 f* Z 卷积网络实战 对花进行分类
1 }& T/ H1 e: @0 G, L' h8 z 数据预处理部分
9 q) ]3 x7 ], u. `& }7 ? U 网络模块设置
) c6 O5 k/ t/ Z. p3 n, e1 q1 L% U 网络模型的保存与测试
# r8 G4 X+ E- Q/ E) o% W, M; }% A6 X2 I/ z 数据下载:
6 D% K, M8 E4 u$ [) ^ 1. 导入工具包
2 x7 S* {' [; K+ L; v& n 2. 数据预处理与操作 g2 \3 l! R6 Q: v
3. 制作好数据源% I- g4 z) O& d# w: ~8 L: g
读取标签对应的实际名字9 }1 S# T- k4 E* v! H* {
4.展示一下数据& j6 |3 V0 c) p2 B7 g1 g. g
5. 加载models提供的模型,并直接用训练好的权重做初始化参数
: T$ S4 a1 p6 H, z& ~7 O 6.初始化模型架构* C0 U0 J3 K* S; I/ `
7. 设置需要训练的参数! C5 x# m2 b+ G1 U+ m2 n
7. 训练与预测& f6 M, g4 p: {! o5 ^7 `% ~8 [
7.1 优化器设置
3 H$ b1 i7 F6 t 7.2 开始训练模型
# x! x% y; F! U 7.3 训练所有层4 a! W9 R- L" _1 T% ?
开始训练* g: O! x1 b5 ?, ? ~
8. 加载已经训练的模型
' `! i* x: b4 R 9. 推理
9 o; q% {( }. {' s 9.1 计算得到最大概率# }& r1 c! I& v
9.2 展示预测结果: J. v9 W" R5 I6 \7 N F% E6 k
写在最后! A/ r2 f$ K% J! g& K/ V
卷积网络实战 对花进行分类/ ~- d7 W- }7 O2 ^0 B; v: C
本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能8 E# M( v) s6 N. j: W
) ^$ p0 K$ L: G$ I4 E E5 s 在文件夹中有102种花,我们主要要对这些花进行分类任务
+ L- B. L- @9 @' A) t& L. } 文件夹结构1 L: q, H# t8 L+ X7 O
4 F9 G* X# }1 t0 Y+ `+ c2 B
flower_data
4 n5 E+ j) Y" t* [: {5 w
# H: @( W" e8 U# g% ~( i train
7 H) \: W6 y; k! X/ q" `! g * ]0 V8 I+ b/ Q
1(类别)) B! _* Z' ? L. ~) _. R3 Q
2
: m9 H" q% K: p9 O xxx.png / xxx.jpg- Z4 l$ ~# V6 O% T: P2 m8 v9 h
valid
# P' T% H2 n q3 u# v, z
f0 k* v' @+ p# e$ K 主要分为以下几个大模块8 q% \' r1 {1 P% V
6 z. V. \; d* ^+ }; Q! e 数据预处理部分
% s, l2 @: x" Q 数据增强
* v1 r; ?* ^- [$ m1 b: m 数据预处理" o& E4 a& r9 B9 Q& A4 f8 v
网络模块设置
# c- Y: ?' W2 ^2 h3 ] 加载预训练模型,直接调用torchVision的经典网络架构' f. x& I2 K P$ X
因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务* T' Z# r& ?8 M
网络模型的保存与测试
, v, p; a" U9 C7 F 模型保存可以带有选择性
$ p3 f) i& f6 e0 b7 e 数据下载:" N$ S/ |; v' y/ w
https://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset
; h- l1 F1 P/ O# i/ j% b" e
" z$ f- b) i# H5 j8 w7 Q 改一下文件名,然后将它放到同一根目录就可以了3 J8 s6 I; {$ U8 k& ^
# p0 v' z7 W D! ~; [! W
下面是我的数据根目录/ k4 U' k" @! w/ ~. h+ V
7 h& ]1 }4 A% M2 S, f7 \% f
& p2 o: w' c) r# c 1. 导入工具包
; x; I& M' W n" p$ ~; ` import os
7 i' W! \( a* r. T0 S# Q- |3 Q# i import matplotlib.pyplot as plt) o! ^# _; q E$ ^# g
# 内嵌入绘图简去show的句柄- w( Z9 n# N$ E: v$ V' J$ Q
%matplotlib inline
& q0 N; h% g7 K9 L7 n# \/ I import numpy as np' i5 D$ \7 u: u0 H" T* b; P& m
import torch
! l0 ^- i. K- D/ j' K$ i from torch import nn' E8 W5 E/ L0 O
7 g9 B% S1 l( p# Q4 o( \
import torch.optim as optim. h- f( O; |1 c$ p" o
import torchvision
$ A( }; o" J" L from torchvision import transforms, models, datasets
$ ~* a$ P) U" F9 `
0 q+ x X( I c7 F7 @) j0 W: q import imageio
) Q. G4 P" _* X; ]# J import time# S* m2 _& I! a# P
import warnings
: @/ N2 x4 E \6 w- ^. z- Z import random
, R) a! Z6 c! _ import sys
5 k6 @3 ?& Z8 u/ F, q, P import copy; u" P3 O& H, A* u: c! c
import json3 S6 p: p0 w/ c$ }- y
from PIL import Image; Q5 |0 @/ L7 | d
, z' r* l; P6 C
# j: ]$ v1 e5 N3 c% f 1
, A7 n" W! M* w 25 G7 Y- i, g) {# x' z, r
3* E4 R8 ?" Q+ o$ |+ a1 }7 c
4
5 l% _7 p: O0 D6 a. l1 b6 ` 5# n9 u$ W s" F' J; {
6
* b, g2 [6 T( M/ X 7$ }3 C9 z/ }1 y; \3 O0 P
8
! B) s) f5 \( G7 U' X) Z* u; T9 \ 96 F O* A# g+ q- j
10
- q- c7 E2 P, b 11
) o9 v& I6 m9 }. P 12
' i- W# e2 ]8 @' W1 q1 z 139 H0 ^2 } G y- K2 M9 b6 j# `% z
14
+ ?+ T, ?* ~+ `. t7 q% B9 l 15
' ?. q9 [( U8 G. Z4 \ 16% J% t! h2 `/ }6 ?5 B' E/ R8 Z
17( u6 X/ y' L! L& |& g: i( ?
18
1 v) K/ r: D0 r) y5 o: K9 W7 d 19
% M5 R+ a9 B0 d) M/ ^( l7 i- V+ D 20
) m/ J* [* N- i$ J% \) w& O; Z1 @ 21% E# j/ F5 J" Z, \
2. 数据预处理与操作0 g' M9 B0 {* _ x0 Y0 e
#路径设置
4 j; [- y) R' K6 ?* U" w data_dir = './flower_data/' # 当前文件夹下的flowerdata目录# W% f6 }% F4 T- P" q
train_dir = data_dir + '/train' U4 `! }, p; ]2 h) \5 s2 A
valid_dir = data_dir + '/valid'" x- X% {) b4 [: x4 X0 T
1( C8 M% Q( d" _) s$ R( B
20 C, S/ B+ [' k* }
3
& Y" ]! z- \4 P. r/ Y% k 4
) [7 @. |/ t$ D( }1 W python目录点杠的组合与区别' r0 P! z7 a$ o G% i; H& ]
注: 里面注明了点杠和斜杠的操作. z0 b0 _. t/ G* e0 K
5 Y( R( k% b) l/ r$ [. @ 3. 制作好数据源
3 x/ G$ |1 k6 Z: a) `3 T) G data_transforms中制定了所有图像预处理的操作
+ A/ p0 E9 I) j2 D v ImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片
+ h3 |% C9 b" l data_transforms = { X$ c+ d$ p7 A/ \
# 分成两部分,一部分是训练) \, ?4 Z3 y4 ^# b& Y* G
'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间; m& f$ E" J6 P% r& |' B
transforms.CenterCrop(224), # 从中心处开始裁剪
7 H% r$ B7 z. b: O# V+ b% l1 L # 以某个随机的概率决定是否翻转 55开
) k- ?6 [+ n" m( X( C transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转
" t6 s8 A' x* Q% j; F( O transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转
4 t3 [$ P+ m2 [+ M- b # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相
; M. N' j6 b7 |" Y transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),
; ]# r& N/ `& A9 G transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB2 M' \1 S& c: v8 m4 f
# 灰度图转换以后也是三个通道,但是只是RGB是一样的3 _# }# _- _ O7 f5 u* j5 X
transforms.ToTensor(),
: i4 ]% u' \/ ~' D transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差' I1 h3 a1 n( Z6 Y, w8 I1 q3 L
]),
5 R3 h, H! a4 K( L- ^: m; U/ f # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化" w! N E/ m* c7 N" d* b
'valid': transforms.Compose([transforms.Resize(256),
5 y x, e7 B p& d( g1 y6 U transforms.CenterCrop(224),
: s6 y% r9 i. h' O. P transforms.ToTensor(),# i2 H* B5 [/ ?1 r" p
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同
, j, ~+ L ~9 B" \5 X5 p ]),: Q6 T. I' j/ \8 r
}
( S8 K# E2 t. x" j+ u' z, a. M 7 U! Q) }9 [* P. v
1
$ v7 q6 x2 i$ x% e# C* I- Z2 F 2
$ i6 Z8 C; z/ y6 _% q, X$ P 3
2 g+ a0 x: ~& ^4 h5 p0 t1 A |: c I 4
- ~4 C8 E3 C' f3 r6 M 5
I, Y, v/ k# h) B6 p 6
- z# s& l( J( T+ \3 L5 U+ l 7' i2 r/ R t& Y8 y
8
4 }/ s2 k$ G! s0 q 9
" ]) `' Q9 s! u# I, ^# _& V' w 10
/ I! \- \' A9 v- T! ? 114 P9 O' r4 Q% Q8 _
123 A- K/ Y D+ z5 |
13
8 X1 c$ e0 i5 n( t( ], k+ ~ 14) v+ K2 Z/ D5 A w' U2 N
15
5 P8 [. {4 f* R3 C& g. C 16
7 X- \3 j- z% q3 x6 ^2 }# r1 q# ]- v 174 N6 \0 V! w5 A d8 t" h3 @* V
18
3 t( f n- H G1 A6 [7 w/ r 19
8 a' I4 D/ Z1 i; L 20
0 |$ P' @9 ?, D, Z( w 21
; X4 I ]) m" r* }2 a; m3 q# B batch_size = 8
- }. F' L. H7 O3 f: Q image_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}2 M$ g6 p4 I' }( B
dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}5 X3 }! F( K$ V$ M
dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']}
( {/ `2 N; e% a$ h, E S class_names = image_datasets['train'].classes* B( }( W Z3 G3 ?, x
& \& }- e; Z; ~( V6 ?9 u #查看数据集合% q( R' B% X+ ?! {- e
image_datasets
( J# ?# N& P$ t% u7 V J
+ E" `" T; a+ E 1. B# M& u( v4 J* @$ B! q
2" \) n5 p+ o5 Z4 m
3: z1 K+ y5 ~% Q6 @' u: _$ ?
4: p- V" I3 C% ^1 I1 l
5
% ?" H- H& O* [/ F5 z+ ~% P 6; e8 m5 `* u% y6 ^0 u
7; h8 _$ _! }( x
8
: ] a1 p- V7 P1 m- M 9
3 r9 R5 W% m F8 X {'train': Dataset ImageFolder7 e; h5 E2 B1 E$ l
Number of datapoints: 6552$ a9 p8 O, K' \! w
Root location: ./flower_data/train
5 {8 c: K. C; F Y4 m: C s StandardTransform9 V5 Z( z$ B( h1 H7 [( s
Transform: Compose(
4 @% v7 W% G* r0 x" @( f RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0): o5 l4 J9 E: X( f: O: T% i
CenterCrop(size=(224, 224))" j7 x+ H, o+ L; a8 G5 M$ \
RandomHorizontalFlip(p=0.5)
: K( Q; J4 p. c RandomVerticalFlip(p=0.5)( K; c0 Q. \# ?' v l
ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
' \8 o' z \1 G6 h" X. A RandomGrayscale(p=0.025)
+ \; t7 _2 L3 m ToTensor()3 A& i' l3 _! w/ M" ]" S
Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
W' a* o" ^0 F- r" d ),2 Q8 E2 T+ @6 S! n2 m. V
'valid': Dataset ImageFolder& d. F; e- `- p- K
Number of datapoints: 8182 M* @, ^0 |, h6 C
Root location: ./flower_data/valid- _$ R% ~& X6 D8 M8 b
StandardTransform
* j$ v8 O* x5 L5 }. W Transform: Compose(
; ^* N0 x- b. a2 H5 |; f, y6 R Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)( X: b* j2 F5 u* ]
CenterCrop(size=(224, 224))
0 J# M- a0 k$ A+ h, r2 \0 X( C ToTensor()
. ^4 f$ y5 U! y# d0 r: a2 M( u Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
% p$ G/ D6 P$ S+ u% m. S )}2 x# E4 y7 G0 ~- y+ D
. [; }9 T: y5 ~8 G+ @+ [! T1 [ 1/ {7 x. j% @* A, \; Q5 c4 h B
2# S z9 `: m: [6 z; i$ P& u/ i- T
3
# W) Q/ i/ \4 U& d5 ^! N 4
v5 w$ K# ^7 M J3 p0 h, o' x 5. f" u( y( h% i: D, N: Y1 E" |( t
6
( E; O. \) C6 B 7# k- ~6 ^' N9 P n! c; l! p. S; [
8; x; `$ T" _1 x" F
97 v$ I6 G/ e( U- E4 o) O' v
10) p% l4 o0 n1 Q6 u, B
11( J8 S7 }$ B6 h6 p$ L
12
' `* G. X4 m; d4 ?: X# M! z& ^ 13
T8 k8 A2 ?; B 14
8 }! l7 x/ |6 i2 y5 [ 15
* V* X+ p& C5 X$ T/ H 16
! L* B- u+ X+ c, ?; K% M- c2 a& O 17
2 ?, h& A2 L6 E9 K6 P" y 18, J& R7 R1 g% k4 \
19
& A$ G. M* T, ^ 20
! k1 ~/ }' O2 ^; N2 s4 @% v, t 21$ B: T( [3 F% m8 Y( H# z
22
* p# F. P0 R8 C5 `6 } 238 e5 m' \. l( |3 Z2 [. j) Y( U
24
9 n, B/ \3 }4 B5 D- |6 c # 验证一下数据是否已经被处理完毕
& h$ z; f: ?& g; X- S& Q& x% I0 v dataloaders- e, I6 s) T/ S8 J5 }0 P/ I
1! j. i1 ~: F# S }$ \' Y
26 m7 ]" ^, w/ D4 P, { p
{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,
8 u0 w) T0 e% j6 A+ x! L 'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}0 q) L, {2 \% T4 {
1/ z/ Z X, c' c4 F; }: {. A
2& T% b1 G! F. Q9 G5 ]
dataset_sizes
4 M) q' t$ q& s) Z 1
0 W4 _) P+ T2 B& S. x {'train': 6552, 'valid': 818}% z4 S5 N. \+ b# ~2 |+ J# |
1) ~, G- \( J3 D; b
读取标签对应的实际名字8 ^# {! j2 B( K: h! K
使用同一目录下的json文件,反向映射出花对应的名字
; z8 B( e3 B- v# M p% ^
& p7 t+ \0 t5 B. U3 [$ c with open('./flower_data/cat_to_name.json', 'r') as f:
. j6 L% b) f; M+ @# C3 ]5 V2 K cat_to_name = json.load(f)
4 y8 ^% d, N: C+ n6 S 10 M. _! O' e) {. \2 Y
25 f# U$ R* B' P4 }! d- ~8 @
cat_to_name1 _$ e: Z( e2 H3 z
1
7 T8 Y0 ?6 e( t+ Z3 k3 y! g {'21': 'fire lily',
, C3 {9 h6 ~9 q, b6 U* h& t, v '3': 'canterbury bells',
4 n- b7 C! u2 J) }2 f+ } '45': 'bolero deep blue',7 d9 d9 d8 @, @! Y$ q. }
'1': 'pink primrose',
; | M' G5 ~% H$ Z: F '34': 'mexican aster',
; l7 M( a3 s6 K9 h1 y '27': 'prince of wales feathers',
6 ?, \( l7 c7 {' N9 ^1 y& _ '7': 'moon orchid',& v$ F0 W6 K# `. o9 V! m
'16': 'globe-flower',
" U) s/ q0 Q/ w5 z+ T '25': 'grape hyacinth',
! p7 l9 k/ Z! N0 i/ @ '26': 'corn poppy',* K0 k! Z$ W) R
'79': 'toad lily',
& C& W$ `! [- n1 c '39': 'siam tulip',4 l I9 w! r1 @" M
'24': 'red ginger',
5 `! d. p$ s2 b6 b '67': 'spring crocus',7 u% P: ?8 B8 A
'35': 'alpine sea holly',
7 I" c+ z" V$ G. f! _ '32': 'garden phlox',
7 P( q m5 x3 A( c6 |5 Y/ ?& i '10': 'globe thistle',
z& m7 F- E/ r! ]/ n '6': 'tiger lily',6 @8 u. D6 p0 q, z% e; U
'93': 'ball moss',
$ k! ^- C* Z% c h '33': 'love in the mist',
2 B" ^5 s" d) t '9': 'monkshood',
- k y' e9 j2 q' s; l$ T. M5 D1 t '102': 'blackberry lily',
* y2 F# V' n/ u; T* ~( s, r6 s' r- F '14': 'spear thistle',
8 Y2 k1 f! R0 m) K/ x8 c! m" y '19': 'balloon flower',+ |( B# Z ^$ V1 d. S8 ^
'100': 'blanket flower',
. y% x9 G" z# }( `5 ^: d '13': 'king protea',
1 t2 D- V; H7 H4 b6 p7 _: W- u; Z '49': 'oxeye daisy', u9 d3 U, l/ \0 A, q
'15': 'yellow iris'," X: L/ X9 Q* q- P$ ~2 m
'61': 'cautleya spicata',
. c4 r. X( b$ J% i( D$ W+ m F& j' Z '31': 'carnation',7 p3 _! c& J2 b0 E; X
'64': 'silverbush',
1 r. H9 [7 D4 d8 U; H8 I/ z" |$ S& | '68': 'bearded iris',/ Y3 y+ v6 k8 m6 O) i+ d% S) i
'63': 'black-eyed susan',3 K# w/ R+ e* b7 D* V
'69': 'windflower',8 M0 \3 G7 e `! b
'62': 'japanese anemone', j: f+ u7 J1 F/ |
'20': 'giant white arum lily',
2 J! R5 B6 }$ g( y* X0 x '38': 'great masterwort',
2 m' D7 \9 \& d- v) U" [ '4': 'sweet pea',3 P8 l/ U$ N2 z! q
'86': 'tree mallow', k2 N. x/ B |. F" n. C) r
'101': 'trumpet creeper',
0 n" m. ?, X- I7 V3 j" X3 P$ w0 V '42': 'daffodil',
, \8 w! o3 ~+ Q: x: L- c9 O) l '22': 'pincushion flower',
1 w+ n- d2 r4 O3 A; V '2': 'hard-leaved pocket orchid',
) F% l+ c+ f4 Z$ w3 H0 K T& k# I '54': 'sunflower',8 \. F) m( W' `3 W
'66': 'osteospermum',
2 C. H+ n! N, u r- j '70': 'tree poppy',1 w5 e* Z& d' n9 ~) `
'85': 'desert-rose',% d+ L, p6 M# N' Y: n8 }4 b
'99': 'bromelia',- |3 R% f3 d7 Z0 R9 Y/ ]
'87': 'magnolia',
* P; g2 b' U& Y# N) H) q: E '5': 'english marigold',, |7 @4 [! p! `* S9 x9 m3 Y
'92': 'bee balm',9 g, N* j+ }) d7 G( l' _
'28': 'stemless gentian',
1 X' b3 _3 ^" e) ^3 C. d '97': 'mallow',/ G5 K$ E: C, J4 \+ o
'57': 'gaura',+ W1 s. ]0 m1 k% e, S- B
'40': 'lenten rose',
! p0 e. A3 v9 H( A1 z '47': 'marigold',) u( s I; [- U2 W; l; h
'59': 'orange dahlia',$ x ^- ] i1 Q7 D
'48': 'buttercup',+ d; V6 J7 f) H! F; F# B+ O/ n
'55': 'pelargonium',
5 T I1 Y; Z0 U& G4 @ '36': 'ruby-lipped cattleya',
2 O) ?5 f( i# t C/ z5 M '91': 'hippeastrum',: |; ^! y7 Q2 \- D
'29': 'artichoke',; g ^; ~7 V1 }8 }/ T: V
'71': 'gazania',
# t0 @ `! ]& P0 ]. p4 H '90': 'canna lily',
- F3 ]( _) G ~ R8 N0 Q1 Y) Y% \ '18': 'peruvian lily',' S1 k" j$ A& l3 F8 V6 \+ S
'98': 'mexican petunia',. P' c/ V% s$ k. L. b/ n- N
'8': 'bird of paradise',
1 E8 }) A8 }8 G6 Y( T) ^9 z- i '30': 'sweet william',) @0 h X# }* c
'17': 'purple coneflower',
* `# K6 C% k! M4 [# \4 f# X% c4 p '52': 'wild pansy'," M8 t3 Y+ e8 F( `
'84': 'columbine',! Y5 t! A3 R$ `1 i$ v
'12': "colt's foot",
! B; y$ l/ X" u0 j* x '11': 'snapdragon',1 J) f! q/ L5 L4 s0 h i
'96': 'camellia',
! W+ x+ `# T/ ? '23': 'fritillary',; |; Q4 v6 K' d7 |- z
'50': 'common dandelion',
% ~% t- P$ m1 q( p% S '44': 'poinsettia',
0 R6 W* C# ^, q1 |2 Z* B, r' m# d( v '53': 'primula',3 L3 P& ?8 e. r
'72': 'azalea',
# U# i' W8 F; ^' B+ C$ r '65': 'californian poppy',
3 S/ d$ {5 a6 b '80': 'anthurium',; N" w+ k) }# I$ _, F+ r! A
'76': 'morning glory',* f( d B5 K( g
'37': 'cape flower',
, L% w% e" ^, |1 p# { '56': 'bishop of llandaff',
" s- X: L) Z# n5 A '60': 'pink-yellow dahlia',
2 r3 t) Y( l# S# {! A) o '82': 'clematis',3 V: ^. B! { |2 W) k0 j" ?" W
'58': 'geranium',+ \; s; X" p' x2 E! H
'75': 'thorn apple',
% I3 l" e2 y8 P* o '41': 'barbeton daisy',% P4 m* P6 v1 `
'95': 'bougainvillea',/ K' j% I0 B4 V
'43': 'sword lily',
' p$ |2 |7 n6 q7 v# y( B7 y '83': 'hibiscus',
5 b. \! i/ [- \ I8 G2 t '78': 'lotus lotus',
" |% |: Z) N( o/ m: E5 ^! } '88': 'cyclamen',
/ c2 s) t/ K) Y+ B3 g4 j' s '94': 'foxglove',! g8 o% f" l( n7 N
'81': 'frangipani', M( W( z a; O# P" y& |
'74': 'rose',, g8 N" @/ I. S }" i1 H5 {
'89': 'watercress',: ^ g, p4 E; W4 d8 q, G% `' h
'73': 'water lily',
9 m. f F- H. J5 i, o- B' r '46': 'wallflower',4 R- J. F2 x- c4 D) N$ b$ y& E' C
'77': 'passion flower',
: V. B9 F& P8 D5 q/ x '51': 'petunia'}, F' ?5 p& E9 D1 O
' z- K5 l) s( m0 C
1% m# M. R* E' u/ x! Y; l; q6 K
21 `8 c1 @0 h* {
3
1 Z, ?" K! W/ x" J8 x 4
& i) w- S! z3 R 5( V. V/ j: t; |, Y' W" k
6
% Z! M0 k3 _( L/ m6 O1 \ 7$ I0 W4 g2 R( a6 m; J3 P
8
& p% t. u8 l0 b8 `; b8 p6 } 9
3 u1 ^8 t' ?3 Z! P2 c2 K 100 _1 B9 ~5 e9 n p, m
11 \0 m" i) u8 I& z- }
12+ c; D/ J8 I, p8 k: J4 z" ^ p
13: R4 v! W0 x; \ L7 D
141 y/ D+ ]$ Z$ J) j( l
15/ F5 A" h3 C) v3 i3 _: \
166 I/ r$ r G- P/ l; y% k
176 w4 K5 w$ a4 s; T! {3 f
18: Z# a( O% t# p; f6 X& i, [0 m
19
' Q/ a* C+ a0 u, M- M, w' J% r 20
& p6 t& a' L$ c% {, ?' H8 R4 Z 214 f5 P# T# M5 v
22
1 J1 s" F* t) f) B, J* V. G 23+ U3 Q B1 ?$ U) z% q
24
$ Y. C* e. \+ o' X 25
: R- L. X1 {3 f- s' v 26& d6 m# ]' K' B' I
276 _' T/ y! P" ?9 \; `$ i7 \( X
28
, o- [& k* t& [3 F" y 29& `# W3 `; S2 X
30
2 f2 O- F% b6 A3 J) I% x |, H 311 n$ X# [; o" l/ G! b7 m
320 |5 N2 ?+ E* h5 T, |' I
33, U) Y8 ?: t7 L. [& Q. a
34
' q- F) { `5 p! D' H# u; {* k' [* `$ N 35
5 y7 g7 l" o* ^ Q4 B 36
) A( v$ U7 D f5 h6 p 37+ o! b- n; Q' _5 Y4 U; O
38
- A Q9 a4 ^% O, ] A 39
: @& z+ J ?/ s5 |. p7 y5 y 40
! j, L) J0 A# ?0 i7 s7 t 41
3 U8 f6 y0 U- L& u8 b 42
5 I( E( P/ Z; R+ a+ e) x! j/ W 435 M4 d0 h+ u2 ]" O" A5 v7 v
44# T6 y. X2 p5 p* ^) }0 X( l
45
2 e ~3 l0 i" n. k# w5 O! S/ \ 46- W! p, j# m* P2 d
47# x; G. X0 a; S! p5 P# `3 B
48
* u; z5 ^4 ]( C, K8 P" k) Z% T a 495 F) e# F$ B4 N S0 C6 Q% @1 K
50
. [5 Z- Y1 _8 D 511 F% X% Z( d7 E2 t' G) `
520 i, T6 {2 K8 j- t- u
53
3 h$ @, O; y, x Q: k; N 546 J% g/ S+ W1 i9 T3 E1 c0 [1 Q
55
* `' m1 `; s3 m% [% q6 @ 56; n& n& O, X2 V' o0 ~1 z/ b) P
57
7 ^$ V: i+ J3 p% V o4 T 58
5 a: o7 Z2 O6 w" ~! s1 O% p 59
; s$ v9 d z+ e 60$ i/ e2 V2 u. f0 H1 m
615 _' q! \7 z0 I" F* G& u; M: ]
62
- x$ R. ?* D! a" p. @ 63) `! v! Z( W/ Q# a6 r) R7 z+ F% D
64
4 L/ y6 G' j/ x5 R- b 65' A6 k1 a( l- ]; M( n% _9 ?
66
- x, \( a7 T9 d4 d8 { 67
1 Q, _* I3 i- t& [) S7 F# ~ 68) ]; R3 v l( N# p
69
9 U5 _* M. W9 k M1 [) Y 709 V% V' q% W: H- K* M- C2 |. f% H
71
& V: Y; F# o9 l) P6 _4 E& _# _ 72
6 W2 p' I+ X) u# T* L* n 732 {" m/ k7 Q. u7 J. ?& g+ p# }& ]: h4 ^
74
. }: Z* T3 {+ E+ N6 Q* a( E 75
# _; \* G+ B8 T# E9 v: Z 76" B/ F' e( `* F# N5 h9 R$ o
77
* p/ n6 E8 z* f' V 78
0 K( s' f' ^/ t 79. D; }" O! K6 c- I% H
80
# D U9 s1 c- T) J! @( X! `; D 81( O/ O, M7 A, X2 P7 K; f
82/ r- e, e; u) r: k
837 [- ?; w/ r9 K6 u( V8 }) K0 Q3 D( {; W
84: q+ \1 @& k2 H- g) J1 n
85: H8 b* |( x- l# d
86! m6 |% k4 `! {5 J, w9 J+ ]- ~) @
87
- p2 b' M8 H3 w6 s M* f 88- ^' d' K" |0 Q# b0 Z0 k! q
89
8 b7 y* t# s, s8 m0 }8 Q4 F' u& d 90
$ E% E+ M( L! i" H, p ^4 N 91( F9 g2 a4 C6 D ]3 Z0 x
92$ K+ r5 K6 ?" t
933 x3 B' q" V* }/ q
94' S% F+ B) j h. L) V7 x
95
1 ^: c- ~1 ], j- m! s 96& e$ n: `6 b4 c4 t
97
& Z/ {2 A6 `. r( ~9 a 98
5 B4 a% r) q3 [5 |% A) h 99
+ J1 O4 M' ]0 p 100
/ W2 X9 e7 b2 g0 ~8 J8 |, s 101# R6 X! X8 @* O7 n+ }, j& F9 |# Q
102# M+ ^: e) k J% T/ }# g0 V
4.展示一下数据0 k) |* N2 V- X: F: \ X
def im_convert(tensor):: ~& u% S- w$ T0 f2 q
"""数据展示""", x, M4 l1 v8 v
image = tensor.to("cpu").clone().detach(): r5 U G. Y5 K
image = image.numpy().squeeze()) d2 r. Q; W$ ^
# 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图- T& {* Z f! n
# transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)
) e( t/ Q; S; C3 `) T image = image.transpose(1, 2, 0)3 a& l ]; a. k7 g- K( J
# 反正则化(反标准化)5 U! }% o! n# [6 ?% {
image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))
( L+ Y! b A) H y& u- H
. D9 x$ `. s! h, a( [/ [8 f, T # 将图像中小于0 的都换成0,大于的都变成1
8 |0 K9 z) z) T/ Y image = image.clip(0, 1)5 T2 m. a5 u7 [3 O; b1 W& X
8 g G+ F: c# j( O' k
return image
4 A \; \* N# R5 g 1
8 H7 K* X+ J4 z* Z( y 2
. g- V* n' G" x( o$ J. f 3
6 }1 X0 z" a6 Y9 ]0 y) E) B* c3 C 4
1 h1 N' i: K1 K1 B }1 r# H1 F: t 5
4 }, W0 e9 i( @$ Z4 p 6
& ]/ T% N7 G8 L2 |$ K) E/ z1 M! y) z 74 [* O- Q4 Q7 ^- N# {. r
8
! g/ {6 G, S1 H2 J: p 9) P% G' N4 t+ A Y# a3 p( b
10
, Z) P# L$ i( f2 q- P 11
% B3 Z% y; x0 D4 o: R2 ^ 12
# z' u8 f- x/ F 13
) V, N7 d3 [% m# ^8 l) m# ` 14. [3 u. x$ ?+ o% c8 C2 \/ p( s6 F
# 使用上面定义好的类进行画图3 `/ o; R. N$ E, |
fig = plt.figure(figsize = (20, 12))
4 \7 J6 g$ D2 {2 E/ X columns = 4
) m6 J% Y4 t$ ^ J8 p rows = 2
2 E0 I9 n" e! f
1 S+ l7 [$ Z2 ]. |# ^ # iter迭代器. ^4 z2 D" m C! R, U
# 随便找一个Batch数据进行展示7 _% j: Z$ f, S* n" ~2 L, n
dataiter = iter(dataloaders['valid'])2 g6 Q3 h9 B! v9 t+ e
inputs, classes = dataiter.next()5 X- o9 t, J; ?7 C
; E) K( X! M. B( W; O, T" R for idx in range(columns * rows):
' f8 r J% e2 K3 Q7 o" i ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])! o$ }: V |3 G# v0 ?
# 利用json文件将其对应花的类型打印在图片中
6 u. w3 y0 K/ \# j7 J3 _ ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])
, t7 t. t/ K" S$ N; i% U plt.imshow(im_convert(inputs[idx]))
# P+ a6 i* t6 n5 [( v& I y; f0 z; L plt.show()
- ?$ F- \$ O/ X6 \% A( m
& ?* V- f/ m5 B/ R, }( B 1
9 J3 _& \4 k' `6 h$ m& i 2" f% R E" @0 N; U
3" i; \+ k- }8 l5 q6 ]) l! J
4
/ E- v/ r. z( _, H" Y* K; @, ? 5
$ a( \. z" }& G3 N, w$ { 6
& s& c! w z5 s, ] V( [9 i 7
+ u1 D& r3 k" W* R' o0 z 8
- C& ?8 {" Y- B, A. x2 c 9$ G+ W* _7 o8 V S0 I4 z2 r% F
103 ~1 V" ?) D* F) Y. U
11
7 B) ^- U1 }3 `. @+ v9 H& k 12
; l6 {' {9 }( v7 j# c- D! z 13
- T1 M8 U7 J+ A' ` 14. {# a. O+ ~$ M% }, n% {! k1 ~* s$ E$ ~
15
( C+ r2 Q/ b7 Z' M2 b; | 16# p! ]9 E2 y7 w7 V
" ` x$ w2 H1 K' C: I
$ N3 A8 A2 s7 E0 V+ s* N6 | 5. 加载models提供的模型,并直接用训练好的权重做初始化参数
) C- m, ?2 ?6 D5 T7 Y+ t2 I model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
2 g; |( l6 m0 y/ q9 H # 主要的图像识别用resnet来做( h `7 E, r& d- [! e4 w4 q
# 是否用人家训练好的特征
. }% N; h" D7 e feature_extract = True$ K) G/ p# H' q/ v0 t# b% Q- v
1
: s! {# q+ g! y9 o 26 D, R5 r, ?% A: K
3
& B- c; j" P+ r* p, W' z 4
. z% r5 E/ x2 l7 M # 是否用GPU进行训练
0 E0 A2 |8 i0 b train_on_gpu = torch.cuda.is_available()- _% C8 V1 G: ?. r7 r$ f
/ j9 P) {0 n8 ]* h" F if not train_on_gpu:
+ ]% ~4 C: G/ I7 m print('CUDA is not available. Training on CPU ...')
+ r7 ~2 g8 T5 E* q' f else:3 c8 T+ f& i6 `
print('CUDA is available! Training on GPU ...')( {9 A( g8 g& b$ R& n- u7 `, f$ d ^# s
. a F7 Q% a" S* v2 ] device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu'); C. e9 o- v# w( ~5 s
1
" ?3 F- z0 w2 ^& E& p 20 Z, R7 }; C1 y6 {3 S2 E
3
: U" ]8 k7 v$ u) L7 g, M3 s F7 P 4& I8 Y# v7 r9 k2 C& ^* k K
5, p+ d: M8 r9 m& M& w9 d4 }0 Z
6
p8 v8 e; _; U9 {* ~/ J 7
- n4 \6 `2 O% J; B: k) V 87 T+ S* V& _6 V l8 r
9
$ \4 l" z b4 e$ k4 c. I2 f/ H1 ~ CUDA is not available. Training on CPU ...& c% m0 D0 ?* v9 ~
1, i) e! D: j4 E, r+ e. j: `
# 将一些层定义为false,使其不自动更新. S3 T5 \1 C0 @
def set_parameter_requires_grad(model, feature_extracting):5 T+ Q8 N; d# L3 q
if feature_extracting:
8 v# @: V6 o n6 b I- s1 R for param in model.parameters():2 Y8 a) M3 P) v) D& Z
param.requires_grad = False6 `/ X7 W4 K% O/ k8 w t3 ?) ^
1
9 `2 V0 t0 S9 z" f0 Q" G 2
4 z# s# p3 V5 c$ ]0 Y3 f9 M4 i 3
l( L/ W- S" F; \ 46 `: ^# ?# U; d
5" @# O* E! p8 v5 Y
# 打印模型架构告知是怎么一步一步去完成的
6 r3 E6 c2 ]+ _( Q+ ^ # 主要是为我们提取特征的" y, G6 R3 P3 O4 {) p k
% \2 t% |, S4 r( W+ u/ ^* x model_ft = models.resnet152()
p/ ]1 N1 S- u! M model_ft6 X D- G3 C2 ~1 l
1
5 _/ w' f+ T8 R3 p) c7 v/ { 2; D- @* h" @: x1 G0 A
3' k* l' W) Z; n% S _! G1 ?" y9 j
4
, D! z+ y$ [: {6 I4 p" F( O 5
8 g3 D1 ?: L2 ^& z4 ?+ F5 _* X" } ResNet(& i* `$ ~4 } Z! {$ N
(conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
5 d3 L# U+ b( y, z' ^0 @! z3 R6 l (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
0 I7 i0 [1 {0 l6 g9 _6 Q. y# i (relu): ReLU(inplace=True); b8 z' K. S/ S8 u
(maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
& Q/ {* r% ^+ \0 i v (layer1): Sequential(
2 W6 ~$ O g, `8 n' b (0): Bottleneck(
' A9 d7 k4 a! f" l! H (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)
]* k0 _3 q" w$ R (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
, l+ V" W1 d, C& v0 D# ? (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)" S' p6 q; s' q! u& c: Q
(bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)+ t p5 M/ \ ^/ n: `
(conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False) K; [# I# u! t. d& }3 [
(bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True); P# l/ X @, F- V0 o
(relu): ReLU(inplace=True)
7 ]5 M# O. j9 @ (downsample): Sequential(
! w3 o. [! p# e. L6 ]6 @7 A# r (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
" u* ~# h/ G/ U0 ^/ E (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
; ~5 J/ J& q6 S! m )
; ~" C; H; u4 x- }2 v( e )$ E: X- Y6 Z1 _0 H7 V
中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。
& C) c" ?1 ~+ P( |% t (2): Bottleneck(4 T' {1 k7 r9 y) e0 M2 L3 F9 s
(conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)# U1 N) \) I. k
(bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
5 K' l/ q3 c$ `# Y (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
5 d+ A/ V! k# v8 J% O, r2 ^ (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
C1 k) ~1 |' ?; D (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
- }* O1 n! H+ l' ]# ^1 Z$ _2 E+ D (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
0 k8 T$ _1 F' u9 L. y (relu): ReLU(inplace=True)' Q# H: i/ p0 T {
)1 I2 a, Z: f3 ?: p6 F
): b9 `+ T; ~9 u: a8 |
(avgpool): AdaptiveAvgPool2d(output_size=(1, 1))
U2 m' m" m! s O3 ^7 \7 |/ d% \ (fc): Linear(in_features=2048, out_features=1000, bias=True)
% V% ~9 ^# i3 C$ h" d4 j B )
~, L% O4 O) \0 U3 @ ( b& G. G' L8 f! Q+ m
1
! x# ~, Z8 ~: z: M6 Q3 @& i8 l" a" Z% y: ` 27 A9 N+ j! g' Y) [/ d' i
3
2 i* s# y% I( k3 e# ^% Z3 K 4
6 a T9 F" s. _6 h5 \ 5
! Q9 G0 e* B+ L/ A2 s2 ~ G6 \7 U q 6+ Y# x4 {/ E1 @1 S" Y2 t) T$ G
72 g( X2 h+ z. {) v9 _
8& J/ l y8 K! S- K( i+ {6 I; p' q. @ J
9
' v- K" M6 v/ v; m# R 104 y O; |9 ~. _. ]4 x, N$ f+ } z
11( k/ [5 w* h! b9 c" v4 [
12! T. \9 X& Z4 ]! e2 _8 N* V
13
* Y9 @) V" Y4 t. @% s: l: e 14) j4 \, R* |* }) V4 ?
15
4 ?& p2 n/ j3 c! t( W' s 16# C) Y& _ L+ M" k2 ]7 @2 S3 L
17
; V% h3 K7 c: }" \% C6 g# e. B 189 f: I* Z5 T4 h( Q/ H& P4 e, J
19
$ D# _' G3 t/ ~, c9 v: c2 @ 206 k3 |/ v+ B, f$ U0 l
210 q( F: w* \. U. K
22! a# X: I# r: ~# D
23
1 W7 o4 ?, X( z3 H: F: x 24
4 m9 E$ d! N8 `/ V4 q 25
4 C. _. ?% R2 Y$ k" o 26* i5 } x# Y0 Q/ R
27
" z) j2 y% _! M( Z 28" ~, t& T) j Y- Q# {
29
- u* o$ S7 j `' g5 \ 30. j% L% X6 ` w# I/ u
31
3 W: R& N( c* w7 r% y/ | 32: F/ c2 }$ H: P' [
33" f* F: A6 G( i2 |" s
最后是1000分类,2048输入,分为1000个分类
7 A5 z. m/ {6 c' w9 E4 i 而我们需要将我们的任务进行调整,将1000分类改为102输出: p+ ~6 ~5 ]1 f" l# R- ]2 [
6 A% \, l# r$ m% m 6.初始化模型架构7 F$ t Z7 r8 ]& ^" Q
步骤如下:
6 S! O) n) }% j. u; }. g( F* a
3 l& O) L% A& a+ N# c( {7 D 将训练好的模型拿过来,并pre_train = True 得到他人的权重参数 C% i% \% T( i9 w2 ]
可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)$ a9 }3 a9 s/ j
无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数
0 g( G# o# q; z' L+ {$ T/ ? 官方文档链接1 R ]# Q# X3 U1 j% e0 w. M% G0 ]+ j
https://pytorch.org/vision/stable/models.html
# E- X% R3 P/ X' H+ a 8 w7 Q; I" R( K8 m. a: ~5 A
# 将他人的模型加载进来
) K4 X- ~1 r; l; ^2 a! n def initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):
- s: q; [2 {% [7 x( S% I& m # 选择适合的模型,不同的模型初始化参数不同
0 q. y8 T+ a3 Z6 z1 h, Z7 K model_ft = None
T, f+ z, k! |3 `) b) F: F input_size = 0
. }* Z5 u `0 }6 j
5 \# y: K' x. u. ?2 w if model_name == "resnet":
F; y* K Z |% a! f% i9 \ """
# W1 H* o9 ]0 F* f# q- ] Resnet152. g; c. s- ?0 J! D# l0 P, n
"""
% m$ f: ^. Z! j% Q % }' @5 f' v5 T; u; ~
# 1. 加载与训练网络& G+ B$ s# L" q8 a8 D
model_ft = models.resnet152(pretrained = use_pretrained)
# y/ B. s# A' R6 J9 n% A # 2. 是否将提取特征的模块冻住,只训练FC层' g7 O o: K3 y7 f
set_parameter_requires_grad(model_ft, feature_extract)
# R2 ]9 V: C9 \1 m/ G # 3. 获得全连接层输入特征
; \; O1 z( I" H; U( }! \ z num_frts = model_ft.fc.in_features V. X9 _ }/ j. g
# 4. 重新加载全连接层,设置输出102% G' M8 x7 n# S8 _8 L s; W
model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),$ m% a0 C' P( o p+ u2 f2 ]
nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1
' I. e- b( G! \8 [5 f input_size = 224
4 ^6 L0 M9 O, v- m. x6 @+ i - D+ m9 W- `, C _1 }. m5 w6 ]
elif model_name == "alexnet":: W5 c2 j, q c
"""4 J3 s( c9 h& a8 t% A) i
Alexnet
1 ]7 D+ d& [: Z7 s- N( c! _ """) p& j [3 F; A+ {$ m
model_ft = models.alexnet(pretrained = use_pretrained)% E R; c4 g' \. W {
set_parameter_requires_grad(model_ft, feature_extract), ?5 {0 f; q# z! k4 |1 D& q
+ S, m2 G6 F4 N5 g1 y # 将最后一个特征输出替换 序号为【6】的分类器
+ B2 { V6 q/ q2 s4 K3 z0 E num_frts = model_ft.classifier[6].in_features # 获得FC层输入
% @; {; _) }3 t, n7 p% K" i model_ft.classifier[6] = nn.Linear(num_frts, num_classes)! o7 W& z2 _& a" \+ z
input_size = 2245 P L- T/ H5 h- ?% v
C. Q7 h, U: X: I elif model_name == "vgg":
' Q% c) P- o; \: J' C* } """
! ^6 s# W- R/ P+ h$ d: z VGG11_bn0 E" |( n" S9 [( U+ Y& E
"""# q3 \, I) \7 r) S2 _
model_ft = models.vgg16(pretrained = use_pretrained)0 ?7 H+ g& U' r$ o1 j1 ]
set_parameter_requires_grad(model_ft, feature_extract)/ z5 U/ _2 Q9 }: U8 C- c
num_frts = model_ft.classifier[6].in_features1 n( L: d& ^) a$ L# \
model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
$ J7 t: [* c, z& p$ a input_size = 224
# h0 B! _$ W( L/ z0 L ; a* q5 m& h! T7 E
elif model_name == "squeezenet":0 s) i0 r& A' V5 h6 h% l
"""
5 E% Q8 B4 O" G' _ Squeezenet- N1 Y3 O( P% G* j: S
"""
# Z- |+ z4 a6 k* s7 O/ M model_ft = models.squeezenet1_0(pretrained = use_pretrained)
& }7 B# o$ c# N( q- w" o: t set_parameter_requires_grad(model_ft, feature_extract)
! w; D$ c- m8 f4 I model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1)): n) V, X' L0 O4 @
model_ft.num_classes = num_classes3 \: T( g+ B$ f6 f. m
input_size = 224/ V( p/ r8 Z& u# P& ]3 o
! ]8 F' n1 P( @
elif model_name == "densenet":8 E! o2 ?- R' B1 g8 O& d
"""
/ }8 U4 O9 K; y* x# O: I1 M Densenet
. F/ {3 u. @; R% A """; Z: ]" R$ G& U( i
model_ft = models.desenet121(pretrained = use_pretrained)
+ |2 f% U0 B& }$ E' A8 o set_parameter_requires_grad(model_ft, feature_extract)
0 j7 Z3 C& N3 a1 U8 k7 N0 w num_frts = model_ft.classifier.in_features
6 [# F8 o4 l! K/ H1 ` model_ft.classifier = nn.Linear(num_frts, num_classes)7 H7 R, c6 U4 m# y" I8 e( Z
input_size = 224" p5 \2 V: @1 H! X* l( C: L
8 O9 t. E1 A* K' ~1 z% b2 K elif model_name == "inception":! i7 y2 C1 R" ` t- d
"""
9 U Y% U/ B" H$ Y5 G" o% ` Inception V39 ?% a2 m2 _8 J k! Y3 t
"""& b) G7 e+ I* A, o
model_ft = models.inception_V(pretrained = use_pretrained), v0 _$ w* I- M l
set_parameter_requires_grad(model_ft, feature_extract)/ y7 r* f. w+ U. W, O+ [
4 t7 q& o) p l, \- c$ x$ Z, t
num_frts = model_ft.AuxLogits.fc.in_features
! Z, C0 v t1 D& Z0 J6 e3 [" r- M0 P model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)7 _) X a y( A5 {
" C+ x- L/ G, x
num_frts = model_ft.fc.in_features+ d4 ~ P# F A- n" p+ n
model_ft.fc = nn.Linear(num_frts, num_classes), h0 u$ i& W$ y
input_size = 299
; v% j' R- \" U0 j. }" z
; {( |) B# W6 \/ k6 ] else:3 L$ H4 w: m" O' A9 s. T) m
print("Invalid model name, exiting...")
s6 z: K( Z. k2 n. ~: [1 M) h exit()
7 G/ `4 Z3 z0 U / I a1 @6 K! t; ]
return model_ft, input_size1 s, h% k9 o% l! t8 L+ J' ?
+ k7 c# c% r9 T) @, x$ c& s% L
1
! D& Z7 H3 K F7 P 2. }* a" @! L! @' w) Q, O, t6 _) Q+ Q
3
# }" X9 z4 I P4 ?1 S9 `; ~/ X8 k 4
, f7 {5 C1 F. y+ R 5
7 w# r- v) Y: f9 b4 I4 K- {( H 6 `) {" H4 ^* {$ }
7
& ?0 U3 s. v h9 K( ~ [8 B 8, @6 |1 h2 i+ c
9
+ Y+ L, a* y) L0 G0 [$ }! _* r 106 i5 v1 M8 L! l ]& C. H {) L
11
+ N! P4 g/ K L& O1 O 12
9 e: [7 b5 x* Q% z 134 c6 c+ j9 v! l
14
1 K4 I1 t' c, ~* U+ g% f \/ G 155 |5 v! S8 x9 A
16
) z/ v, f* F# I1 h$ B) T( o# A# Q 17
) u3 K; L4 o- P* D X% Q3 _5 G 18' o" j( X0 e L0 s
192 M- U- {1 S- Z9 J8 P
209 J0 N8 _( q! M; h7 c+ n) _) u# O
21
) X# D5 K7 H9 |# L$ F- s 22# l; b- B- @+ Z+ `" j4 i
23! d8 Z" ?$ O0 y8 N; T4 U- n
24
6 u$ H1 h7 D: G" A4 C; E 25
# z0 Y- G8 ^. D, } \ Z# ? 269 {8 Y( S/ B/ M4 p4 p& p
27+ k0 V% l$ F! t/ B: Y; a) y8 G
286 n' o+ d9 v5 }: n: Y5 K
29% } ?4 G S0 I9 v, i/ G& S2 F) A
309 S% P# s K1 \% G$ s
319 H- g4 |! b$ x2 H& d/ D
32
1 ]2 C8 ?5 ]6 E' w& S3 \, d7 O 33
% R( t" [2 q3 X; ~9 n 34
! Q( @$ {' T' u' B/ w 35
) I: a+ @$ {- q$ T7 `/ E1 g( e" P% A 36
; u7 r7 Z& ?0 ?+ E0 G. w) a7 \ O 37
+ V7 t9 v% u# ?3 A, _; m 38: U( S& }' r4 ?; A! R
39
+ r* i" U% O* d9 Z 40
3 V: a7 e0 p7 d3 V, s' u! m 41; b: S# |2 Q. v0 a
42" A+ d+ \, }: g' r
43
8 U/ ^7 v" U: o n1 `6 { ~ 44. w7 j+ g8 Z6 ^# E8 S$ E
45
* O" P0 p3 f; s1 [3 H 46# H) ~5 ^. N, ]7 d3 N. }: ^( L
47
3 P! s, O0 j2 K* ]$ V, o+ D) U. q 48
' U. a1 w5 Y. z" j$ L9 n9 J: t1 J 494 {; q) T; u- L. F
50
, z* c( } J4 s 51
& n- M- Z: |( c) x. E, g/ @$ ]0 D 52
" F- A" L: ?+ d o2 Y; h/ y3 H% b 53+ e" A2 X0 }, X: E8 L0 h" ~2 [9 l
54
( E' w9 J& @% g% ~4 d$ e( |1 J5 B 55# h! [: h6 k# m+ `( B
56
* P: H# j* O* i8 _ 57
: F% v& m5 B: d* C# `# q8 t 58
/ N- v& u/ Q6 k. d% _; q& h+ t3 M( { 592 B$ y; i @' P2 m
60
+ f) j# x9 Z! n9 R4 s 61
1 ` }/ O, n" N* a3 j) a( E# h 62: U6 v# }7 \( ?0 V! k* J
63% T3 R( Z1 O! H
64; I# ]& j W$ [
657 n0 M( U- b1 N7 X n* ], N) V
66
4 D7 o' O+ g1 A1 g0 i 67
% C1 b9 S& {9 r& g 68. ]1 ^2 \* h6 B/ Y
69
- U, G! f3 y: P" z# f 70( K: }/ r& p1 }2 W
71( ~! p! I+ Y! y7 M% K3 O) H
72; P/ @! `3 O. a }
73: u9 C) u: {/ r8 e
74- z" p: w8 c# ~& w$ L
75
8 r, e# ?9 ]0 @0 l) e 76
& s+ ]2 a$ ~/ J4 E8 j6 b 77
9 _, N: z, r' Y4 f 788 v& |, B" ^/ V: |/ `
79- q9 @" @ z( f' a
80' G4 s# U, n2 o) [& H/ n
81
# d: u, j' v6 [, ^. b i 82& u3 m, H/ W; q2 W& b0 F6 \/ t
83: |# C L$ f D
7. 设置需要训练的参数3 j, T- d, ~4 V9 y0 ~. d- `9 E' j
# 设置模型名字、输出分类数
+ y" u& B6 y$ l& ]' l) C/ | model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)
& i- ^" C. F+ ~$ z! [ ( I0 H# u2 ]+ r- ^- m1 c }
# GPU 计算
% x* `* m0 w: R p6 K6 T8 ?7 ` model_ft = model_ft.to(device)0 v( m9 R1 g& }5 t5 v5 k
2 ]4 b, K* Z" u$ {, e/ P # 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取8 `9 a+ i' i6 {% k) E. ?" `" ~
filename = 'checkpoint.pth'5 b# J8 K' r$ G# [: d
$ A7 n$ u0 Q# O; a( [ # 是否训练所有层7 N4 i9 z, l# U8 {7 e
params_to_update = model_ft.parameters()! a& `, u* B3 G
# 打印出需要训练的层' @3 b" Q, Q8 X" k5 G
print("Params to learn:")
0 F+ F7 R+ o+ g' P/ K- a if feature_extract:7 S$ o* A! Y2 a( d2 F3 p1 j# p
params_to_update = []/ Q/ N' Q7 \' k* I5 g( I
for name, param in model_ft.named_parameters():( c1 U7 c; Z2 ]! z0 K2 |
if param.requires_grad == True:
: n0 e3 _7 {/ d; g params_to_update.append(param)2 l' B$ j4 z5 w0 e
print("\t", name)5 Y' W% x: B7 i5 n; I
else:" I0 j/ L U, P: w
for name, param in model_ft.named_parameters():
/ L) _: \" _9 n- v3 L& d. X if param.requires_grad ==True: @2 S6 O2 D) v
print("\t", name)
( C6 [$ G" f' E# E . I2 p/ O5 J# J& v3 B
1 x$ `1 R1 v2 j& A6 }
2
0 l$ ^) v* S8 L7 }7 m7 m 3
* ^/ X7 k( v: e" ]3 ]& o 4+ ~) M8 x% d; G
5
1 m+ M% p. r4 O- Q/ V$ A7 q 6& ^5 n X. e7 `2 ?. X
7# [- a) F1 \; ]4 O$ q( B8 O! |( g k
8
$ l) U+ q c$ m- x$ ?8 ? 9% B* q' n7 ~+ Q7 b
10
: R5 P( t; m, h8 Z5 l' f3 e 11
& b8 q$ U1 b4 r! R 12
) M. L/ N% b E; W 13
+ G1 [0 X' o; k& }. [5 q 14
& z9 ~5 \. j- W1 h 15
3 K" m" c* N# d0 k 16
! [, x* j' K4 o% D' O' H 17
0 m4 O' N* Y' N# f 18& j) ]4 ]2 i& s- r% q: N) H9 k: S3 h
19
( I4 l, Z8 Z) t Z, a, d( l$ r6 O 20" b% w( \9 S+ ?5 U( o; x
21
& Z: L- [7 G# Z 22
7 w. ~! ~# v r& I: v4 @0 n* Z+ }6 F 23
; g( M7 [8 r$ t7 D Params to learn:
4 Q: D/ G) j8 m- R fc.0.weight' X$ `5 C4 u+ b5 Y
fc.0.bias2 t5 X/ y, h, T1 u, W0 F; ?5 m* [4 M5 H3 R
16 p+ Z1 V7 z5 E( k" x4 p1 l: r
23 X7 L* {( i' y. f+ f' V5 D
3
2 ^! n8 P; r$ d7 j. x 7. 训练与预测
% A6 m1 u+ g) t 7.1 优化器设置
! H$ U) @, ^' `3 j1 F4 J # 优化器设置) f, Y N& L5 ?8 ?' M8 l' o
optimizer_ft = optim.Adam(params_to_update, lr = 1e-2), z2 R2 b1 Y/ j" Y0 T3 I7 k
# 学习率衰减策略
( l& R t2 a* B2 @* a scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)* ?7 f0 S* L( ]- g( u
# 学习率每7个epoch衰减为原来的1/10
9 b) x$ ?! x/ g! ^: {0 Y1 e4 O # 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算, d: I) }9 |9 R* L: ?
) x7 u( q/ n* k
criterion = nn.NLLLoss()' t7 W3 S4 l9 C' K+ t/ u
1
+ `8 |7 t7 N* M, I& f" Q 2; Q9 e% H; `7 Y4 P$ f
34 [; b( z' ]3 t$ f" m2 S! }
4% w' y& u4 ]0 S8 e
5
: l; b5 r- Y- k6 i3 S$ e 6+ Z- H1 r) q3 F4 t- ]; m
7 ~, B# H2 Q. J6 p
8+ P' f7 E1 D; b0 l2 l6 n
# 定义训练函数
* {2 _" o& D# \4 z) Z1 S #is_inception:要不要用其他的网络+ o$ f! Q, ?) {9 @. F% Y; q
def train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):* |$ Z4 l* i( H9 @
since = time.time()
- Y4 g! w( C8 F R7 } #保存最好的准确率6 Z$ b! X8 ^0 T) M( [! J
best_acc = 0+ x- g$ p5 l, U9 U0 w2 ?, j
"""
( l4 Q: G: h9 w1 n, u9 i checkpoint = torch.load(filename)( i0 r0 U# s1 x7 |: H# h) f
best_acc = checkpoint['best_acc']% ^3 Q* t- M% l0 I5 Z7 T
model.load_state_dict(checkpoint['state_dict'])
7 \* y1 f! p! C" ~8 F optimizer.load_state_dict(checkpoint['optimizer'])
: |) D9 l" `/ a Z6 k+ q' ?5 v model.class_to_idx = checkpoint['mapping']
/ j6 z7 y8 o t) n9 y """! E. ?6 F2 P' ?) M- G
#指定用GPU还是CPU
( E W, v' H( I/ |5 k6 T model.to(device)
, \* \- l; Q7 F4 J" o( L" Z* M7 W& c #下面是为展示做的: a, a) U0 R3 j0 h* N
val_acc_history = []0 S/ e% n- q: ~" m' {
train_acc_history = []
+ j& D. j# ]% \7 j) b( j/ X, U% z( | train_losses = []
7 O% G$ J- }0 H) E valid_losses = [] ]6 Y, A1 H; o# P, Y
LRs = [optimizer.param_groups[0]['lr']]8 d5 V7 V$ o7 w" H- e! ^0 {
#最好的一次存下来5 X6 Q/ K. V% c3 Q* h
best_model_wts = copy.deepcopy(model.state_dict())
2 j& u2 ^, @2 G' G 7 i& Z7 Q: o8 Z" u( |7 O
for epoch in range(num_epochs):
# K! _( H+ \$ o) b. N print('Epoch {}/{}'.format(epoch, num_epochs - 1))
* N8 T2 ^1 M7 N4 }3 h" M2 s print('-' * 10), k8 N( s. j/ [3 l
/ }8 y5 Z# k9 Y8 T6 c # 训练和验证$ X7 u4 B2 R9 V9 u4 U1 a
for phase in ['train', 'valid']:
. Q3 i" U( y5 {3 O1 @# w2 V if phase == 'train':
" B, ^+ Z( v& w+ V7 L0 V model.train() # 训练) @% u6 j8 X6 T' O" L
else:
7 K1 o* t0 J' h6 U+ e3 \ model.eval() # 验证- h7 R6 u2 G7 I1 M. }$ }
( [, p0 L" G4 N U: P5 s/ J
running_loss = 0.09 G ]: T: |9 c3 a
running_corrects = 0% P9 H0 S3 e; e9 t8 D
6 M) N) o7 j5 u% b6 A; I$ N # 把数据都取个遍+ s% `6 e; e! \. M' q( Z8 B
for inputs, labels in dataloaders[phase]:
- J: J( r2 p8 B H( c! r. o% H #下面是将inputs,labels传到GPU
: Y, z; q h! W2 i, z& p: c1 z inputs = inputs.to(device)3 V6 l6 U, t* p, I2 P
labels = labels.to(device)
6 t* D! {, B% l
6 n/ Q+ w/ n6 D. k( X! F9 X # 清零5 X) R( h2 R0 R* q8 L& m* Q9 P
optimizer.zero_grad()
( R. i% i7 \ t' T: Q4 r. k! n7 F # 只有训练的时候计算和更新梯度
7 I+ @& d2 h5 `) k6 M0 ^/ C with torch.set_grad_enabled(phase == 'train'):8 H: r- p- i# d2 I; h& M% O$ a
#if这面不需要计算,可忽略+ m6 ~) S+ [- B0 h1 }( _2 ~$ N$ z
if is_inception and phase == 'train':9 R+ ?* K4 w$ {3 T" l
outputs, aux_outputs = model(inputs)# a* }& S" ]. G, {
loss1 = criterion(outputs, labels)
: j X/ k1 \2 o loss2 = criterion(aux_outputs, labels)! y( ]# |) ?; y$ w; U0 {
loss = loss1 + 0.4*loss2
* v0 ]8 }6 c. V7 ]& @) e. p. r else:#resnet执行的是这里
" v2 s+ }! L$ H4 C5 m! s, g outputs = model(inputs)( T, T( c5 X3 W# x3 J
loss = criterion(outputs, labels)
+ u. d. f$ S* @ 2 U2 j7 Y/ U1 G7 w9 H# V% r
#概率最大的返回preds
4 Y5 A0 s% X G1 }- n$ U _, preds = torch.max(outputs, 1)
) B1 V( p/ y! c' L* E. x* J. w
* u3 D8 a* L: R5 ] # 训练阶段更新权重
( K: ], T# S1 l+ d0 T" N8 Q if phase == 'train':
% e7 o7 `3 L8 `( k) P# \& p( ^ loss.backward()
/ _. ~$ ^( E! B, \2 j. h# [6 q optimizer.step()
" q8 Y% G- ^& R4 `5 k n; {
, n( d6 [% ]6 _: `' ^ # 计算损失/ x( k1 r4 F( c; ]3 Y: i
running_loss += loss.item() * inputs.size(0): T) k/ c: `$ v& \) L
running_corrects += torch.sum(preds == labels.data)) E. T e: |2 q( Y+ s, t
/ m: Z) d( o2 e( Q2 Q( s #打印操作
t1 l" i* {. p3 G. h- T epoch_loss = running_loss / len(dataloaders[phase].dataset)
! p5 r; {+ z! p, Z' K epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)2 E4 x+ a1 H& z# Q7 `9 K
7 t8 n0 c b8 w& n , M2 ^: {( ~2 w* H
time_elapsed = time.time() - since
1 Y# j/ Y" [ w& ^: i; G print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
+ C- r! w& {; {) d6 I1 b+ Z( Z print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))& `: L2 z2 f. K. E F1 v, l/ ]
8 t& a) C T3 W7 h: W
+ R6 F& M; j Y+ f0 N0 ^ # 得到最好那次的模型
; g8 ]7 p- a, ^3 O! k if phase == 'valid' and epoch_acc > best_acc:! r. P: u9 s4 X' X1 g
best_acc = epoch_acc
6 M1 U- w/ p4 V, Z" c #模型保存
" W8 {! ~! n& b; c9 @ best_model_wts = copy.deepcopy(model.state_dict())
8 `* \0 \8 X% E9 | state = {: f1 ^5 r2 L: c p5 u) q! T! q
#tate_dict变量存放训练过程中需要学习的权重和偏执系数
: y* H6 e# f% x5 B3 s5 ~ 'state_dict': model.state_dict(),; ?7 g2 i! m$ P$ z( F* g' h
'best_acc': best_acc,
3 B" P L9 y( C" E# C L H+ s& o 'optimizer' : optimizer.state_dict(),# }5 {3 c/ G9 j5 d2 H' j
}- k4 J+ T3 V0 {* [; D0 Q
torch.save(state, filename)- c2 _: l1 o5 J' D' T* T6 N
if phase == 'valid':" _; K/ e( a# H) Q! i
val_acc_history.append(epoch_acc)# k6 D4 p% l* _7 L" c' [2 Q/ v' {
valid_losses.append(epoch_loss), W: P, M& f: I' Y
scheduler.step(epoch_loss)6 r. q: A. S8 }- X0 c/ u c
if phase == 'train':0 _( x @. D* u. n4 P' N
train_acc_history.append(epoch_acc)
2 ^4 D6 I1 e4 i9 [/ ~9 ?8 ~2 v3 u2 N train_losses.append(epoch_loss)" Y7 D* v' q; Q: l2 |2 T
3 d! v! o" }: w3 o
print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))
3 o: I! X8 R- ~* w# N LRs.append(optimizer.param_groups[0]['lr'])
1 ]; G" a, R! a print()
. T. R$ S w- [8 u g1 e |" q4 ~, X# O' u! I: `4 {
time_elapsed = time.time() - since/ P; z$ e" r# G: i/ Q& x
print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
6 _% f) P2 Y0 l, B" E0 _5 P0 c% V print('Best val Acc: {:4f}'.format(best_acc))2 @% \, | ~% t: z) i
* P: G4 o {# L/ p9 C" S1 K # 保存训练完后用最好的一次当做模型最终的结果
. q6 s0 s. a' t2 ?3 A model.load_state_dict(best_model_wts)4 c, t9 d9 T* d
return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs 9 C; G" b4 t+ j2 G3 T$ k9 V
: A5 [4 G0 s# g: u& R+ E
) P- A6 _* x+ q 1
* E0 e5 q+ V5 t; p9 { 2
6 g r# q- N" F0 c4 F) z% `. x4 Y 3+ m r/ [: T( y `2 [4 u! ?
41 f0 u8 ?& S) R+ j% G2 S
5& A; ]( r! d0 [' u) `4 z8 V- d
6
* N# v. V+ `! h8 \; B" z8 ?8 [ 7
' j5 `2 i0 I/ l 87 i$ X+ q( h8 @# O
9
/ {3 L) ~- \# W+ @8 o 10) U9 e6 X5 b/ E5 L- P
112 o7 N; r& e' s; _+ ^
12, A# h, d2 f0 y" ?( A7 [
13& c3 O$ M% K! r( x0 s6 E
14
5 \* n; {5 c0 N5 P3 Q5 v' I3 a 15
; Q% w, T2 U- H4 u: \, Q$ O s 16
z/ T4 C9 j2 c( A+ }, p9 L 17
* a+ y8 q" L- z3 U 18, Y( _* a1 ^( G0 M+ M
19
& T8 R* V+ n' l0 L 205 I. B5 E0 W! |& e+ h. _
21; _$ R5 G7 Z3 G& Z0 r5 n
228 {2 x ~) j& n$ a4 ~0 @' y9 x( ^
23) e: W; ]" b! k1 B9 r& s
24( x& l u9 u, c7 b+ O+ U. j
25, L' z! s# l" x
26
6 r' n7 p1 K6 p+ A7 ~ 27
6 J! ]. E3 O* m3 x" r& D* C) ^) y 28- i- K k, o0 S! X
29
/ n& ?. U% k1 y0 v! G8 b2 h 30 ]' D7 g: q2 M) u; ?) o
31. @5 j" r" X* Y& {9 `3 {( _
32
+ o* U+ b( E+ s/ t4 z 33
+ U. B: A. o0 z; m: k( g 34
- R/ G; G# \) @8 G( u 35
4 ~2 f6 S! z( J3 G. U7 l 36+ }' p: E. i2 g/ M6 I
37
; _& K- d* u" O* A4 T- g" M 38
, e* n e8 Z$ W/ ^& l& y/ }% A- t, I 39* |1 |1 X( i( L+ | `2 W8 {& g6 T
40
/ f9 H3 d1 S6 D* @* a 41
7 d. z2 r9 f& d' c6 y; c$ L, u: ?, C2 z 42
7 p( E; n6 V7 H& @, w 43% Z- k, B6 y+ d+ i. H$ j
44/ v! x# g5 A- K3 ~$ O$ Y1 ^- r s6 B3 o
45( ^5 e- @; A1 k( Q! J6 [
46
9 t4 D+ U1 x0 Y& g2 p 47
0 i) D ?/ e U7 w! O4 J$ ]) G2 F 48
! B$ l% e& P# u: L7 {$ E 49' N1 R8 N6 V$ _% M
509 W) s3 b, Y6 `5 Q; ]8 n. D5 L
510 B, B( v! N7 x0 r- q$ L _8 v
52
P( [# c$ [1 u0 x 534 z. L5 D* z! p' u
54; ?* L" J7 m* T1 a
55
" T U$ X. s6 [2 L# p 56
4 [- @8 D6 I5 N0 n* z 574 {5 X. ?1 M; Y8 e% M. R
58
% n {% i+ S' U0 y 59" }7 O6 [, j- K2 N2 o' J
605 M2 G* T7 R2 m' r$ W0 X% D
61
p$ @: {6 f% [ 62# L4 I9 V% x3 h. B+ c# Q. o' ] e
639 W0 j$ n }: ^. k) q
64
' o5 l3 G% o) V" Q' P# X0 p 65
% Z2 z4 M* q) v3 L% { 66
# u3 `& @3 c' H' ?. x- W- E/ m# c) P' F 67
: M( Q* ~6 c% o 68
$ g7 N& ]7 |8 v9 k 69
0 z) _- D( e2 J$ j$ f 702 _* X7 U* d2 k7 g1 B" C; J$ Q/ i
71
7 w* f+ _4 Z n. ]4 j4 ?8 V$ I& R 72
3 ` a# a% O! \3 z; n: G 73$ l& A7 z1 c$ U+ E W. y9 J
74
( U) N: B- n$ r0 k* Q* Q4 V& m 75
2 [5 V- L) m8 ~3 y+ ~2 [6 w0 y 767 ]7 e; F0 J( e# |! b( r
77
" ?$ V' Q3 Z% P9 T' O( u 78. |6 Y" M" @% k0 V
79! e( q8 O/ e) S' b
80( x+ F$ U4 h6 t! O0 x* g0 S' S
81: I1 i* V3 r3 Q, g' F# D3 Q
829 a9 M) _3 T0 ]' i' v' X; j
83; G$ I( h3 l7 p. c; s
84! J* l% R& G. }0 m4 }2 A0 f; _
85
' i' t, @& t! R+ n 86
$ t! o" r4 c8 e4 r0 u' V* d 870 r- D0 Y. R5 @' L' `4 u) X
88
+ I+ u' \; M3 X* l6 c" | 89
+ ?; i8 l$ y! U2 W4 y' [8 N& R 900 a# B, n8 U. F
91* w% J' }( m" {# `& q
92( q/ T! t! Q6 B, ~
93# T$ |! J! j: S! Y
94
8 j6 c0 N3 o) Q) @, ?" I# ] 953 I# @6 ]0 `) f, s/ i- H* {
96- Y5 }/ O Z/ F; b! _- t- Q
97( Y0 t P9 I0 A7 ~" I" {
98- Y T; p9 ~0 x( ^) L% L5 r g; M
99
7 G4 x' X6 c0 r( } 100. A& V; S3 T7 R" m2 m, u
101
5 N c) p9 y& }4 X 102
; M- R1 Y0 |' Q 103+ Z$ O! d9 F8 T& M* L. w
104% s* |3 E$ V m$ t6 ~. z
105
4 i; r. E7 B1 `7 g& g 106& b$ \4 E& a- Q, K
107) @4 k* {( E+ ^- v+ T2 P7 q) g
108
' j: i) W+ e# k/ D 109
- h. V5 J5 t7 L; Y6 H 110
3 V5 j+ g( Y# e 111: ?8 y0 I' _% C. ?$ @
112* L+ ~; ?3 }) U! W2 R! p
7.2 开始训练模型
1 T$ ^4 I2 P% r& e/ ^; c 我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次
8 u& n5 u! o4 H" ~ ' _* p6 b$ u+ W" k! Q4 Q" J
#若太慢,把epoch调低,迭代50次可能好些
: g% Z# `9 p8 v* k. }6 h8 Z #训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合
+ [8 Q3 s2 Y" V( {; [3 V: Q 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"))
' B/ v1 T' O- V0 E- l4 q / U' r7 N9 B, g x5 c3 S& S
1) K5 ^- Y' _; g$ O6 @ D; y f
28 h x% a ^. s
35 k/ `5 E5 ` h7 V+ v8 J
4: R; ^ X; ]4 _9 W+ L6 V7 h
Epoch 0/4$ s( P3 c* E' Y& N- a
----------
- K. j5 U* G' z4 a Time elapsed 29m 41s
1 O2 K, T, {% N5 h0 y9 G train Loss: 10.4774 Acc: 0.3147
; B& q% V9 D8 L' { Time elapsed 32m 54s" \1 x3 |$ Q# ^* o S, @
valid Loss: 8.2902 Acc: 0.4719& q, C/ g, ]9 q% e$ \" m. T& W
Optimizer learning rate : 0.0010000
9 t9 |* o2 Q# `/ z2 l: I - B( t1 ?- t4 G6 D, g
Epoch 1/4 Z! ~( e5 L$ u
----------
* u9 y4 |! M- \' J9 I& T Time elapsed 60m 11s: i8 I, p& b$ K9 C; X9 L
train Loss: 2.3126 Acc: 0.7053, O3 G2 V, n# G3 J* F4 W/ C$ I! p
Time elapsed 63m 16s) O- f+ O& t- I/ \+ V9 y3 Z% j
valid Loss: 3.2325 Acc: 0.6626. e( } D5 I" D
Optimizer learning rate : 0.0100000" u6 |- P U A/ N; \$ x
# _5 d( w5 |2 |; O5 p" m
Epoch 2/4
/ d4 S1 `( S! | ----------
# Y; {! S6 o6 v0 [ Time elapsed 90m 58s
' b! F3 X3 h9 Y6 l% k0 Z7 s# K train Loss: 9.9720 Acc: 0.4734
0 t1 \- [* S/ I& ]; T3 Y; @ Time elapsed 94m 4s
( Y8 U9 w% P: s' V! f valid Loss: 14.0426 Acc: 0.4413
2 w% z4 \: O8 w. g; v7 w Optimizer learning rate : 0.00010005 }2 P1 T! w3 n7 u J2 r2 s B( Q
: a: h" Z" T7 K1 V! [) `( h3 b Epoch 3/44 p' g' T/ s a& y3 v
----------
& {% }- X. W3 N: o$ X9 F Time elapsed 132m 49s
" _* u3 o3 k+ z( k0 t% D0 j train Loss: 5.4290 Acc: 0.6548
: k! U4 |/ s5 V% z) m Time elapsed 138m 49s! H' ~/ y7 h1 D q* P4 b( C* b
valid Loss: 6.4208 Acc: 0.6027
* a- i; `$ f" K H0 P+ Q9 t% K Optimizer learning rate : 0.0100000
4 Z% V* P3 c( q5 ~ ' O+ p1 f2 z( x1 S K3 ~
Epoch 4/4. k+ W4 _- r% [* h5 v3 ^9 V
----------/ _; @$ Q1 `6 w- y' O9 U7 l& a8 c
Time elapsed 195m 56s% ~ X7 u3 ~5 ?, @0 O+ ^- A5 g7 _
train Loss: 8.8911 Acc: 0.5519. U6 u% G! ~0 Q4 P/ S
Time elapsed 199m 16s- W {9 Y% e. r- ~1 f3 t
valid Loss: 13.2221 Acc: 0.49143 |6 |9 `) J" u! V
Optimizer learning rate : 0.0010000+ S! W6 b' G g8 m# B) q" O
8 ]/ {4 L; d4 @# J# {% f* s4 H Training complete in 199m 16s! D+ }3 j+ [- \! M! N8 h$ v
Best val Acc: 0.662592$ ?- c: A8 @/ f! y: V6 X7 L- r
5 G9 F7 P0 n# z
1) M2 ~& _& L: g* w: U" _
28 Z0 \+ [) `$ O$ t- {& A+ _+ W Y+ y
3
1 D/ X; V% \/ [$ q 4' R( f+ [/ [+ g+ K
5
1 U9 Z* i4 f" ^9 S& e( ] 6
# @5 {) a3 y$ R 7. Y) Q1 `( {, O! ]5 ~7 X+ N4 v
86 J0 T! V2 Z" O7 d, w, ^
98 T# {6 Z5 j! m0 y$ F, x, Z
10: w, a# H' }6 j0 _8 \. Z
112 U" R/ g! I0 y, ]0 N( j
120 P) h, M0 j1 L. Q4 T2 [5 ^" n
13
. `; E4 Y7 g% k# V! k 14
B b" {, A( E" q 15
+ r5 D$ h) y8 R, v" j7 P* U H7 A) A 16$ h! L' s% F- {4 T# Y+ h2 l
17
/ X) w: y& _* ^8 w 18
6 t7 h \2 L! q8 c+ ~ 196 ^+ F, p6 [" g, Z& L& ~
202 S: K8 v! v5 q- ]
21
$ S) c. U4 Y& W6 [; l' l; P; R 22
* ?* K* g8 J% a- l; i4 A2 c 23
; `. s$ p" q3 f: ^7 l1 W 24
. x; k) l( \4 B8 H5 u 25. M6 a) \ S) S. C" p
26+ M; x: u! s4 d7 g3 M9 K& z5 G
278 u" U, e( o1 ?) R
28! x; |9 o/ ?7 G n, v
29
8 j. l/ D2 P4 ^- B 30
4 Y3 n V; A) J5 Z- A' j& F 31
$ E8 O! l3 t7 ]9 l; h8 I H. c$ I 32
5 h0 P$ V" H1 h. g, P7 v. u1 ~- T 338 K3 G% O2 [& a; G; U
34& [/ X# v. K' x0 L% G0 x
35
7 E& Y) j8 ]- C! H* M2 P 36
) S( t2 i3 X( `9 l. L 37# v+ o }- ~8 z
380 e5 x6 e" g1 Y# @
397 y2 b9 H1 }5 k8 w; q% Q
40
$ m( \$ e( e, k 410 O4 O! M, U l; z* {
42 J. n3 O8 H9 f; c b1 Y/ b Y/ g* }
7.3 训练所有层
0 t7 s. i4 ~8 J q6 T; ^. p # 将全部网络解锁进行训练
* J; Z( g& i9 W1 y for param in model_ft.parameters():9 F0 g: L9 Z( Q1 x0 d( v/ C
param.requires_grad = True3 l# G( A" Q& C3 d$ U* A% m
2 b( o% W! ^8 X6 x' w. G # 再继续训练所有的参数,学习率调小一点\' K, w6 N9 {" W j$ V
optimizer = optim.Adam(params_to_update, lr = 1e-4)' U5 z8 C/ ^! C1 S2 k
scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)
. I& W3 f& O `- | 8 S8 N Z9 T& \
# 损失函数, B; c2 s( ]) _+ L
criterion = nn.NLLLoss()
7 _* x' i- O5 V& R) b2 a3 f 1
" t$ A2 e- h1 Q) q/ V 21 U& a$ _7 E1 r0 w5 m* K
3
0 r# I1 f) h* A' L4 F 4
* R/ o4 `+ }/ C' J 5) `: I0 y# o) ^5 F5 s& f* Q
6
. ?* |( n+ E6 A3 B9 J/ [! ? 7: G/ g5 |! ~+ E4 b% H' t" x( j" s
8
) F8 J# Q; T/ h1 y- J 9* T1 h: [6 ^8 c* d6 u
108 f; `) B/ j4 |2 d& e
# 加载保存的参数
/ u: \. @9 j0 U2 t # 并在原有的模型基础上继续训练7 A$ C% @, j, U; J* k+ Y
# 下面保存的是刚刚训练效果较好的路径$ b5 Y- c! K/ l: |8 E. _
checkpoint = torch.load(filename)/ p4 k( v3 N* t* I' ?& p# Q1 s
best_acc = checkpoint['best_acc']
: _% K h! w7 ^, W' y" H0 d m7 o model_ft.load_state_dict(checkpoint['state_dict'])
, I' Y' T8 h' v/ D optimizer.load_state_dict(checkpoint['optimizer'])
i5 g: m0 _. Q0 `. i0 B 1; {" y! \4 X# H4 R3 {: t
26 p+ ^8 f$ L3 G. i6 b( o
3: C- Z; L; S3 V% u+ q
4
: [$ y; ?: ]) z+ e: |9 Z 5# v P; w ^& D5 ]% ~" T
6
" t% \$ J* x2 K' Q" V 70 u9 W2 k4 x* ?& [$ L' k! g
开始训练
! I% z& @/ m$ J" N9 s2 D( O 注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考; L9 C5 s; p) G
) A* ]8 d; f" y. \ 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"))3 F. m- j& K: _ ]
1# `3 Y* |- R& Y' D& m8 q9 D* D8 k U6 j
Epoch 0/1. e; \4 O; \% v8 I
----------" l' w0 S( e# S; c* m
Time elapsed 35m 22s
3 D4 Z) [3 V1 T7 l train Loss: 1.7636 Acc: 0.7346% P. A( X. }" b P! T# _' l- u6 B
Time elapsed 38m 42s
. t1 [: H% `8 X- y3 o valid Loss: 3.6377 Acc: 0.6455/ v) |; E. Y) s3 |
Optimizer learning rate : 0.0010000; }# J4 Z3 O6 b6 }- ]2 X
: A" U& E1 E/ R8 U- n! }$ Q
Epoch 1/1
: P- S8 |1 J# T ----------% K) ^/ H7 B' v& w( ?5 c, O9 H' ~
Time elapsed 82m 59s' Y" i: S) u0 \
train Loss: 1.7543 Acc: 0.7340& C J/ z% I9 d6 v2 _
Time elapsed 86m 11s8 `( h2 l! H; B: ]- j
valid Loss: 3.8275 Acc: 0.6137
/ w! m, k5 g: S) \' o& X Optimizer learning rate : 0.0010000
: X$ p" Y4 j7 r& l, S+ F8 ], y- m
# L- V/ j5 d+ F% Y1 U; U7 H) e Training complete in 86m 11s
, @3 S9 k% Q$ x% l* d Best val Acc: 0.6454776 _/ Y5 q6 \( P4 J
) p" T7 e) H$ v# D! T 1
/ R7 N9 ^0 |, B 2
7 j, W* F8 u3 {9 X3 J4 C9 q; C! H1 S 3" U* }4 x7 V3 k/ `% Q/ G& j& y, Q3 _! }& L
4
6 W& M* _/ k/ i# ` 56 g* E! N3 C6 n$ a' G
6
# h( Q4 |+ N0 T2 _1 Z9 N; i/ j* H 7
, T' u; N+ g+ y 85 K4 w+ X G8 Q4 O
9
% V8 Y2 ^& ~: D6 F 10
: v+ Q# n3 M- |6 c! M 110 ~- @: m6 V. c" {1 E+ [2 u
12
+ |0 R m; a: W9 i( j# n 138 q& h, v6 v( z6 Z' M7 y9 N
14
) E$ c5 C( V) o8 j# q 15
2 |# X: v- q: ^3 s3 i. v. A6 } 16
5 ~) V) x+ K( M' \8 P& Q0 s 179 z, k+ j5 x. [
18. ]$ a' {! Q/ L7 A
8. 加载已经训练的模型
2 O/ i2 i4 S3 `; f q 相当于做一次简单的前向传播(逻辑推理),不用更新参数) y- p: G0 r3 X% A T0 Z
9 A" J0 O! ~* M9 B7 z# g6 x model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
' t* h# E% z3 Q6 v% r5 C4 f
6 m3 @1 f1 q) Y y) c" I" i0 T # GPU 模式! D# P) y( [- h4 }
model_ft = model_ft.to(device) # 扔到GPU中/ |' O( p. v: W5 i) [( F
- n* m4 x2 w+ T7 [% z5 d9 R # 保存文件的名字* j; ]* o' d7 b* N
filename='checkpoint.pth'+ p5 m& @( c* ^ |. R
" o* {5 G8 P* H6 c7 Z& C6 s5 G
# 加载模型
( m# t8 L: ^! p/ r5 G& c' M checkpoint = torch.load(filename)3 }, j& I3 J6 H4 b
best_acc = checkpoint['best_acc']. a& a6 g. m6 \- S0 U) S0 r
model_ft.load_state_dict(checkpoint['state_dict'])
! Q9 J2 g% q# [+ h 1
Q4 j4 Z" l* q 2& O8 n9 \: }! A" e. ]) ~8 [
3
' v* b1 `8 O, W6 e0 ~ 4: G8 F7 ]* u$ ^+ P: @7 y8 E
5
1 ~0 r. g# g9 G 6
2 U8 g4 K' q7 e) `/ g- T; T 7
# U7 \- R% Z7 {7 A6 u 8
0 X2 s' O2 s$ |% _( ` 9
1 r) T) z0 K6 j: v 10
6 }* |. |, b7 A& \' C- E2 ? 11
% @! h& P F* j. }! c# r. @ 12
; d3 @# N% j5 B <All keys matched successfully>
: X; K$ u- A, A/ F3 V- ^* o/ ?' }: F 1
3 Z* N" F$ j" a# e @/ p' q def process_image(image_path):3 B$ R9 Y% m/ J
# 读取测试集数据
! _% y7 k, A0 u! o' r, e( \ img = Image.open(image_path)( U: u" x# S+ e4 }' w7 w5 I8 d9 l
# Resize, thumbnail方法只能进行比例缩小,所以进行判断
, T; x' g3 n+ Q # 与Resize不同' t7 @- }7 E; ?# C/ b+ B
# resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小& i3 R; n- L, v6 K2 c. h# \
# 而且对象调用方法会直接改变其大小,返回None
4 l9 {) i) ~2 q1 x+ P) ~ if img.size[0] > img.size[1]:
% m0 ]' y0 J+ e6 E: r img.thumbnail((10000, 256))
$ w4 [+ J( i" N9 ~ else:; J1 S: m4 t7 P: j$ v& D- d, [' V' c
img.thumbnail((256, 10000))2 T) [* j* z! o, z7 Z* o
: \" ?) y: N5 t' _$ e # crop操作, 将图像再次裁剪为 224 * 224
" {& j9 V% e3 E+ A+ @ left_margin = (img.width - 224) / 2 # 取中间的部分# A2 Y7 l( d* q% u
bottom_margin = (img.height - 224) / 2
1 n3 o z/ N0 I- I# y+ m right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度) m* W+ {( U' v, S4 _
top_margin = bottom_margin + 224
; M$ \+ K1 H, m# ^4 R# ~6 Z 9 L6 V! u f% E' \8 _
img = img.crop((left_margin, bottom_margin, right_margin, top_margin))
, E+ n4 g3 t: Q7 S - X1 P) q6 \& N' n
# 相同预处理的方法) O3 H5 e+ o8 \9 n
# 归一化2 ]( o- M0 X- a I6 M) y5 m! c
img = np.array(img) / 2554 `% a! z$ s5 M4 U8 s4 S( N
mean = np.array([0.485, 0.456, 0.406])1 i: I. r( h& v$ E8 H6 O
std = np.array([0.229, 0.224, 0.225])9 C' j1 ?( l$ |9 F
img = (img - mean) / std! p7 h9 J# M1 S4 @0 L# Q, _
" l4 t& O* W- X3 p/ R, ~, ~* d3 r # 注意颜色通道和位置
( ~& Z; L' G" @: U img = img.transpose((2, 0, 1))) _) v" a ^2 c! o+ ?# d: I
' {- @ |; F$ I3 P6 E8 Q. B return img* n4 [ `' j# n
& a& L$ s4 Z2 ^
def imshow(image, ax = None, title = None):/ N% b) o" Q* ]. N
"""展示数据"""* M9 v, Q4 j1 w! j0 C- M
if ax is None:& v: Q. H, _7 R( ?; F
fig, ax = plt.subplots()
3 b5 A2 ]; `! {
" _- H" j3 ~: v$ P0 X/ L # 颜色通道进行还原
: j; |& A9 m3 G4 o3 _$ A0 e image = np.array(image).transpose((1, 2, 0))
7 P% F( f/ _9 G9 V/ A
4 r+ b, m B8 ?4 [' W: R- k: L # 预处理还原' Y7 |7 a) Z! W5 a, `
mean = np.array([0.485, 0.456, 0.406])
4 ^& E0 Y2 s/ W& S* q std = np.array([0.229, 0.224, 0.225])
; q; U( U. E1 Q( R0 t" k; _ y! \ image = std * image + mean
( M& z4 s3 I; G# g; h image = np.clip(image, 0, 1)
( x4 s7 H$ o* p& K& Q' X6 Q - |2 K# E) I2 ]! v' Z0 _
ax.imshow(image)
, a W8 o3 Q- s2 o ax.set_title(title)
4 {9 Z6 p0 m- w G- N. h. A9 T3 ^3 E
return ax X0 k6 a' j8 K. C
7 d# z- f' V6 J* `
image_path = r'./flower_data/valid/3/image_06621.jpg'( v# f2 W1 D, Y* L8 D& @% L
img = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理
6 f* {* n( H/ @9 |3 b imshow(img)
' C3 A; k; R% m & v( G/ y# K6 B0 P/ i
15 v z" }( K( p# d
23 ?. H3 p9 ], C% X
3
5 g. f, ~: u2 v8 t 40 K( b9 ]1 w' T4 s7 R: \
5% l* x7 i( {/ w2 K) \: R1 u
6
5 N. a7 k, v9 _5 M/ g 7
* H! L. O$ B! E0 Q& Z+ Z 8: I" u& [9 w8 Z
9
) {0 T0 ~* |/ s7 k, a/ r 10! A. B/ n, l/ L7 @1 P! R
11
v' q, k* Z* O1 m& Z2 r/ X 12
; \* u) I* x2 K' Z1 X 13
5 T& L4 s) e7 @' @ 14
: i- l) }& T' U0 { 15
5 v# }. c$ h: ? 16
) x8 i. b- \6 O- F* q, ? 176 ? s( T; ~) j! b! e; M
183 v9 S% i: f3 h% p9 z- q4 h% w
19
^' |' G! o2 y1 E8 e g 20# s" u; o3 s3 f) ~( R
21 a# }& W9 A* Q) ^" F5 f. K
22
0 W, H) I W7 J7 f 23
0 u* K7 A) g7 j- f7 N, u 24
9 t; x3 e% S7 A) M: c* D& N 25+ x( j; l' d9 ]6 J( o( k
26
8 E ^" L% U: z 27
' x0 }/ c. u3 D* [ 28
+ t3 o7 m3 O6 S: f 29
( h, a0 S( P8 c, Z- p2 J" ~ 30: D$ D% [2 _6 m5 L$ E# j
317 [, c# K7 I" W7 X# k2 V- j+ A
32
7 Y9 r/ |$ {$ |0 \2 k- W' x 33
6 h5 @ v1 P/ G% o1 X7 ^! {# I8 R 34: w) M" z+ s: b' D
35
8 n1 W9 x" V* y! p3 q& A# c 36
5 @1 ]1 I# n! V% G6 K4 \2 A, _ 37
$ l9 D. C5 X# U |2 I 38. i, w" e8 ?9 V8 p1 Z# G
398 f2 C' P2 r u
40" ?. A, n9 M: f* @
410 T0 O! f' A/ G- B- F
42. [2 }3 Y- |1 Y9 r0 c8 F7 @
43
: X/ }( r3 M5 v2 U 44
% h/ P# Q/ r. R# _* J# d% u7 A 459 W Y4 ]7 b0 ^/ p
468 w' B2 p9 [& x/ S6 ^& G( n$ s- N
47. p5 x2 K* u8 T6 Y$ r
48
+ r) k) N0 @8 i' D# _1 [6 L2 K 49
3 c* O# F; d2 H/ v- J* T 50
; ` H h" m1 B4 a! ^- B) w 512 _4 v4 B: q$ y1 c9 c/ F
52
6 h5 @/ T: [8 J; ^; S* a+ h 536 n% O1 |$ ^- D4 k5 o
542 n7 D! n9 L* \1 `$ o
<AxesSubplot:> L$ Q2 l3 ]9 b, W+ k
1% y T$ h# i4 F) x
7 t9 U( T7 e ? 上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确& [& j5 K2 D1 t2 `2 k
; b( N( ^5 S1 _/ z: e5 M+ q0 d8 p
img.shape# J3 v _% N# V, [2 v) L m
18 N- V* y; E5 M
(3, 224, 224)
! B8 I+ Z( q1 D' l3 Q6 Y" C0 k 1
: J% n: P4 j* r" B 证明了通道提前了,而且大小没改变
) p* _9 }- A" V$ l4 C7 T; T- A
$ [2 n3 Z! ?% f+ h, |' G 9. 推理
* W- k' ? y# ?* I9 m img.shape
" [2 Y& z0 c* z. P3 v
. L3 D9 M3 h# s' R* l # 得到一个batch的测试数据
q( s. Y2 R9 }8 D- M; L6 V dataiter = iter(dataloaders['valid'])
- \/ s0 v3 G& |5 N5 T! O9 j2 E. m images, labels = dataiter.next()+ D8 z4 ]9 e& _: i5 W
' |6 v3 q/ R# G model_ft.eval()% D0 D7 v7 _ g9 X! O
3 T Y7 d) q3 S* N' [" X if train_on_gpu:' {+ F4 q2 }& ?
# 前向传播跑一次会得到output* P/ Q+ y! b$ `. F( P" ^
output = model_ft(images.cuda())/ ]) n8 A; n4 l
else:
: t& X6 E1 d* Z6 ? I5 Z output = model_ft(images)
7 ~ ]3 |1 C4 {. z9 f
1 c- U1 G2 p# z) p: S # batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值) K6 U; f. D- @6 w `
output.shape
% N! o8 q( @' v- U" r * p/ T& b) P$ t) N. D
18 g8 `4 j0 f( d
2
2 W8 h+ j6 R" n( L; o' a 34 s' B. a2 H4 X9 S, e$ x' C0 z
4, r% N. y5 N+ K
53 C) Y; x& Q3 b
6# ?. ~4 N+ s, m. f
7! k, I. O5 F$ {4 C5 P1 z. v
8
* z" k5 S# C4 g- z6 o/ V2 S) X 9& }4 _9 N( J0 ^+ V% h$ ]
10
; I& P' J6 n, g5 K 11
+ q" n! N% j" J" n' K. z 124 O, l! |3 C% l" _, U: L, x
138 y: L* j0 P6 _, ~
14
. @$ }+ D3 j/ ?1 N M" i" D 15
, E/ i) p: o8 ?1 N \, F 16
2 L- c- d0 q! F8 S2 L torch.Size([8, 102])4 ]0 i6 t5 D1 z" b$ Y
1
; y- M5 G' F; } H0 J0 e 9.1 计算得到最大概率
2 s8 H- L3 Q2 `- L( @ _, preds_tensor = torch.max(output, 1)
9 m" W. w+ s0 F# O" p0 n& E 4 @% d& M& ?6 @0 _
preds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量3 ]" |: Q8 {9 x
1
) {: o d% \$ \) l; U! e 2" R; @2 u& W. Q5 C; W7 S+ h) a
32 ^6 G' y" I* }! B9 ^, N
9.2 展示预测结果
/ K5 f$ C( U2 H( {5 k fig = plt.figure(figsize = (20, 20))
+ q; h$ i( @9 f columns = 4
2 b2 b' d# ]& ?# Y rows = 2
5 U- n/ h# H5 `
& Y# s/ o }: U for idx in range(columns * rows):0 Y9 q5 s0 A* P/ l3 C' V
ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])
: s1 s0 j' Q) h plt.imshow(im_convert(images[idx])) d! O8 M) _- p, R; x; c
ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), 7 Y1 o f. p8 I% `! T) C) a) g
color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))" M1 n: o$ n5 F9 |, A$ g
plt.show(); y g9 _9 G" w) S4 _& j) j7 x
# 绿色的表示预测是对的,红色表示预测错了
X% n. s# I. ?5 ?7 e 1
8 Y% l% B- C% m8 Q* } 2% c; E4 Y$ s+ H( o2 @7 y ^2 R; s
3% L2 j2 P8 P: L3 Q; X' y& Z! |! s; k
4, D4 L& f' d: e7 ~8 c1 M
5
. q p/ b# B; {) Y 6/ }2 c( g+ X% w9 ^" q
7
3 w& p+ O3 b) Q9 V, j$ i 8
. }6 ^& |7 w" ^/ V9 y1 Q& w 9
( d% d! V7 _7 L$ L 10) `2 b$ a& P3 T# X. |
11
8 l ~7 Z* Z9 k; z' W( z/ ^
4 Q; R) e, z3 b& l 9 q2 f' z' K5 G# n4 v. K$ l: B3 D
. `! e9 v' C4 P; s
————————————————9 A$ E, v6 i C: s
版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。+ ]. b8 p! r7 W* W3 Y( O( ^+ E
原文链接:https://blog.csdn.net/LeungSr/article/details/126747940, W1 U$ l% U$ e' |
0 o* g- ~0 ?; C
" j, T5 }& |' ?; p# x5 ^
zan