+ ^' J9 C5 a) e( L9 d # A# K0 R& h2 t/ H( Q1. 导入工具包 ' N9 n K* l) Z% G+ |. R) @import os # ]5 r& F) Q' X4 J: u3 e! ]1 Eimport matplotlib.pyplot as plt( M% ?9 k7 x+ Z) K L
# 内嵌入绘图简去show的句柄' b8 `, G$ v" j* x1 }7 ^
%matplotlib inline / u6 X& [2 ~. X5 U' I/ L2 G0 S* _import numpy as np; X3 g- d8 g0 `( e
import torch 6 s/ v& s; f7 u0 ?& x5 \' l4 N0 pfrom torch import nn1 m0 u7 m; b5 S- s/ j1 ^4 c# I
. D! W& H: i: B5 W: x5 ^- Aimport torch.optim as optim3 m" F5 B9 c; x+ S% R: t
import torchvision& v, h% Z9 L9 X @9 Q# D
from torchvision import transforms, models, datasets5 m( i' N) H% w
. ?, ]( {! p' D+ Q0 O0 L$ wimport imageio , M8 ]! c1 p0 |2 K, Yimport time, J0 Z6 b: U; h0 D; I5 }" r) m5 g
import warnings" O9 p' _9 s3 C
import random 6 B; g$ k0 K6 s8 m4 Jimport sys' e2 Y- Q4 H" O
import copy/ B8 n$ l, h- G5 X2 T: C
import json9 C$ X8 S. G, P$ u v1 j# {
from PIL import Image: j) y/ U: M r/ d, U
* I1 _7 v% O' d1 v& v2 k! g$ V' H% T$ h3 W3 y4 L, X& g$ d
15 d% C7 _+ x6 }9 ^5 F; J, Y
2 * f; L) x& p. @1 R4 W7 D3 c* k; \# r& Y" d2 p4 + A) a8 \2 p; x7 C& Q& b4 Z* \5, H% q! y3 l- w
60 A, G: h- q( g" i
7 q' m, }- G0 @3 \
84 q( j' V6 H4 C+ F; c/ }: Z
9/ p Q r+ x6 B5 z# K/ r& x
10 4 `' N# T- \8 _9 f11 2 u* S, e: x7 `2 q1 T: K4 q2 l' q/ d12* ^1 O" V, W3 J: f( w1 Q6 p" K
13 4 k7 t' t( Z5 b; A9 c0 h14$ W( y- L/ A1 [( W- L: Q
15/ Z7 `6 {4 E4 S: P+ M' j% v
16 e: D2 a* ~9 b; l$ ?/ \0 [17 ' s! s$ y! k, @6 [6 b% R- T18 ! V! E: m' U9 d) B! ~19 - o' D$ [, z- d: l- [( U20( d4 M Y( y9 S6 ?
21, `$ h) u, F, W% k$ p
2. 数据预处理与操作9 g/ ?8 h0 ^0 ^6 `# \1 \0 z
#路径设置 5 S& x1 X+ Q# A- ~- Cdata_dir = './flower_data/' # 当前文件夹下的flowerdata目录( j; s* ~' @4 I. O* x! x
train_dir = data_dir + '/train' - i- B' u$ y& g3 D5 ]/ z+ ]6 g9 a/ yvalid_dir = data_dir + '/valid' 7 W4 x3 Z: E& C! |: |1 C" h7 \ p1/ b. f9 ]8 e8 H
2 , I1 V" d/ j* t+ Q# T/ O- u* ]9 S3 . ]% U3 P- M6 E& H2 o# s2 H$ `4 ! _* Q7 V5 o( mpython目录点杠的组合与区别 ; Z: \! m1 @* Z6 U0 W8 [8 o4 Z注: 里面注明了点杠和斜杠的操作 ; @/ {' b* J" f) p, L% W" Z9 S7 L8 ?) t6 y, @
3. 制作好数据源8 I6 o! B' r8 ^
data_transforms中制定了所有图像预处理的操作 5 N( T6 L, s6 f. ]8 t1 Q4 e: QImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片3 l" J, z0 }6 ]( t% E
data_transforms = {% g1 P+ k4 S( H+ Y. u* O( W+ J
# 分成两部分,一部分是训练 1 a5 m$ a( @: W* C: f3 j6 b0 T) G 'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间, W/ ?- ?5 a4 ^; J( S5 g
transforms.CenterCrop(224), # 从中心处开始裁剪 7 ~) D" t& d5 O S% q8 D- P # 以某个随机的概率决定是否翻转 55开. J0 R$ K D4 ~6 P1 o* K& n# w
transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转$ P2 k3 ~& a" B' f4 b
transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转 3 F& D" R5 J- M2 l% O& _ # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相 8 R1 `# L0 K; f8 \4 P2 l! z transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),6 J* t+ z+ J% U" |1 c" |
transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB5 `* Z/ q3 C+ B: w
# 灰度图转换以后也是三个通道,但是只是RGB是一样的2 W" S' ]6 a7 L# Q- s9 e
transforms.ToTensor(),# u$ L! b* I1 N! b5 E
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差 ! p4 s O1 D" ^% S( `7 ] ]),. @7 _1 T' Z5 |* j% K9 }( e! A
# resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化 7 i: i1 J2 e; p. D+ [/ v 'valid': transforms.Compose([transforms.Resize(256), # x% z3 D$ Y1 J& P$ b C transforms.CenterCrop(224),$ z1 k- S& D$ A |/ l: u, A9 ~
transforms.ToTensor(), * m( T/ @# {/ ~8 W5 \3 f+ w* k: H transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同( H. w9 p( @; b; M$ v# {# Z! P# A
]), $ z% o/ R* `% t& ~7 X! `} , A) @% _1 u( X) L+ d" ]9 X! F . i- ^: M2 u! C8 t8 F1 # p0 p' v5 a" ^( j& _2 0 O2 z4 { E7 s6 _30 L! C. ]7 v v0 [# M( q
48 o* e3 F) Q6 e$ \4 |8 j) }/ c
52 C E- D/ o: r7 Z2 J
6& M$ a+ o) s" S" _. R- l4 o3 }
7" B" @/ i. u1 I! H# o+ w6 h! ~- p' A
81 {$ f. I' U( d2 J
96 P% j c, D; x5 } A, y! V s
10 4 Q0 m( U, Y, }. Q& }! V111 w6 S8 p8 h; A4 B% Q! x# k# L
12 - u: V+ y8 [# {5 I( L+ J, d13 9 q" s5 t- q4 q( C14 : l) V! Z9 r( D! H9 C' x" D3 _& M15- Y ^7 r; _4 v: S5 q) [3 {3 ]2 c
16 p' B* R. R/ S3 q: P
17 ) ^$ U% L" w' ~9 z4 L7 K3 l$ V18# }4 q( M. ?0 k( u# s
19 5 a3 s c6 Y8 c203 q; ]( H( N$ @. x
21, p( K. e3 S7 ^$ A2 E$ {
batch_size = 8, F$ w9 k. `6 H6 R/ H [2 \ j
image_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']} 0 A2 \$ n) d( P% S5 |+ \6 A5 j% y2 \dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']} : u( j; F8 G! ]0 z- L: A! idataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']} * n I P: P: \* w
class_names = image_datasets['train'].classes + F. n' z) u, T# h3 C [8 C" ?9 F2 q( N# |" t
#查看数据集合$ m5 x" B) p- U7 ^5 y1 n5 p! G
image_datasets o/ j4 U3 \# ]: w( v, T& I , G4 h" H7 A) r% ]* Q Q1 \1 k1 5 l2 m7 S4 A& Q' S' k2 R r' e2 3 v7 @5 Z; r6 A) C5 t3 1 T; h- V8 ? X4 G0 N7 h" B9 |58 f2 D6 }! i, f
61 f O# A! E9 E2 ]; n
7 + |+ o8 Y: [- l8' `4 F; ]7 N `7 I
9 " X% V ], P% t l{'train': Dataset ImageFolder 2 c8 S) ^$ F9 v Number of datapoints: 6552) K- a! P/ l" X* G+ H
Root location: ./flower_data/train4 n; c) o2 D4 z, X
StandardTransform ; M4 S5 T; |2 }7 U( E# M6 l4 X$ J ? Transform: Compose( * L# j4 w' M& L& | Y# K- a RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0); L/ `; l3 Q F# u: r5 l* L/ [
CenterCrop(size=(224, 224)) ; e: @5 n. _* ^' o. }" o5 V \4 ?1 O RandomHorizontalFlip(p=0.5)& z' V+ s: n G! c5 p3 V
RandomVerticalFlip(p=0.5)2 w4 S; }2 {' ~; U. _) _
ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1]) 9 T% f4 a( }$ ^/ z: B5 Q$ @! K RandomGrayscale(p=0.025)9 a4 @! c1 {6 P, T% P% a: y& S
ToTensor()" I* Q2 J: P" S/ v
Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) # n5 D2 e/ B4 x7 z3 D% { ),1 a& _% G' l! c& F1 K
'valid': Dataset ImageFolder 0 d" Y/ U+ N2 R: v1 a: ^ Number of datapoints: 8182 a1 k4 Z3 ~% |% p+ E
Root location: ./flower_data/valid 3 W; T0 r1 H/ D6 z. u0 x% w4 M, `8 { StandardTransform3 W# q: w7 v) J8 p# k
Transform: Compose( 9 E. l1 r1 P* [$ ^6 R- w Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)# |) Y9 c% X; p" e- ?" J; I
CenterCrop(size=(224, 224))' U5 h8 Y3 V$ ?& T+ [* K( ]$ K1 n
ToTensor() f, M- T* K4 z
Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) 4 | M8 u4 _- c5 `2 J# y1 v )} 4 i7 d' W- k! W$ E0 ]7 |! t3 P, E; \
1% w" a4 _6 c" \+ M
2- [2 u! g/ u" e$ u2 P0 f
3 r8 d0 |( n7 n8 `" { |; P) Q4 4 y2 }% B: J! F I4 B5! r- T* _6 R. }. J
6 ( p% X5 n. X: c X2 N/ \79 b7 z" M6 l4 J' E+ A3 S
81 }7 U- o6 T- S+ J# m% \: y
9 ! b7 M; D' m3 X6 M& p, ?2 Y102 j" W# ~- C9 H) ], i, m
11! i2 m9 W& W9 X. a4 u3 w* s
125 {+ E$ [( |" a. F6 `, x3 F
13 2 D. z, H; W* r9 x9 E2 J. g( T14 " f; ~, L6 W, ^2 k% {156 A9 M% j# j5 u! ]1 U3 S& q6 `( ]6 p
16 3 y9 ]" K$ o9 e! j17 5 Z+ A5 v7 j6 }2 i/ q18" z- f9 |- a- `3 u% \
19 % F( p3 c* }) ^$ [8 l3 `( L p20 # ^6 Z& V C% Z3 @ r213 f: `' ^3 f! x+ M
22! H; l! U; ?5 x* {7 o5 M0 F
237 X/ a3 N/ d0 Y7 q, o
24 " H1 C1 t( s2 m5 k# 验证一下数据是否已经被处理完毕 6 t* C. k, e0 T4 `' Pdataloaders 1 u: r7 N( u8 e9 U0 c' U/ T9 |9 M q1 * O0 C+ e& `# U20 [: Y( z; X- I, [8 a& |
{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,: P: M0 @- X6 J( v9 D" n" {$ I6 E
'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>} 3 t6 [' W t- Y/ o& [ }1 ! p h9 O) e0 j- |3 o+ C: S2 , J* @8 O" I3 L p6 K3 B- ydataset_sizes9 W6 `8 d* z( C: X
1& B& O* G2 [5 M: C) Z1 y2 i
{'train': 6552, 'valid': 818} ; s( E9 X% J i. K- x13 ]7 b q" o& ]; u" N" d
读取标签对应的实际名字 & J1 q/ E% u( C) G$ y/ c" f使用同一目录下的json文件,反向映射出花对应的名字 % C) ]# t+ G$ K, u0 N' }2 Z & J2 j, o* [8 i! _1 k# Xwith open('./flower_data/cat_to_name.json', 'r') as f:* m: m8 L- k9 ]( O* Y1 H: `
cat_to_name = json.load(f) . K+ w2 z0 J3 N8 F14 |) L; ^5 p2 Q6 q
2. O; k6 ^) l) d
cat_to_name2 E3 {8 L' |; O1 h, @* a$ G% Q
1 $ q9 g( E% b- D! D! @$ v{'21': 'fire lily', : z# j5 d7 B6 a9 }6 L% Z '3': 'canterbury bells',$ O E5 I; I6 a- D3 N+ K
'45': 'bolero deep blue',$ K) w$ k8 n" e& {: y
'1': 'pink primrose',& s5 v- R* C( i8 ~& s' t4 Z
'34': 'mexican aster',5 R9 Z+ e' r! p) U, X" d6 y; z
'27': 'prince of wales feathers',2 A: H' s$ H- e+ J
'7': 'moon orchid', 6 c# c$ c4 D6 ~! w9 q- p3 P2 N '16': 'globe-flower', + Q/ w0 L1 O/ O |8 z& q9 q '25': 'grape hyacinth',+ i5 ^" F! K! ?3 s- X/ _8 |( j
'26': 'corn poppy', $ g8 s8 l, y) D& s '79': 'toad lily',, @( j7 p8 D# f4 B
'39': 'siam tulip', 2 r$ [0 \9 S f' | W H# G/ h '24': 'red ginger', % h6 {! P, U! H: E) j7 j' U j '67': 'spring crocus',8 K$ v" k& W B
'35': 'alpine sea holly',+ @' r& [* j) t8 T
'32': 'garden phlox', * I( i6 h7 B. d4 i$ G- y7 \ '10': 'globe thistle',; S' P+ A9 Y1 j P, T, c6 _; }% z' R
'6': 'tiger lily',! A) p+ M0 K1 f. w1 C6 x$ ^
'93': 'ball moss',3 o- I8 C( g, ?( A6 t9 v2 e
'33': 'love in the mist',& v3 Q" h9 ^ c9 |/ c s
'9': 'monkshood', : Q+ c7 L! R) c _5 g '102': 'blackberry lily',. I* |( Y7 H+ e: ]8 I
'14': 'spear thistle', 3 Q9 `) q/ d' @6 e( G( p6 Q/ e '19': 'balloon flower', ; ]2 W9 C. e" s( S1 e, @% M. j '100': 'blanket flower'," R$ D6 F0 p$ U) a! p* d# a
'13': 'king protea',( O& B1 v) f `9 U& v
'49': 'oxeye daisy', ! I V5 b9 r) {3 E$ | '15': 'yellow iris', ' i2 T0 n% u! J* g: r' x '61': 'cautleya spicata', ( U4 b; g% ^2 p7 [ Z '31': 'carnation', 5 x4 x2 A' `, g Q; J; o g6 q '64': 'silverbush', # ~* |* P1 g7 l F# f9 i '68': 'bearded iris', * T4 z: j$ d; f5 A- d3 z* ?7 \) ? '63': 'black-eyed susan', 5 C h* }1 h; ? '69': 'windflower',6 H- ?7 E, }& Z) { A
'62': 'japanese anemone', ) N$ ]! H6 z' _% g" e5 { '20': 'giant white arum lily', ' u- y( w& P* E9 X '38': 'great masterwort', 5 R, X" a' U1 F/ z( s '4': 'sweet pea',7 Z- H* T+ ?$ l. g7 \! Q* f
'86': 'tree mallow'," m0 y8 i* x) a; d2 k
'101': 'trumpet creeper', : y) X9 d( T1 f# V" i '42': 'daffodil',7 V# J$ B0 a x4 P
'22': 'pincushion flower',2 @* N6 ?$ |) O& ?
'2': 'hard-leaved pocket orchid', 3 n. e5 G2 G& f. V '54': 'sunflower',% f, g6 ? u$ N7 g, b/ T" c; m
'66': 'osteospermum', 6 G+ E7 ~2 Q: P6 | '70': 'tree poppy',+ Y- t1 C; |4 g: V
'85': 'desert-rose',% T, ~- F! I4 u. x7 o, l- [' V
'99': 'bromelia', 3 }( F* o! S$ V2 l$ E! Z$ U1 V0 q '87': 'magnolia', ( m6 C( w6 v- d$ m/ a- I '5': 'english marigold', 9 U4 n8 W2 F' e: |7 `! r '92': 'bee balm',9 v/ A& X; _, X
'28': 'stemless gentian',0 `+ o. {) V& l
'97': 'mallow', $ b6 ^1 B& p" U '57': 'gaura',, X5 B' u& ~( t1 c* X
'40': 'lenten rose', 1 t0 R" a- o5 Y/ h( _9 ~ '47': 'marigold',: ~* u7 \2 y) }, w0 D& W; a; `
'59': 'orange dahlia', 6 ?& N- f/ y6 `7 i8 Q4 T) E! F( X '48': 'buttercup',; G, b( N3 G/ h, }6 O' v: J
'55': 'pelargonium',) p4 `; Y6 m# u" H
'36': 'ruby-lipped cattleya', . c' n( b7 l: A, {& n& z* x '91': 'hippeastrum', 8 ]# e# g3 @: @, J- s" a0 F$ p '29': 'artichoke', , x2 W3 q. Z! M1 `1 j4 ?1 k '71': 'gazania',' U0 C6 ^& C# T" m0 {
'90': 'canna lily',% w: ^$ @; r H. z2 B
'18': 'peruvian lily', h) `, X; p' L! S
'98': 'mexican petunia', & J0 R& @7 \- _1 E8 O; y* i' [" T '8': 'bird of paradise', ! E! T. r# |2 v2 o. v2 i '30': 'sweet william', , z5 q$ E6 p1 C! g; a6 B( f: k '17': 'purple coneflower',- ~: I: B9 @1 ^
'52': 'wild pansy',+ g& c9 ?; x3 Z1 U [0 Y: k2 }, \
'84': 'columbine',- g2 @( I' f0 Y; D* R, \8 W
'12': "colt's foot", 4 c6 y* y' K, c4 Z1 e# ^ '11': 'snapdragon',( Q! T+ d- C; X, b8 H
'96': 'camellia',3 g" H: R7 W- j+ I% p/ K
'23': 'fritillary', " e. u2 v1 g! O '50': 'common dandelion', 2 t" q1 }& o- K1 d2 g '44': 'poinsettia',( D1 t6 J6 @3 V
'53': 'primula', 2 q @8 P$ B5 } X& o; W6 Z; j/ p% ~ '72': 'azalea',0 L# K, Z+ y. e' Q0 ~, c
'65': 'californian poppy', " s6 K3 s0 [1 [ '80': 'anthurium', ! s0 j: S B) L. R9 O& O '76': 'morning glory',1 h0 S& D' C2 \
'37': 'cape flower', : W! T9 s7 m, w4 c2 M6 D '56': 'bishop of llandaff', 6 t$ ~# K9 [! ~9 {9 g '60': 'pink-yellow dahlia',2 u1 N9 C! J: j+ r* ]& j
'82': 'clematis', / n, x1 Y6 L' S$ M4 M '58': 'geranium',; \# q0 `- X1 \, W7 s; f/ l
'75': 'thorn apple', ! B9 F$ [/ C: G, \$ j '41': 'barbeton daisy', 2 M& T4 w2 P: V5 t% C5 Y/ c5 `; K '95': 'bougainvillea', - X: P3 ]1 V; y: ~ b, U '43': 'sword lily', 4 A: A0 G+ P' u& t% m X3 K& Q( w '83': 'hibiscus',+ u6 R" {" f* [1 n
'78': 'lotus lotus',( d0 _- R# a4 E' \- S
'88': 'cyclamen', * o% K6 G+ y; j '94': 'foxglove', & L8 y5 k- H/ o8 j) a '81': 'frangipani', * c! `: ^# g; H; H/ O" G- P, E+ ~: ^ '74': 'rose'," l/ p% y4 v3 F- ^" S7 n# g6 B
'89': 'watercress', e0 P7 j: }$ z2 u- h& R '73': 'water lily', ' v+ E+ S0 X! B( k5 | '46': 'wallflower', ) d/ @2 A( r2 `, O1 N; k. r '77': 'passion flower',0 U* j3 ]9 a1 i( X4 q. R
'51': 'petunia'} ) ^/ l. i( V/ ^ 2 H$ a3 x! N0 |7 v4 W13 g5 m6 t/ m ]/ s! ~' z5 t
2+ V5 S( D) P& w6 u: y: e
3) e* q& N% f9 x2 ]& W2 F0 [! X& D2 l
4 6 H9 O; c$ l# _: b: x) b5 ' A4 V9 c* K7 H7 c0 [5 Z3 T3 \6 3 L) @% m$ M4 Q7 m7 " {7 G/ a& Y- i% i8# L, T9 S" _% j4 l. Z5 S1 |& N
9 + P7 @( K2 X* u" o10 9 T9 S* D( \2 p) e11 ' K) z; L9 @8 T3 `12 ) r; G$ T5 M0 V' T- J130 J. g( @! }8 p
14 : d1 X, {4 K$ |2 I, H( B9 X' ?15 8 Z+ Q3 } U: f, _1 _% z0 j16 1 Q" O+ d9 J5 a, W: {17 ; x4 N6 o3 W# m- n6 C18 4 H- \8 L5 ]7 g1 S- f19( y) N) ?1 v7 U7 K8 G
20 9 l5 L- t4 _! x* d1 E& x21 $ J4 v1 i* M- |# k" z! F) m2 T7 u22 5 n9 |. W: m( h5 t2 Z) M23* _, v7 R, {2 Q# g8 L. J
24 8 I5 L- N! \ k1 V25% K# N$ e2 O. |5 [6 Y9 Z2 o/ G0 h0 Z* ]
26 ) u# ^& O7 M8 i+ T) y27 ! ?% A3 `4 C0 l# y1 p1 {28* Y9 ~% P$ a9 t7 v* ?& `+ L
29 / F4 q+ l7 ]7 M: O% }309 }$ \, X6 W7 p# F; G" t% r
312 f ?# D- ]/ |" f; `1 ~* T
322 m0 v/ j- s, _/ I0 g6 Q- h
33 ! x# H( i: K P# I4 ~& [34 , M1 Y6 Y* S& A* V) {, K* k35( s6 Y% U$ w3 b) @% X: T2 [
36( h* E2 P* k9 b5 z8 i( c" Z; }) Q
37 A) T+ a; ^/ b4 n/ f% q38# z! s% J, ?* A8 V
39, x- ]& N. M/ d8 l
40 / f2 V, C8 P: b I41 5 L% ~; q. a, d; W2 h42+ e, h8 Y, N' }7 p3 t; f# q4 n# Q
43 4 I5 s/ @4 {( |" P$ B% k; B44 ; E4 N4 c( K6 H Z* C9 D0 f# N45, t6 B& P! K' `
46 7 y6 p" w3 }: q+ r+ h0 h# k470 y! } {' @! g3 _- c, b% w! O: k
484 o. H3 Y! |: ?
49 / l# Y8 y$ _, Y6 l50( |) p: P, e& T( v
51 z- D; p' z1 [, e2 m
52" ~; n2 _* J1 n
53. E2 W+ B/ K4 l# N5 H" }
546 ]1 i7 I+ K1 H9 J
55" i* J: Q+ V: W
56 7 q: n; h4 z9 f) I) [576 V6 g9 L" V; c7 W. k0 ^0 _; v
58 ' a; Q: s/ r- Z5 ~59$ L: ?* S0 u" {5 T
606 K2 w4 O* ~) q0 _. g8 b- o5 u0 @& |$ G
61! j- a9 r: H6 N; @" Z
62 5 Z1 d* j$ ]) V% B2 |9 Y7 m+ A6 b63 % a: o$ [/ t W. W- U64, C. p2 r! Q, y" O, L
65 6 F) s" S% w; m! G! ~+ U661 a4 _1 K2 q- e& T. R* [+ G
67/ g/ e( ]4 a4 D1 F6 M1 M0 v
68 1 f% d4 k) A# j) i+ Z' S69 ; r2 k/ g9 p% j70 3 U/ n7 C- o& d z+ [9 J71 @ G9 e& C' z p2 |; l% M
72 . r- J5 w7 [3 ^. E2 d* C% g73; y' Y ~4 W0 p; T+ ]; `
74 & G2 ~+ x' H7 S- ?75 5 `( Y: q2 s! D0 q76 ; ~( |7 z' d4 ?1 v+ g9 Z2 [4 g771 w5 N; Q# Z0 B5 n6 a. w
78 7 \2 N h: x5 f/ b5 D; b; A79 ) I; o3 O* {/ T/ p803 i9 k( S& a% @
812 C3 A9 s4 Q# y5 t2 m
82 H; m G6 F$ q7 A I
83 $ o$ h h$ @9 F: Y0 {; q84 1 u2 X( _, i/ o& }+ C85( w4 D) e- i6 j& ?$ R* [+ j
86* r8 u* V+ s+ Y7 t2 r4 d/ y9 U
87% `$ m0 N" S6 x5 v; R* u
88 6 `4 e% g% Q( P6 q' c( `89 ' c6 S" w/ T1 |7 c* w90* u! y) y5 H3 L, L i, C2 d! ^6 i
91' A+ z# c" R0 C9 B0 s" r3 s+ t& X, A
92 # X- d+ |0 @3 F# R93; F, i) ? `( l5 ~+ I
94) i! C- k( p# w9 S& f# F- h
95 " X: \: ^& |! k4 m96 W5 v- N1 B C! Q/ o97 ; D0 m J9 ]4 n4 [3 `5 q9 R3 c5 D98+ m. B) s$ k( Y- f
99% @: Z# [2 M) I9 e
100% \; G" l& Z0 J& U2 f
101: t- s z/ G' T
102 , L. X- u+ M7 k# n% E2 N' j6 E- d4.展示一下数据 + S' W( W- x4 W7 E. @: g$ ~def im_convert(tensor):% r" P& Q2 u; k* }/ L% k) m
"""数据展示""" 1 e, E$ \$ {+ {/ G4 ?) m image = tensor.to("cpu").clone().detach() # Q( t% z7 z* S2 i9 I2 ~$ n) E+ [ image = image.numpy().squeeze() ( M& k+ n9 c V3 X, }! \1 t # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图) A i+ l( t0 O" F) A
# transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)* ^- G4 g0 e) j* e# h. }$ S0 I
image = image.transpose(1, 2, 0) : B4 ~+ H5 V% t6 t& \: X # 反正则化(反标准化) 1 R9 ~- x+ l# U1 G image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406)) ! t/ ?2 Q0 I! O) K3 @9 L; b+ u' I
# 将图像中小于0 的都换成0,大于的都变成1 3 ?9 |3 O$ d! S image = image.clip(0, 1)" X5 {; ?% D$ ]' z& B/ G
$ Z# o; f, J F$ D+ ~$ j1 } return image* J' t$ h# }$ C. L& x" ?
11 v2 `& `( R# s: T- W& O
2 * I5 k# }1 c6 L: O5 E( y8 C3# j% O7 y, [0 F+ \6 s4 Z
4 & j1 X9 F A) b4 z6 K8 X5 1 C# |9 P7 B& H0 j: g0 v7 g" s6# H, w4 x9 P* O0 f2 a, J# {
7, Y5 R& X% A# u/ |5 g! L
8! r; A: F) d5 m+ t
9 $ h3 y, J' I4 l8 ]0 m: S8 ?10 7 A" W) Y$ f# J- m9 }/ O8 G11 ! }5 K) `, h' h* G3 x( i12* V4 M2 J6 ] n- l
137 j9 d' q6 B( k' o
14 # ^' P/ a4 R" E4 N. a5 n- |# 使用上面定义好的类进行画图. K: C: r4 w/ s: `
fig = plt.figure(figsize = (20, 12)) # O5 n T& n; ?2 M, |0 X0 Bcolumns = 4 - O; V# \$ U' H ]rows = 2: n8 {& A H" u& ^) X3 ^
8 q6 Z b( P% N) O# iter迭代器 ! g3 R3 I: [! Y# 随便找一个Batch数据进行展示6 S4 a/ Z i% u' G: y4 v
dataiter = iter(dataloaders['valid']) % N. x/ ~/ r7 b1 X6 Rinputs, classes = dataiter.next() . M8 d4 M1 \; j' l) ]5 }4 h5 Q3 o! Q7 g! U3 i/ S- P# o6 o t
for idx in range(columns * rows):6 |2 b4 @1 h3 G3 b7 P; A1 M
ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = []) * L# @, }% T1 w# v # 利用json文件将其对应花的类型打印在图片中 ' V Q* q6 |/ z, a7 G- ]& [9 c1 ` ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])/ q4 r; w' M5 f( }1 L3 q5 u
plt.imshow(im_convert(inputs[idx])) 9 d3 Q2 @8 ^/ W; Uplt.show() " H2 [. s4 ]* [4 \4 r: |2 D ( ^/ L' l8 ^" O7 o# r6 c1 * D3 D5 a) o; T- D2 M# K+ v5 b! D, O9 E3: i" d4 Q' W8 w; N) }8 E. `5 ?
44 x" l, t5 s* o0 e. p
5 4 Z' J; b9 V! T, u c6 + X3 `& \% ~+ H; M" b7 # x; R; C# y+ S- U2 O2 z/ i85 x$ |& a; h; a( w" E8 C
9: }; N5 z5 `; x- A! ?
10, v% ^, D8 p$ D8 p* o
11 ; q% T4 l. R8 R9 h0 F! C123 ~8 l1 m) D( J
13 3 I5 O% j$ a8 a9 A14 % u8 y! |% o0 b) V6 J7 p) g( P15 2 e1 J! z$ g/ Q16! ^2 ]' V, H5 p: ]0 L
1 I- c0 V/ y: O$ x# g
/ D; T! v8 l$ s% U: y) x6 m6 d
5. 加载models提供的模型,并直接用训练好的权重做初始化参数' k+ b- W: T0 `: g5 n7 p
model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception'] ! V Q# m& {. ]) y# 主要的图像识别用resnet来做% p: O4 G) M5 _2 _& m k
# 是否用人家训练好的特征 * o' G! x$ p; a5 `feature_extract = True1 Z! j |/ B% I0 o- a8 I
1 " r1 C! |, A: g4 n! R2* M; s; ~% G: k: e2 L
3 * b# W. O& D- U! r' V4 * M" j! x* l6 O6 a# R# 是否用GPU进行训练 ; i V' ^& e# o& m9 x; B- R1 Vtrain_on_gpu = torch.cuda.is_available() [9 m: A( a, j3 Y6 L q' |! O5 y
( u0 g, r2 `( \if not train_on_gpu: 4 a- [! l, D( M1 g' f print('CUDA is not available. Training on CPU ...') 2 z/ U. B$ p% S& L5 {$ F: i+ felse: 9 i' J" z2 L3 x1 N1 _8 B w {6 A2 f print('CUDA is available! Training on GPU ...')5 e1 H+ w+ o' E( U$ d0 A
& e# f/ t9 X$ ]! Y* Z. Gdevice = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu') " [" y' M: R: |. v7 {1 * c& p, x' B5 a7 J9 W) B5 G22 r& V( I! _( |. p
3 2 [" C+ i: O9 ^ V4 ( i% _! ^$ y' u; s; j3 k, D5- i: C* x* ^7 w l* D
6 + h/ g0 `+ X$ C( g6 q7 4 Z0 D; F5 D" A0 O) z6 e% [" L3 W& E1 P8 2 d0 q4 H- L; V# C& M5 V/ X9" [0 K& g! @3 K, C" E
CUDA is not available. Training on CPU ... " d, N; ]! q" Y8 _. x, [' q1 ' @1 ~1 {# f6 f7 V) U% O* ]# 将一些层定义为false,使其不自动更新' f1 n7 k& l& a9 a
def set_parameter_requires_grad(model, feature_extracting): ' N+ o- F" A$ i/ Q2 O& A( _ if feature_extracting:0 h; }* c# k7 ~/ p) ^. l7 y
for param in model.parameters(): 3 k& r' {& Z- J# ~ param.requires_grad = False I3 H& V% ?. {- h1 N& j) E
1' l+ T" g3 B+ }* M+ j2 T: @
2# G" |* W. w; i' x
3 * e6 [& }6 T, J( J, X' X4" u1 u5 W& a4 r2 k4 D( p* R
5 ' f4 G. N3 C& a# 打印模型架构告知是怎么一步一步去完成的6 [% \0 P' I _7 D: i
# 主要是为我们提取特征的 . u. V; g: @! U# X4 w+ W' ?, ?# m9 K! G! Y* v' f* [! F6 @
model_ft = models.resnet152()$ }+ V% `7 s. U* I/ t
model_ft + D, g; ^6 ]$ l. Q7 x4 X: l. `1 . L3 X( J" \6 T- _2 ; G8 C5 U4 T, b3 7 M2 G7 w% }. W8 j5 s5 P41 c% ?. ]5 ~# P. ?" J
5 , T# |9 V/ m* e. |ResNet( - y# L4 u1 [0 c' U% b- T* N' a* N" v" s (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False), A1 M! A! P& j( g1 c5 g. D \
(bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True) , o" B; v4 G7 }( Z- d" ~# Q) S (relu): ReLU(inplace=True)1 O2 d5 \0 o: v+ t; t
(maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False) ' d% h1 q2 U" N( x- K (layer1): Sequential(6 P* L$ t- \( i$ J6 G! y6 s
(0): Bottleneck( # f: F. r: @5 i4 N (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)% }4 D- `+ Y/ @1 M. {
(bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True) 5 e5 ~9 a; _$ \9 ]4 K8 W. _ (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)) R1 m& C2 T% Y, p& Z
(bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)$ u0 F+ t/ Y4 D" }( M
(conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)4 N' Z- k$ i* b0 @. G$ g
(bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True) * s) E7 J. z& R/ M (relu): ReLU(inplace=True) 3 B4 A' ?7 f9 v( C' o9 \; z (downsample): Sequential($ |/ v0 l$ h5 b/ V4 s
(0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False) 7 F6 s6 A2 P" ~- ^8 Q% u' K/ r- | (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True) : y C, U. w/ b& G2 q; n )- Q- r. R, X- K& I
) $ V9 A' M" D+ K% P$ M0 P/ ?; A8 a中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。3 m4 h4 h6 q6 ~6 P$ P; f
(2): Bottleneck( 4 n% Q- X# g* h3 |# M (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)& w" M' D6 B9 u
(bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True) & @# [# b. s/ `. T+ b& V (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False) 3 A+ [8 ]; ~2 u3 \+ x (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True) # h* A( s0 o1 Z- d5 | (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)0 a! ~( E% g2 m+ K: h
(bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True) 7 h4 [4 x1 P5 M% H2 p3 r6 { (relu): ReLU(inplace=True)1 x9 ]/ U) k& p8 K
) h1 o2 X, @1 g1 r' b: R+ Z- } ) ' ]5 t, S0 |9 R# p! x* [# M" j (avgpool): AdaptiveAvgPool2d(output_size=(1, 1)) 1 i7 W: C& s6 {6 n# h6 { (fc): Linear(in_features=2048, out_features=1000, bias=True): S& j3 r( E. z
) & L. e3 e( P4 e9 I. T/ _0 F% U% K4 U c, D( F+ x$ d. S
1 $ {( K, T3 f! A0 M27 L2 S) E) g8 K, k( \
37 T) D% a r( V1 N8 u
4 / M3 m5 o, C. W: ~' j8 d5( O! U! k0 G6 q; D% p6 m+ {
6 ; t. t% D# S. t' a9 |/ l7- b6 Q. ]7 v+ M4 Q, V5 B1 B4 ]
8* v2 x Z" v4 A1 H# Z
93 b" G9 e* U; I3 l6 A. A
107 u% s3 w8 G- S4 d( {/ ] D
11, o) Y4 _$ _$ Q5 ?3 o
124 m# w$ e4 W3 c6 i$ d) o
13 8 {4 o% B. E8 g+ L l14( U& I+ b" F- m1 T O' ?
15 0 ~3 c z* H, [5 c$ F- g16 # H: m+ m1 K' ` t0 N; i9 x$ Y17 ' I8 p% c5 ?$ z A' _6 c7 ~18 3 t8 Q2 F# ^0 U193 w7 j3 I" _, U1 H4 Q/ A) n; s
20 7 E* H. j. U. X& M21 " @: W: G% B2 @# F; w, z* w22$ j) l G' b1 `: V
23% M8 P, X; ~; D9 X
24 9 M: N6 ]' M9 y& d' X, f259 z$ ]% T; @9 b. Q5 K- O8 m& u+ j
265 d4 x9 G; V% `8 o$ b
27 3 n* B2 U3 [" a28 9 v, v* z8 w( a% W- E; N6 {3 U29 6 K. G$ g3 j4 u$ R308 P; a! @4 A5 V- h2 `3 ^. c
31 # N \7 i' ?% x' u* U) v# B$ P$ X0 |32 ) K1 [; l3 W) l1 Z: B W" v8 B8 u5 T1 ~33 2 s: J4 d; s7 Z3 E. F最后是1000分类,2048输入,分为1000个分类 ?7 d, @, b2 Z, j而我们需要将我们的任务进行调整,将1000分类改为102输出 # P1 ]5 |$ s+ I ' `$ q3 Z6 y" x0 p. G4 U6.初始化模型架构 ' f) D/ c+ @- ]步骤如下:5 S: o7 |6 |$ H/ q6 V' ~
& t4 v' B3 s9 O$ B8 X$ ^5 ]
将训练好的模型拿过来,并pre_train = True 得到他人的权重参数 6 G! T/ T4 U/ y, q8 Y可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)0 j1 [2 m. \& T5 _* P( e* c
无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数0 G. N" a* G! C( K2 J6 q: O
官方文档链接: ?3 A; h9 v! D- f
https://pytorch.org/vision/stable/models.html. O3 x0 H4 P0 V3 L& G
, d: T0 v' m7 y7 C) N8 |
# 将他人的模型加载进来 n S' s- k+ W/ G
def initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):7 J0 } _+ ?4 P- g* L$ B
# 选择适合的模型,不同的模型初始化参数不同 : c" V: G1 M/ }* P model_ft = None * t- n# U1 z, v0 n7 K input_size = 0 ) O$ b# H9 a+ n! F/ l0 h1 D5 U# E% N
if model_name == "resnet": & y# L0 i. M) F) P5 L, @" c$ D8 g """ 9 `- V/ `8 p$ o$ ? Resnet152 % y6 H. O/ f S """ " s, R5 J3 g [+ ]8 s; U- t ' c) D7 J8 P8 Z# f1 `$ x6 X, t # 1. 加载与训练网络 # y" Z! w4 K2 k* i" M ?6 K model_ft = models.resnet152(pretrained = use_pretrained)6 ]0 m4 E% r+ x: H$ _$ L- M" u
# 2. 是否将提取特征的模块冻住,只训练FC层 ( _ V/ _3 S& n9 e set_parameter_requires_grad(model_ft, feature_extract)" E6 F, Z7 }; t! T k
# 3. 获得全连接层输入特征 4 z1 K4 `8 R) `# w num_frts = model_ft.fc.in_features: x7 F! {5 `3 h! u
# 4. 重新加载全连接层,设置输出102 ; e7 n6 N% r! R# s/ _0 w model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102), 1 U% _ Q8 ?) G8 P0 D7 b. ^ nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1 & {" T& |+ D- f/ N# A input_size = 2243 k X4 C, I! \4 ?. p" l8 I
! ], r; m$ I9 g' \) }/ V% M
elif model_name == "alexnet": ) e5 N% A. s N5 ?! t; Y """ ! v- X$ ]8 [. g+ L$ Z2 T6 A Alexnet * r8 T' k- c4 w$ ^8 M e, b """ ' N' p" \* O: E" B) X: h model_ft = models.alexnet(pretrained = use_pretrained): \5 K' f; I) B8 x `
set_parameter_requires_grad(model_ft, feature_extract) 9 c& D1 l( E) h, @8 U1 _/ p* p6 o" B
# 将最后一个特征输出替换 序号为【6】的分类器 ) t: D1 T7 h2 n' t% V1 I ? num_frts = model_ft.classifier[6].in_features # 获得FC层输入: k* o7 z$ T( h& p: ~- n
model_ft.classifier[6] = nn.Linear(num_frts, num_classes) ) o% L0 x4 U9 G6 D: @% i$ h input_size = 224 " }- {( K6 r. z: y5 o. w6 w m2 D4 v* j& ?% t
elif model_name == "vgg":- W0 p+ Z. S' y; n/ {
""" 6 Y L* t: f! N VGG11_bn. G# M$ r+ S2 \- `# n G
"""- G' H* c% ]) F
model_ft = models.vgg16(pretrained = use_pretrained)# H1 E' _' S3 v; G) N
set_parameter_requires_grad(model_ft, feature_extract)$ G8 `8 w3 S$ ^! g( V
num_frts = model_ft.classifier[6].in_features7 w& ]* |0 u: w5 U
model_ft.classifier[6] = nn.Linear(num_frts, num_classes)- d! \# ?# N8 T) b* u# a1 e4 d
input_size = 224 ( W( i& N& ?; l% a " q5 n* G9 j; x: i7 V9 e. [ elif model_name == "squeezenet":9 }8 E/ Q# D" ~! G0 K
"""7 k! t; D) _& n8 M
Squeezenet$ t3 }5 u3 G1 Z
"""5 I8 p$ ?) q2 n& l$ h; k6 J( u
model_ft = models.squeezenet1_0(pretrained = use_pretrained): H; G; M: _" x! w3 }; p
set_parameter_requires_grad(model_ft, feature_extract) ! v5 c) w' s" N+ R" f( I model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1)), ~9 L$ q$ X2 r2 ^ M
model_ft.num_classes = num_classes ' t. P+ S6 w4 V2 O, B" @ input_size = 2242 q* G0 t: w' `
! ?0 R& t" o% t9 v/ i elif model_name == "densenet":! h8 _! T: f! V' n4 Y! L
""" 1 R# P3 I9 `6 c- _+ Q9 x Densenet , J) M5 x8 F! a" k" C! { """ ' C* V" L5 j! k: k: C! w model_ft = models.desenet121(pretrained = use_pretrained)5 r5 O9 j/ k( `8 e
set_parameter_requires_grad(model_ft, feature_extract)+ E) n. o$ T6 V) v4 U. Q7 B
num_frts = model_ft.classifier.in_features 0 y# h2 Y4 H y model_ft.classifier = nn.Linear(num_frts, num_classes) + p c$ i. F" t9 x! y2 o4 [ input_size = 224 . y2 Y4 \4 V" P {2 y! v& B, ^/ ]. K1 [
elif model_name == "inception":6 x6 b, s( O7 j! _- T
""" , w4 @, Z( s" a3 P: S9 { Inception V3 6 Y% f2 D1 l% g/ g. { """) C y- H; k/ E
model_ft = models.inception_V(pretrained = use_pretrained)6 H4 y R1 I( K) G# d3 l1 q
set_parameter_requires_grad(model_ft, feature_extract)7 i3 u# W: ^6 a, Q( A4 W/ ~% V
: y0 n2 u% i, \. H% t8 ]3 I1 U num_frts = model_ft.AuxLogits.fc.in_features 4 r% }/ p' ^/ V, ^8 K% H model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)2 Q7 |; J/ p) B% q. B* H& w
& x1 v$ d2 t) E7 _8 |2 U4 L
num_frts = model_ft.fc.in_features , E* S! R! `; j4 Y0 {6 e# H model_ft.fc = nn.Linear(num_frts, num_classes)6 `7 b! a7 ?2 h4 c6 m8 T3 |
input_size = 299 4 E6 v; |( N& Y) {# J* D s) i# M2 y& }4 v6 t
else:9 A: c" v, u& \1 e4 y* ]5 Q
print("Invalid model name, exiting...")& s4 C& C8 }! e7 g. X
exit(): @: ^% g: ?& ?
! M' A, |( y4 u T% p
return model_ft, input_size5 U$ Z2 l6 L k, T& C2 J; }+ ~9 ~
+ F7 K; @- e6 X4 G0 q- H' y
1 9 Q/ B" o( I# O {: v; ^: u2- y7 J+ V5 C) ?2 u" x7 q7 H
3$ K, J) h) ?& q7 ` i! c' _1 K
4! I8 o0 ]) f# P ^& `7 m
5 x6 \( Y0 k" r6 { T8 B6 & V& U6 a7 E) s! e78 ]1 N, e1 ^5 N" H- ]% E! D* y7 ]
8, L/ L4 k/ Q6 |3 N9 F
9 V) B. B o3 j9 ]% ?+ c105 W6 ?9 Y: v5 y$ o
119 @4 ^$ L1 x$ B
124 D& g: P- F' A; U4 P* ]
134 y8 q% o I9 b0 {" _. k0 a
14* F0 l5 O. ], K1 {) |
15# e! m( [. t. B( ^
16( C% S1 U2 ?! y# O5 T( c/ ]6 y
17 * z2 E, z$ N. I6 b" i3 W' W4 U18 " g, l0 ?3 G2 H$ X9 P' u( o# c19 ; F6 f* Q; V0 d, ], \' t# P% \- `20 ' S; [0 m( M5 ~" [ P+ a; ?3 P21; G' d. e4 X8 M3 j+ V- b
22! S- t' S& {+ m* D5 u" [
23' j9 C& ~$ g4 K+ |5 ?: w
24- \" Q$ T& k/ }$ Q8 A
257 ?" w( y: p; O3 n8 m
26- N; q# ~6 {* M; G
27 - t' P7 l6 L3 }+ Q& t& ]28 $ _; t% `9 j% `# V+ K/ x& J29 9 s6 K: N0 v z0 s4 O( g- T5 |# B* ]30" _9 P# e3 h5 Y T" O
31' Q4 @1 X( q5 t0 o
325 \/ ~! g3 y; s! A* |0 S7 J
331 W* Y3 Z7 P8 q& h# D f
34 ) I- i6 q! ~" M: f X35/ U! \- O& c" I e$ [5 o( r; i8 j
36- V5 _: ^* U$ k( K6 d
37 * D( r; {- {& z L% g: N [38 . o7 ?9 I) o+ A, v. [0 F39 * v* L: w5 n! h! ?. V% m$ l; c: X j40( m) z! e1 \9 Q
411 r n8 j Y1 d
42 ) z' u9 J. E4 i. s# m& A43 ' D% a1 r8 ^5 x. q2 ?# S' ^: q448 t4 _5 N- z8 W* G$ n, t: G1 c
453 D( r! V, F4 @3 ^1 |& U5 j
46% [" C$ T5 i% Y% Z- f, V' s
47 4 M2 L" N/ Q0 L6 G( A2 D ~48/ i! n& b+ n" e+ b! ]! e+ U
49 ' _- Q5 W! U# b6 V. @1 r50 ) k: ]# m& ~. l# B+ c51 ; h; ~& M5 H7 I2 {$ h) H7 w52. ^# ]& h$ P1 M( r- i q- Y
53 7 H! I$ d7 y: ^+ G* q54 $ P7 G+ P0 V/ N0 B1 l554 k% [9 E( n& d' {
56 / w. P9 \0 r0 u4 x57 E; \- K+ M4 o0 d- B( V58 + L- a% i$ ~& x5 e% q59 # K& `3 @6 ~/ t8 `# q- _' t60 ! g9 u6 L" m+ P) ^) p1 |( A) D' i* H61+ T! p0 x0 M: C
62& H# ` b9 ]! ?
63 + |! W& E9 X; g7 W8 Q' \64 6 c- P% A9 ~8 u0 ?" h) d65 : n! k' O' ~" g7 s# y% N66& ?: U$ }8 Z$ q2 k3 w( H5 m
67. I( j3 Z9 I3 ^: ?( H
68* C4 `- R U* K* Q
69 2 A: q; O9 o% y8 D- W8 c70: g. p) J( ?& H% x0 ^, O# |
71# `) R9 i4 z! Z4 g
72 ' o0 v! q/ T" k" f: `. V732 g6 b8 G) S+ d+ T& q
745 w& u. n( u& r5 p
75: K# W7 U* P7 X8 [+ L3 P3 S5 O% U
76 . X3 \4 E# W# p: r0 ^77- N* n7 Z( V, U- ]1 [. q* I
78+ v: ~5 _- h3 L$ U' C8 S' J
79 ) K& z9 u8 Q, h1 s {% p: r) }80 ! b0 g( J4 S( [3 ^; e Z. d81 + p( J) Q: U9 x8 Z82, x/ E" \2 u3 O. ^. y5 |, c
837 y# u+ d8 B0 U) ~
7. 设置需要训练的参数0 a' Y& {0 }/ `9 J
# 设置模型名字、输出分类数1 c6 n/ V2 E' Y1 Z- Q9 d* t6 j! N
model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True) : F. j) e+ [8 G, v! O5 P' _# h+ Y" ~+ f, N) o5 z
# GPU 计算 ; ]: U5 O& r6 q9 w. Zmodel_ft = model_ft.to(device)5 r$ m# x a/ f; J& Q3 `
$ x2 y# K& W( Z, {3 m( }& Z0 z0 r. w
# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取7 W, m( u0 ]. E7 G; M' R5 }7 S
filename = 'checkpoint.pth'. N- P; F( t0 U# t
& A; f) W9 z( T v1 D# 是否训练所有层 # t0 `, I( B' s! Nparams_to_update = model_ft.parameters(). j7 E" T4 S; L7 @/ A( t4 Z
# 打印出需要训练的层 5 L) k/ S( X9 z+ n! `0 U2 Nprint("Params to learn:") 6 n. R+ F$ ]7 O( C8 l& g* S% iif feature_extract: - ^8 Y9 Z1 s( z. P/ H params_to_update = []' y6 k/ y' M' P1 `+ C9 g) b
for name, param in model_ft.named_parameters():9 l* ^! _, d# }9 R6 H" P: N
if param.requires_grad == True: 8 U, w& G6 o; m3 D- ^% [ params_to_update.append(param) " ]0 N0 E6 C7 X( O% l0 {1 | print("\t", name)7 A" l1 N) `# I$ C, f; l/ d
else: / Y0 J# U/ r0 H( l* o' e for name, param in model_ft.named_parameters():+ w! [4 {, c/ O9 D( p! G V: p/ `
if param.requires_grad ==True:& A, m7 Y6 j; k. Y2 t; _
print("\t", name)7 {1 ^7 m4 \7 s% {) s) G
0 |' w l0 `+ t/ p
1 & q p1 L4 B" ]: l2 4 a: T! }. E8 S# F7 c" g0 ~5 B" [, v3 $ t* B7 j- o1 n- C& A3 O+ ^4 0 a3 O9 \2 k0 T7 L; J$ _5" x; N# p$ b, u' S, x4 D* u. E
6 9 |& P5 n+ W+ S. |. }7 % M' Q7 D7 C/ c2 E. O/ U' {" z" B& x8 & m6 k( r# A/ D- R6 C9 0 E1 t" W1 E9 j1 S& U8 R10* Z0 x) [0 t) x. ?' t' ?. o6 x
112 \( h) q( |, C/ U4 p" R
12 : U- {$ I) K) d6 N& \13 7 ?, c: l% e2 U) D" ]; Q14) X- b8 C& q5 o! V5 C
15$ X8 D6 C4 m3 O' _' w i+ h% d
16 1 O/ h% J' b9 X3 N6 T- k; a o17 2 R8 B G6 ?1 {18 0 M: a! P. |3 r8 G- ]194 t5 ? O: ]7 m( j
202 f, t/ S! V- {' v: l
21* n/ g# q* m8 Y
22 6 X {& Z: i& }) x& p233 B9 Y, u8 }3 h# v1 E. R
Params to learn: ! Y) }( b% _; |( k- ]8 X fc.0.weight& t9 o) g: @7 X( [
fc.0.bias2 A3 G0 m8 ]) n0 R* d' j
1+ R0 [7 `! g! |
2( g, z2 x) i: Q3 C0 N
3 % Q5 [) t$ n9 ]7. 训练与预测+ V8 U0 z# f5 v; n6 Y# K! l7 L
7.1 优化器设置 2 I) j: |3 J3 b# 优化器设置( P; ^$ K. i/ @3 K* r: ?
optimizer_ft = optim.Adam(params_to_update, lr = 1e-2) ( X" x6 e k) k* @# 学习率衰减策略 # m" t! R, }. M- M* oscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)9 ~/ T! x: ]8 ?$ E9 J3 Z
# 学习率每7个epoch衰减为原来的1/10 p" L& g* x$ S9 n' F* d
# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算. g4 }' ]0 S* h( \8 C
* S) F7 q2 K# s% g5 Qcriterion = nn.NLLLoss() 6 g' ^7 m N: `1 T7 P4 ]1 ^, x4 M/ Y R1" N4 \0 J+ E B# L7 v
2 9 g2 Q$ }' n: R3, J( m4 q. J ]1 B/ I( q8 t
42 f# p; Q# h' f1 ?
5 2 j y0 D7 O# l+ A63 |+ {/ m3 J+ |
7 , c/ e/ I }: U6 j& G0 t5 Z1 e8) H) q7 F* j9 d/ i! t8 |0 }
# 定义训练函数 * ^1 J: @4 w% b4 M#is_inception:要不要用其他的网络) [/ b: q8 _, _. q# A( z2 S& l: b6 Y
def train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename): 6 ~0 ~% |9 M1 S since = time.time()8 _2 c# i8 t, \7 D D
#保存最好的准确率) R( w9 @, F7 C* h @& a
best_acc = 0/ F" b6 ?6 X" `
"""; M# H* c2 v$ V) z
checkpoint = torch.load(filename)$ C; w9 b$ ]3 ~! E' W! [
best_acc = checkpoint['best_acc']$ j6 d* a7 g, N0 T7 h2 z. Q: {5 B
model.load_state_dict(checkpoint['state_dict']) ' H) `' d5 k5 J. f) j1 _ optimizer.load_state_dict(checkpoint['optimizer'])4 a% I; D" H: ]* P; h; I. C$ k
model.class_to_idx = checkpoint['mapping'] K* X! q5 E2 [, F ?5 z """ * d9 H( I) ]& H8 ~# ?: ^6 ~% X #指定用GPU还是CPU" J' L! x) A2 Q) o- C. c2 | g
model.to(device)0 v0 I, J9 |8 w- V1 l! a7 ]2 X: D
#下面是为展示做的 ) J# K/ c; W, h2 c. M: l val_acc_history = [] ! i' N7 a/ c5 b$ F! j4 C. W train_acc_history = [] 2 _4 U% J! ^8 Y, p4 J train_losses = []. {8 u& s3 [, s: S- y0 ]: N
valid_losses = []5 o! V1 j& @& c( _# p0 g5 ^
LRs = [optimizer.param_groups[0]['lr']]$ d2 y1 ?4 J; I# J) e
#最好的一次存下来 & N% R7 b# C" {) E1 u' y7 w7 T6 I best_model_wts = copy.deepcopy(model.state_dict())+ ]) O* X* z+ Z- k& z% T9 F" j
8 t# X4 S4 f* Y& y
for epoch in range(num_epochs): ; ~# J9 R( X8 Y7 t% A print('Epoch {}/{}'.format(epoch, num_epochs - 1))9 Y3 p- ~& J3 |. k* x
print('-' * 10) l4 e z2 k% b5 Z/ I; R% s/ f
7 I5 h9 N: a8 @
# 训练和验证. F k4 K I9 k; b8 ^
for phase in ['train', 'valid']: & M5 G- z+ g2 V0 D' m0 f) Z if phase == 'train': % [$ R6 Z$ Z9 l( Z5 V; T+ u model.train() # 训练 ( p- ~: `8 W3 i+ h0 Q/ k! k else:7 {' O& }$ Y) i7 z! q0 e: O
model.eval() # 验证4 x2 w, f4 C$ k. T, D, V9 @
! L, e7 y6 t6 F
running_loss = 0.0/ q$ B H1 u" H" d4 I; M: Y ~
running_corrects = 0 # G; P7 b3 W5 n$ U. c& y( _1 I5 y% c; P6 z
# 把数据都取个遍9 q8 q1 m9 N7 K, y6 x) B
for inputs, labels in dataloaders[phase]: / |! P" _/ O% T' Z, } #下面是将inputs,labels传到GPU$ m! Z7 h7 r. C, A- H- [' H
inputs = inputs.to(device) 3 l% V9 S: N5 D5 ~* r' F9 c" x& \ labels = labels.to(device)# ~, z4 C9 b6 i) F% I O
' n- ]7 W3 b1 U) f- {0 U! d # 清零4 h, l+ d! n- h# S7 ?
optimizer.zero_grad()7 G# D% G0 ]/ `# A" f
# 只有训练的时候计算和更新梯度1 u6 f- X7 |% }$ M: @
with torch.set_grad_enabled(phase == 'train'):9 T4 ?; n7 O+ D
#if这面不需要计算,可忽略- b/ m3 J& ?; A0 N! J1 Q# p- [
if is_inception and phase == 'train':: ]2 w: F2 F* }7 Q7 K
outputs, aux_outputs = model(inputs)' y. R" l) O. B# _8 M' f; _
loss1 = criterion(outputs, labels) , A _% i8 X* H& ]" q( { loss2 = criterion(aux_outputs, labels) 5 t( Z7 [) \# Y: `5 G loss = loss1 + 0.4*loss2 ) L/ C U0 m1 Q7 ^& h% o else:#resnet执行的是这里 1 `) C( [! I# r6 c& U9 {( s outputs = model(inputs)( [3 I* y" z6 U8 c5 d
loss = criterion(outputs, labels)/ R0 M$ k6 F' W1 |: K
5 K) W$ f# W$ l* C2 [' r; L' w$ [& S #概率最大的返回preds % a. o) C; v" ^8 x) h _, preds = torch.max(outputs, 1)( C/ X/ [. ` {3 l# ?$ ?* H9 X
% x# I ^1 x5 S6 c# |; K # 训练阶段更新权重/ ~7 Y- S- e5 ]7 y$ V
if phase == 'train': # I$ u8 [: W: `" @* W _8 q loss.backward() - |3 G" |* C1 `9 l o optimizer.step() b* O' k7 q& W/ M2 G* ?5 _- c) J3 `6 ]% D" h& A* j
# 计算损失; ]+ }' i: _; y; H$ \
running_loss += loss.item() * inputs.size(0) 9 t; R) D/ o# e0 u0 W running_corrects += torch.sum(preds == labels.data)$ Q0 _& W1 z7 `
; Y6 l3 n, h* b) ]& T- R #打印操作 % r; o+ K! Z) T! K9 L h( v epoch_loss = running_loss / len(dataloaders[phase].dataset) 1 [5 x: I3 W1 ^* A) L4 D% O epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset); O+ J \3 S: r5 K; [; G
) q! F5 n; b% ^6 {7 s/ E * R x8 a! L9 K, C" N) x time_elapsed = time.time() - since ! r7 A% ], y, O# l# ` print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60)) 1 O( m' [: f# i: o print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))' K) O& P* {+ h& o# @& x/ i- @
( i! T$ h) a. ]2 Z
3 ?8 e8 L: T" C) D7 O # 得到最好那次的模型8 M; p; `- S6 T
if phase == 'valid' and epoch_acc > best_acc: G( Q2 B) L3 w0 C: r4 v best_acc = epoch_acc 6 y: \/ C5 \" L6 d; k #模型保存; v! C" G, d3 Q3 ?1 s! Z
best_model_wts = copy.deepcopy(model.state_dict()) + c$ D& A( j0 @8 n( u" d state = { # K' A/ O8 [. t. j #tate_dict变量存放训练过程中需要学习的权重和偏执系数$ a" g! f4 }+ q6 c; y1 {/ y7 ]
'state_dict': model.state_dict(),7 V9 I0 c' J F$ ^! T
'best_acc': best_acc,; z* e, f* A) W
'optimizer' : optimizer.state_dict(),) w* t ?( ~% e- A' S
}" h6 z! R$ ?1 U
torch.save(state, filename) # h, H2 ?( X( b1 w7 k" a, s9 M if phase == 'valid':1 I% u# l2 V+ J: E- R
val_acc_history.append(epoch_acc)9 u g" m W; m* U8 A# |9 Z
valid_losses.append(epoch_loss); y. q" @$ ~' i! t
scheduler.step(epoch_loss) 4 I$ _$ [; ^& L if phase == 'train':' S" V# O, a0 F+ `' h3 I
train_acc_history.append(epoch_acc) ) W& m7 @' w7 N( H train_losses.append(epoch_loss) 1 g/ N' D, o: C: u# ~% L+ Y8 m( r. l4 H7 \, T% _
print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))8 B2 M: g- Y3 r" X" w u
LRs.append(optimizer.param_groups[0]['lr'])9 j* p: J" a' N. `& s$ D m6 ]
print() # Y6 _! _5 c, S/ d6 i: W+ c, X2 o2 u, N/ _( @ w$ x. T4 b0 C" [
time_elapsed = time.time() - since6 O+ r" S( [7 M9 t0 J+ M8 @9 R
print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60)), j' Y; I2 r8 J' _2 e( u4 W. R$ L
print('Best val Acc: {:4f}'.format(best_acc))0 V) G$ A: s0 U
9 h+ E/ S6 P6 V0 U$ C
# 保存训练完后用最好的一次当做模型最终的结果 y2 _# G" ^5 b( K. d7 J4 L
model.load_state_dict(best_model_wts)) y9 Y, ^( o4 Y, ]7 ?
return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs , X- a3 A7 Q: n2 }! _* j' D; |
7 S% Y* ?1 {+ h& }& F$ v3 B