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