( M4 [1 \/ E* a _. `5 Iif not train_on_gpu: 5 h! ]% ^8 `# R" W3 c4 t8 m print('CUDA is not available. Training on CPU ...')1 j; i( w4 E% D6 I! P
else:1 f% Z I& ^3 u" ~$ J
print('CUDA is available! Training on GPU ...') 2 v& j# x9 e4 s: D. ~ 9 Q" ?& q% K$ c* xdevice = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')" S/ P7 t7 C# h
1 C; ?- E) y6 M+ P" W
2- H( o3 w& j% B% l" _! X0 s8 l
3 % N q6 |$ _6 R, u @* J6 Z48 U x1 e; T8 J$ y" N" i9 Z
5 * B4 k; r% } y7 [$ {/ h4 P$ B67 V/ m }+ x4 a, K) y2 L5 ]% p
7 ; T0 Y1 y! k: u& |6 u, E+ y9 V, h8. u9 F! i0 A s* e0 @1 G0 Q
97 M# B1 @2 R/ S% a5 W. @( k
CUDA is not available. Training on CPU ... , @' r9 {: A+ Y3 y1+ s: w( A! X$ |% _# Y8 I* z
# 将一些层定义为false,使其不自动更新 8 ], h* Q9 H2 m7 H6 |0 edef set_parameter_requires_grad(model, feature_extracting):) @4 b( o" I8 l/ k! X1 `! @
if feature_extracting: 0 V: c. s0 ?1 d' l& h& u3 U for param in model.parameters():9 T) }+ k# l$ T
param.requires_grad = False8 S2 [" O# A* t! j. |) u! p
1, e+ \6 Z Y* k6 Y5 _ d" f
2 7 b% k `* C& b* E# S' ?" V. u3$ ^8 N+ t" @; C. G
4" K9 C* V- P, v
50 o) P/ E$ a% f" x
# 打印模型架构告知是怎么一步一步去完成的 % M7 h7 @7 @& a7 D+ e# 主要是为我们提取特征的4 H, ~- E$ c% y9 W$ H0 f
! Y, K$ r$ S m8 C6 S" J: cmodel_ft = models.resnet152()& B$ N9 ~! b% B, _
model_ft# h/ P$ l* m+ S7 m) n3 h
13 W0 F$ E1 a, Z& }4 S4 w5 \
2& }* ^2 J9 X7 ]& t2 Y2 t- \8 \( T
3 / c$ j6 W r2 v1 D n3 }7 r2 x4 / T" D# Z Y! Y- a$ U3 A3 f" x5! O6 M5 n8 G* ~. _: V
ResNet( 9 l3 i$ |" s. r' t& W (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False) : ]: O ~, x3 {! c7 i$ d5 V! a (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)2 a- }# @/ e, U" G: w3 b4 t; M4 J" i
(relu): ReLU(inplace=True)4 Z% z3 r2 m- n1 }
(maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False) . i# a( o; C0 P$ K (layer1): Sequential(3 M, j8 w" M& s' _
(0): Bottleneck(2 i3 d1 Y$ ? [
(conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)6 v- v5 Y! U. v! }9 D$ ~
(bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True) G9 k: G1 P, B
(conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)4 Y o, y `7 ~
(bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)+ F+ V, Q3 d% O2 H2 c& S- X
(conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False) 9 D; L& i+ P a% y2 Z$ w+ k- x) ~9 u (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)5 V; S, x! {" H2 s- d- [. J8 M# ?: I
(relu): ReLU(inplace=True) - S' O( V0 ~! {! z$ E ] (downsample): Sequential(7 M i) l' p' N6 S# P8 ]
(0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False). J* k0 V& L/ d# A% ?
(1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True) 5 p2 b) ~7 u: O# Y7 p. j )2 Y9 G- Q6 y8 C$ U9 i( p
) . X. G+ `: P* D/ X. m1 Z# i) @中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。 4 H4 |6 T, F& ~) J: R% u (2): Bottleneck( 2 B. S0 U7 R5 W; d (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)+ w" w# L. l4 i, B$ D1 x
(bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)% S, m9 n4 y8 |
(conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)$ Z* i% ]! p& ]7 q5 s
(bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)' ]' W+ \4 Q/ D8 v
(conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False) ' _3 m% B v' {0 ?. b! l$ ? (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True) 1 d' j, s' \! x; G# _ (relu): ReLU(inplace=True), Z) R9 z m9 W2 n# u% {! _
) / G3 R2 U0 M U H. _ )0 U! q' K' ~, j5 m- \3 a
(avgpool): AdaptiveAvgPool2d(output_size=(1, 1))2 S4 r' |+ H3 H
(fc): Linear(in_features=2048, out_features=1000, bias=True) $ ~5 j# _0 h" r& v+ r) / F2 h( @- ]3 D5 S d8 F5 v* J$ r* y4 { \# X( \1 / F8 L, g+ f+ {7 C0 \24 ?4 u2 P' C2 ~. z
3$ a, e- m) Y3 t
4& C2 a2 A1 T3 M' N8 s0 m3 D
54 h% W) r0 ~1 w" Q. `" w
6 # O* g% J M2 r. |7 Z0 m7 / z$ n$ k& j7 c: Z' N84 ?7 ]6 p' x1 f* N: p( g1 ?. D- @
9$ s A+ w( `6 ]
10" J' x( i" D( y
117 c4 b/ B/ \- Y& W7 p8 H" h, |
12# z; B/ v8 ^, M! V L7 e
13 7 m+ T" G" L. c0 D148 ?% N( G9 M9 Z2 j
15- C% ? Q" b+ V9 u
16 1 `( d. v" B: v% P7 g ]) C9 E% A17 0 x, H( K. }* [" p1 q187 q n* y% }6 c" N. y( d4 O
19 1 `+ L$ a7 ^7 X: C9 a20 ( O+ X& z3 A$ C$ `! n* u21 & i7 q: I, J- d9 D9 g# f221 M9 O+ @( j+ R5 k! W
23 ; n: Y1 O2 ^% u7 _- P24- m( o7 p: K8 h6 ]- u( i: i
25+ K% v3 D8 t- O" N T9 C5 ^
26 8 w% x1 Q8 o! q" @$ q0 c277 U2 x# p9 t3 q$ |" s1 y8 b/ S
28; o! h# i7 v }: G7 s/ @
29 1 Q4 o8 ]7 L1 g& u7 f30/ g s( V/ k6 T# o+ r; h J& ]
31' f$ a+ e/ D- | L5 A Q2 ^" G3 S
32* H9 Y8 R! g5 y; m; t3 M) a
33 + x, Z: Q: s/ X: Q6 x8 l: i9 L最后是1000分类,2048输入,分为1000个分类3 q O% G% x* P4 P
而我们需要将我们的任务进行调整,将1000分类改为102输出 / ]7 J r" J% C G) v: M# r # x# U) W4 f" x! U6.初始化模型架构5 \1 e, `* q" z" l
步骤如下: : v- V3 ^$ x8 q' K$ M3 W3 M- _: x, T
将训练好的模型拿过来,并pre_train = True 得到他人的权重参数 + g8 m, l% U3 X! ~0 t7 w可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)( W1 i: M, [9 }' |. v6 a6 `
无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数 @: F, u& c( z+ K' H
官方文档链接 : Q/ l$ j C) N7 f) ^4 `4 B& ], Dhttps://pytorch.org/vision/stable/models.html D& k8 G7 T4 V$ T
( V5 Z6 K: [2 x* }
# 将他人的模型加载进来 0 M0 R% ~5 H. ]1 ^) M# e# Udef initialize_model(model_name, num_classes, feature_extract, use_pretrained = True): 9 ^! C; E" Q/ W+ \4 I # 选择适合的模型,不同的模型初始化参数不同2 E. q6 R6 e. n
model_ft = None : g% N- R/ {! g# g, s1 v+ V, S$ i input_size = 0 ! N6 [% L' `' j% M# ], v1 w' u/ P
if model_name == "resnet":/ i1 V9 v/ v! d7 e7 @0 I! H
"""% d8 R' b% W( D: Q9 v" ?$ y
Resnet152 ( H& ~6 Q0 t! N """/ F6 Y% F3 Y8 R* T& K& P) ?6 P
# ^' N! h- A5 v Z
# 1. 加载与训练网络 5 q8 ]7 x- U G8 S model_ft = models.resnet152(pretrained = use_pretrained) $ X6 L7 p/ Z0 u+ @' s4 ? # 2. 是否将提取特征的模块冻住,只训练FC层) i$ y; p- s5 Y. m' I
set_parameter_requires_grad(model_ft, feature_extract)0 i8 K. D$ {1 x9 J. c; w0 ^) _8 U
# 3. 获得全连接层输入特征 / H# ~) Q5 b+ c; Z num_frts = model_ft.fc.in_features 3 q% M1 O4 T. `: B, s # 4. 重新加载全连接层,设置输出102: z# p2 F% k/ a7 b$ ^5 o
model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),, G# H u- m! O/ b
nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1 8 R1 a/ S$ L9 M' w9 { input_size = 224) y! w5 T$ Q8 [9 q, J! n& @0 y" F
3 u) b/ U7 e9 f" L: f elif model_name == "alexnet": % M; F7 S' R4 a T """ 7 l8 f% S9 B0 s, I6 b0 ^/ Q* V Alexnet ! Q8 {' Q7 Z: ~% L2 K$ t( ~3 ^ """/ T b: n* ]! L5 f
model_ft = models.alexnet(pretrained = use_pretrained)1 k/ o# \* d E$ p2 C2 E5 a
set_parameter_requires_grad(model_ft, feature_extract) 7 @3 x& |% V. H9 S T- h* Z/ b+ h* g, V/ M
# 将最后一个特征输出替换 序号为【6】的分类器2 H- u6 y9 t$ i% |
num_frts = model_ft.classifier[6].in_features # 获得FC层输入- I9 K4 |' N k! S, i/ _
model_ft.classifier[6] = nn.Linear(num_frts, num_classes): p. z. \* j& ^4 }2 F& }; w
input_size = 2249 r& p0 S2 J& p4 _7 Q( Z- g; t, x g% D
. T% H) C& s, S* V% m' p elif model_name == "vgg":7 P; C: l: |2 m* `8 N6 A
"""$ e+ T9 ?7 @0 [ u! |" _
VGG11_bn6 @+ [- X4 u' ~' S
"""3 E7 ?% J1 n8 s, y6 }9 G
model_ft = models.vgg16(pretrained = use_pretrained) - ]% }+ o& E2 G6 y+ o7 q set_parameter_requires_grad(model_ft, feature_extract) 7 H3 z& ^3 t* W, ]+ W' H: O num_frts = model_ft.classifier[6].in_features/ }$ W, S1 T0 g# [3 s2 c. L
model_ft.classifier[6] = nn.Linear(num_frts, num_classes) 8 P5 t$ h7 H1 x7 n input_size = 224 8 ^" f% @% U. o5 O 8 t7 I# v' p9 B+ n" G' t; W: V elif model_name == "squeezenet":1 E" N% }5 n H+ p
"""! x- _) A* x: A+ {8 |
Squeezenet4 f$ C( R0 D. ~. E0 z! ]- e
""" ; F+ ?. j/ u0 G+ E model_ft = models.squeezenet1_0(pretrained = use_pretrained) 4 ^5 j2 m9 p9 H3 G6 z- w set_parameter_requires_grad(model_ft, feature_extract) $ y. z0 U( T- O& Y model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1)) , ~$ i( a3 b5 U' L* k model_ft.num_classes = num_classes # l5 l& ^# I3 @4 |5 e {$ ^ input_size = 224 0 x: i2 n& H% G9 k6 {% r - w; w' ~" f7 }3 S% d elif model_name == "densenet":$ K- i7 q$ N' e6 n( q
""" : a8 ^1 x% D8 `- i, g% h Densenet ( v7 X: K: Q/ c2 @6 x. d& h """ / R* V* j% m) p' e; l j+ ?/ ? model_ft = models.desenet121(pretrained = use_pretrained) + j1 f7 t H5 o4 F- V8 M/ O set_parameter_requires_grad(model_ft, feature_extract)! ~( z1 t, f, P: `4 j$ y7 f
num_frts = model_ft.classifier.in_features }) l u& ?) Z, U1 w8 { model_ft.classifier = nn.Linear(num_frts, num_classes) ! X) q: [( l3 {+ @4 K input_size = 2243 v; @4 N5 T% V" _# B
: u/ x$ S9 P$ v% O, a# d, |$ t6 M
elif model_name == "inception":; b1 X1 X$ ~9 e: X8 `
""" 7 q' N4 ?& @+ n! n/ L9 R! { Inception V3 ( q% V, k; w/ t ?3 N """ 4 `/ l5 K' X* ?0 g. Y o model_ft = models.inception_V(pretrained = use_pretrained)" j9 a& S9 q; q- E9 b) y$ y% l/ X
set_parameter_requires_grad(model_ft, feature_extract)' N+ \$ h# h+ ^
- P0 D9 o4 P8 ~9 y6 D+ |5 ]
num_frts = model_ft.AuxLogits.fc.in_features2 a6 x; b& p4 Q
model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)2 j2 @2 q, a1 \
" N9 E( G, e2 q$ U num_frts = model_ft.fc.in_features ( ^/ o- [, i: j+ i, s model_ft.fc = nn.Linear(num_frts, num_classes)- u6 b1 p6 v) M* i8 }' s& z
input_size = 299 , H' \7 H) z; c9 ~( B ?- O , F0 B8 b8 N4 O7 ?2 K+ n/ N else:$ S2 Q" A& d3 |/ q0 x
print("Invalid model name, exiting...") $ V) R! t8 b1 Y$ R) F, J7 s! s exit() % {- R6 ]# V( `% h5 q, A/ o, J& Y4 x! I7 j. X& w
return model_ft, input_size 0 n4 @/ v: P3 ?# W% l( T" A: [* r) k* k" ?% k3 C
1 D' h8 D6 N! y0 |' p' K2 ! j: s( M: a6 D3 Q6 c8 W32 `' C5 F5 p" y6 Y1 |4 C# }
48 R! ~( n7 ~7 |1 G# o' V& j
5+ R6 i! a# c8 h
63 |1 F, j F+ {8 Q5 j. J
7 . r( ?6 e# I) h( T+ W9 c8 $ l' f% d$ i4 W, Q; n; _9 * ]* U) ^. }; T( u1 w& K2 ~, M& I5 ?10 ! X4 s' H" J* \, C11 8 o9 I) y9 R4 q6 J, |1 v) v12 6 v' ` X( X$ l; S132 G+ D( t, E M$ L
14 1 b' S! M( G' }" l' |; ?& c# W15% t# L3 @) C2 O# S
16 * @& D' I# s1 {8 R17 - |1 ~0 t* R8 [. p+ S5 q! _2 j. o182 _9 d) r: E+ s
199 ?2 ~7 b" o7 N1 Y
201 g: T: s$ u* ?2 X+ e
21 7 }% P9 ~* ?1 n+ i4 W( h' y22, `; }3 v; A3 Y+ s y: t
23 & c* p( r7 S1 {' n1 z i24 2 T* j& F5 _- S, U2 G& b25 / t" c7 |& h% }2 I26 5 i- A$ Y% s8 F4 q4 D6 @3 l27 - q# ?+ B' i* E) p3 W* {; `28: [4 O8 o; A- @9 u# c
293 ?" v7 C9 w9 U5 ]" d2 @
30 : Y/ r! z* u# F3 y1 y31 ( A/ x P# e- s1 ^; D; S32# b( f; y$ U; v
33 / ]% M% V2 p0 n' G6 N; `5 t1 t347 ]" [6 l: v. d' S
35, N0 y+ K+ z q6 r/ ]1 I
36+ G3 {; o, o, l7 \" r" I7 P* A3 ^, f
37' u2 O' V- j0 X; M; K1 C: B
38 - k. m" I9 o% l9 V& z39! R3 v9 l/ W7 T7 h' n
40 & T* h+ @) x" C2 w/ Q b3 d+ X" K41, J5 g4 E: s: K: w6 {& u& M) G( ]9 O
42 H7 W( h' S# @! c/ ^8 _% ]; s! G43* S3 @! _ `, R1 S5 B8 r9 w0 L
440 Y* j2 H2 R5 m/ a% p+ b
451 Q& @% e! f5 @! w/ e$ N
46' d4 k1 R5 }* \2 c; q+ M( @$ i
47 + ^, E% X: o6 |* [4 ?) a48 + L, m5 R: }" K0 z7 z8 z1 \; X49 9 X; L, J J; E50 ) F D# E8 L @4 D( ^518 K$ M" H7 l, C6 L1 O1 t/ O
52 2 B* e+ N0 V( ^: J K) O( \6 g53 ! _1 ?# p! \7 v% E7 X) Z54 ' g5 h- b% s& P/ `; E/ Z557 V( _5 L0 g$ l4 l* ~% M9 O/ b2 J- k; x
563 l0 \& Q1 o9 ^' T5 \5 D; i
572 a$ X0 A8 w4 s) m
58 " o! d7 J- T0 M) O/ q1 }2 ?59% c6 J/ b- O* l% V% w l
601 D+ ?7 G4 i% @( n9 P
61 8 Z$ G& P' p9 B8 t+ N62 1 o: ]9 {9 y! s+ @4 t3 b$ d# K4 M5 n# ^63( y3 X) [2 b$ a" w; x& `- y/ G
64" D# t* P3 t; h1 j& U, `& A" q
65% B6 Q( k8 p, e0 G
66) T6 l: p: n: s$ m$ ]6 t7 [
67 5 ]+ t5 m+ v; o68 0 i% z7 D1 H' t696 q) }6 s& I. E! q6 z( O2 Y
702 B, |0 ~8 ~1 x1 H
71 8 |: }; P1 w* [# S+ Z" j72 3 Z/ r5 x2 z y, q$ z73' E. A0 y _* |# p7 v
74 | o) L7 {- a8 V6 M% x6 `. b75 $ z, p; _4 f5 X' I) X$ ]+ ~# |76, O- ], w5 ^* Y5 }" Q
77 4 j4 C% g& K Z: {& v78* [; a6 e5 t. I4 F- d
79 ( \5 @6 s& T$ c/ ?7 J# A( K* l# c4 {80 % C2 E* n9 p1 v81 # C+ a* ?. {! _# l7 s, ?. |82. R( K# M' \5 H" W
83 3 M! B5 J7 f5 q" R, f( v7. 设置需要训练的参数5 U8 A! O+ ~9 I: U
# 设置模型名字、输出分类数5 X: k% A* ?' g
model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)+ z3 W" A5 K3 Y. A f- _) m* m( ]
( ?% I& [0 k* N! c* V
# GPU 计算 ) j3 G5 @( x8 l' z# zmodel_ft = model_ft.to(device)/ Q) T; |1 g# R$ M7 ?: V) ]
9 X9 Q4 g$ L# ]0 P5 j( S0 }
# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取 1 r* r. l& ]# D/ C4 c S3 Y& [filename = 'checkpoint.pth'( H# |9 o* C) [7 J3 t1 P" \8 X
- o S# D5 {5 j* @
# 是否训练所有层 J2 @0 F! p% }. R, d
params_to_update = model_ft.parameters() , W' V) l4 T7 T' |% f; z# 打印出需要训练的层 ( K5 T" T) C" D" w: N, Rprint("Params to learn:")% b$ v+ u' V6 r b$ T8 W# n
if feature_extract: ' }/ ]9 p' t2 u3 ]7 v3 h params_to_update = []/ ~: O7 H8 g5 t/ N, b% H
for name, param in model_ft.named_parameters():5 y: w* h" R; X4 {% D+ h5 v o
if param.requires_grad == True: + p! Q0 x% _$ {2 \ params_to_update.append(param)+ l5 w8 O! z. d- t* F( g9 ^1 ]6 d S3 n
print("\t", name)* h4 Y& V7 l. d+ t# e
else: 6 D" _( b$ M! R# s% i! r for name, param in model_ft.named_parameters(): + ]+ Q0 [ W- t3 Z6 v9 Q) x& L if param.requires_grad ==True: 0 p5 }0 M6 c% `' M; s print("\t", name) # Z) H* _( u8 F% Z+ O% i O% B& K7 K: K4 n% r8 g
1 ' D- F. r+ ^: W' H: r* P8 o$ f H# D2 ! s& E9 j6 M! n O- B! f5 M. v3 & O# M! _0 r8 h" T5 o+ Z. ]+ e4/ Z9 a/ \5 F4 A4 P
5 % y2 e Q4 M# X2 q Q6/ ?( q; U; R* C: x# V
7* n! g, k$ b( |; g
8: k8 q8 N x f- V/ a: {8 m3 t) X
9 % R- R/ H5 P- t10 " M! h$ t& a( C' K% {- z/ j1 y11 & v& h' [# I# }3 \; ~- D8 ~12 6 E U8 I6 f5 i2 k( `9 I9 e13 # Z0 ?4 [9 f0 N& a14: l: t- @1 R4 n$ }
15 0 _* P& f& k$ y: ]" c; W16 - ]0 I+ c) s+ T1 E& G171 B% h$ ~1 ^$ M
18 & y: |* k! w6 ?+ X) x19/ ?( u8 @% a9 g
20 n- m1 |, [, R6 N21 , Y6 O) w* O, {7 }, ~22 + I3 x5 g2 v3 Y& [23& i% F# _" x: C0 e- z) r# ^: W e
Params to learn: $ F3 B) j$ J% }5 w5 f& a" @ [ fc.0.weight2 @( U5 w6 @' P- }/ r W/ [
fc.0.bias: p4 g# K3 l: u$ T9 b
1 & S' B% q% i i$ C+ q" |& Z2 & a }6 L' T( X2 E3$ d4 [' W0 E3 T' H0 ]* {% S
7. 训练与预测4 t- {% ^* ]! J5 u; N* H K
7.1 优化器设置( C0 G, h( t9 l
# 优化器设置 , S0 I( m8 D2 M- Q: t" Foptimizer_ft = optim.Adam(params_to_update, lr = 1e-2) " [! V0 F: m% O% s# 学习率衰减策略 8 E$ A! l) V& k8 b8 ?# s3 Xscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)' \6 q# G5 v" N3 o
# 学习率每7个epoch衰减为原来的1/10: Q# h" K, J' X2 z% Z
# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算2 S3 G. Y. J, C0 L, G: o
8 P* g5 X- @! P" J0 g# Mcriterion = nn.NLLLoss()& V, B* u6 l/ }$ h# P! ]
1: f6 t, o% e) ?+ u9 T4 Q4 M
2! v7 d8 l8 a7 I; d$ r9 Z, O
3 . m! a% |. ?9 G4+ m' V! o# w: P
5 7 g5 }- g( R, _' A2 w: S! |' ]6- U0 I% M( V8 F, g/ {
70 M+ [6 y4 k Z. B: W
82 f, l& x+ z+ A S5 V5 y4 M
# 定义训练函数 7 J2 s1 V8 q4 H#is_inception:要不要用其他的网络7 p. ?& a: K# K
def train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename): 3 Z3 {7 e% S. v) U$ X: n7 _2 `: L since = time.time() Q4 ]: D0 V; i, P/ g: Q2 H #保存最好的准确率5 |5 v8 ]4 r# i5 w; S
best_acc = 0. J4 ^. C' `% K6 p' y
""" ' O1 y( d# \" ~( _# O checkpoint = torch.load(filename)& D7 `, _, g* I/ F8 Q
best_acc = checkpoint['best_acc']# c8 ` d+ U' \$ {( G1 R
model.load_state_dict(checkpoint['state_dict'])# Q, |8 q7 j2 U& C9 W9 V
optimizer.load_state_dict(checkpoint['optimizer'])8 h) R' ~/ Z# s) n3 C- s2 V0 N
model.class_to_idx = checkpoint['mapping'] # V% R( K4 j+ }+ [* [6 U """ " G! L0 N" {" A #指定用GPU还是CPU! v: ^) p# t- B: ^
model.to(device) 4 q7 \9 c% y! N( M. ^ #下面是为展示做的9 n/ C+ f8 m( {7 c
val_acc_history = [] $ ]. i4 I7 |2 U, k* u train_acc_history = [] - a7 b; t+ X' `2 f0 @: c train_losses = []/ G; t! [) @9 m$ C; [4 H' q
valid_losses = [] ( y* N7 |( r @ LRs = [optimizer.param_groups[0]['lr']] : p( z3 ~! i3 J" p# C' v# @! T #最好的一次存下来 & `7 y2 l3 w$ Z! S- j g4 d best_model_wts = copy.deepcopy(model.state_dict()) k; b2 T* ~/ c
2 H% z0 A/ O$ b- t/ \" P: A for epoch in range(num_epochs):0 O( X! N+ m0 d% I, k# y5 ^+ V
print('Epoch {}/{}'.format(epoch, num_epochs - 1)) / A( B. P7 r: S$ |9 N, W% u print('-' * 10)$ D$ E1 c" D, D
5 L" q5 c; A B W # 训练和验证 ! x# b! I3 M" q! v! }6 j1 ?& m for phase in ['train', 'valid']:6 {' F+ b2 U+ f# ?( j7 R9 y i
if phase == 'train':* r( b+ D2 G0 M
model.train() # 训练 2 W& V( T; D! P' t$ o/ p else:8 f. q4 q& i& M$ S1 o
model.eval() # 验证7 F Y, K; O( b: Q h( g8 m; L$ G
" C: {4 x4 {9 a6 v running_loss = 0.0- R n- q6 w6 b$ h& L5 Q
running_corrects = 0% v" U( `( @( [
, U8 c% G; E( {' r* e. ]
# 把数据都取个遍# e; d* t+ I9 d9 O. `5 z5 f. G
for inputs, labels in dataloaders[phase]:5 x7 t' l* Q7 ^2 W, D' ]- A
#下面是将inputs,labels传到GPU! x) N, o s0 E* w9 z2 Z: F( A
inputs = inputs.to(device): {6 p6 G' x8 A$ [+ J: d
labels = labels.to(device)" f4 N" P6 ^/ g: K* d$ ^$ {
! y; w/ A+ G$ k# `' M # 清零: x( w* ]- N( y, B
optimizer.zero_grad()- t5 W2 r5 T# n5 j
# 只有训练的时候计算和更新梯度# s3 V) B0 ^# P9 O
with torch.set_grad_enabled(phase == 'train'): + M* Q% n: t# H6 t. z; C #if这面不需要计算,可忽略 + N1 ] `6 o8 F5 S if is_inception and phase == 'train': 7 n. P# @6 V% [2 R* l/ e outputs, aux_outputs = model(inputs). q7 H" i9 H$ {& m! v
loss1 = criterion(outputs, labels)5 k2 ]7 l: k' M; ]
loss2 = criterion(aux_outputs, labels)9 @$ i5 }# m( Y! _) R$ y* e
loss = loss1 + 0.4*loss2, a4 X9 w( p+ h8 c0 d4 R
else:#resnet执行的是这里7 m) P+ \ O1 M$ w# B
outputs = model(inputs) 9 P9 ]- f+ h. {; I; U loss = criterion(outputs, labels) 2 L2 J8 i, r _, h7 t0 Y; B- w7 U$ K [* f- o
#概率最大的返回preds ( B, p7 A- m9 j: q* X" j _, preds = torch.max(outputs, 1) . ^5 C1 T7 @ g; D0 J 4 @- J1 ]7 W' R3 u2 q$ w6 c$ D # 训练阶段更新权重 . Y& r+ a, S- ]: n. r4 G if phase == 'train': {( w6 p1 Q! v U/ {
loss.backward() ' T5 B# D: H8 j8 P3 }, i! y E) A optimizer.step() , l; o' e1 @3 S7 V 2 R& k- C/ a( o # 计算损失 , ]0 _' } E( O: ^5 C2 U# d$ ~% u; _9 i running_loss += loss.item() * inputs.size(0)1 }$ _7 l! R, Q( }* E
running_corrects += torch.sum(preds == labels.data)* q) c5 B- V% u3 [+ X" ^- B" f$ `
0 C" m% k, P0 |8 @0 T- F #打印操作6 }# @1 H9 H N. r
epoch_loss = running_loss / len(dataloaders[phase].dataset) 5 z, ]9 @" h, D epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset), ^7 b [2 ^ r5 L$ k3 p
# D( ^ [' y6 ?
7 j e1 U8 K% _; n time_elapsed = time.time() - since ! B! e' Y; o. C print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60)) , S" O( \! x, ~, e print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc)) 2 ?; U t& W6 j7 y8 k* V8 q7 I, f " G1 y) c2 C/ }# Q0 ]3 m5 R, `9 p. c6 \
# 得到最好那次的模型 , I0 P. }/ e( x+ X* | if phase == 'valid' and epoch_acc > best_acc:# g6 J/ ~3 m# G
best_acc = epoch_acc & ?7 V* X4 c/ K7 a #模型保存. P' X7 p$ c- j1 ~ a8 x
best_model_wts = copy.deepcopy(model.state_dict())+ J- n- w6 T- O# l$ \4 a
state = {* A. g' i3 w. D2 ^
#tate_dict变量存放训练过程中需要学习的权重和偏执系数 1 p( z+ w( k% U& p 'state_dict': model.state_dict(), 0 B; V- V6 p& L5 J7 e- Q, ` 'best_acc': best_acc,9 N _6 Y' S) Q A
'optimizer' : optimizer.state_dict(), 9 `7 K. ^( k0 |& ?6 D } ! m1 W2 P! N# l9 o) q torch.save(state, filename) + O! s* L2 Q; n6 n if phase == 'valid':8 s- x3 {1 g* R+ B) H
val_acc_history.append(epoch_acc) # I3 ?/ h7 a& ]& Y' ] valid_losses.append(epoch_loss) . M2 M, \, t* ^5 L3 s1 D scheduler.step(epoch_loss) 9 u2 P5 D. m$ D% u if phase == 'train':5 ^" r/ j7 A1 k& A7 L/ {: W# S
train_acc_history.append(epoch_acc) 0 T- ^& v% m. k, P. U% A train_losses.append(epoch_loss) & w1 f9 n1 U4 e5 o8 W6 f; ]9 V# Z7 b 0 {; z* W9 {4 e, h print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr'])) 1 n7 ?- [4 m$ p5 @, Z+ X LRs.append(optimizer.param_groups[0]['lr'])$ E# ]: e: T: X
print()/ S- B( x; y3 `" ?
# W" A$ { M0 o" f; ^0 U time_elapsed = time.time() - since 0 Y2 k8 e8 Y9 B f print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))$ U# K& L. o: c( r
print('Best val Acc: {:4f}'.format(best_acc))0 p9 s, X0 z1 {) ?
9 o$ D7 a, V$ v; J _ # 保存训练完后用最好的一次当做模型最终的结果9 N& S1 Z c7 u3 L, Q
model.load_state_dict(best_model_wts) / ^ a) h6 d" B' G9 B) ]0 [ return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs ' S( V6 V# H; t2 H; l0 w( d; T: h7 Y7 q9 I
* b [+ Q5 e t* P- L2 X
1 & E: `" F6 q2 b; J$ F N" B2, a/ y: g P: c/ K
3( l- Q$ Q' H0 W" {: A8 Y8 |
4; p& s3 g6 a/ w$ ?- B
58 f I/ l& X5 Y! B. |
6 9 a1 g/ d" }' V C" J( @7/ M- ]6 V4 n( B" q a4 y
88 o+ R/ {5 ~7 R( e4 s8 t
9 / ]! g3 T' q; C9 ]0 a, _10; g! F; D# S: L, b
11+ E" ?5 d, z' e# s& f6 l& O8 b
121 J5 h1 O) U8 }' _# Y
13 9 _+ C& y# |8 R5 C! Y$ y! E14 ; }& [ y4 m3 ~3 @) c8 ^15+ Q, d5 h; |3 q8 D% `4 u
16 ; T. Q0 R1 E# X* E f17: E( s3 ?8 _" _- u- {+ w8 ~
18: C5 J1 p: M+ _* p3 ]/ g
192 ]0 t! |6 X! `9 M+ C- u/ g( q
20 3 B5 ?0 Z+ r' c: D2 F215 p4 x: I4 g3 r; N4 U
22 8 y) O7 h5 G8 p. m. I3 j) D4 j23' r. O* }, L: g+ c
24% D& M6 Z+ m; r6 a& k2 K, f
25 , m# O; b, T7 F Y" _# ~26 ! M- X5 V' _& V, O, A4 i27- P& g( z1 U" t1 }9 w( R
28 ) y [& Y9 f& y, `' U5 `- j3 G& f& f2 J29' L- h/ f, |8 y/ |+ A
304 {7 A/ \' T4 O5 s- u
31 . Z9 P: C, p+ Y4 S, s7 Q* W32; e2 m( O+ y4 C: e
33% @" p# g Z+ e+ i% Z8 E G8 s8 s
34! v0 ?% m9 S2 Q! s) R6 |: Z0 T5 Q
35 4 @, a: {; R; T$ j i36 7 \- o; W7 o4 f+ r1 Q377 k4 o0 v0 K+ |: c& H
388 ]$ U6 b" A8 @
39" Y3 |# e% P$ n( Q; d. K3 Y5 J
40 : `( E; S) h7 U+ l/ N ]- |7 B- b1 ?41 % \- q3 T6 H; y7 d42* h! ]8 u) o/ N! P3 u4 {
43 4 S( z; ?# \1 X4 u444 c* g8 b" q+ X* y# t U7 ~9 c
45* r( G0 a, u1 N, Q
46 " z, j# o/ i. }, ^8 p- P q" N47 & Z: N# h) x% W48 7 ]2 u- d+ W' S6 F# q+ o2 R* l" E49( Y- Z% r1 N/ H
50 1 R6 z. T3 M" t4 `1 N9 I51 , M2 \; g6 u0 g2 V$ U52 " n- A$ j& N8 ?+ I. c53 7 f4 b* ?+ _: g. |0 W) u! Q- Y: F547 `) F2 y3 H$ c' w4 Q1 T8 u
55 2 O1 U6 L, H8 C# m% ~ a56 6 {8 P5 C" M& \+ V2 q- M57 5 ` f+ e! Q# a4 y4 `3 D1 U580 i: g( Z$ c T0 t7 s+ f+ H; b; h
59& C e G* z" _7 ?' [% h' G. X# `
606 q0 c D. }! e7 i8 y, E
612 l9 v1 m; Q1 r, w0 u
622 n' ^' l! U7 r
63 U: Y# M/ w% e- G, n
64 5 q% r; w, K6 `) N65 , q6 o4 L* ]7 P. P9 a66 7 f; \# R! w" c& h3 a! h67! }$ n- z8 u# D: r/ v
68 % U) N; u, d- ]) h69* w1 a% Z3 G! s2 W' k9 b* ^
70( J& e; G0 T% \) @
71$ Y7 c' U3 o! }1 b
72! \9 {' ^$ m' A/ P& P) r
73 & Y( x3 Q( _, ]8 c74$ M" J; J( V) V0 u
75 ' P0 h! m( k- |76 ' Z! e8 F2 \, C/ b, g4 g* n9 r8 \+ q77 * N7 w2 }. U5 v" P, v4 q78 2 x' ^* u% g+ s8 j79 2 @4 i+ g- t& {+ f3 Z80 2 `) f2 j4 _4 u* Y# @81' u( q' l# N! K& `7 Z
82 * {/ ~# E! [9 |838 ^; V( B/ l; x* ?" [( w1 }: m
842 k9 t, N" V" t
85 ) V' h. \/ ^0 B1 K6 S4 p86 ' `0 b, Y& x% b' q2 J1 L2 K( D( \87: Z4 C r' X" c. s- Z
88 # G, a w1 g8 K: J6 h- j1 g89 ' k* B* J( l; g90 $ R& v; Y; f& z% L K91 ; L# @8 o2 S6 e/ z( Z# ~923 `! Y1 P) ?9 D3 `% i
93- |9 Y0 O# X, f Z6 B" V P
94* ^' C, M; A! E
95$ ~1 R: O8 @! I/ j, q
96 & k/ p5 p6 ~) j! ~( J H! i971 U- d! {7 `% N3 J- q- N
98 . ?5 t6 l& x, U' d& D' t7 U99( `$ ?( E% e6 z
100. B- {- L/ I+ p0 u, R
1010 g% n# e. G7 D0 i
102& Q* _) x8 q& [: s6 Z+ ?
103 / [4 _, L- ~4 Y$ X5 Y104 0 N3 X: h. G" w$ _' C105 ! o$ i2 x; ^ R8 i8 a2 d9 P% ~106 4 b; Q3 M: b* L7 p3 k- E107) b9 _* }. ]( m% @" [/ x( s
1085 ?* R+ f0 G7 d+ y" n
109. O( h5 L3 n! |2 X
110 " T' d1 ]9 |9 l& V/ N, l l) L111 * u2 ]1 q" C0 Y4 R/ V' M# k% V5 z112" I5 @* b" @6 K1 Z1 M
7.2 开始训练模型 & ^7 U4 I8 q2 Q7 U我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次& b: z" R: R& B5 F