- 在线时间
- 1630 小时
- 最后登录
- 2024-1-29
- 注册时间
- 2017-5-16
- 听众数
- 82
- 收听数
- 1
- 能力
- 120 分
- 体力
- 565609 点
- 威望
- 12 点
- 阅读权限
- 255
- 积分
- 174906
- 相册
- 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)实战案例
* z' p' O1 g. f' B- m1 k: Q: G! I. j7 ?
文章目录; p3 x: O) S- d/ x! F) O. R
卷积网络实战 对花进行分类! G* s0 I8 i( [/ o
数据预处理部分
* }. N" K' t1 V9 w网络模块设置
' _4 [8 o0 j9 {( X网络模型的保存与测试- p% Z7 g- l# E) R' }
数据下载:
0 Q* t5 ?+ F. c5 ]1 q1. 导入工具包
0 M/ e) I, n7 V) z2. 数据预处理与操作
1 Y7 c7 Y0 s! ~ e6 s' C3. 制作好数据源. v4 h( U0 j: U: F o
读取标签对应的实际名字
, K5 b& H# e' Y- s7 Z6 |5 ^4.展示一下数据
7 f1 n4 Z% R! w5. 加载models提供的模型,并直接用训练好的权重做初始化参数8 }* }; ~7 M4 Q) s X
6.初始化模型架构/ }# r4 G7 _* R' x9 d4 U5 Z
7. 设置需要训练的参数
: O' n5 N4 k) W+ u% l7. 训练与预测5 W% E7 P, n# K9 a W+ g
7.1 优化器设置' X2 b; Y) \" S3 I* k+ ?
7.2 开始训练模型
! C7 w- A) h8 Z6 }7.3 训练所有层7 e4 d8 G2 `. |* E5 A5 v- ^7 ]; F4 i3 d
开始训练% _8 ?4 B3 U0 I9 C V
8. 加载已经训练的模型 n( x: d$ H2 Z2 g, Y# c2 Z
9. 推理, U9 A( A! @4 |
9.1 计算得到最大概率8 c) l% U% s' h' L
9.2 展示预测结果
) G, T/ k U* [5 Z: g写在最后
6 K4 m8 Y- {8 I卷积网络实战 对花进行分类. h4 h: W/ L; @- E& g5 q
本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能4 l- B f$ T& @& A8 S) C
* u% x2 u( _; F6 @
在文件夹中有102种花,我们主要要对这些花进行分类任务
5 I0 T2 W8 \$ J文件夹结构& k3 W" V( l: u% `
( B+ C/ k8 Q1 H6 Q/ a1 E; I
flower_data$ m7 L, v2 j# I9 ?% B$ {
) \( b+ d9 \- ~
train) ^) `4 [) w8 r- [
% k e# m' B5 {1 t6 r! U, w
1(类别)
4 r. t: f" X6 R- s4 a2* v" u7 X5 D0 [3 J9 ]( i/ h
xxx.png / xxx.jpg: U: T I/ [( y/ H5 J
valid! Y6 C1 ^" J0 U
/ Q3 w, Z9 n3 F' r, w8 P) | @主要分为以下几个大模块! F' p- m) Q5 I( t
" M& e7 u/ D/ C/ x2 f. m& D
数据预处理部分; R9 o/ |( |# G
数据增强) I- Z: A3 ^# s, q; I& t8 [
数据预处理% v, ?+ I- S' l: a
网络模块设置7 c, \( a( m3 E
加载预训练模型,直接调用torchVision的经典网络架构1 j' K: l8 M s+ ]6 I" f8 H* c
因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务
: c6 b3 D; Y" [9 j网络模型的保存与测试
2 d& t( m5 `) N2 E# e; [模型保存可以带有选择性+ @( @! c; C; L: t4 f; N
数据下载:% }+ D( A; w0 R8 p
https://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset5 K4 ~$ e9 D Y. W3 D, V2 L
- _% ~; f$ K+ H! B" N7 V" I
改一下文件名,然后将它放到同一根目录就可以了
0 K; r; i0 [( D+ g6 B2 e# H0 z% [0 b. }0 E
下面是我的数据根目录8 J `7 ^+ _* ?
8 x% P! i9 q' ~; S T2 k/ j( c( M/ z
1. 导入工具包
* I; E, g$ I, E; `- n- f( @import os
$ B: S$ x% M8 a$ t9 `, [( q& _import matplotlib.pyplot as plt2 R# b; l$ T8 @; F
# 内嵌入绘图简去show的句柄
" l6 {! f7 M7 j%matplotlib inline
% w" ?7 T3 S& p% b) b* @2 ^6 Gimport numpy as np
. D# n0 `! \9 g3 K: v1 {7 gimport torch
% y; I( v2 f6 j* {9 m& D) V/ ~from torch import nn
7 ?8 N/ g" U0 C9 W7 z8 i0 i4 G5 @0 O4 I5 h# ^2 }* X G+ ]) A
import torch.optim as optim
, m* t% U# X# A0 f2 X' Vimport torchvision+ g7 b) G. J7 m/ ]+ e
from torchvision import transforms, models, datasets9 S1 O; j2 w7 `
0 a" J1 |3 u. ^4 V* f: [ {, Himport imageio% D1 r1 ?2 E k
import time
( ?* v/ G) G% F! T. X0 yimport warnings
: B( k/ x; ?) @6 q# \. Zimport random
% I$ |: Z2 C8 eimport sys
! i, |4 }# G+ W# D5 @# nimport copy
, Z `% u' K+ \" gimport json6 O3 Q) z. d: t/ y& r& G; O
from PIL import Image
; X/ o% D! ^ ?) z2 d! f; ?9 t
4 c0 s2 D2 ]8 j, e' g) U
7 c7 d& S5 @+ ~+ e/ z. J& K a* j1
( k1 P, p+ I( O2
) @0 z2 G5 i! x6 `2 w9 S3# B/ |; W+ n/ d$ Z
4
( N2 G# m. U8 ~! X53 R: s3 x8 X$ A# L$ i$ |
6 n$ H# c/ N, n5 @
7
8 I `/ W4 h8 ~8 [% `5 k( R8
2 p4 Q, `; x' A9
8 `' c3 a2 g( s2 B: U10* B2 K- z7 O# X# M' b4 Y. h7 k# f; r& w; }
11
0 Q2 }1 n+ ^8 j! ~12
. K) |7 q2 Y! L6 k% H, I13
/ {' d9 d! U2 k! N+ ]4 t# J, B14
7 s6 R r' P, L1 R15
. \5 E4 X# J2 z0 ^2 i& Q+ e16
% H2 }6 X d% H" O( P+ a17
# g) c; s% K7 y9 a9 p18+ I. B V. |; |/ h; m8 N
19
( M% s2 ~# s) M1 Z+ U209 }6 H5 X/ w! r8 o- C" M7 W
21& {4 D* H* U; t! ?
2. 数据预处理与操作
" \% P3 i; }& @& v! \#路径设置6 u }+ k0 {0 W$ S9 Q3 k& R% B. [, l
data_dir = './flower_data/' # 当前文件夹下的flowerdata目录
5 r0 p8 m" K E2 m: Ytrain_dir = data_dir + '/train'
* \1 L" o0 ]4 y5 kvalid_dir = data_dir + '/valid'
5 p$ h6 L5 v7 D! @2 ^& B1/ q$ j. q$ U3 H" Q0 r, @4 w
2
. L6 E* A" d$ Z, l! _) t" G, r+ R3
! t7 D+ M% r7 E" T5 [7 Y47 J Z; @! `7 C9 K2 c5 i+ q7 V" o
python目录点杠的组合与区别9 }4 |& \0 s6 Q( h
注: 里面注明了点杠和斜杠的操作
; @8 W7 m* w5 e3 p9 D( ~6 K. N6 e) T/ H
3. 制作好数据源2 X* l$ A6 P+ c2 K! L9 P9 s
data_transforms中制定了所有图像预处理的操作2 ^! a; B& U3 O8 Z7 ]5 }* V
ImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片# P- ?! q( j: L2 h3 r
data_transforms = {7 R7 q8 [2 X: Z. H
# 分成两部分,一部分是训练
9 K. s h+ G [7 R' e5 D 'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间 N. O- C# A( f }2 o7 ]6 j4 w
transforms.CenterCrop(224), # 从中心处开始裁剪; N; @/ N% E' N! t
# 以某个随机的概率决定是否翻转 55开
" U5 S, r4 \9 i transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转% r1 p* @9 k. O( I* L
transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转4 A# A- R* E) \
# 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相/ D% d4 R1 {1 F3 q( ~7 Y6 U
transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),
$ z- r; w- s, p5 x* h* ^" B6 { S# y" S transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB' }6 }% Y' d3 G
# 灰度图转换以后也是三个通道,但是只是RGB是一样的$ k! C4 N& Z. J$ f& L* O; o
transforms.ToTensor(),, S. e( s2 m* a6 d. Y
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差
) h+ O- `" u8 [+ H$ `0 B8 o: q ]),
& F4 O U: v$ k% I4 i3 u # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
5 t+ v, \, m& H; t9 u0 I7 ~9 }0 `( r 'valid': transforms.Compose([transforms.Resize(256),
) r/ {: d4 h0 U transforms.CenterCrop(224),
3 j( x/ i- p+ N transforms.ToTensor(),
" f6 ]" b- h; q7 z( k transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同
1 G0 Y2 a# |% g5 c2 {: r ]),/ A" S9 T7 M, H) x7 `1 t
}0 \' d% G& @& K3 Z ^- [) r, }, c
3 A! a! _; Q7 a2 i
1: q4 m6 ]% [% M+ {( V
2: L" G6 G4 i1 q$ [
3
% [+ u' z6 s4 Z# Z5 X7 g$ L+ n4
+ ]. D4 @0 ]2 b! f, f$ \7 [* L5
c" S- ?+ V5 J/ | t) q Q# s, m# o6
# d& W j* y0 g) W: W" c3 Y7
3 m# H- w9 U9 X' A8
& q7 e; a( L) d0 p# o9
^& T# b2 ~; ]" ]7 ^/ d10" s- T @: X- _
11
" I; d& E% N- H- l3 g12
1 e% _+ c" M1 h! }139 r$ z& a( e: G V& p/ z
14
3 p. Y t0 D0 R+ a W15
$ D z" c" z3 d3 g1 l. m C165 {, g/ z& O3 l4 Y9 |6 |( [
17
5 u E: D+ ]' Q, e3 X. c18
2 R5 g' i! {' ^, w5 ^19. g. {" X4 ]0 f% h- v1 x/ S3 O3 q) x- C
20/ F" m3 W4 z& g d2 G2 `2 y5 X
21
/ p+ J, A. I+ e7 Q4 G# Xbatch_size = 8
7 }5 O! E: a" J9 b/ U2 _2 yimage_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}! I* X- s0 B: j5 L- i
dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}$ u4 a5 F4 [! y
dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']} D7 h4 t+ k- I6 ?
class_names = image_datasets['train'].classes
+ G, i ~% T8 l9 w1 }
5 ] K% h4 h: @$ T Q* X+ \#查看数据集合
! u, F, e' T# k- L* f8 Pimage_datasets
4 ^+ Q# e# ^8 I3 ?
) r) L7 b F4 I4 }* n; Z5 }1
! m4 K, M7 S, N2
& w; a) J8 N3 j& f5 h/ r36 Z4 k! g% N u t
4
4 o5 s1 A0 W7 l9 v6 p56 R N' N: a' c( l2 ]
6) P, U* H# U$ u4 ]1 L
7' S, [# S* w' q% S, R1 O1 h
8
6 a0 j- u$ C0 p4 t4 J/ E$ Q95 q9 O3 |0 L( `0 E3 [2 B
{'train': Dataset ImageFolder/ W8 ~: c6 p2 w0 j
Number of datapoints: 6552
) ?5 G8 S! B% o5 F Root location: ./flower_data/train
0 Y/ N# R+ t2 T" ^+ z; b StandardTransform
/ F4 R: g2 u' K. V) q/ z: Z o. ?' z Transform: Compose(
* Y, k9 u6 D8 c# S5 Z, q# D; c RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)' i$ ~9 x: C) S) f5 a: X
CenterCrop(size=(224, 224))
0 ]* R( e1 N& F# A RandomHorizontalFlip(p=0.5)4 @3 ~8 D( Y$ u1 U" _, Q" Z
RandomVerticalFlip(p=0.5)6 ^# B2 W% E- S% p* \& s) \ b
ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
+ m' I2 ^5 W" s$ W) Z* I RandomGrayscale(p=0.025): @8 E j: S6 A! T1 D3 @1 \$ R
ToTensor()
' s( x# g7 i* @5 G) `; @3 {/ F1 | Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])+ X4 d G2 \/ {% T
),
; q* r8 W) I2 ], k/ j, d 'valid': Dataset ImageFolder
$ x: F; J( j( l0 w; R Number of datapoints: 818
% z+ F- L0 p& r" B; f Root location: ./flower_data/valid* {7 m: ]- D3 z, |7 A
StandardTransform
u2 }2 k$ T9 L7 E4 z. a& W- E8 { Transform: Compose(3 \- ]8 ^, y" {
Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)
/ f+ [) e# K7 M2 ]3 [. K3 C! t CenterCrop(size=(224, 224))6 ]( i7 i3 }6 f" x: t! H
ToTensor(). [6 N' p9 z& E9 ~0 @; t
Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
2 I. S, b+ r7 r5 a {- Q$ U: H; F0 v" ~ )}9 I% ^- I; B/ ~; s, {0 |$ l6 L
* y' Z3 F- `. x8 M& _* \
1, T! R# Y: O$ _, }, e
2
7 R1 k) [# D- z w; p3
: D% G/ y6 t1 `( g% v0 Q4
: M! _ K+ ~0 e, _& i. U5
! n i" v$ c! s* u# D6
& C8 m1 {' S' k/ B- t8 y) p# c7
; l R1 j x9 l% g8
2 T6 i, y: ]% L94 q6 _8 R4 Q. u4 @) }* r
10
$ W- L ^9 |6 Z0 z: X& I11% [5 D$ \! r' {9 p
12
) R. z# e$ U6 f0 F @( k. U13
/ J9 ?" Q9 u0 Y P% H14 L) f( H# H( v; V( D+ I
156 s; I- j; f. G
16
' `+ M( k, m6 x/ ?6 x171 z+ N+ \1 P" p0 B+ ^9 `9 T9 }
18) _- P8 \/ \7 T% `( Z5 Y: `2 M
199 h h; [: x" O9 ]
20
3 G1 Q1 _) b8 f9 k/ y21 G5 _$ V4 s2 R# t2 ^5 h$ w
22
$ e" c- {0 u9 {- K235 e+ ^$ J7 {' G5 v* [' ^% I; ^, F
24
4 Y7 }6 ~, ]) S/ O- t# 验证一下数据是否已经被处理完毕
8 e2 h e W- W, B; J+ r, O- ?) ]dataloaders
4 Z) W' ?# I6 T10 V& z H$ @5 D2 E! E6 K
2
' w/ K' Y. p$ S, A# R{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,5 S' v, K8 y- z0 O; M, v- M d
'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}. [2 z0 l: a$ Q N
1
8 l" v) e3 g* I6 @9 ~" m, Z2
6 S7 P: N9 L J4 j0 W& J! A" w+ Rdataset_sizes( ?- h2 z O1 A. [
1) o# Z* T9 P) {' E3 U
{'train': 6552, 'valid': 818}" o c' |, c( x0 G
1
1 o4 o: z, }; W( D读取标签对应的实际名字" O1 B& [% }( a8 U6 e2 f. Z, Y
使用同一目录下的json文件,反向映射出花对应的名字1 K0 F" D6 N I. X
0 L1 J9 F) ?. ]. k% P4 zwith open('./flower_data/cat_to_name.json', 'r') as f:
I; M# ?$ E4 [; g# r% X cat_to_name = json.load(f)- E% h- h9 |* Z! a% S6 J6 O% y! t0 k
1
/ s# l3 Z- P5 Q$ S! \" A: r6 X2
, M( L8 v7 I+ {+ Kcat_to_name) R- E {. @9 @6 e( `+ _+ Z: P
1* O' u- m$ v, U6 z) X( ]4 i6 S1 g
{'21': 'fire lily',
1 B9 _1 M+ r2 q% h2 o '3': 'canterbury bells',
3 w# i# `; H% Q7 s7 U# @7 S '45': 'bolero deep blue',! p& m& l* |# e2 |# k6 V
'1': 'pink primrose',0 B0 V+ c: }4 \7 k, c z
'34': 'mexican aster',. a0 n3 k1 }0 E; h0 J8 U
'27': 'prince of wales feathers',
, B8 u% K; l8 N# N; D% c0 v '7': 'moon orchid',
b9 G! Q$ N* |1 E) P) O/ a '16': 'globe-flower',
0 d s: a+ Y" |, z '25': 'grape hyacinth',
2 W/ M% q+ b5 M! s9 E; j' G '26': 'corn poppy',
, R1 U/ R) h0 g) k: R '79': 'toad lily',/ @' @1 X/ w# |4 K4 q& a! H* B3 }
'39': 'siam tulip',
" ^9 ~5 m* t' @+ x# B '24': 'red ginger',# ~- V( X0 {4 t8 [+ }8 S' c! }9 G
'67': 'spring crocus',
8 a( j$ e) T8 w6 |, T" i '35': 'alpine sea holly',
: @ H b4 p+ |( h- k '32': 'garden phlox',+ f' z, b3 G( j, C7 K
'10': 'globe thistle',6 O. V+ w6 O& L7 A: Q/ _4 A
'6': 'tiger lily',/ K( W* \/ b b- t U
'93': 'ball moss',
0 X, q# t/ _$ c# O% g" T; r '33': 'love in the mist',
8 h" v0 z) X1 _ '9': 'monkshood',
! r: `# ^$ X, U) k$ v '102': 'blackberry lily',
1 K$ G1 s+ Q! e$ k1 b. l$ J- m '14': 'spear thistle',: U6 \3 v* v& L# L& i2 {4 o
'19': 'balloon flower',
) x" h& P% c* m& Q# r '100': 'blanket flower',
8 a* j1 f: j L0 C '13': 'king protea',! i' B- O% H$ m: R3 T
'49': 'oxeye daisy',- x+ ]6 Z* q" B6 \3 }$ z0 d& B; l
'15': 'yellow iris',
5 v h; t( \2 {* t$ D/ i- V4 I '61': 'cautleya spicata',9 i9 s, I4 d3 g: P" s
'31': 'carnation',
/ @9 x! q+ _3 q, a* P '64': 'silverbush',
( v5 s/ ~ D6 p '68': 'bearded iris'," H1 X. ^" m( g# P: M5 z
'63': 'black-eyed susan',
8 G8 `9 Z; W2 B( z( @1 U# g6 ^ '69': 'windflower',
) V/ Q: F7 G7 ^9 p$ ~1 A '62': 'japanese anemone',! C8 b: M1 n; W9 w
'20': 'giant white arum lily',/ Z e6 r3 `( C1 T1 v
'38': 'great masterwort',
6 v5 ?; v8 y ^& P: e# b '4': 'sweet pea',
7 |/ _+ e d- G* u) B9 i; t% @2 X '86': 'tree mallow',
# h& `3 S5 s6 e" \7 G '101': 'trumpet creeper',7 f2 ?% Y7 D6 v, I/ o U
'42': 'daffodil',2 |% f- R. z2 s6 M+ b6 O
'22': 'pincushion flower',8 U" y9 K! V; i
'2': 'hard-leaved pocket orchid',1 I( C; ~! ]" o/ ?1 @) [
'54': 'sunflower',
3 E6 u7 v" C3 T0 Z! M7 l '66': 'osteospermum',' N4 O. v9 a, p/ }) [- l
'70': 'tree poppy',
! }* f( \3 ?% v3 F; n, Z F '85': 'desert-rose',
( }6 n/ `7 o. Z# f8 r: W$ m: @ '99': 'bromelia',: D- q+ `+ D. D% I! B( _$ ?% _9 J
'87': 'magnolia',
( l; @" [' O+ e6 o! d4 S7 c '5': 'english marigold',
U9 S, k% W* F0 k! |( x1 Y '92': 'bee balm',
* s; B2 y" p' \- Q) D '28': 'stemless gentian',; m$ M, z# ]+ R6 |4 ]7 H
'97': 'mallow',
" E2 {: Z! r# v9 H; o( \ j '57': 'gaura',4 K# Q8 P6 E$ L7 o
'40': 'lenten rose',
9 x# g l6 V6 s9 g, { '47': 'marigold',
! ^0 }+ w) U) |; |0 L- d" s- v '59': 'orange dahlia',
) u4 F: [9 v9 j% N3 u '48': 'buttercup',7 Y7 p }7 c: G0 i" H
'55': 'pelargonium',; R3 E* U, n V) _ v
'36': 'ruby-lipped cattleya',
( X5 E3 U4 d* ^! H5 L0 J '91': 'hippeastrum',
" u4 k6 b" O; u1 S9 e) _, z '29': 'artichoke',' Z4 V `) r2 C5 W
'71': 'gazania',3 h/ v0 e& I3 w- X) g1 Q" y
'90': 'canna lily',! A2 ?% p4 H4 F7 B5 q5 m8 _) m& T3 l+ M
'18': 'peruvian lily',
3 E2 L* m+ J2 V6 T9 h3 s4 C1 y '98': 'mexican petunia',
# s4 f1 r" O1 F. U1 F9 ] '8': 'bird of paradise',/ Z- Q( n) q+ P3 `. H5 e9 }* C# H
'30': 'sweet william',
( C! E$ E d# A5 [2 K" n '17': 'purple coneflower',) I, u5 X. p8 b6 h
'52': 'wild pansy',+ }/ k- p. V; @/ P, U. U
'84': 'columbine',
% M% f7 b0 m9 o '12': "colt's foot",
9 t8 Q1 P* w4 ^& U7 z+ @7 B* `0 j1 v '11': 'snapdragon',
6 C* H6 |. Y2 X V8 } '96': 'camellia',$ M- s7 G) Y: _4 p% }: Q
'23': 'fritillary',
" s: ~; Q) P, O. L '50': 'common dandelion',
. `6 W& {4 g1 K$ R$ n5 K6 Z '44': 'poinsettia',
! ~ X, J& [9 H# A F& ] '53': 'primula',2 T+ t# v3 N2 \" h- m
'72': 'azalea',/ r- o5 i L; Z1 u/ d* C
'65': 'californian poppy',
5 c C9 [* k* W9 G* F8 ^' f0 ]7 m '80': 'anthurium',- @( l* V* Y$ e' M
'76': 'morning glory',/ K# S8 b* g. c: _2 T& ~
'37': 'cape flower', Q- K- o, S- Y& r: V8 V# T1 m
'56': 'bishop of llandaff',
+ j9 ?( G7 ]3 t# c '60': 'pink-yellow dahlia',
{ l' |3 x B7 P6 G% d '82': 'clematis',
1 y' N* `/ g5 W: ~1 V. W '58': 'geranium',
( E9 |" d) j) j '75': 'thorn apple',; i1 ]7 M" t- j* v: B" U+ U5 b. ~
'41': 'barbeton daisy',
' w. P4 ^, _) S, L4 g '95': 'bougainvillea',
: b, o. _7 p+ f2 o0 @ '43': 'sword lily',
9 d5 `, |! D! A '83': 'hibiscus',
* m" T, ]& V- l+ v- [ '78': 'lotus lotus',
9 ?1 U; y; d: u3 ~; R8 H '88': 'cyclamen',/ I/ D, N' U; |1 o" u
'94': 'foxglove',: {1 w) f. q* ~8 h' W6 e# n# |
'81': 'frangipani',
% x# F8 c. I% ?# J '74': 'rose',9 L2 U/ m4 j# u. I G# G; G
'89': 'watercress',. E8 V3 f8 X8 ?0 B; ?
'73': 'water lily',. Z b0 `9 x, ~3 {9 I1 R+ N
'46': 'wallflower',
2 Q3 |& L- I3 F) {' y( o6 W '77': 'passion flower',
1 [1 P! |( S. j% ~$ v% ] '51': 'petunia'}
3 `' e# b" q; v6 Q
# v) E# Q# y2 w11 o( D4 k1 w7 K. {: J t
2
9 R% e7 r4 |! Z; V( ^( M6 {9 l34 ~; m9 Z# A r9 s0 e
4
- a, D: j3 f3 L52 o9 c" `& O4 h, h5 T- g% i
6
6 t& p, ]% [9 ~) h9 x7% K# i3 B4 ?* B' s* J* H1 _2 r
8
8 s; t1 b |) G, A* {, V91 O* r. l- w+ q7 J
104 y! X5 S& {+ ?: Q& c+ n0 i
11- G; T! A. O! M. W- F
12
$ Y; z0 k9 q4 l136 E; U. Q# D$ z1 x" Z% h! t
14
* e* Q$ ]7 Q& O8 z' w0 w15% @, J6 F. r- K4 R! a5 S
16
! C3 x6 a2 m, n% N17
/ ]) _7 ?( g& I* m& a9 E2 ^2 f18
8 m: z2 J, m; r$ u1 k, E* g19* w; M/ n7 H, b9 o1 E: ]) @; u& l7 M
207 J, p) ~" d1 P7 e. z( L6 A
21
! f( @: ]8 U4 }. ^7 v22
! U9 e0 O! F( |) y: ^23
! @6 a* O" w; P- }0 J24. ]7 Z4 k+ ?; ]
25
+ a+ ?3 f3 }7 V- C2 Q% }26+ F' Y( h$ B1 ?) p, T
27
& o/ \, O2 e; A1 a2 R% }28
9 x! T; J: x, v J29
`1 @" g) T' t/ u$ z( |30
4 R# s1 F- q7 J1 N8 \31
* Q1 b* G! x% @32
" o. A# f4 U. h6 f* X# B6 m2 P# t33
' T" Q2 Q g/ i) G34
: b6 e# v8 C: y. W) K35( y( l" h6 F" X+ u
36
/ L* X4 T5 e, d4 X9 u37
) a. x$ ]* N2 w6 q7 t38
% U1 J/ s* ^4 x! x- s# ?1 p39
" }% ^5 B% c5 P( d; O' t; v40. \- Z$ S0 T0 o W8 d! {: b
41
6 f5 @! T5 N# h; t# U2 B42; \4 T9 A; G# u: m+ @0 ^5 T
43- k* |8 r: [3 T# f* D' f
44
4 H; u. Y- e h45* v0 x! o" [* c6 T
46" H+ u( F( ]& Y) {. X7 ?4 t3 C
47) x" m5 X. V9 M3 ^) a2 J3 u# K
48
. j Q2 _9 A6 s) u5 D0 F& h498 V$ a @: @: P1 g9 [
50& Q# ]1 H: O0 o) N
51
+ j1 s" V4 A, x5 t0 k m2 R$ l52" r3 U; t9 g; z- f1 k
53# b) H- }$ i2 I" T7 l
54
$ o, |' H! n4 C1 K' b55, m, ?" F4 Y0 _3 ^, s, D2 [! u
56
& K6 E- O+ l& T( m8 Q574 t# S( V7 X3 K. e- V2 ~
58
. \0 Z% R2 H9 i$ e59
* C/ h0 g% W0 V) d60+ G A5 I7 S7 F2 ^7 M
61
# n/ s& r: m. {# S9 t- e2 J& S62
7 F1 L) D2 K( y3 p$ w$ e0 w# z C637 N6 X+ T, C9 H# c
64
) d a9 G7 v$ n' Z: [659 z3 m- K0 u M: V5 h3 W7 \
660 \! j# e& J8 m. ?' y% P: d) X; m* B" |
67, A4 _9 R0 Z% A5 _1 H
68
6 R& W: [1 ?" F7 n7 O& Z* H# D69
/ m5 j6 z: {5 o# F; |3 i70 k/ p* V5 Z" z; X
71
8 F# Y% T' d* r5 r; O0 X8 b. M72
. n7 ?: O f$ `# \9 l8 b5 p73
+ Y8 W1 X7 u& K/ M, V& Z74" j7 }6 E7 J0 E/ r7 z
75
8 [/ y6 d+ v3 S! {" B76. R) q* o4 ^' I* [
77
# {3 m# E* ? ~6 a$ p N7 x# i7 K78
/ Z* n4 f2 S1 G8 q0 |0 \/ A79
1 i# p* `* s+ [80! d- l; t$ P# ?
81
5 ?9 C1 U2 g$ @, t/ k1 [82" X" _% ^ h9 H4 d
830 L7 l7 q! m9 N ]% L$ |* m
848 J- @3 {6 k b+ p, g) q1 |3 |
85
$ I' C G& O% x6 M, I1 w6 g86" a. o. h' x; @ @, q$ Q
87# y w2 j- h9 O& m8 `
88( q- C0 b6 i9 Z- q: a! }
89! U8 @. L& N% C& n$ }; c
90( Y* ?. K7 S! A7 ]( w' |
91# I+ [8 @1 u3 T! ^% S
92' m0 | ^' y( |! G2 i
93) t4 L3 _4 ?$ @% s/ N, [
942 ~: Q* S2 g$ M4 O( c
957 o8 y/ L, b. v( j
96
/ P- d$ G7 I' y- n3 S# O97
# C2 M, d) K$ e98
. C2 G; z" P V7 b999 a, Q, {! f5 m* _& S5 E
100
% p1 \8 Q' I V, s! Y" O& a101* z# B/ J# x1 N" R% i5 \
102* S/ _ A m# Q5 i2 l) ~# R
4.展示一下数据' g; H4 i1 a/ d5 i
def im_convert(tensor):
6 k7 J p, N( k% Y& d """数据展示"""3 y. W3 L9 ?/ F% ^' ~" f' `. v s
image = tensor.to("cpu").clone().detach()
, D' r# j' }( x$ q+ b3 ^; V( R. F image = image.numpy().squeeze()
: M+ O% k' U& f& ~ # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图3 U& W3 D- u5 ~: r0 {# b7 ]
# transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)
# Q8 N! }* v* v7 n. X$ b, h! V image = image.transpose(1, 2, 0)% R! W' @3 u* s1 Q, z6 h9 u
# 反正则化(反标准化)
0 m3 ?6 H! |6 ]0 |) W n4 F. m- i; P( l image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))
6 b% i8 U+ ~% L7 z% o
7 y1 U9 N! t1 _. G1 k # 将图像中小于0 的都换成0,大于的都变成14 j) x) c3 K# u( g
image = image.clip(0, 1)) S4 x! _! T, b. q
3 [+ K! ?9 {6 g l- a) L return image: h) c! G" d% |; a. [' v! @7 N- e
1& M' k) c- f5 n2 J
2( d$ @/ d' y: V4 K7 i
3
% f+ e' ^6 Q) J4
) X, F; A; m! }/ y; E58 d# f' x( N3 d) U% R. }6 S
6; b9 z) H5 E6 l) q7 f, [# _
7+ V7 S( \! A: L* y
8: n D- V! M9 S$ r6 n
9. l/ W1 r( f* |% O2 S2 Q
10! V1 J- c ^, n( |) n
11& Z& ?" E$ ~! R c u! v
12
$ J) f& j7 c6 U136 y+ p+ E% N+ h0 S8 o8 S
14; J" g* a6 n/ H; O4 K
# 使用上面定义好的类进行画图
% J3 z" J1 B& l% u* c9 C# nfig = plt.figure(figsize = (20, 12)); v1 l( Y8 o' {! S V" K2 {
columns = 4( G$ q: V, l( W( E2 L& Q
rows = 22 \5 h: M# T, z; T) k8 ?
* t' N4 y( S# _+ W' h; E* q( F. l# iter迭代器
, S8 k, S7 w/ k8 `8 W9 }, K) j# 随便找一个Batch数据进行展示6 J3 e8 Z6 F, X/ _: S! B" k$ b1 _
dataiter = iter(dataloaders['valid'])
: d0 n9 E5 B. B' Q% d; e( Xinputs, classes = dataiter.next()
8 Z) e7 d1 c. K6 m7 o% }* [5 L
4 _, ]2 N- [3 u' Z; a5 w: afor idx in range(columns * rows):; B+ ?4 a8 x& A$ m3 Y5 s. p
ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])
: _( o, H2 n0 W6 x # 利用json文件将其对应花的类型打印在图片中' M" m( ?- p: A |, W3 u+ }
ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])
$ j- O* f8 m6 v- ?0 v plt.imshow(im_convert(inputs[idx]))% L% W, x1 ~5 l1 `; O) B
plt.show()
6 s. ]( }2 v3 t4 }$ g% q1 G$ ^- b. M1 Q$ H' p. o. I' s
1
; r1 e: d+ D; l/ r: m( |2
- |: e4 B9 ~0 u! E3
3 B. ~" O+ h E4 `2 z2 b4 C5 D4% {% K2 E5 J7 ^4 t
5! \" [ \, ~+ b4 N5 ~
64 a; F; f, M5 c8 h; Q; N9 S
7
. h1 T3 v1 s( ^4 G0 N. k E8! i# ]1 v& ^7 X
9
( i9 K- _* V* @3 w: [: y10
% l5 y- E- C' W3 m' M11
! H) z2 }5 C' Y I$ p9 A124 X) @/ A, G0 O+ X1 e
13
* S! ?' L& {. ^. N. |2 U1 k0 s14
/ ?- Y( k# g: K; d% R; s0 }) i5 G15
& _) [2 @. _ q, D16
; }, H& @5 R9 Q t; V+ U) F- y4 O
1 I3 t9 H+ M4 u- y; C$ n1 {9 u3 d' N9 F8 S/ E5 j* r
5. 加载models提供的模型,并直接用训练好的权重做初始化参数9 N* o( Y2 ` r1 o
model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
; e G- l1 A" P) w- ?# 主要的图像识别用resnet来做
3 f! Q0 ]2 S5 _2 X# 是否用人家训练好的特征& ?! I9 q( U6 a
feature_extract = True
" u7 P2 E6 N$ \1
- e {! s/ R9 v* W! {2
0 p& H/ J! f) K! c' U: y4 W3
, u6 h4 ]/ }; h; y% a/ L, X& w& l0 J4
1 l' e2 P: Q. F. M" \# 是否用GPU进行训练3 A3 V$ S/ i' N; t& g
train_on_gpu = torch.cuda.is_available()( N/ V2 A( t+ H9 G+ g) m2 L
- X: j+ v6 u" Y" Aif not train_on_gpu: L1 Z9 U, I, z0 f/ ^
print('CUDA is not available. Training on CPU ...')
( Q0 P! x# {0 g/ }$ Xelse:* K1 p, A7 i, ]6 o& L6 w7 [2 t
print('CUDA is available! Training on GPU ...')
4 e6 @' V7 c, e9 _2 B/ h6 ~3 J; o4 {9 c% G
device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')4 i+ R1 A" R1 M
1
. E1 Y' n/ u+ R+ N3 F; ^! }& H$ |2
- \8 ~' c5 M; H( C2 Y3
; B7 u. c, W2 A+ ]4) C9 Y$ v# Y+ c
51 g4 m8 d" z* V. [. `, @4 K
6/ J& F$ K) E% g" A: }- I( m
74 }6 }1 R4 y1 M0 P+ M' Q
8, m7 y: D( I) a% B0 e* W9 B8 b4 ~
9
+ O! x- u+ x( p& Y3 t5 W7 UCUDA is not available. Training on CPU ...8 G+ C, O9 k/ ~( M0 e( X% r
12 e7 h+ H- M( F# b
# 将一些层定义为false,使其不自动更新
1 D8 n$ r2 }# O( V! N) f& odef set_parameter_requires_grad(model, feature_extracting):
8 s+ O7 i* E' }( h# _ if feature_extracting:
$ k( Y9 _' A: D for param in model.parameters():1 n% g+ X7 }8 ^- s2 v9 d3 z d
param.requires_grad = False
* P2 g) g+ s# P* r! i, D) X4 ?1
K3 W) i* |2 I0 h1 }& z2& D* f H3 v* c: x* F
3
6 ?# ?3 ?* @9 ]+ Z1 \+ ~4
5 i" d9 H# q( M7 E5& h7 C. Z% f5 L7 p4 B3 F
# 打印模型架构告知是怎么一步一步去完成的
; d/ {$ Z. t, b0 Q8 u# 主要是为我们提取特征的, v" Y+ h5 [1 | s% Z' D
0 o$ {# N4 \( Smodel_ft = models.resnet152(), d8 d% y+ B) r) K6 b6 V, \
model_ft
6 R" _3 u3 R* Z7 Y o5 ~' Z1 Q1
6 P, H; W. I' a4 G9 M' i8 k; d2' m9 i( r- x2 h; y" B) f
3' M/ E" B" |9 [$ h0 ?0 F0 P
49 l- J$ M u; n1 T" O6 W. Y+ T
59 o" {, S! p& ^/ R: N
ResNet(
4 A' \& I( n7 B' S2 o& W/ ^- ` (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)9 s Z( k. C( C7 M$ v# f
(bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)9 F) y% J# V8 K/ K; e3 _2 J) ^$ R$ ~
(relu): ReLU(inplace=True)
/ p* z! z6 ]/ d/ U5 M (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
1 v6 q6 N# \: S& @ (layer1): Sequential(
( @. {% d P& t$ y" G- d( l (0): Bottleneck(
) X; e; F7 j. r; w. ~8 t" P4 s9 f0 Z (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)! }: Q4 _. A/ P+ G- ~4 b
(bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
, f2 o" X2 T/ Q. a (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)0 e1 J$ Q6 H z. s4 u9 J6 w
(bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
V+ x: V8 m: D# p* d1 A. ~/ p (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)9 I# c9 q5 b1 l. v# x, f
(bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)0 H) z6 ?. x; j" H( X
(relu): ReLU(inplace=True) X. }" {! G S8 T
(downsample): Sequential(
+ H+ D8 O2 B$ B2 q' u6 c4 m (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)- d. q* E+ u m: {' i
(1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
1 r0 o/ F- w1 K. p, P- u )
' A1 M0 M, K# ~( Q: n6 j )( u( A7 D$ O7 ^6 m
中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。6 j, M. v' ^! P) X/ x* H/ Q# u
(2): Bottleneck(/ r# o0 r9 T. ~' |
(conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)( c1 K2 E2 j- E5 M* m% I( q
(bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
! T* K, a9 T& X* q. G (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
( U6 F1 l- j! v7 s9 A9 ^- o$ A; a (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
/ r2 R4 ~+ h% A6 @! w (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
* Y- {- I, L% N) O+ j. g8 Z (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)( \! T) c: ~% G# h+ Q& G" i" Y
(relu): ReLU(inplace=True)6 W2 b+ ]4 F2 N
)
' O4 x5 _& ^* I7 X* Q7 e )7 p" x; \' ]: `
(avgpool): AdaptiveAvgPool2d(output_size=(1, 1))
$ K) d& G7 G0 K) {* Z0 u (fc): Linear(in_features=2048, out_features=1000, bias=True)
6 g4 u* e5 }" D% ]/ u9 k0 i)
* k5 G# C O2 P6 G. _' a# {
$ o) L7 Z. r, q/ |! A" v/ {1
' S( S* b# y5 z2 Y7 [9 z2
8 t+ W. z& e$ @/ E, E% b1 h- g35 W% @1 J2 [7 q1 a
4' i- U+ g3 t% a5 L6 \2 H# h5 B
5
7 K6 X+ ]6 y5 p1 o" S. j& O69 ]4 b: z) D" W R
7
4 y" C. Y7 k% X m8
* f( C# G7 r' p2 n% O" H5 W# N1 y94 I8 b! G) \$ j4 X# X, ?) W
102 _) Q- A! Y: p3 m$ K
11
+ V! C2 D! c$ |& w: D12, E" w' i) ~1 Q* t
13
: ^9 i2 h j! y: r: H9 v, `: F; R/ e14: k/ p" o0 K# {& h7 Y& {
15
# B) ^4 L9 _) J3 j6 |0 S168 p- R, b7 |9 Q2 D
17
! M, o" n& @3 O( ~18
. S. y+ h0 D \) F. B, C( j19
8 D) Y. D9 J7 h9 o5 N" S; X) l20/ G8 C' A9 v2 V/ l
21
! p$ P4 h1 X- H22
4 F0 z% q$ I$ r$ z4 A6 o- m23; A8 } b' r3 u' h2 f( i9 J2 l
24' J8 u, o, Z) j1 `( Z) i
25
# J$ t4 b3 I+ g% ]263 E1 J# M) G$ B1 s! g1 G: `' @
27
' {& B( Y% j, c# @$ W283 @/ ~. u0 ?5 }
29' ]3 Z3 m k. D# }: C# B! B: K* [
305 t* m# _6 c" q0 Q. `6 Z
31+ w4 v4 _& E2 {# `1 r: C+ @! S1 d
32 Q5 H" r) D: y1 R
336 t; O) e; P+ c
最后是1000分类,2048输入,分为1000个分类/ M' m# V* @6 V6 R% l
而我们需要将我们的任务进行调整,将1000分类改为102输出
% z1 ?: \& K9 b5 {2 v: [) m% M& u7 M$ k3 X, o( _
6.初始化模型架构
+ A& j0 z4 i$ G. v2 u步骤如下:' _' M8 h. ? _) D) W( x
1 @$ n9 i8 ^$ m% [* m
将训练好的模型拿过来,并pre_train = True 得到他人的权重参数7 R+ n: i# |$ L- M- W% G* v
可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False), r% }* p& [3 o% e/ ~- Z
无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数) x x% L. ^! `- r) D
官方文档链接1 D* j4 r9 K) g0 O
https://pytorch.org/vision/stable/models.html% ~, E- q8 {( m- Z* A _6 \" P
/ r3 v) R4 j- e/ O0 R% i M- _# 将他人的模型加载进来
& t. E- u" Q- g# g5 O( X# ]1 idef initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):
* h' D- L# Z7 G6 V* \0 Q5 J # 选择适合的模型,不同的模型初始化参数不同
- T# J0 r/ S- k model_ft = None
6 a6 U+ E; Q# z1 z; @6 x input_size = 0
) c' h: k- V3 v3 b" a% ?7 H. N2 m) ~$ n# e' O
if model_name == "resnet":
3 }! N1 a# B/ v$ @- | """7 c* ]1 l4 E8 h+ Z* Z
Resnet152" p( A- c$ q4 F6 I
"""
: U( l2 a' Z% i& h! D! U! M! i2 c0 _) B( \, o* W
# 1. 加载与训练网络0 K' c5 ?% x: _
model_ft = models.resnet152(pretrained = use_pretrained)
; h' Z9 G) f- C2 i& T # 2. 是否将提取特征的模块冻住,只训练FC层
3 Y7 T# u1 W G7 B7 j set_parameter_requires_grad(model_ft, feature_extract). F. B+ Z( j. f% K0 p$ j8 J
# 3. 获得全连接层输入特征2 B0 w: P6 `- Q ~* z" p8 W
num_frts = model_ft.fc.in_features
3 z# b# }; ~. s2 ~ # 4. 重新加载全连接层,设置输出102% A6 S, _6 w0 f
model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),8 y" u; J& }8 j2 m2 n
nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1+ A6 c! n7 ]7 \& P/ T5 x
input_size = 224" a) B; d' {( a' u: C
$ |! U I R5 }3 K8 | elif model_name == "alexnet":% f! Y4 `6 u$ s3 p3 r+ y4 ?* n
"""
, t; n9 r$ l, L" z) D* Q( s9 _ Alexnet7 \1 M. v, k9 T+ V+ C
"""
! n& F( Q& t- Q' x7 x) c8 H$ L. d model_ft = models.alexnet(pretrained = use_pretrained)& x" n# D2 b5 k# a$ N; z
set_parameter_requires_grad(model_ft, feature_extract)
- ~* {3 V4 h4 A& ?, j
' o- O/ s3 N2 }+ r% j4 t # 将最后一个特征输出替换 序号为【6】的分类器
9 m4 I% ^$ y5 K* i" L! ~. c1 ` A num_frts = model_ft.classifier[6].in_features # 获得FC层输入
+ K' l7 c! _9 N2 M4 A4 w5 S! u1 L model_ft.classifier[6] = nn.Linear(num_frts, num_classes)& Z5 f/ R+ o8 c* ?7 Z8 a" u
input_size = 224
1 E* g& [, c0 p, c6 ^0 n/ m8 K
. S I* z! E. ^: V: f* u7 S elif model_name == "vgg":
! I) ~# V7 W0 Z* W' z: | """
/ t6 O# P4 B# K! d1 a! h VGG11_bn7 Z& V1 |3 s+ m3 ^/ p1 c) i/ w
"""
9 X- d7 v. e: U9 r% U# Z, R- u model_ft = models.vgg16(pretrained = use_pretrained)
# p2 [0 D1 H2 r2 w/ v6 J( R. l set_parameter_requires_grad(model_ft, feature_extract)
# d; t! l1 P. X& g9 d2 [# A, r num_frts = model_ft.classifier[6].in_features4 X( t1 }/ Q; z& {: k3 L; h
model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
- X* a5 ^+ v3 j input_size = 224
3 v, d4 ` u9 x+ \3 T7 ~1 w7 `3 n" r' U* E# W6 L6 A& Y1 G! c
elif model_name == "squeezenet":1 f4 w2 M5 m; l& X% ~1 H5 E" z8 x
"""
5 i% ~) M0 m/ } Squeezenet
; ?; P, q2 d/ i" N& F7 { R+ o """& Y- k. R; F P: K3 b
model_ft = models.squeezenet1_0(pretrained = use_pretrained)
% f8 l A) `5 n$ y set_parameter_requires_grad(model_ft, feature_extract)
; ?! I. h! l) O) J! a* `" C model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
$ J D& q- z9 \) S2 y U: `8 { model_ft.num_classes = num_classes9 M: U: I& t2 Z7 h3 U
input_size = 224
" J! w, ]' f) s9 {5 s
! U% Y' q; `2 z0 J elif model_name == "densenet":: R* f8 t- y+ U; P
"""3 X8 p Y; l" T& p
Densenet& |6 ]1 K1 T0 B
"""2 a. o# z z- z2 c' y
model_ft = models.desenet121(pretrained = use_pretrained)
, z+ }% d$ V9 }' u7 l% }1 A; N! k set_parameter_requires_grad(model_ft, feature_extract)
' G3 `, W3 \# Q) O& A: A num_frts = model_ft.classifier.in_features4 A I+ ^3 i1 N$ S- M/ y6 f
model_ft.classifier = nn.Linear(num_frts, num_classes): C# N$ h5 r1 }3 d* R
input_size = 2241 J9 p4 C: G. y) ]: K- X
5 G. B' e6 I4 b* p1 d" ^/ Y elif model_name == "inception":! G; {" d* F8 ]% \6 X; w$ q
"""
9 G; ?! @5 W- y* P Inception V39 B& H$ b( j2 _3 u
"""- @! q- Z' m. n
model_ft = models.inception_V(pretrained = use_pretrained)
! I6 [* X1 r( ]. V' U/ ~% F! x set_parameter_requires_grad(model_ft, feature_extract)7 q/ v2 b# L! P/ M9 n
& S1 d0 E" l2 f! I" h; l% V& E: ? num_frts = model_ft.AuxLogits.fc.in_features
' V" T5 t' ^) Q4 @; p! X: H model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)
& L5 V: L) e5 j K+ t" S" ^( e4 ^- Q- w- `5 ?2 m6 Z. V* G
num_frts = model_ft.fc.in_features
+ a( ]$ I2 y: c) a% W model_ft.fc = nn.Linear(num_frts, num_classes) ]7 [0 B# \+ x- `! d- ^" o7 Z
input_size = 299( J% f0 l: T: o
' A0 n4 M; q1 P& S0 e1 i/ V+ I( o else: K& |4 } U1 D8 g: Q r9 P( f
print("Invalid model name, exiting...")/ @" A0 S ^/ i9 a# v$ z9 A+ U) C
exit()
5 k+ d- C9 f, H) F! p) z# c; f/ W6 B- ~& ]0 q" Q
return model_ft, input_size
- @1 W8 Y1 c) _9 K+ b
- [: V3 m. l! D* F- ]- A/ O6 y1( q$ a m% I/ q
2- }, A# N& U/ j0 W( t4 h+ f
3" M# @ P7 U ]8 y* v9 W8 [( y
4. u, K, z, h" C$ }$ L
54 U# m3 V+ z) @: G- V* t# W
6
+ ^9 l& f: X$ ~77 _5 n0 {2 L' a: n2 _( b
8
) P# a6 t) y0 J/ g X6 q1 s7 D! X8 U9" U! b5 w# b. _9 T r! ^& B
10
* J: X/ ~, c1 [( a11& ~! k) E; q4 `& b* s$ \; l# Y h
125 j) _* m0 e8 R( }' ^) ^. p/ B
136 P$ n- M; R7 o5 j! A8 ^' f' v
14
8 Z; {. a6 O4 z; H15( J6 D0 A: Q; g4 T" B: ]& p# p
16 Q. b: w8 T4 a T$ Y' R8 A
17
8 m% ^* J4 c9 R: L5 J180 w& b. m7 b8 X* d1 r: i# p
19) O; x2 u3 ?* X- E% R
20, |6 r+ L. Q, W
21
; k# _2 }* M& b" O( z228 L* T/ c) o1 t; ]! ?
23
& G* x# z- R8 v) A+ O; s+ ^3 H24
6 n9 j. f2 {" e# c# a25
1 ^; c4 Z: A/ w+ E26
2 i9 B% _4 o8 R3 l3 M273 ^3 M; m; K1 I1 @
281 @- d* [+ O! q! p/ F* }, T
29& B2 M3 q# t7 P0 k: W s
30
3 P5 m8 x; i5 K/ s- `& N31/ c$ b" P( ^8 @/ W* ^* t
32' j8 l, Y) V# N
338 b; p3 x5 n$ W6 [( I! M
347 e2 k# r8 D& m! E6 i" t
35
0 _4 P& G) o4 q+ E36% ^. t+ N8 A+ o1 g. z0 f4 a' ~
37
4 f% q# P6 N) P B38
( y5 r; |+ |" i* t. y3 N39% y5 L# k+ v, U) _
40: s1 Q" ]+ s; P! O$ ` y5 |8 T4 J$ W
41
8 p, ^( \5 d- a2 O42
, e! j5 N& s3 n9 O43
. }/ H3 R4 C6 w' n0 N44# y; u: V! L0 s$ \# M
450 T2 K+ Y9 {+ u. F
468 R" }" E% W, y$ [) O9 ]# x0 X7 v0 q
47
. }/ I4 j3 g' O% s# v8 c- T G48& N0 p. j! P. j& W: Z& ]. p
49
! A6 n) V, k: H1 U- h( ?50
0 b* @ G5 |1 a1 @% [+ x6 Y51
" B$ q+ L5 q$ X1 O i52
; R: S% M. T$ O# e53
8 D9 M, t( e( \+ x54* l5 Z( ^* ]' |/ p4 N
55
( T# D# Q% l5 C5 h C% v8 ?561 O) ?; D7 v j5 c3 C3 h* m; N
57
' N1 _" Z0 y1 N$ v1 X58
8 ]4 Y( a1 C m+ O' O59# V. m2 C3 q" d
604 O/ W2 B! F6 p4 c3 t
61
- Y% @0 B0 ]+ Q0 K2 }623 @( D7 o- y( V6 j
63% O0 A: v& _. D
64
S, K. s+ I, u65+ a6 q0 I8 g" A7 G# D
66
- ^! ~1 B4 ?9 F3 G# i0 [$ i67
% W. `+ k& L; b, ]. q1 D683 t2 W q, L( x8 y
69) x' C3 ^* Z$ X9 ]$ r
70
6 Y& z; k }2 K" s) O( h& y71 n; N9 u3 ~* l4 d! D. ~
72
0 @# \& C. {$ k1 ]73
! p" {. _" g% l* \0 D8 P74& b7 X& q4 p! `% N" a# z& L9 G
75$ r7 x; w0 |" j$ [# _7 M4 c
76
" b( K0 ~; N) c( s4 C" D+ H77/ u- L1 k2 X* k3 q# s5 E9 t
786 d @& Y1 Z+ a7 }: K
79
% I# \) ]' L) q5 ~0 K5 X. r" n7 P806 n7 ^" c! s3 p
81
8 J; g% p h% h8 z7 R6 d( V82
! }: f7 y; K3 W. y9 }# |; B3 U83
$ H. F/ ]" L, B* s) v* G- D7. 设置需要训练的参数
' e' g5 Z Q6 f1 T# 设置模型名字、输出分类数
( C9 P$ ?: i9 c: `, X" Umodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)
% X7 y1 S, Y9 `7 K0 H5 @% @ y- W5 i; s. U, Z0 M1 d! ?* i
# GPU 计算$ U4 i0 m4 p [% n- y- y9 S
model_ft = model_ft.to(device)
( w* R8 {0 J3 P: b- A. w
6 ^) R9 m0 A; e* v9 c# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取
# j3 y$ Q& K7 u( pfilename = 'checkpoint.pth'
2 x R2 x, \+ L' H1 t6 Z L2 D2 ]3 n5 E4 V( x0 [9 H3 [ R( X2 L8 n" ]3 Z! h
# 是否训练所有层5 Q N) _# P- ?* A
params_to_update = model_ft.parameters()
) x ~6 Q7 B2 f- q# 打印出需要训练的层
5 C" N0 @/ Q8 K# B& @print("Params to learn:")! ~! y9 F P& E
if feature_extract:: ^( c, Y) U. X. T7 z
params_to_update = []
0 P0 Y& w% c0 N6 L7 {# \/ x for name, param in model_ft.named_parameters():6 N+ |. H0 x4 H8 m
if param.requires_grad == True:
- u& G: S- b- V Y- D& m4 r/ I params_to_update.append(param)! b* D% B* l5 `7 e
print("\t", name)8 D5 c# W' c3 O4 h
else:" X2 f1 c8 k" p$ \4 j. d/ \2 O
for name, param in model_ft.named_parameters():, C' j, i: t8 K8 T/ y- [
if param.requires_grad ==True:
' }! ]. |% ?" e! U$ [ print("\t", name)
/ S: [2 {6 B5 U* [4 X% ] E
5 b/ C' {4 c5 I3 T7 i) ^6 \7 o9 d4 W1& f) h$ g8 f5 s7 o/ L
2
7 w4 o2 G' I0 K; ^3
! N$ A7 f0 x) Y" G0 D+ o @4
$ V9 f' r1 P" A* v# K0 v) P5
% H. }# H* u' d- Z. ]) h; I3 v68 f2 j' Y( J4 @; h
7
' B6 b7 C" M" v5 a+ c+ `. E8 {/ V9 `- L% c- L! F; g. G n1 X7 Q
9
1 \$ p" N9 q3 m5 f10
0 o( f& E. D) z' T1 p) _11# L' G Q6 L* f' t' W, G: y7 Y+ K
12
6 a# T- T( A; ?6 y# S/ V Y n13 \% T4 T2 _4 r6 e2 X- `9 v
14* |1 D A: `2 f1 d4 H! ^8 C' m" M
15
, C; U. v9 S- n: k3 K$ z161 [3 Q. O" d2 y. s
17
0 D9 e: y( _% p3 L, p3 f# p* z% r; m18
2 T' V4 A$ v( q19
( I3 p( m6 C5 e8 ?20 q4 W8 A! ?7 f$ b! K/ b
212 [: e% ~2 e& I6 K7 t
22
3 L' L; N2 |2 P; U0 f23" T- o' ^8 L9 K: m F! U
Params to learn:" z3 M& w& i) f" M: i+ }( _
fc.0.weight6 _- j8 Z4 C( X2 M y' t* ~( U. E
fc.0.bias6 ^7 [# S- H$ }8 ]
1
l$ E5 W5 G# G* x22 v6 h$ C) i: @7 f, @' F2 C. B
37 D" e) x1 p5 F8 \6 H
7. 训练与预测
* A! |5 v, ^. e7.1 优化器设置
0 E1 \( B: l6 X+ j; i2 @3 ^# 优化器设置
% H: d4 s A; Q& O. g$ W0 @! soptimizer_ft = optim.Adam(params_to_update, lr = 1e-2)7 s% V: Z5 g3 h: k
# 学习率衰减策略
$ ]/ e' M* j# g( Z) U" m8 lscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)
) N* `: Z% I3 n: q. h0 @. A# 学习率每7个epoch衰减为原来的1/10
2 B7 L, s4 Q7 r4 r3 x( ?) f+ u7 N" ^# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算+ n% B+ o) G" y" ~5 C5 O- h
% z0 v! N' `2 e# o
criterion = nn.NLLLoss()5 ^& J" |% J2 V! t3 H% Z
1# x( q9 F) n- B# X M
2
7 a$ B! W J; J# T6 y3
- V! T& q% Q8 e) a/ }46 K3 W' J, L8 e5 b, a7 B
5
# Y) `& p9 k' p) R6
# m' C( [) G% e o" g" B7
6 z Z( j7 ?- k/ g- z8
7 T) ?6 s. A( X( ]8 b# 定义训练函数
2 K9 N1 v) ~) y#is_inception:要不要用其他的网络
: t @9 z3 J5 bdef train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):
, l' i8 h: l& G# M) _- m: u since = time.time()+ L5 z1 A9 T) k( @- O( e
#保存最好的准确率: \8 U) S- `9 t( x; m; H6 h$ E2 K
best_acc = 0
, J' ~( G3 M: [; k """
1 {# Q; W6 `7 L3 [8 f2 B' H checkpoint = torch.load(filename)- a4 @8 z" h- e
best_acc = checkpoint['best_acc']
, Q9 U1 n) n4 ^ F# ^8 n9 D model.load_state_dict(checkpoint['state_dict'])* n5 v2 b( I; `" N- z+ M' \
optimizer.load_state_dict(checkpoint['optimizer'])
! x. ^4 N$ ^8 p6 W4 b& M model.class_to_idx = checkpoint['mapping']
+ O4 n2 Z0 F" G """
+ x1 D- P+ z& Y; x0 g5 F1 D" V/ I6 i #指定用GPU还是CPU: I/ W0 f8 Q6 k
model.to(device)# b: V8 t {/ ]4 r: ?: t( h" A
#下面是为展示做的, o5 n( c4 I/ p1 v1 ?
val_acc_history = []
) s: f: p+ N7 ?! {* e/ D6 B6 _5 B train_acc_history = []$ W" ]" r# j6 ]8 y# i8 b0 E
train_losses = []& M3 ] S2 x* I$ `1 k6 P; T
valid_losses = []
9 { x4 i1 B( H. L( _, b% `- R, W LRs = [optimizer.param_groups[0]['lr']]
; V. p& f! {: A1 U #最好的一次存下来4 S h, J, \7 z( J; a
best_model_wts = copy.deepcopy(model.state_dict())' P2 E `7 X; \5 I5 ~
$ k5 s6 t+ @4 {
for epoch in range(num_epochs):
- z1 x4 r, @+ t5 D; {! z3 A print('Epoch {}/{}'.format(epoch, num_epochs - 1))
* H& V+ U8 F. f# o) d, g print('-' * 10)
, e0 s5 y# ?" p. \: d: @* k
6 Q. V g5 d% c/ w # 训练和验证# C% t+ m5 ~* y: q3 j r
for phase in ['train', 'valid']:& B" }# m2 [9 z9 r) ~# H+ s2 P
if phase == 'train':" N0 B' r1 y# T- U2 m# w) w
model.train() # 训练
7 ~, j/ f- l( m else:; O/ L; B+ b+ C j
model.eval() # 验证
! M$ o9 u7 J2 W' Q2 c3 J
2 z1 d! Z9 X6 Y+ y* X running_loss = 0.0
7 w L! @2 J0 n7 e running_corrects = 0
5 w q( j [, T8 D# w! Q+ ^+ x6 M7 @3 E) @' x- h
# 把数据都取个遍
" T, n+ m0 J" Z- t5 w for inputs, labels in dataloaders[phase]:
7 l7 S7 d8 v8 c, k #下面是将inputs,labels传到GPU
$ Y/ J9 ?% U* ]/ m+ B# v2 q inputs = inputs.to(device)4 s$ u- L8 @+ [' h% Q
labels = labels.to(device)
" q2 r: \2 z( {' Z0 w& }+ T: K1 y6 o
# 清零
% p) |; M9 S8 n& J g; K5 [ optimizer.zero_grad()# j8 w+ D0 R3 ?/ g9 l1 s* |- h
# 只有训练的时候计算和更新梯度
/ g( A. v4 o1 ?% i6 \( z with torch.set_grad_enabled(phase == 'train'):% ]0 B5 g/ A! G
#if这面不需要计算,可忽略
3 n8 Y5 x; l2 D, C9 K8 w: k if is_inception and phase == 'train':: ~6 t) |3 H1 y7 L8 F* A) J7 g
outputs, aux_outputs = model(inputs)$ _. m( B2 R4 i
loss1 = criterion(outputs, labels)
( g; \. G/ A9 ~ loss2 = criterion(aux_outputs, labels). X0 D. A, h5 I ?
loss = loss1 + 0.4*loss2$ ~7 m5 h) P/ b. a$ K% k# r
else:#resnet执行的是这里' D: i9 \2 J1 U3 }. c9 B1 ~# u. x
outputs = model(inputs). T8 N/ ^8 l- O8 Z3 a2 e6 V
loss = criterion(outputs, labels)
' u% r: P7 Z x* |# G% v: x/ |) t$ G( f; G3 w6 e0 y" ?* {' L3 C
#概率最大的返回preds" n8 `/ x6 J- t w
_, preds = torch.max(outputs, 1)
4 V, A# Q( l" V% c, l& ?" S
" O3 z2 F( @. t+ J # 训练阶段更新权重6 M9 ?+ D5 x/ E# @
if phase == 'train':
: U" Q/ I9 ^) P loss.backward()' ^, k h; g% g% U
optimizer.step()
5 u( Q0 U! Y/ e$ B
+ \9 @, C( `0 g: o' S # 计算损失
/ G: k7 C* v8 n2 Q+ t' z! E& W running_loss += loss.item() * inputs.size(0)
# }1 b+ e: {% X4 ?( B a' G running_corrects += torch.sum(preds == labels.data)
9 T! p; Z- U1 I6 {8 A- t! R7 ?5 @/ _1 ], j$ p
#打印操作
7 o1 Y" J" T' z- d9 p7 q5 @; l epoch_loss = running_loss / len(dataloaders[phase].dataset)
$ p6 `' P" Q: G& L* X$ P3 ] epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)
" ]! k$ O" x$ X. t1 R% k0 t7 E9 _6 f
1 ?3 @% g2 A" m. { time_elapsed = time.time() - since
# K0 \, a6 t1 E8 m( d; m print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
, s |( t; t* P3 }$ t print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))% [( a* ^$ P; W0 o/ X' f
: r# @8 @3 k$ c6 D/ ~1 X7 o# A6 I' @" V; m: w
# 得到最好那次的模型) s6 Z/ [& I5 m0 c) N( k
if phase == 'valid' and epoch_acc > best_acc:
- h" T: _; p; | O9 N- r best_acc = epoch_acc
8 z8 A" M* G" X3 ]8 k8 m( T #模型保存
, ?5 q" b% V: `; d5 X best_model_wts = copy.deepcopy(model.state_dict())
9 H8 i5 `- Z0 d state = {
. K4 O" v$ [9 m! x2 \ #tate_dict变量存放训练过程中需要学习的权重和偏执系数
9 Y2 |4 ]! X8 [5 j 'state_dict': model.state_dict(),' a }( \2 B; s4 j7 \
'best_acc': best_acc,
2 T; o9 s1 j* m2 ^3 L* @ 'optimizer' : optimizer.state_dict(),1 o* I6 z, @& I% y! e( [
}
" X$ J! A" n& m- Q1 o torch.save(state, filename)7 P5 F% v9 L% |1 A% {* D5 t+ D
if phase == 'valid':
' G" k. X9 ~4 ~# L# V* `, F val_acc_history.append(epoch_acc). K5 M6 s8 n( H/ Z7 v8 B2 I: n
valid_losses.append(epoch_loss)$ u- C ^% r C. y
scheduler.step(epoch_loss)( J- \; x7 ]/ O$ G
if phase == 'train':' w" j4 i! F7 g4 [! |8 U
train_acc_history.append(epoch_acc)( k9 p9 H2 I# u
train_losses.append(epoch_loss)
5 \ t7 R0 M: z# h4 `8 _# w! `- t) N. G" x% J0 q
print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))
: H) r* H3 r: R. @ [, Z0 _ LRs.append(optimizer.param_groups[0]['lr'])7 n* V$ {( h6 H2 X
print()$ K% k" _! L/ Q3 S
. Y. g( A7 c( U4 L" W! H+ P1 }# I- s+ | time_elapsed = time.time() - since
7 i% B. S. g( w% [" f1 Z print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
5 s3 X, b/ K7 e5 A) [1 c4 f print('Best val Acc: {:4f}'.format(best_acc))
9 s( {2 T3 a8 o: T H# J! b
^: q8 x( c" }8 _* S& S # 保存训练完后用最好的一次当做模型最终的结果
( D S" J0 n" y1 O# T& ~ model.load_state_dict(best_model_wts)! ?+ S, T% |0 D- r% W9 {
return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs
5 d2 d5 b3 n7 v" ^9 L
: N8 L0 m& r, G
! ]# k j5 c1 o' {! z: F19 H' c3 g7 S- p- @2 l9 X
2
: E* S+ d7 v# k. U- }3
3 k* L( m. n0 w5 s4 K. i4 V4
+ q: @# ?4 r$ k- Y, O5
" z2 X D. U: n* ~* @6. |; M- I; c% o, W- d2 \
7
9 j( a! | S8 G& i- s8
( G5 _0 E' H& L: D& A+ I9
6 G3 B* [$ c/ }10, x0 _' A. X' O" r, V6 F8 q6 ]
11
0 v P# C% E2 B: I U12/ F5 k$ u8 U( h/ W/ e3 b, \
13
) M s- W) j) \2 T3 K& f H14* j. `+ N0 a7 B4 r, [& ^
15
+ G1 D& N& i$ D/ T" n; @16
: m% N9 l' |$ v1 n171 S5 V9 _4 C) K& e f
18
, R, I% }$ d& y5 n. R19
0 b& G. n; H6 {! U203 i. r% [% ?$ q/ @% I" O
218 \& t* n% s1 q
22
! G7 @8 {# r* @ M8 R9 {1 I1 V; A23
! K# D0 ]1 J, l* N- C9 o8 |1 v24" C# q. y1 x7 j
256 W/ T& Y; W% K( M2 G5 b" V
263 f1 [# L; A. U& |4 P
27' W+ f" V5 s- ~( ^5 }
28
% s. m/ \" r' S& n7 j3 |$ a29
' b' i& b0 |! L4 x/ w* q30' H6 G+ q8 |4 ?+ R; A% x
31) o, M0 c. ]8 Q; G' E6 F: F
329 p! ?- A5 e' E' g& e8 N
33
7 G+ C; X5 C; b B5 }34) v4 z5 H4 W/ g5 C7 C) p& {, ?8 e
353 h) l( c3 f% E) E
36
$ t; G' G6 j$ G3 x1 k+ _+ [37( ^; j6 p7 F7 Y% f/ \
38* ?: L4 e4 T8 F4 P s1 p
395 ]0 w4 u+ Z' {0 }4 m
40& M6 }2 W1 e8 @- V* ?
41
( ?" p2 Z) j* P4 H5 _. Q+ S; N42
# k1 N; K, [& G3 i+ s43
* I4 Z8 m0 o: j- @! ~( ~448 j4 H8 J+ q" V9 B
45" r7 s0 d' U0 n3 Q! e5 Z7 \
46
5 c# C: |! u7 z9 b" ~+ p47) d" K. K; \* E) J$ t: l, m
48
- c0 V0 |% F% P' J5 h49
5 C6 a5 N+ A: \( R" o) a50
4 ~0 O, w0 x2 B8 V# E51
6 G0 i4 u4 [" J* r7 s4 X4 @: A0 \52
! V$ q ?" f% R53
+ e: l! s! v" c- O/ z8 H54& h7 T$ _- P( d9 W2 D
55/ Y% }: l+ {: \. |! E6 f+ s
56
* S! Z% e% k6 ^- v0 A57
: X+ _2 F6 X) ?. T& ?1 g581 M; S6 E; o* Q- }$ \) s: B
59$ p2 n0 z, R' {' f2 y
602 c0 N; V, [. I6 K
61& S; B3 V" H2 \/ ` }* F
627 n' f& P* ~( v
63
/ g& s1 a }4 o+ D64
5 e9 `/ |9 T5 V$ K3 g' C65
2 D$ g* a# U8 P0 Q/ w66
* i/ @" x! e8 I0 ^' \" K" E67) z5 J8 W0 A0 j+ R8 x- l
68: K! P: G% j0 z
69
' S" {/ ?7 U2 n- k703 y+ |& D* a: Y+ v8 _/ W
71
, s: O2 Y( E( [5 \2 D2 D. X% J' ]72
8 `* e. j7 ~% O7 z5 I73
2 ^0 s. W" x1 U g: }74
& _; G7 t5 c3 x, J+ z75- m( [# b) I1 o. I: v6 c4 f$ Q
76
" O2 I! x2 a# @9 m0 Q77
4 |0 [0 J M+ F, A- N787 P: r' p! L7 @; a
798 ?% r% }, U' U2 c
80& X# E" w' l3 p
819 @! p3 l* ~ ]/ X) y/ r% q/ _. R0 w! \
82& ^8 d. m. B& r/ A @8 F) @
83
4 ], Y( h- @5 e) T$ ?# Y840 N9 } J: R* T0 K) S, L' ?" c
85+ _' A# D8 i3 Q, G
86% B/ Q K+ |9 } l9 X7 X
87
' O; j' c, u M# n* q+ Q! O88! {" U8 p8 w( ?, r
89
) U1 w% q# t1 {* B( H90
: I+ P9 B* M5 i @912 Z5 Z2 q; K( C {2 K5 g2 I! |
92 A/ ]$ L9 a6 Q8 e7 M* \9 f
93: h* \/ U( j5 y* D3 X! Z
94 u1 H9 V! Z7 b( t
95
* N, c, a, A7 v0 {4 D2 }( P1 e96
( d% y& g0 `) Q( ~0 {) r97& b3 y) r4 S% R/ u) I
98
3 L; q- H: X7 f7 K2 g99! B+ p K) z c. L- Y
100
+ x! H4 f% n' J4 m' x101
7 w; d, M8 L( f" a102" f" Q+ a3 i( a" `5 F, u
103
( Q& _" G2 U) T! o2 `. D/ A104
* r$ y3 T3 ~' Z- Z; i8 j105. U. ~' S6 o4 Q* I; g
106
7 t3 s2 E- ~8 y4 U: I ~4 |107$ m: j2 c0 f m0 ?8 ^
108
. U; W2 h9 x. L0 M: y! F% P- t1092 h" F6 c! V9 A* b; d
110
6 j3 \& W/ a" P* l) E& w$ c) \111) }6 h/ Z% }5 q4 }5 x# V6 r2 K
1121 |7 Q0 B# k8 d. Z" D8 q5 ~& J
7.2 开始训练模型
& o: I& Y/ @( F' R! p" S我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次
% B/ e p! F1 m7 b- z( c' j5 C2 }
8 b$ d# J. A6 b0 V( ^#若太慢,把epoch调低,迭代50次可能好些
" q# j# R8 J1 K% @8 f0 Z#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合
8 D9 ]% y" X% `, o* d0 Jmodel_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"))
6 J' M3 G, P0 V2 v6 @; W1 p8 e/ R5 M7 T# S) I4 a5 U6 H
10 S! i7 @0 E2 r* R) ]
2
' Y! c5 U1 f b( w, c7 P/ O3
7 P& e a- g0 a6 u( m47 a5 Q/ l8 g7 R* a
Epoch 0/4
% w& ?. l: ?5 J/ Q, D& [( {----------
z/ I, S+ q- I4 ^: aTime elapsed 29m 41s }+ u# V3 U U6 u8 K
train Loss: 10.4774 Acc: 0.3147, i4 U# x: {) G: y c
Time elapsed 32m 54s+ Q! W& c! e/ }" C/ M$ S; ]
valid Loss: 8.2902 Acc: 0.4719
2 T# b! o' N }0 H& SOptimizer learning rate : 0.0010000
8 O8 C; c, o1 f5 } L! ~8 K8 t" e' y
Epoch 1/4/ ?" @- x1 e1 G% ?
----------% |: L' S" q+ Q2 B- |. N
Time elapsed 60m 11s
3 r0 ~. a3 T( O5 r4 Vtrain Loss: 2.3126 Acc: 0.7053+ n9 k9 A0 z, m7 ?7 A; Z# l
Time elapsed 63m 16s
- l: S- Z+ B9 f" |9 ~valid Loss: 3.2325 Acc: 0.6626* `2 R# h) s) c! V8 I
Optimizer learning rate : 0.01000004 _* s |/ j$ o- I4 Q# T2 q1 a
4 w+ ]" n, y, ?3 ^6 @ g9 E$ J; NEpoch 2/4; Q; V7 Y* J% {8 H& M* l6 @% K
----------
" q5 O( ~/ } h5 o+ C$ XTime elapsed 90m 58s
5 `8 }7 d$ V6 _( Y+ B9 Otrain Loss: 9.9720 Acc: 0.4734: C" Y) \* S6 g
Time elapsed 94m 4s% F- Y6 x% G! K7 X) n4 @+ l2 h
valid Loss: 14.0426 Acc: 0.4413
& h# c2 @ t* K! g- L; COptimizer learning rate : 0.0001000
2 \& h* @9 n, I$ Y: g6 m2 ^, t
% B; }* f- U( f" M' y) ]Epoch 3/4
4 O8 i' x7 ^4 O" ?( [- A----------; ~; Q- N |2 J
Time elapsed 132m 49s9 v3 g2 ?. ?: `4 W. g( H Q
train Loss: 5.4290 Acc: 0.65489 h/ m7 B/ A8 W1 ~# Q
Time elapsed 138m 49s
! G8 J( H$ y, \! M; @valid Loss: 6.4208 Acc: 0.60272 ]- ]' Z9 s0 ^# X0 Y
Optimizer learning rate : 0.0100000
/ K( ]7 P. A* [8 |7 u
: \ f7 f! `' ?% v. C- h* kEpoch 4/4/ \8 i' j y0 g
----------
( d4 e# ?) ~, i' x1 zTime elapsed 195m 56s
- h" m) q$ ?/ v' ]' ltrain Loss: 8.8911 Acc: 0.5519
- s6 H9 m2 W: ^) I, kTime elapsed 199m 16s
9 U1 a4 B1 i0 |4 P2 svalid Loss: 13.2221 Acc: 0.4914
# K) N4 n. H0 d0 j, j* COptimizer learning rate : 0.00100003 V# N4 D/ h( C* W6 @
V9 c8 m, F+ I! c# ^
Training complete in 199m 16s
0 E I- c4 o* w/ X+ n3 zBest val Acc: 0.662592
5 X* E2 N1 R) X) b3 |& W o) r8 `" A Z7 n+ q
1% A( a# i* d$ ~( h
2
1 A. X' O+ r. w38 L2 h9 R ]+ o5 D
4
; ^# X; L% d! _0 G6 |5
7 l0 M- B" {! {* `" U, g3 f6
. |. J- ]% u# t7 U, Z7
1 T" x5 }3 A% B( E8
8 |( x2 a7 _9 J. y- ?6 S9
2 W7 R: n0 K/ k109 g; ?' j, U; `% N
11 c* [: A$ O( l9 x
12! J# |( s3 i7 j- h" ]
13- R: P( h4 Z6 n1 F
14
8 I4 Q- v' O0 s' h. ?% G$ Y15 R, I- c& r% i+ P' _
16, M* \0 @1 @# D% q* h
17
& m1 R; W T+ h& E9 {1 `18; U' D2 D! C' l+ \4 {
19+ X# `# e/ _4 ^ L0 O" t
20
4 m" B8 J( H7 a4 p) f21& d6 Y, M/ z5 Y" ^
22& D4 j$ l8 l# b- F
23 P5 b$ [+ L: l J0 \& S
24
1 v0 f9 ]/ I9 e1 K( U$ c25$ w1 ]+ l- O* X Z3 n! q0 B
261 k3 c* Y5 `4 `' }; `. k' q
27
/ \ `8 i, O) O8 H3 t6 L28+ a6 v# K7 Q' G0 @% b. P3 f9 m/ S
29% U" u; ?( W$ O# V. A5 R" N& ~4 l
30: x- M; z; R- e4 a
31
. [& ]2 u: m) Z* z) C( K9 } F32
0 I) Z X& ]% ]4 F+ o& v333 X! W# d7 ~( U- ~ Q2 ^( h4 V6 k
34
5 N% k6 O3 J9 [ t35' J" Y! s5 ~6 M9 Z& w% z
36! T# E" T M/ R$ I- Z& n* W
378 M" O! v* R/ D
38
9 c9 I% b$ B- s# E399 M6 y/ _/ [- F# f9 Y$ e* |
40
+ H- J( ^, p' [. W8 ~41& X* ~/ C _2 e/ q% {% A
42
4 s, e, Q- q5 q* J* v7.3 训练所有层+ _6 |$ g3 r$ m6 Z a
# 将全部网络解锁进行训练
: {$ b' b8 Z+ B8 Sfor param in model_ft.parameters():
/ \6 x+ J" R! v \) z$ P, a param.requires_grad = True* N! e4 t) O: }- L
b# T$ X! |2 E! j5 h1 f0 j5 B; Y9 I
# 再继续训练所有的参数,学习率调小一点\& F( r4 C9 v) {% D8 b! w0 a: O
optimizer = optim.Adam(params_to_update, lr = 1e-4)0 F8 R7 g: Y) x6 F
scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)% Y) f E! v/ V4 G* G' P+ t
h9 X1 V/ G$ ]& S4 S' G' N6 T& N' N; `# 损失函数
/ |3 {* `+ E) l+ Y& C0 B3 U' a0 jcriterion = nn.NLLLoss()
/ z P% B& a A! H1
3 r1 [) R+ q6 B) U; k( c21 R4 ^7 o$ i# L6 W) r4 h7 G
3
5 |" V2 R! p8 G% l4
( G& X, u8 D1 B8 B4 w+ I+ l55 C+ D" a: K6 \9 J2 [
6: V4 a1 ?! w4 s$ j$ N
7) ?) {1 O% E, t
8
0 l: {2 _: @3 a O }+ e* O! J9. ]0 M/ j* A4 `( n6 c
108 k+ U0 T6 F$ z+ a
# 加载保存的参数* X3 x5 J: {* j5 {
# 并在原有的模型基础上继续训练( E; s' O9 q. ]# s; [' I& y
# 下面保存的是刚刚训练效果较好的路径
0 l- ?+ i1 [/ o0 T6 k6 l5 Z" vcheckpoint = torch.load(filename)
! W; H s. q2 j- o6 Dbest_acc = checkpoint['best_acc']
0 q/ w% z! y# v4 E( m- B4 Zmodel_ft.load_state_dict(checkpoint['state_dict'])& u; l3 m5 U2 M1 ?
optimizer.load_state_dict(checkpoint['optimizer'])/ n9 I) e8 e- V5 M
1
1 d+ c; s" p$ P/ {# Y2
5 E4 {% K# }( {$ L) e: ? M+ Q3
0 n$ a, U5 s5 A6 I4
: }: R3 C" U' E4 s/ N9 E5
/ s0 t/ w* ^, Y5 u! c6, v" O3 L; w7 Y3 v
7
, v! s. {2 g, h" t2 f0 p开始训练
/ A3 j+ ] p7 x1 _7 ?9 A' F4 a注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考6 V% k6 a. z# {& u9 u9 ?
/ u. l" U+ ?/ A( T. Gmodel_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 b0 y7 q0 t9 t( D1
/ Y- i2 B0 |% {. E& Z/ LEpoch 0/1) j p) n7 Z0 g; K+ k) z2 e( l
----------8 h2 y& Y7 S) |; [$ R- \9 {' w( U
Time elapsed 35m 22s3 `0 e/ u1 N0 `: |& ], ?, \
train Loss: 1.7636 Acc: 0.7346
+ i# a4 o' ~* @8 t' h c2 lTime elapsed 38m 42s
8 X6 G9 z) x+ a4 nvalid Loss: 3.6377 Acc: 0.64552 j! B$ I& |5 u: q0 P
Optimizer learning rate : 0.0010000
; H' x- c7 Y2 j2 j* x! c" g$ H+ A7 j2 K
Epoch 1/1$ f1 j! c7 W3 q
----------
3 G5 R, R, }0 `Time elapsed 82m 59s$ S" @* I) o; R' s F* C
train Loss: 1.7543 Acc: 0.7340) ]" U) c9 p# C( G9 _1 U
Time elapsed 86m 11s8 ]% R3 p8 B B$ @
valid Loss: 3.8275 Acc: 0.6137# i1 f9 P' R1 X* @9 q- W3 P+ [
Optimizer learning rate : 0.0010000
) d# Z, S% M4 ?5 q3 T! _! X2 \, ?7 A( B1 O
Training complete in 86m 11s
; @: R7 K' ~0 ]/ i2 BBest val Acc: 0.645477
' D' f( e b+ M* w$ D
# ]( I( |# _5 z! L; _1
, l" i8 K: |8 p/ e7 w! e' A2" I+ V8 f5 h6 B5 t0 P
3/ t5 E' l- L. J! {
4
8 e; ?6 G! z- p: k* |# |54 v: D# l1 i# X' J o2 H
69 f' m; U! e4 \, ]) M, m! f _( f- I
7
5 }) `- K0 S/ A( y& c8 [* N% s8
: ]( s. Q/ J) v6 m" h2 O2 K9. k% }% I' J+ Y. t7 _, J
10) ~0 a4 P1 |/ G
11
k; y: z: g f+ ], |$ [- J ^5 P12
( x* s8 R: n* O13& ~3 k7 E z2 z' ?0 t, A
14
v7 d0 N) B1 N# O% i15
* l0 y4 `3 H7 z7 K8 ?4 W3 K16
5 I/ ~( {0 u y& J: U, ~17# O6 w, o1 r; ~
18
5 i i" N6 p* V- Y" x! n8. 加载已经训练的模型' i# A) y- Y; Y/ a
相当于做一次简单的前向传播(逻辑推理),不用更新参数; |, [) m" _* E: L& m& |* c
) M2 L! u5 X# V) F* _( p( |2 F
model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
, R+ g0 F* B$ c# ^
# m4 V q' f7 _) T1 `# GPU 模式
2 M! n: V! K6 I. |$ o" Z, p- u+ i. h9 imodel_ft = model_ft.to(device) # 扔到GPU中
- P8 u3 N& v6 e7 Q. t, N. }$ }! w! R) V
# 保存文件的名字
, h* {: d2 V8 f3 l8 ofilename='checkpoint.pth'4 k8 t1 F! J+ s/ m1 ^) o
- b4 F5 V" ^2 ]
# 加载模型
6 r: E. u$ i# a2 ncheckpoint = torch.load(filename): k+ W9 f X# f9 L d7 B
best_acc = checkpoint['best_acc']
) C. {1 w: H8 A3 ~5 }/ Y+ l& Fmodel_ft.load_state_dict(checkpoint['state_dict'])
8 C, p1 z: }2 }" u' r1
9 ^* ]( j; g: |; e2
1 `5 z8 c, t" J2 y0 J$ g3 N: Y7 m. \3
+ T& q7 O. m5 e; ?6 x41 G3 h) C5 K* E W
5+ ]7 \1 i& C" _
6. W. |& S9 |& j4 u( A4 u* Q. T
7
& X: m, I) T3 A8
+ U& w n# G+ ]1 N* O* r9
+ [% P# ^: X5 `( k2 `# r) P. e10
6 T; D: p5 k3 H& m11
6 x: O& d0 B9 `12
# ?6 D6 O0 T* J1 S<All keys matched successfully>! `8 {0 f0 g1 f
1: T6 @/ x! @5 h( X+ p5 Q
def process_image(image_path):
9 |5 B- s" ^. k # 读取测试集数据3 {4 z7 V" j X) a
img = Image.open(image_path)
' j/ J$ \6 d5 |1 y4 v% A& o3 C( Y2 m # Resize, thumbnail方法只能进行比例缩小,所以进行判断
4 e+ X B2 F$ {% P' H # 与Resize不同3 F1 o7 h- _* H5 l
# resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小6 }' L9 L) |% k) S
# 而且对象调用方法会直接改变其大小,返回None
7 h' c& R4 U1 e" m' S$ l if img.size[0] > img.size[1]:5 W- Q7 }5 N: V( D
img.thumbnail((10000, 256))$ r& z% \; I0 P1 z$ a r
else:
$ X1 ~' U( v5 O8 Y6 Q img.thumbnail((256, 10000))
$ S/ N L. x \6 d# L! K* R. E- K+ L! X! Y+ ]
# crop操作, 将图像再次裁剪为 224 * 224
9 P' ^! ]0 v. Q" Z left_margin = (img.width - 224) / 2 # 取中间的部分
1 ~- {* A" c- O5 i* W2 x bottom_margin = (img.height - 224) / 2
# A4 Q7 q8 C% v0 G$ a( z* S/ U& }' y0 _ right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度
& m# ^; r" Q( P; N, ?5 K top_margin = bottom_margin + 224
, t, \2 \+ L2 A7 G+ e3 h- Z
' u. M0 S6 w Q/ Z% t- C# ?& S img = img.crop((left_margin, bottom_margin, right_margin, top_margin))
) d3 d; e7 w) _7 j1 D: y- z' {) Z- s! `' x) U
# 相同预处理的方法- x# x# V, ?# _) \) G5 w
# 归一化, @2 o( p' ^7 \( i9 b E x
img = np.array(img) / 255+ _) l4 W [$ }- s Y) y% ^2 L' g$ U
mean = np.array([0.485, 0.456, 0.406])# n" f" F" T3 M' A8 M W; c( @
std = np.array([0.229, 0.224, 0.225])- ]% m( b5 j) R* N) i0 C3 d
img = (img - mean) / std
8 P' [. [6 n3 f' t7 M3 r1 R6 G
' `; x! [' n; S7 M5 U # 注意颜色通道和位置6 `% w; W' E5 m. u: x6 k1 k
img = img.transpose((2, 0, 1))
/ k: Y8 s5 A7 ~+ o4 W% [, ]* m# h% L* u1 J+ m/ B4 W: M
return img5 _) }( |9 Q, g' s, i% D- ?7 g
+ B: A4 w. v) rdef imshow(image, ax = None, title = None):6 o2 B$ S2 Y& a/ w1 b8 _4 E9 F3 t
"""展示数据"""8 ]% `. w! n! u5 C! l$ R
if ax is None:
7 B; P2 `" g# v k( \ fig, ax = plt.subplots()
" h; \8 l" O( m) q4 x G# @# B$ d' h$ f
# 颜色通道进行还原
" M; I3 q3 M3 C9 v8 @ image = np.array(image).transpose((1, 2, 0))1 z- n; G' `5 M( r4 b( m0 ^1 j% Z+ X
) a" X) F1 x1 x! }7 a' p. n
# 预处理还原
; G/ E$ S' C7 L& @5 I5 p9 I mean = np.array([0.485, 0.456, 0.406])
$ I; O [# H& z# h9 k u std = np.array([0.229, 0.224, 0.225])9 P! e/ K- D+ E, E0 t, x
image = std * image + mean1 W% q% Q! f: e9 o
image = np.clip(image, 0, 1); m* A- A6 V/ v5 J; T) g- l+ H9 V
) C* ~, e7 S! C5 d) p2 Q0 v ax.imshow(image)' X) b8 `0 v- n9 l. B
ax.set_title(title)
& p$ ?9 \9 }* f# b C1 ~; {" n6 i# a3 ~! Z/ |; M/ g, @5 m
return ax
; R5 U* w1 H& _7 z) g3 U, a, o5 U: p! v- W; r- \6 Q+ {
image_path = r'./flower_data/valid/3/image_06621.jpg'2 p' E0 M) K& y
img = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理' d% T7 H8 t0 b
imshow(img)
5 y, z4 P7 C l; B3 D# k3 y6 c
7 a' e; O' \ g( L$ B6 x% s13 q! K% l r, a8 W
2
9 S9 v# V$ c- j3
% ^! i; u( g7 X7 K4 \6 h/ u$ o) p: {
5
" j+ E0 L+ U. }" }- Q# E+ {6
! V+ W9 V# {; K% J z8 r3 t7
) W) k& ?+ ~1 D8
) s9 c. H+ A& S' F1 J. x q) i90 k* F9 A$ f+ A" e
10
, {% X- y+ G# x. b- r0 q) [11
0 U* S" @# q, T) U2 T9 V( g7 L125 J. V) |3 W% s: U. L! f% t% S
13
( r L+ l# `1 v, v14
% Y; c- t1 F6 e15
" I% H: q' j) [* D167 \/ ]. o/ b1 T5 h: K
17
0 ~8 }1 `$ j: l186 u9 Y6 R- Y3 V; R l x) |% F" h
19
. s( H- @( j9 W& D20+ ]$ \- J: Q. Y
21- c7 O4 |8 O, d# }/ j8 Q9 e
22' |) t1 S A% N) Y9 F
23- r( k: t9 y- {4 G3 P! g+ \$ d
24+ H" D/ T, Y% a/ ]0 x
25) }" v5 w5 J4 |+ Y
26
! u% I8 z1 s& m ^! o; L27- i* J1 Y- f- c$ D
28
; K: N7 K7 X+ e' v292 d+ I; X" B: p+ }
30
7 i0 o% N; \1 v31
4 {8 f& Q5 E$ C1 z& e# z32
f. f" q7 C% g6 b$ T ^9 G33 B# X! A- _- i+ X
34" V4 l# s; Y- C
35( M! o* L, k% ~" m+ f
36
; D0 |" |' d2 a& G7 h5 y37
; Q6 Z8 D3 l/ I. I/ H3 p: a388 \+ f) @+ h1 N& d
39. J L7 |* B2 ^ [0 Y. N: ]& v" O8 s
40* v0 _6 d# V8 x9 X4 @7 H
41
; Z) k8 O* S$ G2 g2 n/ y420 D, w0 H N2 |; x2 q6 R, s' v* _) [
43) _5 V* z. D M7 G5 {( c N
44- B- k1 r- ~4 r8 r. U6 ?
45
8 b/ y( O+ h: W! ]46, l# T3 B& {; i Y" Y
47; K3 B+ z: A' N0 B" @! U8 z: _% f
488 s# l6 V- f& z" u' P2 O# N
49
+ H7 X! D L9 B- |& D& D7 I50% m. ?1 a5 o7 L+ U/ v: I
51: ^! r# p9 F& I& q1 c0 |
52% K- S* I2 S4 u/ i; R
53
0 x X1 l' D+ w/ G2 c54
% D. s4 B( [8 Q1 `<AxesSubplot:>( [9 G+ h$ n8 @. J5 r& P+ X; k
1+ B a; ~2 e- P8 a
- O: v4 M5 D" @# p6 o
上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确/ m, ?' {0 e. F
. [& t, J j9 t5 Dimg.shape
1 `# X" O# ?* o, V; }1 S( H' b1
; ~' |' E5 w7 `5 c, t(3, 224, 224)* z$ l' a+ e/ O* x) a/ P
1
$ ^4 T8 K: M3 H1 L/ l证明了通道提前了,而且大小没改变
, K* H7 K: C( ?" r
: A& Y1 h y% O9. 推理4 I% P7 ^ r' }6 I6 @7 P
img.shape
8 E/ P& i; r4 ?2 n6 z4 n/ c- g
V4 @( V! D+ t0 x' p; y% G# 得到一个batch的测试数据6 F3 p( L6 J3 j) e/ p! p
dataiter = iter(dataloaders['valid'])3 W& o, T2 X. [# N7 s! a6 X
images, labels = dataiter.next()3 W) E) c+ ^# a, e& A+ W+ R. J0 m
5 v, F3 _) h v. k# ?
model_ft.eval()
2 A( B* Z- j3 E4 {1 v& k9 F. w, O2 \- G c
if train_on_gpu:
5 T: V1 Y$ [5 C6 K8 i( _0 D # 前向传播跑一次会得到output
* O3 k0 @4 t! \ m& T( U output = model_ft(images.cuda())
^) q/ U! M9 |' a* {. y& H1 velse:
5 ]( x6 ^0 O5 F6 x9 } output = model_ft(images)
" Y/ C- J _. c) x1 z8 |2 q5 f7 h; k$ L
# batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值
2 p: F5 W1 I- f) D7 `! @6 Goutput.shape! P2 k, z+ V) H! F" p
; N% f9 v' O4 @4 U7 T& V
14 R. m- Y0 n1 _+ f8 S
2
2 s4 f( ^# S/ F/ x" w/ j2 F3& E$ o+ Z. j& K# p, O
4
( g, U$ p. l/ d" Q+ d5
" J/ K' z5 i* B$ W! a0 ~9 r( U6
3 s4 f M, T7 M% g3 F% e# G& @7# z2 D6 U' _7 @7 @8 f
8
2 j8 [- D! R% q8 W4 _/ C9
/ Q- u# l. m* t0 o, G2 t10
f- w) Y1 J( L5 i# g7 |5 O11
0 Q0 n1 D. I$ e12' p7 S. U" f1 U( k
13" }3 |9 |5 [+ Y' b0 U& d
14
7 \; \( a; d- H15
1 ~5 f2 B2 V# F, W16% {' {* K5 v# f3 w& k
torch.Size([8, 102])1 [# |/ K, Z* B; l6 _
1
- f+ t" O2 O" `9.1 计算得到最大概率
$ n: s: {! n3 |5 u/ f_, preds_tensor = torch.max(output, 1)4 I0 |. D& g2 m8 [
) N$ w( @, z8 h# x$ O' Dpreds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量0 C2 ~2 B V" e& y' ]* J% ]% K
18 ], R/ p$ r6 _+ L+ s7 r$ F
2
h* ^) Y" X0 ]% s3
; ^ m; l* y+ D4 A9.2 展示预测结果
' A5 b( Z5 k8 x( o2 F0 {' Y$ wfig = plt.figure(figsize = (20, 20))9 `, r8 y$ l9 D1 [
columns = 4; G3 V+ c) B B/ k. Y
rows = 29 `: k4 g4 {9 L4 b, N5 }# V" K
9 `7 G1 Z$ f# }2 a2 J' |
for idx in range(columns * rows):, M6 s7 }+ @, G+ S- E0 ^ K
ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])
0 q' j. |. R q plt.imshow(im_convert(images[idx]))% }2 `5 Z0 o( R; I. k" A
ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]),
+ l" v7 u6 D- q0 V, w3 g M, b color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))
( k3 v- a0 G( p' [, n% I; z1 Vplt.show()
: W8 T: H' Z( p3 `# G# 绿色的表示预测是对的,红色表示预测错了
6 S3 r, K% [. Y8 w2 C2 R1
% A/ B% d# p) k2 r2 B! B& n7 X2' B& m+ s0 r" u- \& g& Y
31 Q/ q0 J& D* f9 j6 t
4
* S* s( c$ ^$ E+ c' K) O; D59 s5 t( ?% j7 V) E+ A7 H! d% k
69 z5 y' Q3 u# n/ Y6 {
7
& l7 c) R1 h, j% u$ [- _8+ _6 C+ Q5 Z( {9 a: |7 X
9/ \) r* h$ S; ]& q. U4 h/ c% h1 g
105 x6 T7 p7 K N/ [* _0 @1 ?8 x
11
. g$ f6 G! O$ f/ \9 [' G2 @) _5 j. A w Y
" r' d+ q8 ~7 p$ w. Q% A7 T! _( N* E b; q* h7 c+ X
————————————————4 i$ x, X1 K7 y* R- J. A
版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。
- `( u' l$ h2 c: J D原文链接:https://blog.csdn.net/LeungSr/article/details/126747940
" F; n: j3 x- U7 w! L6 L; ~8 c: A$ R: a% \. p3 ~! t
+ O6 g0 l% J+ X |
zan
|