QQ登录

只需要一步,快速开始

 注册地址  找回密码
查看: 2778|回复: 0
打印 上一主题 下一主题

【深度学习】 图像识别实战 102鲜花分类(flower 102)实战案例

[复制链接]
字体大小: 正常 放大
杨利霞        

5273

主题

82

听众

17万

积分

  • TA的每日心情
    开心
    2021-8-11 17:59
  • 签到天数: 17 天

    [LV.4]偶尔看看III

    网络挑战赛参赛者

    网络挑战赛参赛者

    自我介绍
    本人女,毕业于内蒙古科技大学,担任文职专业,毕业专业英语。

    群组2018美赛大象算法课程

    群组2018美赛护航培训课程

    群组2019年 数学中国站长建

    群组2019年数据分析师课程

    群组2018年大象老师国赛优

    跳转到指定楼层
    1#
    发表于 2022-9-8 10:41 |只看该作者 |正序浏览
    |招呼Ta 关注Ta
    【深度学习】 图像识别实战 102鲜花分类(flower 102)实战案例9 A& z4 W5 v4 ?2 u& ~' F6 u
    " q4 f& p( w9 L! b
    文章目录
    - c% j; d7 s8 ^8 g4 d8 H0 K+ b卷积网络实战 对花进行分类4 L2 h3 s& o' E# V& t
    数据预处理部分
    6 f: E3 _; f0 u) R: j# u网络模块设置( _) A- ^4 ?7 y6 W) W
    网络模型的保存与测试
    1 s! l4 s9 r+ @4 \数据下载:$ T) r8 U/ l' R  _+ I: `" |- P
    1. 导入工具包
    5 K0 R: ~+ v$ o$ E9 S! H: _4 e2. 数据预处理与操作
    . d$ N* E' Z# Z! ~4 d* l3. 制作好数据源% [! ]; g4 `7 f/ d+ U) r
    读取标签对应的实际名字9 F9 G6 F# J1 u, V! Y% W
    4.展示一下数据
    ' s- \5 X0 w) @5. 加载models提供的模型,并直接用训练好的权重做初始化参数; `1 Q6 @; F6 y* d0 x2 ?
    6.初始化模型架构
    ; K4 a( p+ z- F: g5 ^: V3 _7. 设置需要训练的参数8 W# j6 {4 Q2 l" U) N
    7. 训练与预测* q! }1 D3 m: f0 {# Q
    7.1 优化器设置
    % M" ?" D, C% l2 l) E: n7.2 开始训练模型+ J/ O' n! z; m: n* U: F0 a6 R1 T
    7.3 训练所有层2 W: F) X( z; G: n  O' O+ I& S
    开始训练
    6 _; ]& Q$ O% Q8. 加载已经训练的模型
    7 j! m" Z# z: l' w* j, S/ z5 R9. 推理
    1 D+ n+ n0 H# B- v* e4 f( i2 H$ \' J9.1 计算得到最大概率# u/ f9 P" x/ J
    9.2 展示预测结果
    ; k0 z/ l" k' b: L: g! x写在最后
    + E! F: h# t0 Z卷积网络实战 对花进行分类: j+ v) T  V% b1 X7 Z# R
    本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能
    0 s, r; y) _9 G4 @- Y+ J. W; q+ p4 _: w6 q. `
    在文件夹中有102种花,我们主要要对这些花进行分类任务
    ) Q, E3 s( |3 m' U9 }: S% E文件夹结构' l% I5 p( N/ h$ q1 X$ a
    / P% N; V7 m+ s0 X! {4 ^
    flower_data& m4 @( r9 Z/ z7 C* q% H* z/ r

    + V1 ^" K0 K" V" b2 ltrain
    / K$ H" I9 ?- T6 g
    9 M. Z% f6 L- u. `: d% z5 ~' {& N1(类别)
    1 H6 u/ E- u5 A* Y2 }; t2
      Q5 [# r- ~2 }6 I5 J+ C7 Dxxx.png / xxx.jpg
    / E' J* i4 p. T" g$ d, u. q' ^valid
    & i) j7 w0 s9 h2 T/ J
    7 E9 \- r! f5 U主要分为以下几个大模块( q3 R' H, G% ]. W6 S  A

    8 d# y9 k# w# F0 a, E7 y数据预处理部分+ k/ q2 y" f# @, T9 m
    数据增强
    ; g$ d/ s$ L8 h( z数据预处理! d6 k0 o% `* q' s! [
    网络模块设置# y4 m# r& m& S4 t
    加载预训练模型,直接调用torchVision的经典网络架构
    # b, w) c) X( ]因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务
    ! o' d; L' i( a' O, g8 O网络模型的保存与测试
    4 _: A. {3 g# N: t( W模型保存可以带有选择性# o1 o' |& `/ T8 Q. o% k5 L
    数据下载:' V# r# F& V% q. W$ ]$ p
    https://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset/ s7 r) {% u/ T1 ~* ~- e
    + }1 e: @+ z# H9 x6 x/ ]5 e
    改一下文件名,然后将它放到同一根目录就可以了" l* }' Z" [. A5 M/ i4 D
    . F  Q1 Z+ U9 {; r% G
    下面是我的数据根目录4 z" J8 |) U) E. v) |

    3 U' A' l" f7 T% x$ t' i4 ]% W
    8 n. P) ~. _: Y- S# B1. 导入工具包! `7 C3 v# |' u' ?& o. M* f
    import os
    ) n8 J) k+ P8 a! I; N% Uimport matplotlib.pyplot as plt
    : B+ u4 [( l* s9 T6 s4 F4 f7 {# 内嵌入绘图简去show的句柄; b4 ], r7 D  j! X9 N" q6 M
    %matplotlib inline 3 J) f7 _6 J5 w: u4 O! ]
    import numpy as np& y. S2 @) t) H( Z2 I
    import torch2 R2 E+ f# g" j( W7 K
    from torch import nn
    ' u. }8 c7 F7 i* A- m8 ^
    8 I+ I. t9 I$ j' O5 h$ N" rimport torch.optim as optim
    . {0 m1 C7 v1 X2 m9 K2 |7 H! limport torchvision) k1 O! v& p/ g
    from torchvision import transforms, models, datasets
    + k4 F+ Z! A# @8 V1 g! A
    ) G' ]/ |2 Z5 S* i3 b% W$ simport imageio) H- ?. J2 K, F4 m& D
    import time
    ' g3 Y7 y& W' |$ B# J, P# F& Qimport warnings
    * _2 Z3 r8 c9 z$ J4 D% a9 Timport random
      ?# q, {# q; o1 T% x  Jimport sys) f0 ]. P& c; h  U; i; p/ ^7 S
    import copy
    4 S' a9 g( \& z* Z  B2 S) g% Limport json
    5 _7 j* w+ b5 Q4 Z# Ofrom PIL import Image
    0 L; R) ~0 W! F& n. J( X5 X2 i+ p& x0 G4 Q8 q) ^

    ' O: ^1 E8 c7 {0 t1
    5 n( s% K! t, E' G' ]8 a# v2 l- i2/ P7 C& O, z  O" \+ b% j4 F( P
    3
    $ G( t! b9 N+ Q! e+ v7 k4  M$ X# n" G: J
    58 z0 V3 M. W) H2 ^2 c$ [' j2 ~7 {
    6
    3 {: {+ [/ u2 H8 A7$ z( L: ]* J: }9 E2 q
    8' ^- m3 _) K+ F7 v8 P9 L6 M+ w: d
    9
    / z( S5 ]4 y9 g( f2 n7 T6 y) ^10! {% b+ l- R5 M& Q% @6 O5 g5 i
    112 Z1 P* r$ u9 `$ x/ G( h6 V
    12
    1 R& M! \; u; d# X2 Z! O6 o138 H# c1 O2 k  H  R0 ?- I* _
    14" E! g0 _0 {0 Q: w) I- R! Q
    15
    $ G% x3 r8 ]4 b( `. p& L8 @6 g16
    / ?/ S( r  M+ t7 V2 i6 V+ S17( ~3 v, g! K: ?4 E5 ^' p2 f
    18( y4 w% o9 c, Y4 x6 a2 w2 V
    196 w- l3 l1 O/ @0 E( J
    20
    , W! ?; X3 b4 E% g5 Q216 y9 c  L; O$ C$ J; k. q' i; d
    2. 数据预处理与操作
    0 y/ e2 H9 a' T0 h7 D: v#路径设置2 F8 m! Q* j: L% m4 w% A
    data_dir = './flower_data/' # 当前文件夹下的flowerdata目录
      y$ b% t+ H3 R6 `3 {train_dir = data_dir + '/train'
    0 f2 L- U4 R; |1 a) H: ?( Xvalid_dir = data_dir + '/valid'% z- O5 ^1 W8 L3 i8 c4 e
    1$ ?1 F& K$ W! ^) I% w2 `8 t
    2
    ! a  {8 Y7 j6 e6 l" F3
    0 t- l6 _7 e8 z- ~5 f; t4+ \; o1 e$ T! C, k: C0 v* _- N2 k
    python目录点杠的组合与区别
    ! C3 w/ R0 y& ~! f注: 里面注明了点杠和斜杠的操作; h# h5 d, t$ E
    ( P3 _! p: T8 g5 Z; |/ p& ~9 ]
    3. 制作好数据源
    % @6 P1 T' e5 J' mdata_transforms中制定了所有图像预处理的操作
    ! D3 Q9 z4 r9 Q1 e( HImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片
    % N7 Z& Y$ y- L; \3 l: ddata_transforms = {3 L9 l$ ]  ~3 J* B
        # 分成两部分,一部分是训练9 ~* X- _6 Z- P+ _' t; C( P
        'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间
    9 W0 Q% y  w# ^" U4 m& M* S, p                                 transforms.CenterCrop(224), # 从中心处开始裁剪
    + H' ?# P: V# L                                 # 以某个随机的概率决定是否翻转 55开$ C4 I, B9 m- j+ A$ I7 f
                                     transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转2 E! K5 Y9 N8 B. f" ^% j5 ]! Q: l3 T
                                     transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转  e5 k, F7 z9 z9 N6 z" H2 G
                                     # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相: S: B" u3 V, a+ _/ N; T
                                     transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),. r7 G8 Q- W# n
                                     transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB4 i" j  j- o$ w& x
                                     # 灰度图转换以后也是三个通道,但是只是RGB是一样的; \$ |' [9 ?" V2 \! Y
                                     transforms.ToTensor(),
    % Z. P! h4 v4 G0 [; o; a/ S                                 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差, t' V5 W8 |/ g* e. q
                                    ]),+ ~* J, w) o, s/ G6 d
        # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
    9 C2 H; Z1 f6 S& @" I- @    'valid': transforms.Compose([transforms.Resize(256),1 J5 L# g( J: {: w2 I) N
                                     transforms.CenterCrop(224),
    2 |6 y+ T/ S( N$ d( ]; t1 K+ d/ u                                 transforms.ToTensor(),1 k9 p+ v6 n) l# s2 z/ A; R
                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同+ `) ~& Y) N, |" Z2 P1 }6 g) s6 T; O
                                    ]),9 {. o1 K! c- }% ]( o7 Z& G0 X8 H
    }
    9 k( B  L2 W2 b+ X0 w+ K/ I5 f( j5 P& |/ c, T% [  j6 D" t
    1
    . S1 v. p3 U! d# h2
    ) _/ l5 v, i5 m% g3/ X" B0 f$ t3 Q# z- ?
    4
    * }/ y: r- b. S2 `# Y; c/ q7 c5
    , H  n: G; J  `2 |) {6
    * H7 B+ y4 g3 I. s- V76 b2 I, V$ f+ U; d/ y1 G; J
    8
    . W. c2 p1 M2 o5 }3 C9
    # ?! I8 ^6 t2 q10
    ! v/ c$ S& K" \11
    * b% \: s+ i: Z& K128 z% L* P2 M9 o5 j% o5 h+ M
    13
    # R: l" s- h; p: l3 E4 N! p# o. h14
    9 [7 F3 J; S5 y( x. @150 [; \0 N, j3 Z* w" k
    169 K% U& F( `0 {9 a& L( S  Y
    17& Y% w7 |0 F5 I# t( P. c
    18/ H, Y1 b4 f: {  o8 h# Z* R
    19$ v) H  a& b, @% G1 P9 w
    20
    ) _2 {& ^. Z$ I. d9 U7 X) ?" d( P1 y213 x  h# D5 [1 E/ h+ g
    batch_size = 8
    3 b4 V) I$ _/ Z" P! ]3 Uimage_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}* F: X! O) U9 P8 ?6 `1 s$ m* L# G
    dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}% V, a9 J* h' ~3 T8 E% m$ D& y
    dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']} * j8 g; O+ T5 b
    class_names = image_datasets['train'].classes6 q$ g7 q( B9 W4 e# A3 M' O- s

    & V( A& U- s  Q$ E#查看数据集合
    0 {1 w* w0 D! e8 g7 R9 a% Zimage_datasets
    8 v3 V4 r  {( `! [7 Z
    0 P1 T1 o8 |9 R0 [+ u1
    * k1 B& s' f! Z. i" w1 j27 m6 {9 Y, k2 k8 V
    3
    ( H$ x8 {4 }. ]7 S, A2 ]( |4
    : }; u$ Z. G# P  a* C' e5
    ) t# l% d' D& i9 `- x6$ Q2 U, I. }1 W* q
    7! I4 J+ T4 o" \
    8% c3 M7 y6 T# C% k: c  d3 a
    9  ~7 R9 R" d' C+ [: t- W$ q
    {'train': Dataset ImageFolder; U  H5 v9 c3 A; N% G( |
         Number of datapoints: 6552: n/ h; H  B3 x% W6 T& p, I5 q
         Root location: ./flower_data/train
    8 x. M4 Z, a' y7 z     StandardTransform
    : j" y' p/ F% `; ]3 B Transform: Compose(
    " F# R" W% \" w                RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)
    : q2 n; \7 J  ?' O                CenterCrop(size=(224, 224)): c, X* N  ^" D9 w
                    RandomHorizontalFlip(p=0.5)
    & m3 E; G6 A* m' g, h/ R. D9 D                RandomVerticalFlip(p=0.5)) W/ \& @9 W5 J$ [0 A" p: Y
                    ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
    - L+ M; S# z2 \. Z  e( i* Z                RandomGrayscale(p=0.025)1 Q, ^# u2 f! _
                    ToTensor()  K; M6 z* f2 T
                    Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])- Z$ C; ?  ^0 v" H' X; R" \
                ),
    4 ?$ t( U# r  D 'valid': Dataset ImageFolder" L2 W  O6 V8 f0 [" Q% K
         Number of datapoints: 818; \: y; T, l- L8 a
         Root location: ./flower_data/valid) L. X* U* h& t$ Y
         StandardTransform
    ! G0 c/ _2 u# h5 p' P4 ` Transform: Compose(. a0 \5 \  F4 d  @6 s8 D7 s
                    Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)
    $ I) {7 a* R% m7 v) v' F                CenterCrop(size=(224, 224))
    ! _' u" j3 s* `2 }- T) u9 h                ToTensor()
    # r! S2 }6 j, t# F: n                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])0 g: p3 v) ^: k3 Q/ ]2 g
                )}
    / ]2 y0 D& f- U: P7 c! I
    * {& e4 V! i8 A3 D1
    & q! M- P) e- c2
    $ i: n2 r. ]! B/ @3; n8 a8 C, {/ ]8 g3 H  f
    4
    ( \8 X4 [2 H: }- K  P) X4 M5
    1 \# ?' w% ]- d) z/ @: b* P/ {6% L% i( n5 V+ }, C
    7. r! }0 e  r" R4 J8 N# ~
    8* V% n" G# O; k2 @9 A
    9- w- c: h8 Y# o5 o; w# n# z* h
    10
    6 s! l! n' m. r112 e! [  p" F  R, B# e
    120 N% _; }7 j. h1 D2 |# ^6 f! r4 h
    13
    ; P- F8 i7 t) t142 R6 y$ u% h7 J7 T( ^
    15" |+ F! A0 n. R- T; c7 |
    16
    $ ?9 ^- r3 x9 o) a/ m9 z178 ~; s. A+ s6 ]" n3 n- Q) Q" Z
    18
    % |0 n# _5 a  A19/ h' j, a& z+ A( g7 Y9 M: n
    20
    3 W6 g+ g& _  f21& D. q4 B' }' N, ?) y& @& S! k0 W- R5 ^
    22
    $ |6 X0 P6 [5 Y0 s7 x23$ o" o7 w6 i3 T/ n0 F3 q% ^; ~
    24
    2 O& h. ]) e# W. f  K# N0 L# 验证一下数据是否已经被处理完毕3 p7 z# _2 j$ i" p
    dataloaders
    6 ]5 b% W. z9 l- ?( R' ]1
    0 u+ S' H( C6 g; H- s; W9 ^4 n8 p26 E3 z- @6 c" v
    {'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,
    5 ?0 ]6 P, ?1 U, `6 m$ d 'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}
    1 `: \+ _* ~! ]( w, Y$ a17 N+ c& @8 b' j1 M' X
    2) {! V6 _3 H* y& o& }: n
    dataset_sizes; R8 Z1 {- m9 R- Y1 V
    1
    4 [5 P* w# V+ s/ w+ X{'train': 6552, 'valid': 818}
    7 ~9 N5 D5 Q) v% r$ ~4 t. }1
    : @% Q; q# `3 J; |5 H& c读取标签对应的实际名字
    + V( U( G" h7 \. p& l使用同一目录下的json文件,反向映射出花对应的名字
    6 [3 o: X8 q1 F- E0 S3 f4 L0 Y, K/ m  G
    with open('./flower_data/cat_to_name.json', 'r') as f:' `# P3 d' J; v/ P) q+ V7 S9 n
        cat_to_name = json.load(f)4 w- `  r% D7 _. |' `6 s: y4 W
    1  d! G( T) t8 ]( N6 n6 p/ P( I
    2
      ]2 ^1 p. s- e6 Ycat_to_name2 N% Q( g6 e8 B0 L
    1
    : {* {& {6 b3 A4 o; h' t2 V1 a* `( {{'21': 'fire lily',- u" ^5 q9 |. U0 Y  e
    '3': 'canterbury bells',' W! n, o  `; V7 f- S2 K, O
    '45': 'bolero deep blue',! E+ E( Q" H2 g
    '1': 'pink primrose',0 g- L2 T8 c, k6 G- n
    '34': 'mexican aster',
    ' f7 |) N. q* g3 u6 _ '27': 'prince of wales feathers'," K$ g% Y/ W5 y2 ~! b9 y
    '7': 'moon orchid',2 g$ u" s! k, S# Q% f( l# J
    '16': 'globe-flower',
    # J( W  f' j) g# _ '25': 'grape hyacinth',4 [1 e% ^0 t) M% K6 ^/ T* l# Q6 r4 N9 J
    '26': 'corn poppy',
    ) N( B; a$ Z" L/ Q '79': 'toad lily',
    $ ]6 d6 _6 A; \' W( N1 ~ '39': 'siam tulip',) v" m- R/ [# F; d) O3 z. F4 L8 e: f
    '24': 'red ginger',
    $ R! I0 _0 Y) H% R4 S7 R '67': 'spring crocus',
    * {6 E) r- j9 h '35': 'alpine sea holly',
    % W1 ^& W: W- J! d '32': 'garden phlox',: S* ]! d% g6 X1 u2 c. ?6 ?, r
    '10': 'globe thistle',
    - k$ ~( P& O5 g# f '6': 'tiger lily',
    ; c9 _$ ^0 U7 Z/ z1 S4 ? '93': 'ball moss',
    - o! ]9 j; c" ~ '33': 'love in the mist',
    / b7 n* w# v( G '9': 'monkshood',
    . A3 r$ c8 \7 d. K1 c( n0 U '102': 'blackberry lily',
      @, p4 x+ ?/ {- a' f( k6 m, O: x '14': 'spear thistle',
    1 V. S* Z. T  f6 c '19': 'balloon flower',
    * Z. [* }: W! ]5 s4 H. @$ a '100': 'blanket flower',
    * n% P3 n7 e; O. U6 T '13': 'king protea',
    0 _: P" t+ n1 f  J" C '49': 'oxeye daisy',% \5 j2 T3 b$ ^9 c/ t) J
    '15': 'yellow iris',
    1 p& ^  D( V/ O; G' r '61': 'cautleya spicata',
    " f" [9 {2 D1 t7 R6 f; Z0 Q '31': 'carnation',, S1 F; B: G, R
    '64': 'silverbush',
    3 O1 ]1 w. y* }: X( u1 k, Q8 ?  c( L '68': 'bearded iris',: i! H6 ?$ J' v* a* W- J2 T! {! y5 d
    '63': 'black-eyed susan',
    ! N$ m, t5 m/ E0 l/ k# d9 R2 | '69': 'windflower',  |6 k( M+ `6 y! B4 j
    '62': 'japanese anemone',
    / Z& @/ V. A- k- G: O '20': 'giant white arum lily',
      K2 R- D& v! P( [ '38': 'great masterwort',9 J, L, }  f; n+ f0 \  m
    '4': 'sweet pea',
    5 [! Z. }7 d1 W  ~ '86': 'tree mallow',
    ( m9 ^. l& Q- P. K '101': 'trumpet creeper',! q+ i7 a. }$ ]
    '42': 'daffodil',
    ) r9 \8 W1 L$ \1 o '22': 'pincushion flower',+ Q) J* x$ {1 c( G% F) H8 e" w
    '2': 'hard-leaved pocket orchid',' }: h, a5 |' X% t; C
    '54': 'sunflower',
    # Z8 H6 @0 g8 H0 ^1 O7 [ '66': 'osteospermum',
    5 m% V4 G- S/ L '70': 'tree poppy',8 a% j& N* L. g8 F9 G
    '85': 'desert-rose',
    ' A4 r) p5 y0 ?& y& U '99': 'bromelia',
    0 M, a5 r; ^2 q' i5 G! _" u '87': 'magnolia',
    7 H. L5 Z( A0 @' a '5': 'english marigold',
      t  Y3 N( o, _; y( j '92': 'bee balm',
    # V& A$ k( R  c' u" V% o1 J '28': 'stemless gentian',+ u) d; w# c6 L2 A$ ^5 n
    '97': 'mallow',3 g- ^* d1 h9 |9 @7 e3 S
    '57': 'gaura',
    0 |$ c* c) j$ u& J+ y, z, Y '40': 'lenten rose',
    8 H' s/ x/ T9 C* {  d4 k: f '47': 'marigold',: s, A! E; T1 E& b/ ~4 `
    '59': 'orange dahlia',
    8 O5 W+ D6 @0 F" R& [ '48': 'buttercup',8 U# u$ H* @9 e4 R6 Y# n# J
    '55': 'pelargonium',7 A, `6 k1 W  Y2 s5 ]2 t! }# }
    '36': 'ruby-lipped cattleya',
    7 {% @: x/ e: j9 x '91': 'hippeastrum',
    : a/ e0 K% W  ~2 W  Q$ s% g; z) X/ D '29': 'artichoke',# Z, h4 f" b# h2 |* @. w3 d% d
    '71': 'gazania',
    " q+ ?& Y# r: F) c1 |* Z" q '90': 'canna lily',
    5 h8 q! K; R7 e, @: p! W. J '18': 'peruvian lily',; d4 s1 w- a1 O: k7 u: ]- @
    '98': 'mexican petunia',% V& R+ k& f: S# l9 {, E7 h
    '8': 'bird of paradise',
    ; S8 V% i4 F8 m) b; A '30': 'sweet william',
    6 V/ I  \  G. p- [" }2 d5 d8 G '17': 'purple coneflower',5 P7 Y% F, L  G+ K8 F  p
    '52': 'wild pansy',& u8 t* q" r, W% ^+ j) n' M4 D
    '84': 'columbine',4 D" T( O+ D) C  M1 g- b9 S* ?
    '12': "colt's foot",2 M" M4 }8 m- ~9 L
    '11': 'snapdragon',2 V0 _9 U6 U: Q
    '96': 'camellia',
    ! _/ w" L3 M! s. X5 _6 ~5 T6 V '23': 'fritillary',5 ~) r  K2 Y  x" I' s' G$ J0 N
    '50': 'common dandelion',& f, ?; b) X9 q# h: l
    '44': 'poinsettia',/ @. o5 n  O' u
    '53': 'primula',
    % q2 |/ T0 h" z4 j" y '72': 'azalea',; ]7 E+ t; R+ |
    '65': 'californian poppy',
    . O3 k+ q# a. |  q! J; Q' X '80': 'anthurium',
    / }/ c( C: Z: v9 E '76': 'morning glory',
    % `( N1 c$ _9 z8 s$ \2 J' @7 j '37': 'cape flower',
    , M$ m5 [) J& D, o '56': 'bishop of llandaff',' E8 P. W0 X8 L+ w
    '60': 'pink-yellow dahlia'," Y; ~0 X: y" W+ \) ^
    '82': 'clematis',' j2 o% u1 d4 c) }5 |+ Y  B
    '58': 'geranium',
    ) s" E) G  V" @# b& s8 f) } '75': 'thorn apple',/ I% g3 k, k1 l, u$ S+ }+ O$ R
    '41': 'barbeton daisy',
    0 W1 u2 f: E0 T, r. e, h '95': 'bougainvillea',
    : S5 c6 |9 |4 I7 _$ a '43': 'sword lily',; i: C9 H& E7 U& Y5 r+ ?9 Z
    '83': 'hibiscus',
    1 K) q' d8 G& q, l) J5 T2 c- a '78': 'lotus lotus',2 r$ \" m  g/ L+ Q
    '88': 'cyclamen',; S! J! Z, P. a) ]5 r- C8 Y2 X4 Z
    '94': 'foxglove',/ E* m% y  |0 e7 o9 o
    '81': 'frangipani',! N) x' ?7 G# x- O. Q/ ]( d
    '74': 'rose',( Z4 R6 [; [1 s" @5 |
    '89': 'watercress',7 D/ M; G2 N" t4 Q& y7 y$ J: B
    '73': 'water lily',
    + Z+ E! c" x! E6 n3 {. m '46': 'wallflower',: C  N% E& U# O/ }7 l
    '77': 'passion flower',
    9 t$ m( n5 A% E( E '51': 'petunia'}# r( K$ A  G2 S8 R$ g# [) ]
    + W+ A# n3 h4 ~  K
    15 Q* I) N1 C8 r9 p+ i
    23 X% F" J* L* U! y1 i5 I. |, [
    3* K' R" P5 [4 ?. k* E
    4
    " r7 T9 o% E: T8 i" X5# I% }- J; f* D8 @1 z* z$ f; A3 b
    6* [. e6 Z& o0 O9 M$ Y
    7
    - t8 K4 W- [: u% l$ m8- g" l' [( d4 l  X. _) }7 o
    99 a' i5 ~" n: E6 f6 b  @& X
    102 A+ \% c5 R, k* G9 |, R% e! a
    11& K7 q$ P: B+ X$ w
    122 g; W* d6 K2 g8 x
    13% a6 Z$ D3 a6 P5 R  B
    145 }4 u+ W2 @1 b7 g
    15
    / e+ V/ }( e0 }16) x! L' `6 h0 A2 t/ |
    175 G/ g! C* B! A" g' Z1 Z* \& o4 D
    18, Q  H% l, t( T9 ]1 W, b
    191 x5 f4 B' N7 v7 N( s" t
    203 c5 d% O  p4 l+ c+ _2 H
    21% a# W1 _, a* a3 W! K: x6 S
    22* v2 x8 @( |' C* k8 {5 q+ D  X/ U) C
    23& ~& B9 ?4 P# p- ?
    24
    ! h& N% ^8 _+ H) H, e25% i, j2 r/ L8 _% S8 Y6 U  g
    26* b6 x% |8 o2 D( t9 P, Y
    279 E; E% m; i* X. C, `1 }* U/ s
    28: @( Q; Y" A; T. }; n1 L! u
    29
      H) e( [& d0 e+ V8 _5 p3 h306 Z9 ?* |2 }. r, y' d7 \1 J0 D. e
    31' t3 s9 _, p1 K5 m# q) v& ^7 F* g
    32
    ( z0 C4 [' _" U) [4 k33/ s' d" B3 c0 ~0 p
    34- w. n1 v2 X; [2 ~/ \9 h/ H4 ^
    350 l0 t- R3 @4 p# Z
    36' V- ~: A- ?) K$ `: f+ Z4 C
    375 z2 J5 Q" N; u0 U4 u$ {- U' P
    38
    & t9 ^( R9 q3 m* O39
    ' ~  d; M& x/ R1 I40
    ) }3 R( G' w6 _2 J5 @" k' R: U41
    . U2 `& d2 t- `4 r0 a1 c420 L( v( w( ?% Y) G; W
    43# p# }2 n0 o7 i4 N! W% j& y6 l9 k0 Q' x
    44" g$ x6 W, D! O" t7 E) V6 w$ a
    45+ N" g8 D3 B9 H+ s6 U" f  {6 ~+ ]/ Q0 @
    468 D  B- H( ^: _& Q
    47( C* t' T! Z- a
    48
    / ?9 ]9 }5 O! o: g49
    + n0 a3 j8 V0 D$ A. }; q# I1 b" ~3 j50
    9 V3 L$ u. \6 D, M/ N. Y& Y% ?" Y51
    . z% Y$ x! G8 Y7 L/ E52! B) C/ u9 p$ `* K
    53# j- A' [9 p" b  ]# k4 Y! j& ~: {
    548 @9 _" p- ]  _- n
    555 @2 {) k( ?& l* S
    56' `- ~$ d+ y! d
    57
    6 K* t' F7 G8 P* \& i! Z% Q: V( ?58  |$ a0 L2 }, g: e0 V3 s! g2 l* g
    59! _0 ?4 E/ [. Z% }* S
    60/ T  ]+ b3 F- u
    610 Y( v/ u/ A. M1 }9 H
    62
    / C& J2 s8 f9 ?2 _- q$ r4 \63
    * x. d6 J' O1 C7 N64
    3 u6 A/ Z& e# H7 \2 Z1 G8 E65
    8 U0 ?# _3 X  L  G: H  J. m66
    + U. y- T0 |: B+ l8 \. r67+ t* y  J2 `( g+ m
    680 c. k; a5 `' J( j& A1 C* s( e! h
    69
    7 |; m7 h6 z6 O& b70
    2 g" b9 O) `! ?4 h; T71' ?- b5 A* L* {2 W$ B* T6 K
    721 K4 J/ b1 \6 B' g9 B
    73# G5 y# e$ d9 D- D& s% u
    74
    + |/ A; Y& k- ]; {# p+ C75; m: d0 h# ?  S' ]2 D2 [0 W
    764 R, l3 }$ Y+ m+ @. ~! p
    77# O2 L% a8 \# r5 i9 a  `
    78
    4 Z  O4 `, n' }; }* f79
    * f7 K2 w- c0 u6 n  q' i) K80
    " _- b' w7 I7 T81
    * K2 G( Z7 M% d  ?( P! L82- f3 D) t9 x, J- O
    83
    ) ]" k, T6 |: ?) k! i3 V, x. L# o84+ n3 a. }( A# M! b. C* U" @
    85
      ?. D$ o7 M8 y0 p- p8 W$ n" t9 Y86
    " y0 i0 N1 m4 {4 w8 j87
      ^! r; w) K# V5 z( p88
    9 s  b5 z  ]  t. C0 A9 G$ o89
    # k; B0 O% l1 W5 I( d2 ^90
    7 Z9 p( f% X& D* U2 G5 m4 w3 ^91, x$ J8 u! v+ b
    92
    - G" e  h) w) i" a' g93
    3 e8 U1 K- ^; G; a; r94! V, X3 ^2 b6 u' P# Y$ b  G0 I6 r( E
    95& B: b8 B+ ^! L. Z& J  T$ Y# X
    960 n6 ^2 [+ U/ N) s) K0 z
    975 ^$ J, a, i4 O% d) m7 G6 r2 ^9 W
    98
    1 V2 j- U+ W% b, ?992 Q1 _3 Z! B! z) `. n: u, q
    1009 F' r4 v, r# o% N$ o: ?& f& w+ X7 B
    101
    / s) Z* Z3 }( Z- M) L102; P& |) P7 ]% {$ K& c% J8 p
    4.展示一下数据
      U) d& T' P1 K' q( e6 Cdef im_convert(tensor):1 v1 W+ [& z; b8 ]' M( u
        """数据展示"""
    . ?: {. b, k. V6 T+ {, B    image = tensor.to("cpu").clone().detach()
    4 B' v2 G5 a% A( I- \    image = image.numpy().squeeze()2 J( m% D, J" ?% O5 J
        # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图; _' [8 R  @! Z9 q8 a2 F! @# Q8 O1 y
        # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)6 k3 ?! o3 Q( O6 k2 F* \, N
        image = image.transpose(1, 2, 0)
    2 ?" @  k4 ^2 I* w! |    # 反正则化(反标准化). g( k( i* k8 {5 }4 `1 u% Y" b
        image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406)). e' N. ?3 w: C  @) n) B

    ( h9 C7 U* r1 i% N2 I; I    # 将图像中小于0 的都换成0,大于的都变成1
    / A1 d. X7 _5 {* y    image = image.clip(0, 1)" M0 h7 Z1 ~. Y& R+ W1 e4 s
    5 x/ M4 T4 b0 O7 R; m3 u
        return image6 b' i/ O( M9 l. e9 L
    1
    1 L& j/ {+ c/ H' F4 R2( E8 J* A4 x8 X
    3) o5 H2 ]' d, J
    4' ~% h- |, F2 g. D$ i
    5* [4 B# G3 J$ L6 r( R
    6
    ' u( X# {2 [& R7 K* e% \; f7( g3 v9 K9 F7 ]5 c5 r
    84 j7 _+ ^/ E" t8 B( T: |$ }
    9" S! x5 k0 @6 h
    10
    ) D7 h; J  g* G. u11
    5 D) D- M( d& w7 G* Z7 V6 c12
    / \/ K1 w' C6 O" c* Z! J8 a2 C# S13- {/ Z0 @* Q; E9 t
    14
    0 `7 V' Q0 f* `! L# 使用上面定义好的类进行画图7 q# p+ X3 X8 F. z
    fig = plt.figure(figsize = (20, 12))
    & H6 N( u. j' @5 F) Jcolumns = 4
    / d; s! a, U8 ?) h, W- T4 Nrows = 2) j! h2 S! \* v; f' h. z

    $ U" @2 b# q5 {3 x! M3 c0 [5 L# iter迭代器$ z3 v1 h* \! \! h3 B. @9 O3 B
    # 随便找一个Batch数据进行展示  e& K5 w/ V' \; k5 [
    dataiter = iter(dataloaders['valid'])
    - v8 f1 u/ }3 B# I/ a: b& L. _inputs, classes = dataiter.next()
    - P3 q, U  T) T& z$ }0 o, M% Q3 d5 h
    for idx in range(columns * rows):& Z" _# I  c3 Y0 {, T. E/ }+ H* H
        ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])% h) g1 e# @. [7 f
        # 利用json文件将其对应花的类型打印在图片中
    ; y& d& g5 {; |4 P    ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])8 I- F  p0 K% u/ ?
        plt.imshow(im_convert(inputs[idx]))3 [( Z! X' W" i( a: m, B) ?  s
    plt.show()! z6 }5 b; S( \- x& a0 C

    + Z# n+ d+ O# f1
    ( z# d9 F' i0 f, n3 b0 r2
    ' \- p( t( k8 f8 x( \33 f7 X6 \+ X" \% i+ E5 C/ e9 l
    4; H: F3 [3 p* k
    5
    2 y* @# R6 w" F8 A" g; Y$ G6 q68 Q3 n1 |2 j+ l
    7" N0 r+ N, F8 z# v/ y
    8
    7 ?! I5 f* [4 S$ Y/ Y% {, x( t, M. z' M9 `9
    ) f4 E; E% ?; Q& T* D10
    " ]+ `3 m& h* d6 C# T" P; @! E11' q" c3 x/ f! d- J& S! H
    12" ^# d; L0 H( N9 q# b9 V2 ?
    13; E  S; c* ?! ]
    145 E. K4 f  G& H! @, T. @; P
    15
    2 z% M' i# u0 s16
    / N( @2 p. J" P5 g( i7 ]2 @2 m) z) K/ ]7 [, p( P7 X  l  }
    - B' Z0 C1 l2 y+ l
    5. 加载models提供的模型,并直接用训练好的权重做初始化参数* g0 F4 D4 v- H& o3 T: M$ [1 m
    model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
    4 ~: j$ p3 M! H4 b& G# 主要的图像识别用resnet来做
    & h3 [% t! K/ q5 b# 是否用人家训练好的特征
    . q: _/ s! N2 o1 N- _# |) Ofeature_extract = True( U6 ]# j5 Q( H; C
    15 i* q+ E) O% N) A$ c2 p& T
    23 M/ y1 ?" {  f- I$ I: o1 w
    3, K' U: v# {0 v
    4$ }9 ~) b& S: ]2 @% @( t; O/ ^
    # 是否用GPU进行训练
    3 a, b8 r  m8 p+ n& C+ ^train_on_gpu = torch.cuda.is_available()2 ^- j" c2 y' A6 R* [6 d
    4 Z6 h; O. V* [) D) Y8 l0 R' s
    if not train_on_gpu:& x, k+ k6 x0 [9 {% l( i, g2 h$ B
        print('CUDA is not available.   Training on CPU ...')2 d$ \2 N; l* h5 C( N2 c
    else:8 v# k4 q! T' r0 q7 |$ ]% n( Q) Z
        print('CUDA is available! Training on GPU ...')
    / W7 T; ^7 C: b" o3 s- t; |2 X" @; }0 e0 ]: i3 T4 R3 @
    device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')
    % g7 }5 k8 n6 O" r# d  m, H1
    ; W; p) G6 d  r2
    ) N% j% J  H3 \0 Y- u6 d( c3
    & g5 H6 P+ c5 N5 e% X# }45 [" k4 C* f6 C
    5  ^/ k/ d' f0 \/ r4 W+ N5 n9 q
    65 t( _5 d; l9 Z8 Z* |9 Z" [' A
    7: Z. {0 A$ K- z! Y0 K
    8
    1 Z3 o, H8 p9 {  o- B9
    & h" o, a: s/ W( k; U5 O, C+ m* q. XCUDA is not available.   Training on CPU ...- ^5 I- k  R( b) B( Q% f
    1+ y/ [% M; J. T7 ^- {
    # 将一些层定义为false,使其不自动更新9 u1 C  C0 u! p& D
    def set_parameter_requires_grad(model, feature_extracting):& o3 c9 b4 x6 y* h  K! v( ?
        if feature_extracting:
    ; I& ?8 u  B' s5 y' J9 ^: Y# H        for param in model.parameters():
    8 {' x: `% A. g' Q: f: _            param.requires_grad = False& f  k* J( F5 @8 n; F
    1! O& t' o2 ]6 W# [8 N# i
    2
    % i! d6 q/ i. q30 E! e/ B; h* e' R5 Z2 |
    4
    9 S. [; J9 {; m" W" Z4 \% Q, ?5
      O' v7 B2 i4 R9 m# 打印模型架构告知是怎么一步一步去完成的
    # _5 [. X% ]+ b+ V# 主要是为我们提取特征的
    / T8 @$ j( a  e7 r8 J% H! S2 b+ N" K8 b4 Q
    model_ft = models.resnet152()
    ( s" P. R" j8 |/ K6 [! O; _model_ft
    & T/ z2 T/ v- Q6 C6 g1- n: }6 M' T. z- ~2 p+ H# {
    2# U. F! H+ E( n
    3
    4 b  a* Z* s1 c- O4; Q! I  W/ w+ X' @0 S1 p. E1 N
    5* `/ g& A/ L. a) a# i  {' Z" V
    ResNet(2 Y; K5 @( w7 ?) d$ v% G# e
      (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
    7 q2 k5 W& F4 D+ |  (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)$ H9 |- ^: [( i- I
      (relu): ReLU(inplace=True), P+ Z2 U1 d& C7 t( I" H8 g8 k
      (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
    & V' F+ T/ z! Q  u, @+ t  (layer1): Sequential(5 F* C: s" h* U9 X* T
        (0): Bottleneck(
      u! c7 M- H, ?) a* ]( c6 Y0 ?      (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)4 P1 x3 \) O6 X% z2 P9 r) I5 E5 J) l
          (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    & k, m- R" ^1 b5 \      (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)+ E' x* [1 {& M
          (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)1 k. K; X5 C# f! z) V2 w
          (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
    ( |: s7 l$ [( U# Q      (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True): `) o. z, }1 q& m9 R3 U" c: v2 p
          (relu): ReLU(inplace=True)
    3 m2 v- a' |- z3 x) E* G+ F7 o      (downsample): Sequential(. A+ W* ~6 S: S( g2 `4 G
            (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)3 f6 o% o1 h9 p1 g, Q
            (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    - ^' |4 S  h& W      )7 G% {4 l) d! x
        ); P0 L+ e# X- q& v
    中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。, {  X5 T# f  E6 g. i: l; K3 U: }
        (2): Bottleneck(
    0 w$ u: K6 m# ^% c, e% x& I, j      (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)2 p8 d! N* z1 z6 d2 P  q+ v' Y
          (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    4 n; ^9 q) i% V  Q3 J      (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
    ( I% ~7 v: A2 r* t/ O/ Z3 @7 H      (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True), p/ z9 ]3 o1 P9 |6 V) @! N
          (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
    & s1 p' H& `# V  B( W0 B! z4 B: C      (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)$ C+ T+ n4 M" ?
          (relu): ReLU(inplace=True)) d5 C: a1 ?) M  E# I! T( ^
        )
    + }+ t8 c0 `4 a( A) o9 s  )
    6 I; T+ f( H2 j5 H, R5 e; C  (avgpool): AdaptiveAvgPool2d(output_size=(1, 1)), Z1 i" X3 v# i# i
      (fc): Linear(in_features=2048, out_features=1000, bias=True)7 ]( c8 w( E7 K: d! b7 ]2 F7 C: \
    )
    # y) C, u% @# r) p% c9 _7 B4 b; S( i4 ^9 D5 ?1 P5 L4 P
    1
    : Y: z* p2 N: u/ w: t2
      e3 q( ^; N2 P3 ]; _, \31 t; v6 b- a+ F+ c7 @
    4
    7 I3 P, Q  V- C$ f9 R5" _# y7 o5 I1 T' V. X, X; M
    6" W! q+ X2 I* p0 I- N  V  x' U
    7: E  E& g' g) g( r8 S, X$ p( l
    8& g& t6 w1 p/ M  v: i5 d( h
    9
    8 Y1 w( [5 |5 h+ \. E10% L# _( s5 V, N% J. Z  P5 y; E
    11! j6 y/ ^7 S" M* Q3 E! ~1 p  J
    12
    ( c' D$ I" J1 q. I8 P1 r- Y139 n3 o9 _- ?$ i5 q1 k) O
    14
    ! ?/ s9 ~  Y3 o4 ]15) n+ g/ h; [  `, H
    16
    $ G& d, L: b+ a2 T7 _2 W17
    # B1 m7 `& @5 N( C' v  \. D! g18; S4 ^" B! _0 N
    19
    ) E' M  `0 X& [0 M9 Z20
    2 q  l9 L, l6 `# c6 ?21  ^9 ?# e8 d1 g3 G9 d  Q
    22# U( }9 I) P! M$ t; i
    235 R: u4 J* m3 L& _) ^% i
    247 n4 L( n+ s8 B& P# q
    25( j8 r. _/ `% Z' _/ D" A; a
    263 D& G% z' j1 h: ~, p/ W. E& U+ z, l( E
    27! v. O, O2 j0 H
    28
    2 E2 u4 _0 a& f4 [9 h291 e5 [- Z1 K- t' p! R* g
    30
    ! H4 _. [9 r( `$ k31$ b/ V- K# j9 Q( d. ?7 b% D
    32  `* Q5 @- o0 w* m! m8 y
    33; r# J& J' f: J) \0 H  v( l
    最后是1000分类,2048输入,分为1000个分类
    % `3 k( g6 w1 N' y而我们需要将我们的任务进行调整,将1000分类改为102输出
    + ~/ X: k' \" v: j/ y. |  y6 v
    * J& L6 ]8 v2 [5 G& _6.初始化模型架构
    + }- \8 N: [+ h9 I5 P步骤如下:6 _2 J: F; c4 K( O* a" P
    " }) j% K) f7 I$ T9 L- ~
    将训练好的模型拿过来,并pre_train = True 得到他人的权重参数
    8 q1 B1 n5 J8 U0 B& {( a5 h可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)
    & N& f& @  a3 ^无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数% d( Q, C  H3 W! v1 C" G
    官方文档链接
    " J, k( }0 R) j2 lhttps://pytorch.org/vision/stable/models.html  |9 D1 Z# C7 `! N* k  P

    ' p, f3 A6 o+ {1 O" v# 将他人的模型加载进来
    2 j. Q1 I1 b  H1 t- sdef initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):. X! m% l1 ^9 K0 _
        # 选择适合的模型,不同的模型初始化参数不同2 l5 h' D+ G7 p$ P3 g: l4 g
        model_ft = None; P4 R% E3 ?, E, V% [6 U
        input_size = 0
    4 H3 o; {/ t/ l& D: E5 _& w2 C" {" w! E/ k4 i4 |
        if model_name == "resnet":/ \1 Q* I& L( W4 E
            """
    - t2 D9 _2 a2 Q# W/ G4 W3 C, C        Resnet152
    ' E- N/ G6 G7 T* u: i; _# U        """* }: u8 k# ?! r1 n3 m; M
    ) O4 ~3 j$ e. @  C3 f
            # 1. 加载与训练网络
    - e9 ~0 s4 G7 D) ]        model_ft = models.resnet152(pretrained = use_pretrained)
    ( _* i4 B$ Z) i9 [% g        # 2. 是否将提取特征的模块冻住,只训练FC层2 }( ?+ i# S* }3 _' v1 i4 }
            set_parameter_requires_grad(model_ft, feature_extract)
    & C- K, m0 G% m: I" I        # 3. 获得全连接层输入特征; v! n7 @3 B& ]
            num_frts = model_ft.fc.in_features3 f: j; Z( B' b/ t+ v
            # 4. 重新加载全连接层,设置输出102: U# ~4 F" w/ m
            model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),, m3 w: E' E4 Q. V+ Y( {
                                       nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为11 z  ?4 M5 `( r  L2 A" j' r
            input_size = 2243 B8 {+ Z- i; k; r+ a1 k- k) A+ u3 F

    8 I9 T4 P* u7 c/ R* |# O    elif model_name == "alexnet":
    * X9 v$ H3 _  p* N7 I7 q/ j# F! Z0 s- M        """
    3 i5 F, X  R7 m/ ?" t; w        Alexnet
    0 x# ^" r7 w' }        """4 F3 z! K0 f5 Q5 T+ [
            model_ft = models.alexnet(pretrained = use_pretrained)! s  J; L& j' i6 A/ y+ w' }& ^
            set_parameter_requires_grad(model_ft, feature_extract)
    ( o' N# O4 |# Z
    ; X) A8 B( E. H1 q        # 将最后一个特征输出替换 序号为【6】的分类器
    , P: u9 D! R) C1 O) t) h  W# F        num_frts = model_ft.classifier[6].in_features # 获得FC层输入
    : B; I% q; S* |1 M1 O5 _        model_ft.classifier[6] = nn.Linear(num_frts, num_classes)# _& G5 a- W. D- f- m( _( K
            input_size = 224
      {1 c: s* O& ^1 W; n5 }* b
    . v( r7 O' E8 {1 S2 U    elif model_name == "vgg":
    6 f5 y: K% y4 i4 ]$ p        """% b* x+ X2 D7 v1 C
            VGG11_bn
    : p/ u, A3 _9 q9 X' N        """! t" |) t+ N' n0 q
            model_ft = models.vgg16(pretrained = use_pretrained)
    3 u6 L4 A! `0 V4 H1 {0 ]  v        set_parameter_requires_grad(model_ft, feature_extract)3 s  O8 A  W& b6 D
            num_frts = model_ft.classifier[6].in_features$ H5 D* l/ ?1 }( u
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
    ! R/ V$ K9 U, ~! }. @        input_size = 224) H$ [& ?; E- F1 V  r. c
    ' x; a3 o7 ^; [3 X
        elif model_name == "squeezenet":- l3 w6 {% m4 ]2 w- h8 ^
            """
    / h% |- Q8 |$ F' }" w) E4 ~        Squeezenet
    + V) n% U# F' w$ O4 r3 H+ V9 u! J        """( _" T. `" B4 b" ]& q% f& C$ ]
            model_ft = models.squeezenet1_0(pretrained = use_pretrained)7 T6 p0 ~9 U: I' S6 i
            set_parameter_requires_grad(model_ft, feature_extract)$ D2 p' S! R$ j6 f' ^$ q  N; d
            model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
    5 x0 ^3 O9 X8 Y        model_ft.num_classes = num_classes& N; |. B& Z+ T' Z
            input_size = 224: z, }0 {( m6 W- s( W/ I) U" I

    4 {( E9 I# T4 M5 G+ y/ b    elif model_name == "densenet":
    " u8 q1 U' v- U+ }+ r        """
    " w4 b. c7 z4 x        Densenet" G9 `) m* ^/ V, \# d! R; {
            """" T. U% Y" c2 H2 R, c( O
            model_ft = models.desenet121(pretrained = use_pretrained)# Z) G1 x  t8 B# x' x
            set_parameter_requires_grad(model_ft, feature_extract)
    / q, y! y" O& t" z        num_frts = model_ft.classifier.in_features
    8 Z8 ^$ B5 m. h2 S        model_ft.classifier = nn.Linear(num_frts, num_classes)" O8 J- T& M( t0 t7 g2 r
            input_size = 224) C& x# K- l4 I2 ^+ b, Q

    . `& z' K$ F" S    elif model_name == "inception":8 [5 {0 t6 o1 u, j1 f; h& {6 Z
            """/ @- L' ~7 |- B' ]" \; h: f9 k
            Inception V3
    " P+ ~" u+ n7 S1 v) |1 |' O4 q! P3 I6 H. |        """
    + J+ `7 ?, l$ y+ A5 `$ Q9 E        model_ft = models.inception_V(pretrained = use_pretrained)4 i: {5 Y4 c1 i/ G
            set_parameter_requires_grad(model_ft, feature_extract), D1 W; {  @  W5 k

    : B; _: h6 u7 H, h$ z        num_frts = model_ft.AuxLogits.fc.in_features
    ( ^; M8 j$ b  ~) J9 e        model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)
    4 N9 y1 I, ~" X  [5 p3 P
    0 p& ~. O$ B; h( F0 U        num_frts = model_ft.fc.in_features8 f4 Z  S4 Q; Q- K
            model_ft.fc = nn.Linear(num_frts, num_classes)
    6 A& F# P6 O1 k) Z' [        input_size = 299
    + \7 V; x2 W! T0 R. N% Q, R& `# P/ k( @  X
        else:
    5 T$ l& h! X- r# O0 n        print("Invalid model name, exiting...")% s# l$ ]& ^, a( h. l
            exit()
    6 _. C. z$ g' }) ]6 o& U+ D0 S  U) C" r5 G$ R' l
        return model_ft, input_size
    8 _% c" X$ e  L: S9 K& E+ R. M7 X" B8 [6 d' W! ?1 r/ n
    1& M0 \- m2 f& O+ U. x, |" f
    2) a. ~  Z  ~( O% j6 V' _
    3
    / D9 b* `3 n+ d$ i& d) h4 S4, W8 d9 q2 e/ G: v  X2 B8 r
    5! v3 X/ J, k' f0 Y1 j: O: e4 ?; P
    62 ^' H4 a( X# Y" }' N& C$ I
    7
    ) ^- z! w6 U. g' f; i1 J! S; F8
    5 P; h8 X$ }4 Z/ ]1 u1 l- a, a. \9, z/ A& j+ A3 w1 @5 |" Z+ d
    10
    % B6 _4 Z( _# `) f112 {/ _4 A6 N7 J9 L
    12
    . v* x- ?. f6 r+ v* J' F& U13
    8 f3 V6 F* Z, n% d; n2 n14
    4 V! V$ v- }% D  v15% I2 J2 h' o/ k- U
    16
    3 I" Z+ J7 C4 w. H17+ X& E; Y/ W% \5 x' z
    18$ l% x$ }/ s) U* e7 j$ ~& b9 U9 w/ d
    197 w( z- J+ Y& H
    20
    & V' W4 e" f, _5 w# V21
    ( L! C' s, P" e9 L" M227 |( Z7 ^3 C4 W; H1 r2 v0 `
    23
    , Q5 u% X' n. k( @* Q24" |1 S9 @8 V0 g; H$ V! D0 }  F
    25
    + q$ C2 K, w' ]2 \# W9 w26) o* ?/ V. v1 |4 C, j
    273 \) k9 Y0 |* @3 f9 d
    28
    ! A( d) o$ g$ c* U! u3 C' {295 B) U, D8 `! [' m5 [& h
    30! o0 N! z+ z) j5 R0 p- x) X
    31
    4 w; \9 `" C; ?32. S) p( {) f" O! `+ B, ~
    33( ^1 t  w9 P( _% y3 M4 f
    34! Q+ a+ p, g5 A4 ~+ m
    35% I  Q4 x) K) M( t+ _0 M4 E! l
    361 o: y+ W- E8 a1 K6 ^
    37
    4 P$ M* w3 @+ f/ k& b0 B38
    4 M0 N' L) n6 x, M39* ]* L) h0 \# a$ @: }, |
    40
    5 G# R) ^- y/ M41
    9 P; |: E  p8 @4 m: H42. N0 K/ n9 l( b) c, L3 s6 b
    437 z8 t5 k9 U! w! F" o' M2 R* @
    44
    8 j; F; n0 e% k: F6 F455 s( P9 N/ C: u* Y
    46
    # @: r, _- x- Z# |47
    0 k- ?+ R- L; A! D48/ w( z7 {. K: e: q
    490 r6 m( k3 q- ?' }# p2 f
    50$ F9 ~" z% g' O( s7 Z. W, e8 M' `* G
    51
    ) d3 E: y2 s& p, f1 @, {9 d0 x52# R: u8 G" b  l# b/ q, N
    537 @8 y6 M# p8 G3 c4 B
    54& E5 R. r& @6 e
    55
    2 G$ |& d; P6 l) {6 w567 E; {2 E$ x* i* `1 U
    57
    0 ]  m& X% F- y3 c* r6 v58
    ( \- \- \' q& E3 z0 x* O59
    2 A: I( z& x, x% Y5 v) a60) b3 X7 K$ {! m, q/ v9 v
    61' n6 Q: u; w0 F6 W  d, T2 Z
    62
    3 P, h7 i! J/ h. F* \632 j% i9 z, L( \# ~
    64
    ; Z% @, v( J- m/ b6 i8 S% ]65, c7 j9 {! m, n7 B* ~# O: e3 v9 B9 }
    66
    3 o- u! n4 ?' U: x, _' m* z( k67% o% `0 l4 _! ~
    68
    + |& R/ }% a& ?7 P* ~! Q69! J" S$ y1 ]& p( c2 L% c
    70* q; ?* o% a. U2 ?
    71
    % x, c5 _. v: a, b2 Z72
    1 b$ {0 m6 F3 m/ B73
    $ i$ V8 ]; C; B% H74
    6 H. a( w2 f, t( {0 s75
    / o1 A( w/ K2 G3 I1 L2 t% {76; g3 m8 x; B' v" r/ Y
    77: n! @2 q0 ~  ?) w' _) U, }& i5 S) j
    78# g5 Y# X  x/ \( }* V$ ~
    79; v9 ^. u& t9 m: \# x3 F- n
    80$ _+ }6 x& t9 M3 _& N2 b$ @
    81
    ) T6 ~  @! t7 K' l8 z$ F  V82# `+ p( b; x$ s
    83
    1 Y1 e( J6 R" @" [' C7. 设置需要训练的参数) E" }/ H0 m- V: ?1 s
    # 设置模型名字、输出分类数
    5 Z8 l# B9 x5 n; ymodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)' t8 V" {. F/ K' V4 W& {! E

    0 z; T# l$ J9 ]" y' z3 R4 Y: I# GPU 计算2 f# p& j" h8 G( D. b/ y- X
    model_ft = model_ft.to(device)* b/ r0 q# j; V9 `" s9 |2 c

    , C2 k2 _' I7 j. L3 R# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取; e6 j6 {3 S4 J3 ~2 I5 B- q
    filename = 'checkpoint.pth'
    $ j+ l; _" h0 |4 |( l' K- _* o
    * u: C8 `8 J. d$ D% P' w3 i# 是否训练所有层
    % C. @, ^8 r' `7 v- A8 I8 T4 _( tparams_to_update = model_ft.parameters()
    # q& _* l1 C- K0 U' H# 打印出需要训练的层6 n' _) b& k8 }. \6 [
    print("Params to learn:")
    ) p) U3 e4 I" ?& j* jif feature_extract:
    5 E, U& x& j& p6 u- ^( y/ O1 L  k    params_to_update = []
    $ S; s1 ^7 D: m9 `2 ~! Z6 ]  l    for name, param in model_ft.named_parameters():
    8 Z$ n* S* G2 k! d2 N/ }9 Q        if param.requires_grad == True:
    5 N/ N4 F8 i$ c* P; _' o) H6 p            params_to_update.append(param)
    7 g, F8 ]; `( U: s/ E+ w+ v            print("\t", name)
    0 L7 u" |6 x6 belse:
    + [5 n; w# \* \" {; ]) i    for name, param in model_ft.named_parameters():
      a4 S- W" f# }6 U. I+ d        if param.requires_grad ==True:( E$ B3 W5 q8 Y3 `$ o# V
                print("\t", name)
    , p+ g* F" ^7 z5 H
    . [4 T/ z6 }. `& W( W& k10 @/ @  C4 e' k. m$ W
    2
    # U" s( k' o, t7 d3: c' z  ]5 h# d- z- h
    4
      Z" O7 V: a& j! h* L) X5# D: C( J( ~9 k9 I# F% ]$ \7 x
    6' t+ ^- }$ j5 S! f
    7
    & }% e; W) z  t1 ]; _! u- P7 l8  W: H2 r/ m* t) g# M
    95 i) e/ }3 D( r3 i
    10& b. A/ z9 r7 h7 C- ]6 u
    11& C2 W5 @* ^, {9 u$ b/ _
    12" {4 ?) s) i. ^, n# v. Z
    13
    * `' N0 ^9 d8 U' X) l1 L14
    ' }! U% b( ]' h6 C$ n* H15" u5 w4 j/ O5 P
    161 p+ H5 |0 F2 s6 F' z6 j& G+ R
    17* c8 A1 G' M0 O& Z' X, d' m8 x. K  b
    189 }0 e4 y# [  I) X$ ]. v! |% C) [$ l
    19
    ! r2 r6 K! @' F! Y2 ^: t  q20, f# c2 f6 P1 m$ p7 x8 @
    21
    1 [# V. D3 W& N  O22
    - H6 T# r  L9 H  i23- V% Z0 M4 w" R$ a& z$ [
    Params to learn:
    1 E9 E8 f  a; y1 @+ P( T1 w1 x1 ~+ D         fc.0.weight( a+ g. `7 X. E6 `# z
             fc.0.bias) C7 G5 ?, N6 Q
    19 Q/ b* E3 @2 A' ^
    2
    0 T7 j1 \# Y# b; P3$ ?/ z4 i: {: n$ @/ v4 ]/ R
    7. 训练与预测
    : @. ^2 n) j$ ^" {) w5 t, w7.1 优化器设置. g& f) `. `3 [3 X- @) K7 [
    # 优化器设置/ I5 S7 O6 p0 h4 O) u, _
    optimizer_ft  = optim.Adam(params_to_update, lr = 1e-2)
    . g9 [; M5 J% C) z* i1 ]$ m. b5 j+ a) j# 学习率衰减策略
    4 `0 U/ h+ b0 `  K1 Z/ dscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)
    ! h1 s; {: c) ^6 |9 t# h7 _1 {- ~# 学习率每7个epoch衰减为原来的1/10
    " G; T3 k% b5 q7 s$ ]  i' \. O# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算! B) M' G* A+ q( v% m
    ( S# c4 Z  ]* e3 {- y
    criterion = nn.NLLLoss()% N% ~3 v7 k! `+ l) d  h
    1& Q7 m( q7 \" B' n! \& g. F
    2, d' t; }  k& m: ?4 N
    39 ]/ O$ w# T: c# A  U% n
    4
    % F% ?( b. K" R$ _5
    # E6 K6 @4 s5 M# s  I2 e61 Z' |+ E: a5 a1 S! W/ s
    7; m0 `! H, F$ p7 v( \2 a: E
    8( @9 H% |7 F2 r: K
    # 定义训练函数9 t; y' v" B, G, ]$ n. R5 L+ G% c7 ^1 c
    #is_inception:要不要用其他的网络
    . K- C/ J- p/ C6 F2 n" ndef train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):% W/ q6 s9 i1 G0 i# f) W
        since = time.time()
    ; `& y, t* X" u$ c" a    #保存最好的准确率
    9 T' F+ \+ L. N  |* b& k5 h    best_acc = 0
    # A: ^, Q4 `3 m4 L    """1 W$ G, H1 w8 j( r4 X* H
        checkpoint = torch.load(filename)3 ?" ^; w8 ]$ x8 C6 ?( Z# c
        best_acc = checkpoint['best_acc']0 b) p+ G, o! ]- G  ]# V! ?* L+ a$ v
        model.load_state_dict(checkpoint['state_dict'])
    " F! j/ }, _% n/ X; U+ Z    optimizer.load_state_dict(checkpoint['optimizer'])8 f: X. k" t) q, A$ b- G: Q
        model.class_to_idx = checkpoint['mapping']* T/ y) I! u* q4 ^0 u
        """" D+ w% l3 L9 R" R
        #指定用GPU还是CPU8 N9 `" R% z5 r2 q0 E3 P8 j. i
        model.to(device)  G. J  z4 i7 K% c7 U* ^
        #下面是为展示做的# b" B: t8 u. G: O/ X2 Y; i# ]
        val_acc_history = []
    ( b0 N; Y6 |0 G! v6 i1 x    train_acc_history = []2 u! B& N' t% {9 c( q
        train_losses = []
    9 f- o9 [0 _9 d    valid_losses = []
    - \2 u' u  o! ]( k1 ^- C3 A    LRs = [optimizer.param_groups[0]['lr']]* e8 C) Z: `3 o- H& e1 t
        #最好的一次存下来
    $ U& y! ^/ M1 w  A8 D% G. B5 f    best_model_wts = copy.deepcopy(model.state_dict())
    2 Q4 X$ \+ C# x( _, Q  ]
    ' p# m9 J( e6 Q/ f    for epoch in range(num_epochs):/ G/ y/ N2 m2 I( S3 A' C
            print('Epoch {}/{}'.format(epoch, num_epochs - 1))) K% ]  R$ m5 K  C4 Y; M
            print('-' * 10)- C! C3 q; A% R2 T

    6 b9 _9 a& G9 f$ b; C8 }# M# G8 s        # 训练和验证
    8 l$ h; m6 s* ~0 N        for phase in ['train', 'valid']:
    3 ~* q' G, ~& ?3 h# M- i! o/ ?            if phase == 'train':- t# Z- j5 Q) t" t1 j
                    model.train()  # 训练3 |2 s+ \' z$ T6 a) j6 K' e! y2 T
                else:2 D9 I# E% }7 v& L" A5 t
                    model.eval()   # 验证! [2 {( i6 g8 U9 C
    0 B' e9 c. W/ h
                running_loss = 0.0  j; S& O5 g3 I: y+ r. @
                running_corrects = 0
    / o7 b2 v7 L) J/ w2 R
    $ J! F. s. \, B# ?6 U& `5 u7 I            # 把数据都取个遍
    ' [  d+ D% i5 w2 b, F0 B: k9 S2 R            for inputs, labels in dataloaders[phase]:9 c* Q) E- |3 D- ~/ m
                    #下面是将inputs,labels传到GPU
    1 r) Q, H6 ~3 S6 u" [' }0 U. H3 v' e                inputs = inputs.to(device)
    ( A4 W1 e" q3 g/ h5 E                labels = labels.to(device)/ j: ]2 H4 `8 I& A
    : u* F- |4 U! Z9 S2 ?/ W& M
                    # 清零
    + ^7 ?" H& l; a+ {                optimizer.zero_grad()
    , Y9 }7 L4 F  a* R+ o                # 只有训练的时候计算和更新梯度
    # J! B, {' Y. H                with torch.set_grad_enabled(phase == 'train'):
    . T$ e  Q1 L% f: k5 I$ J# n6 h                    #if这面不需要计算,可忽略
    : Z' x6 [' N. o1 O, B                    if is_inception and phase == 'train':$ W# `% \. K+ n2 G4 k# w
                            outputs, aux_outputs = model(inputs)
    & H5 X% i' u8 D                        loss1 = criterion(outputs, labels)
    * ~" e" I) P7 M0 D                        loss2 = criterion(aux_outputs, labels)
    0 B0 R3 i) {$ g/ d6 J" [                        loss = loss1 + 0.4*loss2, W2 H# u; h$ T
                        else:#resnet执行的是这里
    " k; P! d3 f' a9 {+ l                        outputs = model(inputs)7 W: A3 l  Y  F# q' C
                            loss = criterion(outputs, labels)
    3 D( G, F0 |( h, N# \
    & v: A0 f* i8 E* b                        #概率最大的返回preds4 W6 Y+ S0 z% j
                        _, preds = torch.max(outputs, 1); y* ]  z& d. X% p7 C0 E4 K

    6 Q( c2 \3 d9 e' J) U$ C0 e                    # 训练阶段更新权重
    ) ^4 D- O& Z1 i/ G/ L% k) N                    if phase == 'train':
    8 m- L$ Z6 z. V; y, }$ c5 Q2 [4 e                        loss.backward()
    4 p) P- u3 O, u0 Z! A                        optimizer.step()
    ) T* U/ d& [/ f  s7 _0 f0 _2 f
    & t5 I/ I: B9 x4 \  R* Q                # 计算损失+ ~, p  h7 F* H
                    running_loss += loss.item() * inputs.size(0)/ f9 `! D( |+ o2 d
                    running_corrects += torch.sum(preds == labels.data)
    ) a9 x3 f! Q& o, R' j2 K8 m; k  y& q
    1 s) Z! c( {: N6 J; H9 @7 V( P- D            #打印操作) B" X) y0 I! M& B: i
                epoch_loss = running_loss / len(dataloaders[phase].dataset)
    % Z3 C" d/ ?3 k$ ]' Z$ y  T5 e# Z9 [            epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)' P: G6 F* N! H

    ! L% ]: n! e% Z2 d& w' n2 n, y' r. V$ ^" q! D0 K1 a9 v
                time_elapsed = time.time() - since' M. t4 O" _5 A6 k7 K+ k
                print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))1 h" `; N& J' b1 E5 ?( B
                print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))# Z5 S$ o) C  g

    2 H! b0 n3 u' T. j) a  r  o; H7 J1 Y5 a7 @3 ?$ e
                # 得到最好那次的模型
    ; q& K) J9 \, |+ }+ B6 M            if phase == 'valid' and epoch_acc > best_acc:  P, T$ _4 f4 C# X4 y
                    best_acc = epoch_acc  m: i- H3 P5 h  R- H. j  f+ L
                    #模型保存( n2 q2 w0 P- k7 b! b3 ?* s
                    best_model_wts = copy.deepcopy(model.state_dict())
    2 y) S) X7 B* _4 @) i3 z5 n                state = {8 g0 K! P) p$ L; C
                        #tate_dict变量存放训练过程中需要学习的权重和偏执系数
    + n4 Z, e. L! n8 R+ O7 I9 Y                  'state_dict': model.state_dict(),
    1 O' w6 w3 e" [$ \3 Y% [  `                  'best_acc': best_acc,% {6 W4 h1 f  N" O( |
                      'optimizer' : optimizer.state_dict(),* C9 Q& X5 ^5 G7 l% p6 d' I
                    }
    ( c* y9 a& U: }% L                torch.save(state, filename)
    : S+ ^0 k1 D( p            if phase == 'valid':; X; ^; z$ m( g) H% F( B1 d
                    val_acc_history.append(epoch_acc)6 {+ ]4 C& v1 `/ m* n& q4 ^4 k
                    valid_losses.append(epoch_loss)* f8 q! @% T8 x. x: H
                    scheduler.step(epoch_loss)1 w) K# f% A! f" v$ W! ~6 w5 \
                if phase == 'train':0 V# M- M0 U3 u. t5 g1 D$ Q
                    train_acc_history.append(epoch_acc)
    # c4 t  C) Z: N3 [                train_losses.append(epoch_loss)9 b$ m4 |' p+ Q  X* S* i
    1 D) L- h1 |# x2 c% \
            print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))6 W8 c; B2 I% [( f# _" @
            LRs.append(optimizer.param_groups[0]['lr'])
    6 L" @) u8 I  \* @4 D        print()# @3 B: Y3 y; M, v* \5 l) l5 H( @

    , q& n, `' z- M1 L+ @1 L- M    time_elapsed = time.time() - since4 W+ V; @  u2 l8 \+ h4 e1 L
        print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))3 y; A+ L" t; `6 M9 X2 t8 g  h
        print('Best val Acc: {:4f}'.format(best_acc))
    ; H7 p  ^  A, P) [9 U+ ^
    - f& ?& D& W- c! m. l+ u    # 保存训练完后用最好的一次当做模型最终的结果' q: d: I+ u/ U% S/ y# ~
        model.load_state_dict(best_model_wts)8 r' v, \5 C+ @$ M$ u: ^/ g' {: r
        return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs
    : |) D# c! `. L7 O" k
    0 R6 p4 \% |% H: ]8 t
    7 N* Q! ]2 b& S. y" W/ u1
      a2 ^) Y4 [( K, X1 j7 G7 \8 t: ~) @2% @* q% j( D2 r( e
    3
    ' R- l! A, }8 ]: a1 Y. @4
    9 Y. q, ?' e+ j; k. P53 }/ X2 [* D  [: h: y
    6& Q" C8 y& T: U! G
    7/ |9 U. Z# |) v
    8
    4 ?' V& g, K' E' k$ o$ ~: \" r9+ ]2 }  ^5 }" S: O5 v- g
    10* |* M. c; Q4 U" ~2 c+ D
    11
    : w* ]  v5 o4 k4 c, k5 {" s120 i9 \5 T9 B& ^6 I
    133 c, d( a$ X5 e5 P8 o/ O2 x. g
    14
    . b2 z) h3 |) y% d9 }% a15
    * ~: m5 ]  {- m; M/ X  ]1 Q4 U16  x/ p" L  R' q8 r- e1 C% ]9 b
    17
    ( d" @) l9 Z& U! A9 d! o18% _4 U$ S. ^! b5 K
    19# x& r" v: _) Y
    20% V# c' x0 `$ a1 f
    21( }; i. u0 O3 B& y8 N* o
    22; B# x6 @+ G+ j7 L+ w
    23) C' e  F2 `8 F4 y, e
    24
    + r0 ^. i; t' A1 H25- E" ^, g  x( R; J2 Y
    265 z* k" o. R" L1 c4 q2 H
    27$ r9 m+ S/ w2 O+ t% r+ r8 A# F
    28
    3 G- J* H, e0 L29) {" L- T) z9 U$ G, p  Z$ k; x
    30% m  ?6 e  c' ^* c# Y5 `
    31
    9 l1 x& G3 G* @4 C* {0 _32
    4 D# `- C- e/ p  U33" m' d2 k3 g( X+ X6 E& e% G
    34- d9 Q4 \9 j0 K/ d; D$ C) Z$ w1 P
    35* _4 ?6 |- u8 m# Z! R! ?
    362 B; T6 S- @* U
    37: d. Y2 B( j# Q  h
    38; O6 M& I! O* L+ j" Y9 s
    392 n8 Y" R. B8 i$ z. T+ ?' f4 u
    401 T6 N9 B9 w8 @3 P' W$ X  D
    41
    & q) V& J" H# B% J1 ^4 A42
    # b2 B: ~/ k& k- m: p& J! E, b) V435 x8 p' U# t2 @* e
    44
    " j* Y/ _; u; L& T1 ]% q; A45! }7 \  }$ Y7 z+ E, B% }. x' j) n5 i& g
    46" A! }: o5 x0 q) @2 Z9 c
    47  @4 Y2 g9 o) ?1 i2 o/ @
    48
    2 j2 ?8 C! v: O( b  R# j7 |49# \! c3 u# ]) L0 K0 Y2 \
    50' z; W$ @( \0 M, {$ y
    51
    7 M/ N  P$ Y+ V6 T7 a52
    : y: ?6 v* m$ G; o53
    . G3 c: v/ R3 q9 }  J54
    5 p4 J- r, x9 b* o( H! B55
    ' B# e1 x$ w' l8 C+ L56
    % y$ f. @! X( w: r57
    1 _2 r' E2 D; H' R58  s* l3 a0 e8 X1 I. y' Y
    59
    - _% F2 ~9 F: Q3 p- U60
    & Q3 ]+ K, g: b' i5 l4 U* w61
    / H: P" C4 Q: T62
    6 d5 ^7 W! c% M4 Y2 t63
    % T( E1 n! g  P- A64/ h& e6 [+ Z8 ^& G: I  W! w
    65
    % Z* P- |$ i$ [! L, M# X& ^66
      p. i- x* _" a* n674 X+ w; y3 u5 \4 {; ?7 H9 M& K* l
    68
    ; Z  W4 K+ M0 h0 @6 e" P( [69
    & C1 b; t5 `0 }1 X, X! C705 A/ E8 Q2 L  S+ D7 s3 o
    71% w, V! x' U( q) G- R, |
    72
    ! F$ V8 |! g# ^- |9 u737 b4 A+ d8 k9 ~
    749 S3 o+ b: P- V- H- Z- c9 q
    75
    $ V8 [' d0 Y, {" d2 [6 K( u76
    & W7 H+ g+ P. H4 B1 G776 g0 \" c0 s, w9 c
    78
    / N  H# ^. P( q) B" J79
    ! S* e' D+ \, e- a. h; }* X80
    6 _5 e% I( o) Q" ]: ?: N& x81. c4 [3 j7 A% }, h% h; G
    82/ f# a7 j4 N- X& C0 R, i! n. Y# o
    83
    4 W4 K$ j: W( d: E844 j1 Q' P' Q, x# W/ j+ o# }
    85, S/ q. w+ ]6 D' O2 R: t1 h" f
    86. \. H  L3 d1 I. l
    87- s. ?$ F3 R! Q9 m) w7 n4 C- T
    88
    2 G8 R; X( |4 O8 `890 \* {8 ^1 z& C9 Y8 O
    90
    2 u3 k( Q1 J6 z5 S7 ^( R91
    1 y7 \5 [9 C5 U  Z92
    6 F! G; }2 e4 b( D/ F93
    1 u) {- S# c6 C& K( V' S/ {94- w3 j& ^# n. y/ Y( N5 J
    95
    , @" m+ L- O7 K, W3 o7 S4 t96
    2 T: g! w+ I/ Z, C2 B7 l5 q97& R" V; O. }1 V* U
    98$ [* C1 F% [$ [! h: h4 Z7 }
    99
    4 O1 w: @' [! d: c- ~- l! R8 n+ k& P( X100
    & I2 N, s7 l- l2 f& ]; |101
    0 n0 O$ k9 A" d( ?: I102
    " B/ B7 l& {+ P* C" T103# d/ I; H0 D% g: s2 Q! P, X8 q
    104
    0 C9 A# ^: X# k105. l- p! l2 @: w8 `3 M0 @
    106/ n. I) q" c6 V! _# j
    107. n4 ?0 b3 a) s2 s
    108
    / P5 u+ w. X7 f6 ~109, s+ b" }/ ?. p( w, J+ M
    110
    2 V: U$ b" P4 ~: c$ I. v1113 a& h  ^7 C/ ]$ c! Q
    112
    % o7 \, }2 L: q" C7.2 开始训练模型2 Q5 e0 ?1 u& |$ f8 b
    我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次
    , y3 u, ^  V) {  w
    $ u5 u# k' s9 \/ f/ c  b3 t4 F#若太慢,把epoch调低,迭代50次可能好些
    ) ]5 i& k/ w/ o$ H) I) L: C! p#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合9 A& O$ w+ \7 D- ]
    model_ft, val_acc_history, train_acc_history, valid_losses, train_losses, LRs  = train_model(model_ft, dataloaders, criterion, optimizer_ft, num_epochs=5, is_inception=(model_name=="inception"))
    9 v% s+ t" O: }. p
    : \4 }$ h# r$ I& |" q1
    ( B; _- w2 y+ ]( A4 a7 s- V# b  [2
    & f/ U) X8 Z. h7 U# Q5 [) @32 r( e1 x7 @# L8 W
    4
    % O# X3 u) G2 H+ NEpoch 0/4( t* i7 b  L8 k% f) E
    ----------! m4 ^- I" O9 d
    Time elapsed 29m 41s& u8 |4 Y5 P, s2 R: c) E' i2 D
    train Loss: 10.4774 Acc: 0.3147
      ]4 _& r  j4 S! a0 ?2 g7 O. \Time elapsed 32m 54s
    + T4 ~" G) H9 q+ Wvalid Loss: 8.2902 Acc: 0.4719  G' \' r: f% g3 X% i8 A% z
    Optimizer learning rate : 0.0010000
    0 Z8 ?! g- R  ~' o/ T( Q8 v4 f5 [  _; q. F1 m0 d; u' O
    Epoch 1/4
    & L/ N- c3 t" f----------% L. L& z# Y! L9 L3 W3 H
    Time elapsed 60m 11s
    # b4 h" C' \6 L8 q% \( R9 h+ Ztrain Loss: 2.3126 Acc: 0.7053
    4 s! Z  U! B4 s4 J, O4 eTime elapsed 63m 16s) O; |& @5 j0 f3 y/ Y
    valid Loss: 3.2325 Acc: 0.6626
    4 E" N& g3 W2 eOptimizer learning rate : 0.0100000
    - q( v! w$ B  _0 H2 z7 a4 w, M7 R( G+ \' |% Q$ Q
    Epoch 2/4  I/ R# L0 j+ J- s$ v: O
    ----------
    1 c2 E' ^$ H' m0 H" u9 ETime elapsed 90m 58s: M4 S- W# C; @/ T
    train Loss: 9.9720 Acc: 0.4734
    , W$ R# u5 s# DTime elapsed 94m 4s' i0 y# f# M0 \' M$ m' |
    valid Loss: 14.0426 Acc: 0.4413
    * \* }3 ^2 ]6 A: i; o  mOptimizer learning rate : 0.00010007 G: Y' U. q- P. X5 m+ J
    - O+ \, p9 `# f( \* B9 p$ a1 R8 R  d# ^
    Epoch 3/40 q, a% q4 o( Q. `
    ----------
    / V2 H+ i+ p* C+ Q  {% O% M7 ~: iTime elapsed 132m 49s
    ( q* _8 ]( e: }& E/ ~8 d# Utrain Loss: 5.4290 Acc: 0.6548
    4 c' B/ w: a) VTime elapsed 138m 49s
    + Z7 O. w, R- O8 D5 j' Vvalid Loss: 6.4208 Acc: 0.6027
    # o. r& l5 }# ~; R1 v0 T! DOptimizer learning rate : 0.0100000
    4 }+ k5 A8 W/ V0 t# J& U
    1 m1 i4 C. z/ [" Y! {1 jEpoch 4/4. l" G9 I- |! L+ i. K: H1 u( N1 ^
    ----------+ t) `7 o1 s  m5 |1 W& T. Z! v
    Time elapsed 195m 56s( P, L+ r0 C5 q0 f& e1 o# I
    train Loss: 8.8911 Acc: 0.5519' |$ D/ I9 ~4 Q5 g/ J* U; H1 [) I
    Time elapsed 199m 16s  S0 F6 f! }5 I( U, u0 W% H( k  b
    valid Loss: 13.2221 Acc: 0.49148 p5 h% N/ B, G5 N% d9 |
    Optimizer learning rate : 0.0010000
    2 h$ u: H- T: n! ^) m# o3 Q2 G/ `# W. Z3 m( X, B' L0 n& D
    Training complete in 199m 16s
    ; c* d$ K; P( i- |) |$ UBest val Acc: 0.6625922 t$ A/ [2 a6 G) y& h
    - Z- @- k* T" o  H/ @1 v0 l3 Q( d
    1
    ( s6 u. U/ ]; l2 x% f% a2! B$ s  \' w2 W# s
    3
    ( b3 ^4 X  X, K" B1 {% d2 @4
    : Y/ H: g# N2 ?+ Q* {3 g5
    : }  r9 {3 t$ |! b6
    6 k3 R; I# O, v6 Z: L7* z) a, _5 N4 ?5 h4 }
    8
    ' V' i+ f8 s3 ]1 y9
    ! U4 [5 q: X, w# o0 c8 x10/ R+ v* x  e1 ^: i
    11
    & }9 ?8 @/ b# r+ \2 r" Q5 H. J$ U12
    7 W4 N1 M: Y* `13
    4 x' i- |7 @) P3 D; C4 j14
    7 W; }" x  v8 W* X: z" d15/ l8 C# v0 e$ {1 o7 U+ A( m
    167 [5 ]8 r6 X# _' H0 d' T3 S
    17
    ! d- C* T1 h0 u18
    + r+ l3 K5 U9 i7 q19
    1 i0 ?. G5 h' f$ O. |20+ t2 [+ e/ S1 o6 k' U! P# ?
    21
    ! ~& C2 c; N7 y. Z: K/ n% C# z22
    ) F& v$ R' k) F; a23
    5 `9 V9 [* w, c1 C24
    8 G  s# D) \( E  u2 a4 |1 t25
    5 D8 y. J4 w8 D, c7 O0 T) y268 i' m& N" _' ^) K" ?
    27
    # y% K( W/ h; E: }1 D& [5 O9 I) t280 q0 F8 N" F6 P2 L* x( O1 ]
    29
    5 P/ J' ]! J1 Z/ M30& m$ ]- }* V, u# g0 H# a: b  A
    316 T2 P% J- R$ _' u2 f/ v
    321 S3 q. f: R0 O% i, `
    33+ F. ?3 }8 W" m5 s
    34
    & Y/ a" k: a$ a1 L, T5 D35& ]6 V9 ]9 h8 ~% y3 }9 P
    36: S. X9 B2 i2 x$ J
    37
    * k% G: W- d( d/ n% R2 ^( j5 j38
    : w5 F& ?9 U7 T! d0 G( Q$ ^5 q396 W" G6 u/ z/ q' r
    40
      |( p( f+ G: N4 o8 R6 u41
    " }6 d' O: ]/ S3 H" |; K42' u  @, o1 C6 [8 x+ j
    7.3 训练所有层: H& D0 Y. Z2 I5 \; R5 d6 v
    # 将全部网络解锁进行训练
    ( a; ]6 t' n) Z, h' K+ Zfor param in model_ft.parameters():
    $ ~+ m& k$ V& s4 Q    param.requires_grad = True
    , A9 @* A& R3 c2 {- h, `2 }2 o3 I& M- ~& Q# h4 E! B) h
    # 再继续训练所有的参数,学习率调小一点\+ o2 O; m) s- U1 `
    optimizer = optim.Adam(params_to_update, lr = 1e-4)
    , s3 H! z7 R: f( o- c. U, j! jscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)
    0 W0 n' d& T/ B7 M% x7 v' i
    , l' ]* k5 j1 q7 O8 k( [# 损失函数' }4 d1 [6 _5 L, Z: O: @3 x
    criterion = nn.NLLLoss()0 P, o0 z7 X) L6 \
    1
    8 Y8 ]0 u) b0 X- u" ^. s1 G2
    3 w$ M  \$ h: K' T3
    , _2 E2 r3 W3 O# W3 W  \- X7 D4
    7 p! |3 A" R7 \: G3 @5
    ) h. k( c( M8 l! t6
    0 B$ M. M* k; x77 b( @! X1 s& H" U! J/ |
    8
    2 f! x4 }8 e. J  y4 u) g4 m9" S% d* l; B5 ^. K6 s% G9 `, G9 K% S
    10
    * q! R3 A% E0 e+ H  [2 w5 I# 加载保存的参数* o! }( h0 b- R1 U' a6 Z
    # 并在原有的模型基础上继续训练. p" v$ v3 d- ?) k0 s5 Z6 a$ d* v
    # 下面保存的是刚刚训练效果较好的路径, T$ }4 |* }# S
    checkpoint = torch.load(filename)
    * u9 L5 d. y+ X# v9 Y  t) sbest_acc = checkpoint['best_acc']2 O  E1 e# ~+ _
    model_ft.load_state_dict(checkpoint['state_dict'])
    % G0 H% b4 v1 K* a9 }% Xoptimizer.load_state_dict(checkpoint['optimizer'])
    , P3 @' f, T# x* L, C9 h0 f1
    - S; q! E$ X. T2) n2 L+ ]( o9 H' u7 k6 J5 U
    3
    ( z6 R" O3 M" y8 Q' u# r4
    ( ]8 L4 W, T! i! w/ c; P/ i* r7 E5
    : y. G" p) c! o  @# t0 O& K6% W3 ^. `1 @# D
    7
    9 A6 ?0 G. {  [$ K8 d) R开始训练# |2 Q% X: h6 d1 F8 k* V. `* B& H) f
    注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考
    2 q* `4 y1 b( o4 k) M* L; _5 `6 A6 i; O
    model_ft, val_acc_history, train_acc_history, valid_losses, train_losses, LRs  = train_model(model_ft, dataloaders, criterion, optimizer, num_epochs=2, is_inception=(model_name=="inception"))
    % }/ @( [6 E, Q7 {9 c$ A# k/ R4 K7 n1
    4 o: j# w2 Q- ZEpoch 0/1
    & d4 @! F2 s* e& m$ w$ N; s----------" P  d% g6 ^! A& E
    Time elapsed 35m 22s% R! J# L& o9 o: ?: Q4 S4 S
    train Loss: 1.7636 Acc: 0.7346& Q# n1 Y. ^( f/ v7 r
    Time elapsed 38m 42s
    - c8 _* R  c# P9 z+ K. G0 Q2 Bvalid Loss: 3.6377 Acc: 0.6455
    5 r* P0 T! V/ S. |+ Z0 vOptimizer learning rate : 0.0010000
    $ D3 P2 M8 ?6 w! }8 D: _2 M+ ?0 H6 P, s4 r! M4 g
    Epoch 1/1. u* b3 }# Y( f. F; r
    ----------% T  _  q/ U4 N
    Time elapsed 82m 59s
    % }$ o& [, g0 ~& {& Z; U5 f! ytrain Loss: 1.7543 Acc: 0.7340( m5 w3 a; I- K. K/ O7 Q
    Time elapsed 86m 11s- m" F+ g/ r" e: O* [* A- U  y
    valid Loss: 3.8275 Acc: 0.6137
    # x7 q7 W. p! d: ]8 {Optimizer learning rate : 0.0010000
    $ ~9 \9 ^2 p2 _
    5 K  K8 f; ~; Q' o6 B1 g# DTraining complete in 86m 11s
    * G) E8 P8 {: o# P' B7 o2 z" m& e  NBest val Acc: 0.645477/ X! ^+ A1 N# O- D; F2 m5 J

    . n4 f5 m  r: I) \4 A9 v  n19 r$ w! L( t8 x+ Z* k
    2
    , q* `. E" D7 K0 L/ i0 d3
    $ _) _* N( n$ c! c) R' s4
      D2 Z' y# F* t3 S5! M, {# U1 U" l0 `# t
    63 P, B, N+ H$ b5 o9 t, Y
    7; ^, ?+ x8 ]! H" P
    87 ?  M# P; {  c% |
    9
    ! d8 A2 [4 O: j4 @, c10
    ( u2 ]# r) M, p7 H4 S114 O1 P* G9 |' ^  F8 ]) w  z1 t
    12+ S, I9 o' t$ P: o; K$ e$ [$ _
    13
    0 A! t( I; v2 W* ]6 R8 X9 ~14
    4 t! t8 @( e3 E* f/ m! ?, {3 T15. @; e1 T" ]" j
    16
    # Q0 u' j3 c  x17
    : A( N7 n( ]7 R( J7 n, I18
    ) E0 s5 e  T. d* c9 q8. 加载已经训练的模型3 R! h6 X2 H/ l
    相当于做一次简单的前向传播(逻辑推理),不用更新参数  q3 o5 M) \9 ^

    ) h/ c3 E3 q8 C' ]  Omodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)( B9 y9 g) g* O

    ; }9 K' A: |1 Y/ |! H& _( N* \+ T' j# GPU 模式+ x2 i, k2 N1 K8 c
    model_ft = model_ft.to(device) # 扔到GPU中/ U# {2 Z" y) q- t1 Q

    ! a5 W2 f6 l. j' L  u8 ^# 保存文件的名字3 S: X5 C3 M) _( s% C
    filename='checkpoint.pth'$ G* }" `; d) f) d) _# x. ?- W
    * H3 H4 l6 ]1 A5 J6 e: x
    # 加载模型
    # T9 F, `( i, O- l1 l3 P  d0 [checkpoint = torch.load(filename)
      x# F4 ]# S- W  {5 D/ u# h. B  Ibest_acc = checkpoint['best_acc'], R, i% V6 K$ A
    model_ft.load_state_dict(checkpoint['state_dict'])8 o/ y9 x- N( T4 j8 b/ h4 }: n% }
    10 E* I( \2 t* t
    2
    : Q3 N$ D$ w- @! R/ s5 {3
    7 V# j/ |3 ~0 q3 c3 y: ~& t4
    6 M0 u0 Q$ P5 p$ r5
    , B: a0 G. v: W, r62 r: d1 {5 l. t. l2 _
    7
    + h1 a+ j" m. Q& \& I: p3 N! j4 M8
    ! [4 ^' Z- ?: H- Z+ E, e2 c5 Z& O9
    ; |- k) Y0 i! `- K1 K% ~10
    ' x  v( F, g: f; D0 f: V& P, h113 c/ i1 ?' m' T* O" ?
    129 w( ^( ^8 |, U( B' \( V3 j
    <All keys matched successfully>( F, L4 `% D  b  I$ y; P. H
    1
    5 @2 w, V) c& p* O2 `. i$ Bdef process_image(image_path):# s' A* @8 U% D4 g
        # 读取测试集数据
    - s, b9 D: r; E" q    img = Image.open(image_path)
    0 B; v' u" u: v    # Resize, thumbnail方法只能进行比例缩小,所以进行判断5 g/ l7 g* e5 X' @1 Z5 X9 y! _
        # 与Resize不同* z- j9 S. n- P2 Y
        # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小1 h* Q4 J. {# n1 h- B
        # 而且对象调用方法会直接改变其大小,返回None5 ?& B: }1 |1 a- w& @2 [' J% h
        if img.size[0] > img.size[1]:" O- n5 k: `" v2 Q
            img.thumbnail((10000, 256))
    1 ^& A& k: ?' ], A- `9 s! c    else:4 n; u, q, {% U  A5 d( A% [( x
            img.thumbnail((256, 10000)), R/ N0 I& m' Z- ^
    ) Y! M0 ?* ]% M' d, I2 X- j
        # crop操作, 将图像再次裁剪为 224 * 224
    ! l9 h. o1 C9 ?; f: `! b    left_margin = (img.width - 224) / 2 # 取中间的部分& |  |1 }5 |- F/ o/ f4 S. M' `
        bottom_margin = (img.height - 224) / 2
    4 D5 n" [, p, c% B" T    right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度+ R6 H1 I  G+ ~: Z& k: B! z6 s$ g* F# I
        top_margin = bottom_margin + 224+ i* Q2 r4 h! ]. L% v

    / x$ a- V) e/ L8 n4 k    img = img.crop((left_margin, bottom_margin, right_margin, top_margin))
    , P5 e9 V: B0 ~7 Y% `: z: i2 ?. G7 R8 a1 \
        # 相同预处理的方法
    % x4 {% x  @6 q0 R7 [$ R& L    # 归一化
    ; ?7 o" I; J; p) y* B" h% l3 f    img = np.array(img) / 255* U% X0 v( @6 d1 U! C6 ^
        mean = np.array([0.485, 0.456, 0.406])
    + M) B( U  B7 `$ g6 E& x& k# m! ^    std = np.array([0.229, 0.224, 0.225])) p9 `! f! m: a, X) W# g
        img = (img - mean) / std
    5 u* m7 c4 q& P+ A. L( F% p: n: c- t: S, j
        # 注意颜色通道和位置
    . Y2 {0 |+ {5 s' A7 }! K    img = img.transpose((2, 0, 1))( E9 h3 R0 \9 g: o' e

    ' H" e, w% [4 k    return img
    ( S7 i7 |3 d1 B, H5 O' z
    $ L0 p+ a& X, j1 x0 D; |  @( G0 H' Ndef imshow(image, ax = None, title = None):& f$ k+ |' e' E: h6 b! O7 l0 o
        """展示数据"""
    5 i) Y  L4 y  d1 H4 K0 @    if ax is None:3 ?6 K( u9 L  k8 y
            fig, ax = plt.subplots(), P& x: t$ G. N/ q- x( T  N
    - `( K1 L/ v4 X' C: L+ g2 s
        # 颜色通道进行还原$ V  |1 Z1 a* g# t
        image = np.array(image).transpose((1, 2, 0))+ ?9 x- L. s9 ^

    ' S* W, N2 |* G' }    # 预处理还原8 h! R' j; Q( R) A* h6 i+ B! F4 o
        mean = np.array([0.485, 0.456, 0.406])) n  h7 T3 c8 c, m% R# Q+ t
        std = np.array([0.229, 0.224, 0.225])
    ' b( |& w# `; I- q+ _" E9 `0 n    image = std * image + mean
    * {% l# s! B& A' p: k    image = np.clip(image, 0, 1)# p8 e$ L* o3 B( j0 q" @* h/ ~

    ! K% J& k3 a! x5 F+ B1 N    ax.imshow(image)
    . I  H5 m' e, k1 I" q+ O9 l    ax.set_title(title)
    . t9 Z5 _* f8 [) i9 M8 g0 S( R7 O/ T6 x# o
        return ax) U4 t( k& m+ v. b( T+ s! y
    1 Y5 a# g3 \; T4 V3 Q+ M( o
    image_path = r'./flower_data/valid/3/image_06621.jpg'
    0 y, t2 i9 S: Bimg = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理3 M( ~1 d  b( j- {
    imshow(img)( A; n. d' ?& f8 G. @3 G0 x" ^
    ( Z( j1 E8 ?4 I9 Z$ u  h
    1
    % [2 r# u3 j% s! P  X2 z) k" [9 H: J2* l1 E) t2 X: F  ~' o8 Y; s
    3% C; ]" F/ N& h3 M+ r
    4; W: S1 q1 g7 i0 ^* e) Y% F
    5
    + Z: z, f/ E5 b8 V+ W$ h2 {6! Z7 W4 M! R/ [  ?! \# d
    7
    - `* I' d9 H1 R# ?9 b/ B1 @4 x8 E8
    3 h, I% u% x0 W  r7 a9+ J4 m0 h$ P* T+ O3 O8 k6 X+ w
    10) [8 w, @  J* F$ [6 y% v' ]
    11% _0 ]0 z' D8 O' N5 H$ l1 i9 Y
    12
    $ s0 \0 a9 f) }/ U13
    - `; i8 y" Q8 C  C, r. ^  f* o3 c147 y2 q/ Q3 r6 |$ F' y9 Q$ \9 {
    15
    7 T. L: b1 N. Q( g' j16( N6 c8 }" A9 H
    17
      M, ^. e' m  D$ o3 W) R& B* @5 f18' I9 u1 S2 O7 W/ Z/ k
    19
    ( g: O6 Q5 M# B3 V. s" A# U' R200 F7 [3 ~, |7 E- w+ {
    21
    9 e# a2 Q* t# u5 h22
    ) @( _8 f; s  B6 C/ p' H23( `, }, l; z7 @6 O
    243 N* J. n8 p6 T/ m" u5 [  ~
    25
    & t0 l& p7 S5 i# ?. r26/ e) M9 l6 f, u  h* @
    27
    ; E% [! D% s6 E28
    2 K1 d' y0 w2 Y7 R, Q; v2 H3 n& Z295 ]2 Q# _4 N5 j0 k/ n
    300 W$ R5 I/ b* r: H
    31( h8 g& c" ]% H% x( l
    32
    # M, ~2 e' s5 Z33
    5 B4 u7 |( R- b$ y34
    6 s- ]4 C, ~5 L3 ]0 k35
    9 N1 z3 ^( T  ~5 ~36) g' ?( C8 w- f, |5 E( i3 c2 _
    371 h  k7 F1 a+ @! c
    38* G. G. W5 A+ Q  [3 d
    39
    - W$ h8 {. m3 ^2 N40! k8 D* b/ a% Y  g" q& W
    41
    3 |- |* P4 P  C/ E. w& D1 i42
    2 E5 L9 ~9 a: e' w0 F7 J431 O4 u$ S: ?  I
    44) ?% Y! ?5 v# o& K% ]6 z4 q
    45# f3 u) a$ L) m0 f
    46  a% i6 I$ E" G9 L* |& ~+ X: x
    47
    4 W; Z$ \6 @+ [' |- c48' E) i' h/ F( i7 Q8 R# A5 d+ ^. k
    49
    7 X, B* T# p/ B) ^% @& q% M, O- r50' K5 I4 L! X2 i( g" h( q
    51
    8 n3 G7 V0 o0 G+ h52
    & |, U6 B# D& ^( g53
    4 k8 d  K7 b( [, L* ]54  K/ t$ P# {# z' \& U2 n0 }2 H
    <AxesSubplot:>
    # {4 M) L7 H1 K. _7 K1
    5 H9 k# ?+ C, u7 ?) q+ A( z6 V2 s, N& r
    上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确5 }9 N5 o4 Y! b1 d3 m
    - z0 m" ~1 t2 D6 x- a
    img.shape
    % k) n0 n+ M) Y, z6 x; u' F1& E3 \. l* f! |+ M% G
    (3, 224, 224)9 d) b( b0 e1 v& u8 h( `, D2 |
    1  _  |: e; j0 X0 C3 |* A9 E4 F  B! O
    证明了通道提前了,而且大小没改变
    * l7 ^/ w0 k3 }# W1 Z1 ]2 o( n6 K, L1 g7 j0 M
    9. 推理" h5 T- I# u& e5 E
    img.shape
    0 y! {) _' p" D2 f2 I( ]/ ^4 }7 |4 T  B" l. j* g0 {
    # 得到一个batch的测试数据, u) ~  m& n6 c9 C; [
    dataiter = iter(dataloaders['valid'])
    0 }2 u/ J; {; D( Bimages, labels = dataiter.next()% _4 p  _; U6 h

    + o7 n, @" p+ omodel_ft.eval()6 K' ^  A4 ?' P5 @1 N$ v) Y: M
    : b0 b& Y& |1 B0 Z, i' e
    if train_on_gpu:
    5 S# V. U4 S/ j/ Q, o1 k    # 前向传播跑一次会得到output
    2 Z3 [6 T& o7 }7 R/ K    output = model_ft(images.cuda())# ?7 ~3 q, y3 D/ ~& J3 r* c
    else:; |7 ^" p1 f5 W4 `
        output = model_ft(images)! g* @7 D8 b+ e" R, W1 z* t
    8 d! m  E; T3 C( j' r
    # batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值
    - @8 Y4 s! }. b( K2 J% @0 U9 N6 y: eoutput.shape3 e( U% p) a* v

    ) s9 ^0 q$ T) j5 S6 ~8 j1# D/ A; w* ~) X$ J! m
    25 d' @% e5 K/ x* t  s4 u5 j
    3
    $ E( N6 X1 O0 w, _" j+ n4
    % [; h; f9 d1 K# e& t5
    + }6 A1 E2 s) ~3 g( }/ K- s60 w. B, ^% D9 q
    71 O/ O9 P( T" J0 F. U0 p
    8
    % T6 F. E, m) x/ B$ Z) k9- ^; y, w! T, c- J; J% a2 s6 N# T
    10
    9 j- Y2 Y! S  @8 K. y. h) X& O111 K, Z* Z; b8 x: y0 x! _4 X0 j
    12  }* l5 W, |! C* J3 K6 R
    13
    - j- ^- z7 j3 e) _141 q4 j# ]4 x+ E7 [  l# t
    15
    - ~) C% G9 I# l! t166 B4 M; t$ t. W! r' X1 ]
    torch.Size([8, 102])( `. g- A7 s7 z! x) s# t
    13 V, S7 E: ]2 o
    9.1 计算得到最大概率0 \( y( k; p3 }# g' u- [
    _, preds_tensor = torch.max(output, 1)# B5 ~% C9 i9 L' G  _
    ) p. p1 b, c! D+ a9 _+ b# [
    preds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量
    1 Z% n1 y4 l7 T4 x) A1
    0 L; Q. X& z* K. A; V2  [* }, ?* w$ G( d! d
    3
    0 y. s; E  |, }$ D9.2 展示预测结果
    $ {) }0 r  a- s/ D( I, \: |fig = plt.figure(figsize = (20, 20))
    4 X  ]+ B1 w6 f# L. R7 V& Rcolumns = 4
    " S% T9 ^9 g! A" e$ v4 L, Grows = 2- A0 ^' l/ K% [: X

    - f. b1 w7 o4 pfor idx in range(columns * rows):
    3 K! Y2 n8 c0 z+ d# x- h    ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])
    + m1 q/ f3 l* w    plt.imshow(im_convert(images[idx]))( Y2 C5 N) s1 b- M
        ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), . p  ~* q5 j5 b  r- O. s0 O
                    color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))
    % i) n6 [! v, `* _; I/ y0 gplt.show()" |/ M3 W+ N& e* h( d/ M
    # 绿色的表示预测是对的,红色表示预测错了
    " ?  `+ s, ]# G( u: v0 f  x17 \. D8 x2 F2 }3 O2 v1 ?
    2
    4 s7 v% Y+ ?) V3 w1 A; D' R3
    5 j. y. a3 D; F8 e" O; k3 y% h/ n- t4
    , y0 V3 R. I7 q0 |5( q3 ~, D$ `$ |& t+ v+ {7 t+ F
    6/ o  \% A, q7 v4 a, D" k6 C
    7
    * V( N4 `/ X' d3 e8
    8 N0 |% ]4 `$ C9! k3 i/ d) G% k+ n& m
    10
    6 w% K6 m. M; \11
    - D9 w8 K, e: i1 e! E& r3 u( U2 `( H" p: n& J. c5 P# o

    , b$ t0 p& O; ^; z4 p+ C' U7 p3 f. K2 V. Z! N$ V' \$ a- b5 B0 s7 i
    ————————————————9 w" r& W$ H: i0 a1 W' ?
    版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。( c5 v/ o: t% G4 g8 u2 K
    原文链接:https://blog.csdn.net/LeungSr/article/details/126747940
    # I' f9 A) p' e! k% }
    ' W8 C: ~0 P7 y0 f0 J' x  ^5 i* K* s0 L) W. D3 g$ z
    zan
    转播转播0 分享淘帖0 分享分享0 收藏收藏0 支持支持0 反对反对0 微信微信
    您需要登录后才可以回帖 登录 | 注册地址

    qq
    收缩
    • 电话咨询

    • 04714969085
    fastpost

    关于我们| 联系我们| 诚征英才| 对外合作| 产品服务| QQ

    手机版|Archiver| |繁體中文 手机客户端  

    蒙公网安备 15010502000194号

    Powered by Discuz! X2.5   © 2001-2013 数学建模网-数学中国 ( 蒙ICP备14002410号-3 蒙BBS备-0002号 )     论坛法律顾问:王兆丰

    GMT+8, 2026-8-2 02:41 , Processed in 1.892582 second(s), 52 queries .

    回顶部