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