: i3 x6 |2 y4 S0 R3 n5 r4 Fimport torch.optim as optim " @$ W o& p0 u! @! o m& Aimport torchvision8 F$ r4 D0 k% h# g+ A4 H# h
from torchvision import transforms, models, datasets6 n F' d$ K& l" B0 ?
9 _. n! z5 d6 \, E1 x% [% E+ `8 Y
import imageio 0 F* [+ g6 j; K; K6 Yimport time 3 h: g$ p R) U# J8 a: \; _* Uimport warnings $ }( O( e0 u9 p( G6 ^0 R" H- N4 Wimport random# R7 y$ \# h; A1 l
import sys , S+ s- d5 Y) Z0 x cimport copy! V) W6 j6 d' s; @ F" b
import json & o1 @+ i# B A& V( ^from PIL import Image- C% g! k5 h& L; M
% |& J; I1 `5 s: O8 ~- ]. Y, X/ {7 Y. d* r& Y
10 {, _4 h2 i5 N. i) S$ ~; C
2( u+ z/ I' O* J0 u
3 . L" d& A7 d6 z4 0 m, ^- p; w; G; y5 z. N- N5+ |/ p o5 H) c5 z( X
6" Y5 u& ]7 b2 N1 g7 D6 O
7; c! f6 s6 k( C/ W. R7 b) n
8 ! ?5 P+ I% @- |4 f2 \1 s) Q9! @- R3 Q$ `7 B0 Q4 y' @
10 % {8 z# l2 G% g2 W# ?' {+ b; z, m3 P11 / ~- a* J3 x! `9 ]6 D4 E% `' d2 G! m12* X v4 D- U" f0 A v# A
13 9 F& W4 c& q" G7 H+ ^: i14 5 M* ?$ Q8 O1 |0 i6 v. ~, o15$ @/ S% r. Y7 D! v& I
16- H `. K/ W& P3 g8 Z/ k
17 # z* \& L' I9 E2 v* P+ x18/ G$ `) U$ u9 \& L6 p* y: m
194 N n, ~3 Q0 c- m j9 [
20- w3 T* a0 H4 I" g: R0 M
21& d$ q8 j; o$ h1 \4 u& _
2. 数据预处理与操作 . V% N/ u' r- K. k. K4 N0 `#路径设置5 Z+ `$ [3 P- Q9 ~
data_dir = './flower_data/' # 当前文件夹下的flowerdata目录 ) O2 ]0 L$ Z% Atrain_dir = data_dir + '/train' & ?9 U/ U# M# ]9 p/ J/ U3 Hvalid_dir = data_dir + '/valid'3 I6 `7 _( }/ V% [0 {2 s: o
1. Z* a5 M* I* D3 a4 @8 H0 _6 b
2 - m8 r9 l& V$ f! B+ y32 D9 j9 [5 Q: ^; ~
4 4 c% ?' G5 l4 ]8 I5 lpython目录点杠的组合与区别 2 I9 a3 P6 _ W5 V注: 里面注明了点杠和斜杠的操作 M4 L7 a: l; U# @: U$ Q5 N( w: R5 B* Q( ~% P
3. 制作好数据源0 f. D+ A! t7 T1 E' @( @* [
data_transforms中制定了所有图像预处理的操作& b2 C( K: @9 L Q: V+ T
ImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片) b* _: r# S0 _+ G
data_transforms = { b2 e0 g; q& s* g( P4 w; w
# 分成两部分,一部分是训练 ) L, d, ]6 Z. O D, q 'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间 # j: `# J# x1 G$ f transforms.CenterCrop(224), # 从中心处开始裁剪 / ^# E; e) {/ i+ E4 | # 以某个随机的概率决定是否翻转 55开 2 Q) j9 _: |0 i C transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转6 |" C1 H6 _( F9 a) W3 y; c
transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转* x0 s- k; m/ \. T) E2 s& ^
# 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相1 c7 S! `! H: ~0 V
transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1), 3 B- P8 Q+ r3 I transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB : R+ s, l) g) N* D2 C( d # 灰度图转换以后也是三个通道,但是只是RGB是一样的 $ D! ]% O0 m; R7 s' `( q transforms.ToTensor(), " j' x2 p8 k) A2 ?" F$ B0 X! P* A6 H transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差0 q" u1 o( c8 v8 W
]),$ {3 ?) s! Y: Z7 W% T0 w p% L
# resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化& P8 l \' t c
'valid': transforms.Compose([transforms.Resize(256), Z9 d1 }% {6 V5 i transforms.CenterCrop(224), - } [: g9 f+ L$ e! L( K: o transforms.ToTensor(), ( @& ~0 u3 D w4 T) N transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同; X5 b& o3 o+ q
]),3 R. i5 a$ J, C6 `
} 2 q9 a& b6 |0 ]/ y: N2 N 7 a- B8 g: u: ]0 |3 B16 o& }9 |# h7 J6 N: @4 B
22 H; M, f4 D. g* H3 Q
3 1 s% {2 F# a+ m6 f+ ]4 z$ i3 ~4& X( l# g* R/ [% `2 K
5 + V% E2 t- z _' d# o9 k5 Q5 V0 q6 B e1 ~8 L$ C$ F) K9 Q7! @5 _; W0 _$ k/ G( N* t
8 : n9 E+ c" q$ V* ?9( t3 G( b3 z, N. s+ S& r0 ~0 b4 ^
108 H/ m! F, f9 V; W( m
11 % r d, L7 [3 q* \& p8 J12 0 ~7 R9 h; }8 [3 C% x" s135 |% g4 ^0 u3 w M# c2 ^# F( J
14 ( f( _0 b6 `2 m! N& B15 & i* {6 O j9 Y" I; C- B162 ?4 n% k' {$ y: s4 p
17- w9 F- c; J* G+ m8 R+ P& I9 M
187 J' |9 @/ {1 ?( v z" q! J" \
19 & p4 [. `) \! g+ M9 ^20, q4 X4 Q: R: q u! [/ z: ~$ f
21 1 u G6 R. u. g* @; nbatch_size = 89 u% R) L2 E$ ~! o/ W; Z! T! V
image_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']} # k: ^9 T& t5 v5 ?dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}) N B2 F: A1 o$ q& q- V8 K5 w
dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']} , i, j' c/ D! }% E" `# J& O6 L5 rclass_names = image_datasets['train'].classes5 ~$ J/ A% L7 X% c4 T
& l5 [ r3 Z$ E6 m#查看数据集合 F( p( ^5 x7 m$ v( A4 J
image_datasets 1 x, Q: g9 d: h) ~6 X; y- C$ i( m' ~& A4 J3 M) y9 e/ g7 h7 r
1 3 \" N) Z1 C- P% M- T21 a2 ~/ r0 ^* l7 D f& j
31 v( g5 F6 H; j5 o, o# d, {
4* S% K% N$ V7 i( a7 t
5 : J# X: ]' S4 V4 e0 [5 a6 ) n3 \3 b7 D0 s7* ?1 ~/ ~- o7 M$ P- q) z
87 S( e& \; L& b3 X, p0 c3 |1 R# O
96 I7 U* n( b- c0 E
{'train': Dataset ImageFolder 4 c5 _5 k2 |# b' Y) O* P0 S5 L Number of datapoints: 6552 + f9 U4 @: Z) ~4 U1 D$ T Root location: ./flower_data/train 0 ~; @. \% A- _+ x P, U StandardTransform M v# C4 y" t! O Transform: Compose( 7 x9 b4 T) U: ~7 S1 @ RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)1 F, H3 h* [; Q$ w
CenterCrop(size=(224, 224))2 X6 O6 C& |1 q1 R4 _8 T
RandomHorizontalFlip(p=0.5)$ U B) [- a: C( y) n; }" w
RandomVerticalFlip(p=0.5)& w7 R; W+ i: a4 N; O s
ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1]) ' f n9 }" K* V- | RandomGrayscale(p=0.025)9 E# D2 n& h- A7 R2 f1 l8 I0 {5 ]
ToTensor()- h5 M/ v1 t m# O( f
Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])4 [$ d! |. r9 O# h6 r
),2 z$ p* y% B, J* t R: t- a0 L
'valid': Dataset ImageFolder1 ?* ^/ ` o6 k4 D& D* P
Number of datapoints: 818 & _) k- u( z) z6 I Root location: ./flower_data/valid 8 n+ Y7 k2 a4 h; m StandardTransform) {$ f9 j! c$ `6 \8 a {
Transform: Compose( / L0 }) s8 D. w4 h4 M Resize(size=256, interpolation=bilinear, max_size=None, antialias=None) 4 a- N8 t1 N8 X2 N2 a% R CenterCrop(size=(224, 224))( p& N) K3 m+ {3 Z2 O) h
ToTensor()3 V/ Y, u0 ]0 G5 o7 ?' s" T& L0 X
Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) / k1 g; F$ u, R( u( X )} * {( ^& o3 V$ T8 Z9 A$ `% T% N/ D) h. i, e! z# _, @. ?" ~
17 c: s7 q1 M. u0 F' j/ j" m
24 i& p) I) ~; w
3 * Z$ v% X w4 t7 M# S49 ~: g. W4 E7 d
58 z8 n) ^, c$ o' X {. U6 U
6 V3 J! t2 w5 @% E8 F) ^7 % G+ _ L3 y& v. J) |8 6 R; B& H; `9 W0 y9 3 L1 {8 s7 q3 o0 u# V+ k4 g/ ]9 h0 @10# J9 _* H" Q, M: ^! Y
11 3 p9 Y, q5 {8 A$ a9 l8 Y12 ; c& ^# @- U C$ f13; E, R8 `( e& Q" V( j7 Z _% @
14 ! ^8 A x$ b M) G& N, D9 p15' G& Q" T$ P% [/ v: m6 l# M5 ~4 J4 p
16# Z; ^2 g: P3 B2 r9 K1 i( {
17, g0 Z6 ~! v9 l8 ~. G$ p
18 ) r4 X: w5 v' x" B% Y3 U+ b; S19 0 p& \: i* n& k! X: p _, I20 2 _. K8 e x% ] s; l( N, b21 . D7 t; T1 M8 ~ l$ `22 w* T% z6 e) u$ n7 E# ?
23 % V' g2 y$ S9 F9 H! C24# ^4 t% z2 F3 w) a% s/ U+ T4 P
# 验证一下数据是否已经被处理完毕 ! G9 Y' { F, O, I+ N4 m- Ndataloaders/ }9 F# y3 g2 n" D) C
1 % q, A% _" W* S6 Y3 q% P, l1 f2 Y6 y0 M6 k# n5 Z. N8 m+ b{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>, 9 d/ P' }# e0 I$ K7 k! p) A- l4 _ 'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>} M8 y$ L4 b! J5 w" d Z
1 + [$ s& h9 M, i4 G2/ c, j& c O! t% Z
dataset_sizes & K9 O/ {; }/ N11 h: V8 e" W p( }3 L
{'train': 6552, 'valid': 818}- U0 j- @, x4 v, |: P+ W- e& ^, O
10 ^. _: M! h; o
读取标签对应的实际名字! P, `* X7 O3 [8 }* f( ]" C
使用同一目录下的json文件,反向映射出花对应的名字, Q" p4 A! ?! a1 w# f* f& @; J9 W
6 {% d6 m& P: }$ H6 ]3 ~4 U2 K
with open('./flower_data/cat_to_name.json', 'r') as f:4 U, ~9 g. |3 ^& ?. p! M) G; t
cat_to_name = json.load(f) - R* g$ }' h9 n+ ^/ r, L" [1, l! m+ s. g$ r- k
25 a& C( W7 I; [7 }
cat_to_name % n) Q: S2 Q5 b; h% v* @1' ]1 D# p% c w& O" N: k
{'21': 'fire lily',4 J4 T% S7 O, Q# A
'3': 'canterbury bells', 1 X& ^$ F) T* ?+ |6 g2 ~ Y$ E '45': 'bolero deep blue', 5 e# V* G V" D$ g' V7 J* s0 a6 m I '1': 'pink primrose', & V, S9 M4 d0 U( H. i '34': 'mexican aster', , {$ G$ \# U7 R2 c/ u4 ] '27': 'prince of wales feathers', 3 X% t9 R- l/ x, ]- n) `6 @ '7': 'moon orchid',. @1 d% V1 }% c" R$ w
'16': 'globe-flower',* n6 G: @- ?: @" n# H A3 I
'25': 'grape hyacinth', / r1 l2 Z+ r( G5 B1 \5 b; E# O: a '26': 'corn poppy',7 N! I# J/ a$ K0 I
'79': 'toad lily', 4 h- H& p% O+ J% C- b/ { '39': 'siam tulip',: O1 o! u0 ~- c* {
'24': 'red ginger', # p' P) D/ g4 ~7 S '67': 'spring crocus', 6 _8 c# [$ p, ?3 C) r4 K5 w '35': 'alpine sea holly', c* t) M+ ?# H '32': 'garden phlox',6 o# T; E* H' c( j$ a
'10': 'globe thistle', " n; q( T& O- H% B '6': 'tiger lily', K: a; P* E, o+ Q' E
'93': 'ball moss',4 P$ {% ~7 c5 E9 y
'33': 'love in the mist', 3 x/ T; i2 W* i2 ? e '9': 'monkshood',5 a/ Q3 @9 f8 N6 [
'102': 'blackberry lily', 9 E/ D. K' K& r, K '14': 'spear thistle', - k, W5 Q& Y" ^% O9 _4 b. d '19': 'balloon flower', ) Y$ Y3 F4 [" g1 x '100': 'blanket flower',) i* K8 e$ C+ `
'13': 'king protea', 4 n; P# n6 D8 O: M6 M% N '49': 'oxeye daisy', / i- g8 F: O3 i4 o '15': 'yellow iris',0 ]- L7 n' f* n. B+ l0 p8 s
'61': 'cautleya spicata',5 Z3 o/ F; }# n7 p g
'31': 'carnation', / |# Z8 i' m2 \- V% Q6 b '64': 'silverbush',* `7 M% x( N, _' W6 ?
'68': 'bearded iris',* K, l/ m* V: K+ s; j9 J' n5 s+ V
'63': 'black-eyed susan', ' a% s. j$ |0 k: P '69': 'windflower', # [6 `- ^, M+ N& v9 l '62': 'japanese anemone',# X( b" `. e$ E+ V/ l, k$ B( l
'20': 'giant white arum lily',0 S# V; p' j* E: T
'38': 'great masterwort',, i( K' P! y" j( ?& }. i+ t
'4': 'sweet pea',: n( X. k$ Y [1 `2 `4 i/ m
'86': 'tree mallow', 9 L$ ?( k3 F, s" @$ g+ ] '101': 'trumpet creeper', p$ C N% L4 b. P9 `
'42': 'daffodil', , j: J) W) w! Q* @ '22': 'pincushion flower', 4 \: j7 P# B. r% y6 } '2': 'hard-leaved pocket orchid', ( Q$ N6 H6 b. [0 c7 g; g7 u '54': 'sunflower', 8 V: F2 E! ~! N' g `1 p) `+ t: S '66': 'osteospermum',! F% E. ^6 a( @
'70': 'tree poppy',$ b- ^, U: R ~+ W
'85': 'desert-rose', ! q& D& r8 q( o0 A; K U! k5 [ '99': 'bromelia',. Q* n7 O3 m. |. m! g6 E7 J" a
'87': 'magnolia', . n6 q8 |2 j# U/ c2 T '5': 'english marigold',- v. v3 Y- i9 e/ t$ |) P
'92': 'bee balm', 2 W( N4 F: W7 A '28': 'stemless gentian', 3 z( i9 F7 O! R) l& E t '97': 'mallow', - I. j' r$ z( b k '57': 'gaura', 9 D8 c9 @: \4 |7 d '40': 'lenten rose', " q9 k& j0 P; v. w '47': 'marigold', / ~5 f9 o1 V8 {: j '59': 'orange dahlia',+ T/ }. s; x' _
'48': 'buttercup',- u) V- `+ v. V* I
'55': 'pelargonium',4 F5 s" E( n/ N) ?
'36': 'ruby-lipped cattleya', , f# t+ Z& F$ G F- x7 A '91': 'hippeastrum',: k6 k- ?0 }9 f# x. }5 b9 s
'29': 'artichoke', * _. S4 K! d4 T5 p7 J. Y '71': 'gazania', + z: j. b" L7 r* b0 \8 P '90': 'canna lily',7 p; _7 R. L" `/ w9 y$ i
'18': 'peruvian lily', 7 @) E( ?# ^5 \$ M6 z '98': 'mexican petunia', $ [ B3 Y" B+ n. \ '8': 'bird of paradise',! q$ k$ b) m7 h# a! Q& d
'30': 'sweet william', & [9 F1 V0 a J) n; k6 _ '17': 'purple coneflower', 3 F1 Z! h2 U; F- @7 { '52': 'wild pansy',/ Z/ ]. a. l2 B% b0 v. V
'84': 'columbine', * A, ~% J" U* { '12': "colt's foot", - o$ \2 f' F2 ]- i2 N '11': 'snapdragon', , f3 w# X0 r1 p '96': 'camellia',5 r4 a" L0 u3 e# ~
'23': 'fritillary'," w; ]9 _+ T7 p( ]! ~7 s$ J
'50': 'common dandelion',2 m/ Z# o4 F% {
'44': 'poinsettia'," B! Z! O2 L! V/ V8 q# ?) _* @
'53': 'primula', $ a5 `) \5 b) Y! L& ~( c/ p '72': 'azalea', & |5 U; d( A* u4 [( m: | '65': 'californian poppy',0 R8 ]0 _) z: [# s
'80': 'anthurium',% C# n* ?4 j, G! j3 ^5 F
'76': 'morning glory',) l# `+ \" c' K$ M) }3 U# P3 J6 s: L
'37': 'cape flower',& Z5 z& C$ X/ Q& i
'56': 'bishop of llandaff', # `& } N+ o7 Q+ i) ?8 r: C9 i9 J '60': 'pink-yellow dahlia',; ^. F- D8 {. f3 l/ x4 l
'82': 'clematis',6 s5 J. ]4 d; B4 t& I
'58': 'geranium', 7 A7 v- f6 v, p" ]6 c '75': 'thorn apple', # a4 |& f4 {0 j9 K$ ~4 J '41': 'barbeton daisy',! F. o% p3 H% H5 W0 Q
'95': 'bougainvillea',! \* x0 c6 a6 P: ~1 n/ O5 X7 c" f
'43': 'sword lily', ( a) S6 G/ `+ K5 T; { '83': 'hibiscus'," N# H4 y; V$ e; l( n0 [
'78': 'lotus lotus', _: O6 K4 S {8 ^$ n
'88': 'cyclamen', & |6 g7 j& o. p8 {+ } '94': 'foxglove', 9 w6 X) j4 K5 ` '81': 'frangipani', W' p( e. p: j' J- L3 e& W
'74': 'rose', 4 K5 ?4 ^( t9 P9 C- U '89': 'watercress',( s9 r/ I( C# |
'73': 'water lily',# [. g& b# S9 D0 U7 T
'46': 'wallflower', : d5 y/ [* M0 g. u+ p4 t- n! { '77': 'passion flower',* ]) @" n" q, V$ w$ j4 _# d
'51': 'petunia'}3 N- b# `+ T* S5 C+ t8 i$ u0 V
! F# T, n# ?9 P6 f
1; F4 p' t& C+ s3 h0 Y4 b
24 N) u5 q( U% m4 w$ S7 S
3 5 k. a9 B v7 E, m6 f4, q/ v, D* L' T9 b1 m. z
5 + r% `& |5 k+ k9 ^6: M3 v) d# w8 I/ K' m- t$ _9 K
7 ' r( J, d& E" n6 u3 X# ^! o) o83 Q0 t) a! y2 S8 P e
9- g. F2 k; {+ M. }% q2 |
10 7 r" o9 A& b7 C" v, n7 [11* ?6 X8 ` @/ l4 Y
12, B+ |* x. T1 ~4 H3 ?+ r# g
13 ' h9 K9 _! B0 \4 i- U7 c148 z1 G+ E8 K7 n8 r4 g" M K/ X
155 M9 y" v- Q7 D$ u
167 |* `& Q( S: w: p* O0 ^) q! Z
172 I" q% [4 Q# ~) b2 |* b. |! {
187 _ i _+ U% w9 J7 F) x
195 M$ v" u8 ]/ x. s
20 , r* i* f* M7 \& b* Z) o/ u* S21 2 g2 `; X8 |5 a0 T4 `8 y3 {. T: o22 ! n1 p8 o1 `. Z8 ~: w$ x233 ?3 R, o' O4 M8 a
24 / U# l% m8 x/ t' c7 l255 _' y$ u& `1 c, S9 g" ?3 Q! S
267 u5 a; i( y+ B! C5 s
27 ' ]4 S# l% r# C8 q289 S; n( S0 d- D1 \6 y$ B+ d+ D
29- M" A; T. E: R+ I+ Z- }
30 $ X4 u G6 q+ |& y310 p/ Y' l4 p _; T7 Y9 p1 S9 l
327 J0 P* ] [3 [- p w6 h9 {
33 ' T5 p# T+ R0 L/ @0 Y7 B, }343 l5 N0 L$ y8 o
356 Q# C9 y9 @! b! }% ]; V7 g0 w
36, ?6 }4 \% Q9 b' s* @9 y. @$ x9 _
378 K$ U& u9 [" g5 k" B
38% ?" e' c5 g! @/ ]( m- q! q: M
39 , Y! Y% T. Z0 X) N$ q40$ i- y1 R% T; w$ f2 s% @) C6 }
41 7 {6 `* M5 w# i3 ?42 ) R4 d( y* Y5 |5 U) `43+ r' V( ~0 p9 b
44$ Q/ j) E/ i8 z+ W$ e1 T3 ~
45 L& k2 c2 d6 F: r& T. t/ _463 _) |3 P( t# A5 J$ r1 U4 C2 V2 p
47 ' @% [1 K$ G, K' U48 L( a% O8 W& L
49+ o% |" q8 h) b
50 : O) ? x8 M5 O7 Y51 + {1 ^9 S h/ `9 G7 J52$ a0 k# U! Y/ p+ n+ ^1 o) A
536 A9 I; \! ]/ B
54/ H7 }/ V0 O2 w! V* E! Z- ~% Q5 r
557 j/ j0 o& D: Q; O
56 2 L/ h, ^% v8 J, n9 O b57 2 d+ V/ G) L) I t58) f' T$ |! U* _! k1 e1 J
59 3 X0 g$ L# ]" G9 J4 H& S0 H0 T4 r60 * l9 |% x! G# q: j61! u1 u! S+ s& O
62 2 H) D. u1 ]. n9 l$ J1 j639 |, `; o1 g2 H0 U0 \
64 ) M9 ]! R* q0 b65 % K6 O: R) \9 Q5 A0 C; {$ k66 g0 M) n4 I3 N. A+ N67 - X6 h# G" |( P68 " K- q L2 o) D8 ~7 i- z; _ G9 P69 % ~' `8 A* I9 F: E70* v7 U9 \" Y/ S. l3 m
71 5 {7 F. i( Q% Q z$ G0 W2 o728 c9 U: C6 I+ z4 E B T
73 % h+ O$ O- v0 ?# O: g74 ' \ C5 Q+ @6 b3 ~& z" `753 v4 a! d$ J& H' {, K+ ^( L. c
76; ^; q8 P3 |* w
77! B: q0 g$ e- O- l, p- W
78 2 e8 w3 @$ J7 e5 ^79 , ~" Q" K0 o1 f# I! p$ @' ?80 8 Z1 t0 D- [& p5 [81 / X3 H8 V8 U% i- k' Q0 U& B/ O82: O; u0 G8 e( o: z/ \+ v" T; n* u2 x) U
83+ Q( L* l& ^9 [0 l' d; v2 [- h: S
844 T R. e: Q6 B& E; D
852 h9 ^; n; F0 L! F
861 n l) y0 O0 o2 B: |9 m* q* `
87 8 ?# _& {9 U6 g, k* }% }% o, ?88 6 X' D" ^/ l; x) _( d' N+ ]89 * o" [. b1 B0 g% X0 P90 ) ~8 |, j& y6 x- k9 _4 H91 + V- `$ h6 P( L; ^" j+ j92( E) O6 P9 u# L; o
93 # K) y! s, A) ]94# d, ?2 t" ^2 B. v8 K9 C+ s( C" I
95/ c7 n/ m I- h) \0 E- i
96! _ }5 t! Q: V8 g" {/ W: W
97. p1 ]/ v, Y" t
985 W0 w0 V+ s* c% Q
99 , d* G9 x0 D: S1003 g9 } C0 y5 [4 y" `1 v' ]7 I
101& X$ _2 V9 W. |0 v% _
102; K$ l9 ]6 h( `0 A
4.展示一下数据8 B" j1 c+ `# p7 K- [9 n/ P+ z2 I
def im_convert(tensor): 9 D/ E4 x3 T- O. E """数据展示"""8 ^3 e% o9 C( r* b/ O2 k
image = tensor.to("cpu").clone().detach()( l0 q2 r- A* i& B
image = image.numpy().squeeze() & @# K; J9 Z' c+ r M- S # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图7 [' A0 Z0 g4 D3 y) z0 C
# transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c) $ U: U1 T& L4 V. w, O" O% M8 x image = image.transpose(1, 2, 0) ! T- G" i; m" ?- f, Y # 反正则化(反标准化)4 o; k+ ^; ^7 }- T! ], O9 l; j
image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406)) / Z5 X5 z1 d5 | ]) i9 c5 X5 X% U: n' D0 m* u4 |4 ]- A( W4 A
# 将图像中小于0 的都换成0,大于的都变成1 , R l s$ v* h& U" [ image = image.clip(0, 1) / x f# g F' C: \5 L. e, l$ j ]- c c) Z9 r* M
return image5 E9 T: G* V2 S" ~, Z
1 m; g# @* z( J8 q( E" f
2 * Y, N+ J8 e5 ~3! M$ c0 y. M k0 f* Z! L
4+ W& B, E3 v6 J5 s1 _9 J. U
5 - ^6 @8 G$ a& l; W! ?6 ' u6 d( L( X- ^. Y/ [7/ W5 O1 H7 ^ L# }% Y1 v4 s# U
86 n( } i X0 E3 @8 H7 ]
93 b6 V# Z+ P& E) U; W
103 \4 n5 m' B- c8 q
11+ L- Y6 @* j: n2 A. T- C
12# f; I. d% D/ ?0 E( M& `
13 " X+ r8 q. M; r- w14% n& s4 o/ m/ _) F' R9 @
# 使用上面定义好的类进行画图 3 K/ W9 m: U) V$ h% i& [fig = plt.figure(figsize = (20, 12))$ y5 R# }$ [0 c
columns = 4 : K2 n! i! Y: `) M& V* B# grows = 2$ L5 A+ c; g* u
* ~8 u. `8 ]4 p# iter迭代器 ( `2 a: b+ K/ E1 R: n! Q# 随便找一个Batch数据进行展示0 Z0 ^; E( `% e
dataiter = iter(dataloaders['valid']) 5 N9 P' }* s f J' `inputs, classes = dataiter.next() ! D; U. t. Y+ ]. x1 p& H* D" d; B ( A% b0 j) ~' n& q) n5 I4 {for idx in range(columns * rows): ) Z2 ]9 X8 J$ U! e ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])" p) T5 X. f9 M4 F1 V6 j! C
# 利用json文件将其对应花的类型打印在图片中 ' |" L2 T/ ?( t% k( x& s5 G ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))]) : I! _) D U: G plt.imshow(im_convert(inputs[idx])) t9 p. Y# J7 |6 B6 g
plt.show()' M0 e+ B4 P4 S! D/ T
) g3 J4 Z {: j% A& [; T8 O
18 \3 ]# W) m$ j: H e ~6 V6 i
2 ( ]/ M" F: {4 V. ?/ ^# ?3( u6 D/ l6 y% p y+ ?0 _. `$ `
4; k# x7 G8 G7 g* b) @
5 7 A7 y& `" W9 s8 [1 `; g2 A6 + a, f" \0 c ~. y& k0 a+ a7 # q2 Y* l! n+ |2 u+ P& S# W( F8 / R& [0 c% @6 A: |. O4 [9$ e$ b. p! F- J7 U% H& D5 U: O
10 * N% m' y V. N9 k* o- L6 }' u5 [11 1 F: b, t# O# c! b" y% D12 : _6 b, `0 w% Z: T+ J M13 , q* H. O. ~" _14( s. D4 K4 l2 ]2 z: W0 X+ x$ o3 Q
15( [: x b3 l3 D5 I
16! x* k; ~# u: v0 z' @" R5 q
( K1 i6 H. V" G! B
6 K0 z9 [2 L0 M+ o( Q# f8 j, B5. 加载models提供的模型,并直接用训练好的权重做初始化参数 9 [, e" A' F) l# i* u9 \! Qmodel_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']; |% c& J3 G5 {6 V$ H( y4 u+ o/ m
# 主要的图像识别用resnet来做! _+ `8 H% o- Q: ~1 p
# 是否用人家训练好的特征7 U2 T; Y- [3 v9 ?
feature_extract = True+ j( K# {0 X; H' I* m1 v0 Y, I, m
1 8 I Q3 c/ E1 ~" _6 D6 K2 , l. m6 y; A; R$ Q9 t, Q' b1 m4 s3' [1 ^( |6 @4 R6 \% R7 m- G) I
4* E+ g" _) g, f& }5 Z+ H
# 是否用GPU进行训练 5 b. H! I( }: T9 d# p, htrain_on_gpu = torch.cuda.is_available() ' h) d( p% A' l* Z. [0 F! e" q0 L/ i' ^& p C# p6 d- F
if not train_on_gpu: 7 X( U' v6 E: u: Q3 p print('CUDA is not available. Training on CPU ...') ' q6 r7 z% z7 ^else: 9 A2 }" a: F- B6 j. o print('CUDA is available! Training on GPU ...')" N0 X2 Y4 L. W6 c! S6 h
, h8 E+ ~! E9 [5 l9 P; {
device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')8 E( I7 x+ n, n3 `5 k P# J- q
1 7 Z. ]9 k+ _. T/ F3 v( W0 O2 % K6 b# j6 l6 J) U- ~) i9 d% t3: W2 w! S2 o) P
4, A9 b7 z9 \% r S9 X1 Y
5 , z! u, C0 R8 m3 X! S2 f& H% |- h6 . P4 i" [5 n5 ?8 E* T2 D; V79 G( i; J7 T" S3 c5 m
8 " M" N1 u' l0 y5 b2 e' j& R, C9 # z# W0 N; B; e" PCUDA is not available. Training on CPU ... * [1 Z' L1 o0 e% m2 ~4 \1: i+ }8 u Z( m6 H9 \
# 将一些层定义为false,使其不自动更新 0 C2 W0 I* y" ~ _+ P C8 `4 T( ndef set_parameter_requires_grad(model, feature_extracting): * N7 h; L9 u6 z1 Y" C! [ if feature_extracting: & @$ R! L+ ]9 ]; G4 N, }& X for param in model.parameters(): ! `; c. ^, V! V) w9 ]& m param.requires_grad = False , k) V: T# t& R2 u3 Z# U* I. q1 % Y* n) T( Q3 ~2 % t+ |# P; {- B* w1 K3 6 S' t+ E4 d% u; A( `4 6 t# {# |0 j+ H. a2 | y& }9 Q5/ M3 t6 i9 q( n: E. O$ N- _% e! h3 c
# 打印模型架构告知是怎么一步一步去完成的+ S/ d# w$ j' @2 W( R, ?, |- h
# 主要是为我们提取特征的7 F! X `9 f/ x, ^; @! K
7 a/ z( S. A$ R8 ^& ]model_ft = models.resnet152() 6 ]8 G Y* X H4 s$ G( K! Wmodel_ft # d1 M3 @! q* X* O1 X/ \! b# J9 o1+ q9 j8 n! V2 V$ n5 E* f
2 7 w# E4 N" o1 l- J3 ' `' s, w& N8 |0 z; v5 I! z4 ! O- K4 S* U1 F- @! \+ r1 C' T* l5 1 i$ `4 T) z6 |7 W/ cResNet( ) u D# o6 N; j- b# b. p8 e; K3 } (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False) ; H% _, c9 I, h (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True) - }* t M# v8 P7 m- _& Z (relu): ReLU(inplace=True), P4 ^ G: N: |" u6 \* F9 s
(maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False) % P; e3 b9 [: [ X4 a$ ?+ V( [ C (layer1): Sequential(: r- ^* d/ c3 |8 F. I N
(0): Bottleneck( ( B5 \) C* J1 h" ` (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)3 j9 |1 p/ M- o" ?& o3 @6 O' S e/ f
(bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)- q# j/ B4 V9 {: R
(conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False) # N3 S& ^. P% k# y2 `2 ] { (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)" \- n5 w c9 K# o
(conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)- \/ }* O, o8 Y% B- U7 g
(bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True). | L9 I- G+ x9 ?- s9 `( x
(relu): ReLU(inplace=True)- u/ h) Y. e/ X0 @
(downsample): Sequential( % V, F1 {; S6 ~1 A (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False) 2 t( x2 |5 V F( w+ Q (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)1 U% c+ x4 f9 x$ ?; ]. I3 X
); P9 I7 T$ C7 F* l" n9 e
) ) j! M3 ]2 l3 L- w6 X% O中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。; k! b1 [! A; {* W9 }7 d
(2): Bottleneck() ]) H0 y# I! N0 b; g
(conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)- v$ o$ F6 b ?/ Z b
(bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True) B1 _, ^$ N( e6 A5 X/ U+ e
(conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False) * n" q( A3 z! V8 l# f, d: q (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True) ! u2 ?% l" k( Z) j% q0 ?- f P (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)+ O; [; U, v9 I! e
(bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)& ]. B* D) f# l5 d+ O" u, M& r! _; s$ T
(relu): ReLU(inplace=True) % R p+ M2 y: ?/ i F0 C, x )5 y8 y) E6 ]1 a' l% z
) 7 Y- W e, |& D& p3 B5 | (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))# s: H5 K$ D9 T2 f* |0 p$ Z- v
(fc): Linear(in_features=2048, out_features=1000, bias=True) 9 Q) h) i+ |9 s4 @6 o) ; \ W+ s/ ?% K6 m; m ) Z3 f1 }4 z/ f* r, D16 j) w+ U: y7 t- f
2" Q% r3 ] b" I& d3 R. w6 O: t% v3 l9 R
37 |) Z. T' W! l% {$ X( J* i
42 F2 K1 ?. _) h _ P1 k; L
5 # T$ `- B) o/ [) ^% B- Y6- j9 V& [) I" C5 ^/ X+ x& s* b
70 D" J- T9 v' ]; g' u L
8 - z" b3 q/ B7 d3 z9- W% A& d: `, J ]
10 " d" c5 ]+ w) E4 B7 B11 . a* u. m' y$ ~2 I9 G12% @; X/ |" Z9 J3 v% {, s
13. s: b% T% S! Q S
14 2 y9 p0 |* s, |8 K% I3 Z15. e5 l. {! X" ?- ?9 ~3 B5 q9 C
16/ T- }, ?$ E5 M# f8 l, \- A4 t( X' s
17 : S- A, I% z1 R( {0 Y+ C3 z1 @18/ [6 C: m3 n/ j9 J- B
19 9 a8 O1 I- m+ f ]20 & E$ l' k" ]8 W- z! C- N21; c5 r, G A( _. Z. l1 h1 o- l
22 2 w/ B2 o* m' i+ U& a. K/ _23 ( o, m3 U+ q8 c% r. r# Y243 d7 C* o' z1 _! |$ O- m/ _
25 F" I7 e' h" E2 l" b+ E
261 \- _/ Q0 R: p, `6 v
27$ l0 A# {+ u2 @9 _! G' E
28 4 N$ W7 _/ U0 J6 c1 [$ U. M29 ( \7 N0 s0 e8 W+ o' |2 E, d30- ^% u/ k! y' ]0 _. b! j4 `
31. O7 l+ E% N0 ?7 Z( r1 N8 z; E
32 [6 ?) L1 y6 \335 V1 L( `3 j P/ w
最后是1000分类,2048输入,分为1000个分类7 x L$ ?1 T+ g3 G: W6 [- F q. _
而我们需要将我们的任务进行调整,将1000分类改为102输出 ' q% H2 [9 H2 {' N) i; k6 k : P; ]8 j( W- y% U, K% B% f6.初始化模型架构 # ?) S0 R) t7 V0 z步骤如下:- P8 S# `0 Z8 m; o
+ E4 Q; ^2 q/ t& W
将训练好的模型拿过来,并pre_train = True 得到他人的权重参数 ( u* [' g5 E6 f+ j9 J, L( r1 G可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False) ; D k2 O* n$ z无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数; I }# ^6 w4 B/ _6 G" l
官方文档链接6 k% R! o$ ?( m. G6 E/ Y t
https://pytorch.org/vision/stable/models.html7 \3 W$ @% b4 w1 b6 E
& B- @6 ]! H3 @; ?$ j* v8 Q m7 F+ `
# 将他人的模型加载进来 3 b; a! U1 x- v* S6 Kdef initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):! M% a, U6 w9 Z7 Z
# 选择适合的模型,不同的模型初始化参数不同 ) N$ @' D* p" `, ^: O model_ft = None& f- }* {/ \9 q. ~1 ~' X
input_size = 05 y2 w* u( ~' a( A. t
" h7 \4 A9 e( r+ E if model_name == "resnet":; e7 ?7 i$ a& n! F1 v
""" , j' ]& k0 d/ ^* P( Q U. O Resnet152 ' D/ |2 b$ L& e4 {& d """& Y- f7 Z c3 r j& Q$ p) O$ G- p6 }
1 \( ^3 [& W* h ~+ B( J0 r
# 1. 加载与训练网络 1 T6 ?* x# V2 { model_ft = models.resnet152(pretrained = use_pretrained) ) @" m% F& k x: f4 z$ J* _ # 2. 是否将提取特征的模块冻住,只训练FC层 , |! ~5 F- e5 H6 ?3 l set_parameter_requires_grad(model_ft, feature_extract)) d9 c6 N! `; P( f4 @ M
# 3. 获得全连接层输入特征 q+ x+ ~. u! c
num_frts = model_ft.fc.in_features" g7 N7 F$ A3 U0 t
# 4. 重新加载全连接层,设置输出102 , B5 J8 H$ [0 B. ?8 s model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102)," v2 F: C) P' x% G) ?
nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1 # V$ J. u+ `* k" O% c input_size = 224- g8 P% @0 ?* G. W6 i( h
! E5 V: D5 x! }3 f- z
elif model_name == "alexnet": # W# Q2 _% U3 B& _5 { """ 7 X. V& Q( ]& n* d- \ Alexnet 3 Z2 }; i& {, [) s """ 9 l6 ]9 C& ]5 ]+ P$ d9 M model_ft = models.alexnet(pretrained = use_pretrained) 8 A& f: @+ w8 u8 b set_parameter_requires_grad(model_ft, feature_extract)2 w/ K- u" Q5 o+ \3 B
9 |4 a1 F' ]* L3 ~& v
# 将最后一个特征输出替换 序号为【6】的分类器, G% G$ T+ _' A. [3 g
num_frts = model_ft.classifier[6].in_features # 获得FC层输入8 T: A, s: m; D' `
model_ft.classifier[6] = nn.Linear(num_frts, num_classes) / V0 L& n) z/ |7 ~. L C; Z input_size = 2240 J- b% V) n: {8 A
1 r" y- k# k. }8 L" C" J3 }( T9 K
elif model_name == "vgg":2 Q. T8 m6 E$ e4 i: V9 _! B$ T3 V
""" ) R8 i# R7 t6 } VGG11_bn 5 V1 S# C& R3 o' c% Q """ 3 i* x5 K9 n7 c. ^+ @8 Y, e" \ model_ft = models.vgg16(pretrained = use_pretrained) 3 S6 ^& p1 N. _9 J set_parameter_requires_grad(model_ft, feature_extract) . H, [( |6 h% ?1 [* ?1 i2 _; u num_frts = model_ft.classifier[6].in_features( Y6 M4 Y5 Q+ x9 U, J
model_ft.classifier[6] = nn.Linear(num_frts, num_classes)5 t$ w6 [7 h: D$ W# D% f
input_size = 224 0 _/ ^! G! q6 e9 z# n 6 A4 y7 y+ ~4 u8 { e elif model_name == "squeezenet":6 \+ y& m' i$ h5 [. M
""" + g3 Q" m, J7 g% ]$ C3 b Squeezenet ( x, t2 v. O9 ~9 x( ]+ a """ 2 R0 t+ t2 b" S1 k { model_ft = models.squeezenet1_0(pretrained = use_pretrained)( m* f8 D" M; l( ?$ H* p
set_parameter_requires_grad(model_ft, feature_extract) - b; @* N3 ]' V" ]- \2 _3 Z5 w model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1)) 5 E0 G2 F. D9 V) e5 Y1 V9 X! R; `& G model_ft.num_classes = num_classes # D0 l/ V; [7 J+ H1 L' G' k input_size = 224 0 b" f2 y7 }7 E a4 h; I: {! ~6 v# v6 }7 n- @+ Q; G5 u" `
elif model_name == "densenet":! o( }; o! S* j. K' q8 q
"""" V, H+ C3 v9 _) M- \
Densenet 1 |5 i9 |; L# f """# p; d4 B* i4 \8 ?
model_ft = models.desenet121(pretrained = use_pretrained)2 Q K7 p& N9 L0 q" q2 N
set_parameter_requires_grad(model_ft, feature_extract) |. {1 I$ k9 R" N num_frts = model_ft.classifier.in_features0 }/ i5 ^) d8 _/ O4 J$ ?$ T
model_ft.classifier = nn.Linear(num_frts, num_classes) 7 m {' l. c+ E9 `# G% C input_size = 224 5 M2 c+ R7 ^! y' r 5 m6 i/ J: X& e! V. Z elif model_name == "inception":, h5 E$ O& L& n
"""1 @* @; T5 M0 }( _! g, F
Inception V3 2 V7 W4 R4 @6 l* r) l- M """. g/ @1 T" L( j+ i+ u+ @# U
model_ft = models.inception_V(pretrained = use_pretrained) [* R/ K3 @' x) h* u; P
set_parameter_requires_grad(model_ft, feature_extract) 9 x; Y, U6 f' [8 ]' [ 4 \. A" e7 w; V3 A num_frts = model_ft.AuxLogits.fc.in_features . \) G8 k* w0 T. s: s4 o# l6 p, J9 | model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes) & F( P7 Z; O1 ]0 b- h8 _4 x' R5 f) N* l' P! ]- ^
num_frts = model_ft.fc.in_features& v5 L' ]- ]" k
model_ft.fc = nn.Linear(num_frts, num_classes) 6 f% F7 L6 M0 j# s* ], O# @. S input_size = 299, _1 g6 N# c- h, G
: ^9 F; m! c$ U' z M' m else: ( V* m" K2 H! x ~3 r7 d print("Invalid model name, exiting...") 9 N/ Y7 y8 x1 f7 g3 f exit() ( K l" J8 W5 h( A' C+ O, c+ v7 G* ?# G7 { u
return model_ft, input_size$ e: |8 ^3 n9 \4 z8 \+ y
: Y: x V& U9 Z& X h( O4 M
16 j6 r- }( j3 Y& ]8 h0 C
2" g2 _9 N# z& v% D- M' y
3. U$ e- D* U: A! s% v4 F: @
43 \/ b9 a& O( e! C, o8 \& i1 b j8 {
5 / f' x7 U) c* W6: R9 C+ h0 ]5 g, Q, S+ _8 J
7/ D$ Z) \2 R' F$ S- ^# \
8 0 @. V6 ~6 p7 I. W( W0 v$ E9 _' }) f V* U* _; r7 ]4 v10 0 |8 ?( l6 W B6 A+ Q1 I# N11 7 D' i: ^! ^+ Q' E4 l" ^12 & M+ r1 }1 m5 p3 s" M: `" ^139 }6 z" L/ H, P# M. z2 n+ R6 s
14% `3 J+ k0 _3 \# u
15: F, q f5 ^& s. [2 i# Q) s
168 N& M" ~; F) ?- R
17: K3 K5 }" y# h1 x2 U9 N/ \* z
183 ]" L9 R2 \# Z
19 6 k4 f; n- I; c& d20 # n6 \, ]/ e6 T& o2 x21 ; `/ E: s4 ]' K$ I4 B226 G c ?9 b, r& ? B, W3 x
23* P- Q- Y! N# Q
24 ; O) ?' R; L. B9 T9 o5 J! P25 % e9 z. k2 k6 a0 l26! U' a3 `( K( G& m
27 ' _, X& n2 z4 G" G9 x! d- q28 + h8 E3 I. D7 a29 " w+ B% h0 \" K0 X# M30 l+ |# B: }$ a. d- E2 E7 a
31 8 D: t n8 d5 s8 V0 ]328 F+ c) {9 C$ W5 K( B" V- ^9 t
33 / r1 O- z& F6 `* E0 o% n5 G34' G& ^$ Q* g) X) P/ k8 J1 v5 l
35$ P7 U* c% l8 R: X7 q: e1 z
36 ( D0 L* v) L% O* x3 H371 [! d/ [4 k5 s. E
38 h2 B, L" N+ o" L2 j39$ s' t1 J6 c$ l
40 . Y* A4 ]% H8 i' Y417 l W+ O# [! ?% V% V
42 1 M9 g; }- ?! Q/ G6 c43$ l9 M* b) A8 g% y/ @( S! a' o
44 0 l D" |* @. q; m1 j45 { g7 D3 J6 Q+ T468 ]; k$ Y: p; c5 K+ e- D9 O; k
47, p1 g( B5 ?! z, K
48 % J/ k9 i! F: I5 h$ r5 M( Z0 A49 4 [& `( ?% x+ O5 X; A. f) y6 z" y50$ m7 U$ u' k5 Z7 f1 O
51 ) Z" m2 r; g( ?% M' u. {; Q7 H52 3 @9 C* ? ]8 L53& e. N+ ?$ k( S
54, s" }" |9 ^6 H& o* B/ s' M/ T
55 * B* e; W0 m0 s: Q0 f56+ B# u: J' g+ P+ S0 ]
57 * m4 d. y8 p$ }2 X5 x T58 E% l1 S. e3 q. w. X! V) x/ R+ r59 % q- {+ n1 d% ]' \' }. e60 0 d" E2 b) Z0 c6 c7 _& _: \) c61 0 Q2 ?6 F! P( p' j% ?62 ) X2 s% q7 u0 ^+ b+ J- S63 ' n; u6 ~9 r+ i6 l" a5 R9 V64 6 o* V e! X" ? n! z65$ T5 q# T% i" e
66* Q, n5 ^) h: l) f
67 , T7 G4 L: n3 k0 m: s0 T0 l68& A6 R$ i; L. p+ [5 o8 I
69 5 v# n, C& S" y. J2 i& P# j70- D9 e) P# w. l# q8 t
717 R y+ f6 F8 h
72 : c# c. O/ T _' F733 t# X* f; M, [" M5 F9 f
74 4 I; @" S- S& [; ~6 |0 o& k* C' O750 C' s) j) K1 u
76, d$ ]. f$ X8 Z3 ?8 P
77 6 b* W! G" a" b! s783 O, e! J' G3 g# N( {4 L" y3 C
79 0 }8 @3 B0 F5 H' e( P( K800 z9 ]) w- o; L
81 6 |: r3 B+ X+ {7 Q! S: p4 N+ ~82 % R U& ~" A: W" V" z) X83- b. u! E+ w: R. z9 G
7. 设置需要训练的参数 7 U# t3 i2 w. x# 设置模型名字、输出分类数 6 d; H: `& a: tmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)9 A7 V3 [% W& |& t% w
# j% y" Q8 G5 e+ }: e1 E
# GPU 计算 # g5 b4 C- W& Imodel_ft = model_ft.to(device) # }+ Z: ]1 ?' B* \& j' D6 R" a( p; L' x# g; J
# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取, @' x: o# X/ p0 Q* O
filename = 'checkpoint.pth'2 g) O3 X# q( Q; U( J/ F' X
$ M( Y* h( _1 X# @0 R# 是否训练所有层 . z; ]3 w) G4 @1 V( f, d, d" wparams_to_update = model_ft.parameters()# A" A9 S+ r& J# a) _
# 打印出需要训练的层 ' X3 Z8 I7 l! d! |7 ]4 Jprint("Params to learn:")4 _4 I& l) _% p! g2 X% x
if feature_extract: / p" ]% C8 S1 E) O2 C params_to_update = [] / ?4 `( S$ U$ N. d( m. G for name, param in model_ft.named_parameters():! L6 j4 B3 R- I$ m: E
if param.requires_grad == True:1 B1 |3 c& E# O( ?7 E: h
params_to_update.append(param)+ m% D% k3 A* ~1 h
print("\t", name) ) l! _0 l4 b2 b& y+ Selse: 8 ]9 d* D% Y- z& U for name, param in model_ft.named_parameters(): `; ?! R# W. O
if param.requires_grad ==True:$ s8 u2 B9 I* M9 _
print("\t", name) 5 ?- P! j* H* R( s1 H* \0 M3 Z
1 ; z9 `# |- F& D- E- z8 j0 u2; _& [/ u" K! f) H
3. S: e, b- ^' s9 @6 K. T
4$ q' j5 ]1 R+ s
5+ A4 b0 P7 m2 A% S3 q, D, F8 l$ i
6 " D* b/ V, T" K; [7! w0 \8 k. [ [
8 3 ^ i; c: w9 R6 Z! f9) C6 t2 _$ s4 _- G. S9 Y. }- O
10 s* Z, ]# e$ _11. X, C) V o' v
12 / v6 S6 l [' P! `, f13* _2 U6 L0 W. |) m1 E. L
14$ s! [1 B! b+ \* R: `
15+ B9 |! g! T% w: M( w) m7 u
16 $ T3 ~* o3 {: E0 S17 $ Z3 k' b( I' R7 ~& }18) l- [' i2 o {+ b# Q
19 B# X* m/ _, E) A! v }5 g20! ^' [, m0 k1 A
21 . j2 N4 ^! F% K* e& ]( w* f22 1 Q/ A! \' x- X# y. `23+ b3 `+ b3 y# L# I
Params to learn:+ _1 p9 A8 n0 n( \# n
fc.0.weight+ S; d# G7 c* A
fc.0.bias$ ~7 g2 e, J4 R% d( `
1 ( T( `; M) a! @* b! A. b2 ! q, c6 [- c$ \3 ) R3 f6 }8 w5 b; o% m/ p4 q7. 训练与预测 & u" P, i7 g( [9 H: N5 K, t% t6 J7.1 优化器设置 - P* t9 U9 V2 _3 P6 o4 y* E$ B/ M# 优化器设置 # V0 M# \: `% J: W; u1 i, v! doptimizer_ft = optim.Adam(params_to_update, lr = 1e-2) 8 u' U) u, v B/ V* j# 学习率衰减策略! @9 s2 n8 l( F! Q" l
scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1) , s* t& ^; b9 Q9 \9 i1 C3 {# 学习率每7个epoch衰减为原来的1/10+ a2 e" \& H& @ o1 o i' L
# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算 ; W6 Z9 e5 k& k6 r, `% o % d- t' z* a; x0 p: ?; Vcriterion = nn.NLLLoss() - `; n" @9 X; ]# t7 P. J1- O" @; l4 d8 C7 y: p
2 ' f4 W9 u0 Z2 W* M& A6 V5 V* F0 Z3 * q+ C" L e' {7 p$ [4 . f" L" K; z, Q* X; U. D2 ~1 l5 % h, ^: }% X. x# C S# R6 ' y* z- c: v8 g1 g4 U7 v1 t: z3 L7' X9 m1 \/ v1 H
8! b$ c* O3 z: Z+ u- t0 K5 \9 [
# 定义训练函数 4 a( o* I* N0 B9 u8 U9 E#is_inception:要不要用其他的网络 4 E' s9 `0 ^0 S: Q! h1 pdef train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):7 d: X, n* c$ n
since = time.time() # J9 z6 D& K8 ?; ^: q1 ^! P #保存最好的准确率 " t. e/ w. S- o) r% q# q best_acc = 04 R( x! W3 i4 i( C% q7 U
"""+ W& W; c, M: U/ n
checkpoint = torch.load(filename)1 s1 M4 ~! i/ v" I z. W
best_acc = checkpoint['best_acc'] ( v, F O) I6 x' U+ K) @ model.load_state_dict(checkpoint['state_dict'])$ q+ ~8 f x+ J2 ?" ~: \2 C
optimizer.load_state_dict(checkpoint['optimizer'])$ i1 P+ g3 L& A' X" x2 m
model.class_to_idx = checkpoint['mapping'] 4 U& i7 c1 j* |! [6 A """$ w7 |1 Q& S2 B& D0 w% E
#指定用GPU还是CPU + W$ N/ _2 p% y" G1 i, h model.to(device)% g3 r0 f/ f2 t. t- [( _ R
#下面是为展示做的0 h& J8 s3 }) B& O" `7 ?
val_acc_history = []' m* w- G: t' Q
train_acc_history = []8 V ]. [1 u5 |* n: C- ?% |, t: n
train_losses = []1 e' v$ @# a0 |) G+ a
valid_losses = []" G4 S; Z: M/ @
LRs = [optimizer.param_groups[0]['lr']] 8 L9 z" G6 h, b# f #最好的一次存下来 ( T m9 J" `: }/ o best_model_wts = copy.deepcopy(model.state_dict()) & a) Z, m! _: K3 c$ E. a" D 9 y2 C; W. [4 e) ^) S for epoch in range(num_epochs): ! A, ]6 Y6 |# W- X: p7 J print('Epoch {}/{}'.format(epoch, num_epochs - 1))0 [5 y4 }; j0 B! u: Q/ z
print('-' * 10)* x4 H* o) V9 ^0 H* n4 L. B$ G9 J, Z6 T
8 z! P$ u& x% {8 ]1 ~ # 训练和验证 + [* l! ~; g. H0 n2 m( C& h for phase in ['train', 'valid']:4 \! A! k8 j7 R1 f! H* o5 D' @/ T
if phase == 'train':1 q g* I# \( N% n/ T
model.train() # 训练 $ u2 S7 x w$ Q: D5 T/ }0 g/ B4 C else: 0 C! C/ r. q. R- G model.eval() # 验证 - g' z1 E$ K/ N+ l2 Z5 Q5 I @$ Y7 o& ^0 `
running_loss = 0.0 6 C K, b/ f7 x running_corrects = 00 O( h2 Z5 | M- j! w5 O, u# N
. P9 F% t. Y0 p. ~ # 把数据都取个遍6 {# }7 U E+ A8 c" A
for inputs, labels in dataloaders[phase]:# ~ b4 i% p& Y% [9 T
#下面是将inputs,labels传到GPU0 e+ N' c: ~- _( [3 r! j& S
inputs = inputs.to(device) # e/ g3 v3 F" J labels = labels.to(device)$ _6 b) h0 f1 C) Y! X
$ |# K: n7 Y; a0 R) F0 n! @# g # 清零' M% j5 Q2 p& }- a2 M
optimizer.zero_grad()! T+ J# s9 B. E6 {2 B m: ~
# 只有训练的时候计算和更新梯度; _7 U0 {) i6 V! i. d9 X( S
with torch.set_grad_enabled(phase == 'train'):& C5 |. V2 p! L7 R. q }* S+ W& U
#if这面不需要计算,可忽略 4 c3 b; j% Q0 n% K9 Q if is_inception and phase == 'train':* C3 D6 O0 P) [3 L' j
outputs, aux_outputs = model(inputs)5 N- J* ~0 h8 a2 ~' q1 G
loss1 = criterion(outputs, labels)+ [3 ~2 G5 @* ?/ D% M( K( L) _# g
loss2 = criterion(aux_outputs, labels)% F) x* L. L" ^ J# A# u: s
loss = loss1 + 0.4*loss24 O& F- @$ R0 C0 \" r
else:#resnet执行的是这里8 M6 G4 w2 n3 H& f
outputs = model(inputs): _( I3 h% _$ Q. e9 K2 @# w, L
loss = criterion(outputs, labels) 8 B0 U: w1 p- A* y; }* i: w) W- v" u1 U7 s& |
#概率最大的返回preds 5 ^. a1 ~. x2 U7 [1 b9 p {# w( s _, preds = torch.max(outputs, 1) 2 @% Z/ e9 S5 R A' P1 g) ?: f 7 l, q. e* @! E # 训练阶段更新权重7 T- L S2 [+ N* T, i' ^- J- T
if phase == 'train':* Z9 }- [% g8 T; i$ N- ~0 m
loss.backward() ! g6 I% q1 o6 N* ~ optimizer.step() + Y+ S4 {/ b' l( m3 K, @$ u* I0 D& g) |# Q
# 计算损失 - d/ ^+ S& O5 W( I- r2 @& ? running_loss += loss.item() * inputs.size(0) }4 o5 ?/ {* o8 \: i
running_corrects += torch.sum(preds == labels.data); C* ^0 s" g6 L+ E( W& N
6 {2 Q% b5 E' n/ I+ Y #打印操作( j6 N- W. ]2 ?- Y0 Q! x- H
epoch_loss = running_loss / len(dataloaders[phase].dataset) 6 }! Z% \% M3 j' c0 H epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset) `( b5 D; y7 y- [! X8 |* S9 k, A3 p! \- w
, I4 C7 a# |8 s* ?3 N: [6 B time_elapsed = time.time() - since " Q8 W% \2 W; p1 X/ P/ z* t/ Z- c print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60)) 2 X4 O4 J/ ]0 h print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc)) 6 [7 z9 u( V" E" j$ K V. s7 n1 m8 F* K9 k# i% E: c2 B8 E
0 b5 U0 M4 f8 _8 c" m # 得到最好那次的模型 , J$ ^' v8 e& t, `/ U if phase == 'valid' and epoch_acc > best_acc: # A$ f4 X" a% u/ {+ L best_acc = epoch_acc - W' {9 P: ] @+ x) s5 F #模型保存 , |! z9 Y8 q( D best_model_wts = copy.deepcopy(model.state_dict())$ Z* _3 G% W+ e. P& h( C& [
state = { " Z3 d) z: @* O #tate_dict变量存放训练过程中需要学习的权重和偏执系数 9 |) a8 O" R& M" v- P( n 'state_dict': model.state_dict(),7 Y9 b C+ U5 f
'best_acc': best_acc,% U! ~5 K1 ^. u
'optimizer' : optimizer.state_dict(),. D% b5 C$ ^8 X- F6 @- B+ P
}; T; a3 n6 b U6 l0 U
torch.save(state, filename)/ T) I- z. s9 G$ B! O# a2 l
if phase == 'valid': 2 R1 ~8 \' C6 q- C% D' j val_acc_history.append(epoch_acc) ; T. A7 }; U7 U! ^. r valid_losses.append(epoch_loss) ! R( H4 t$ P4 r1 { scheduler.step(epoch_loss) 8 t- ?6 o- g& a6 b if phase == 'train': 1 ~+ c/ y% d, E) j7 ^/ W/ [! ^ train_acc_history.append(epoch_acc)$ _& D) t- d% `% \# ?; e
train_losses.append(epoch_loss) ( M- j- k) C+ O6 R7 v0 D5 q/ m" A0 g7 O# B( a# E& P
print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr'])) ( G6 l, o, N5 A: p$ f5 @5 q* P+ ^6 W LRs.append(optimizer.param_groups[0]['lr']); D a/ S: t* B8 m6 @/ ]1 I# j
print()5 Y8 Y' ^8 x8 P( X \4 S/ f; a6 r
! J. v0 @; J P! v7 x time_elapsed = time.time() - since! l/ i, ?( U. X& E0 w; A
print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))3 j! o! g N9 x$ |- J; A
print('Best val Acc: {:4f}'.format(best_acc))8 c6 [" ~/ {1 r% { a8 f
1 N: t: L: x" o7 A
# 保存训练完后用最好的一次当做模型最终的结果 9 t! y: }/ [+ s: M/ z8 @ K* B model.load_state_dict(best_model_wts)$ x' U( g7 B, k+ \9 w
return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs ( ^1 s) N& j4 d5 p- s, | Y! p" e& }; r6 |- `6 Z$ H4 ?/ {. f