QQ登录

只需要一步,快速开始

 注册地址  找回密码
查看: 2780|回复: 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)实战案例
    & J- y# y3 X) M  \* Y/ T- b+ M8 p& d, a% n
    文章目录
    . N3 s. J$ x. I% G卷积网络实战 对花进行分类2 h. l8 j) @9 m7 u& Y* H
    数据预处理部分  L& u# z; d" m$ x5 W) n
    网络模块设置
    ' k5 d( e4 S  n( B7 M! K' p: D网络模型的保存与测试
    6 Q3 S! a/ I* A" H1 C) B数据下载:# t- |3 Y# T& O
    1. 导入工具包1 w. C! T# p. E( ?. P! t7 e4 z
    2. 数据预处理与操作
    & o! O. m. o  t: o5 q: e3. 制作好数据源
    $ E9 Q1 W7 d5 ~8 F4 \读取标签对应的实际名字
      W$ U, X% H2 d) D3 w4.展示一下数据
    ) K1 ]- ]& i) H+ C5. 加载models提供的模型,并直接用训练好的权重做初始化参数
    ! Q. S3 d9 J1 ?3 [6.初始化模型架构6 W9 |' k% w3 t+ a' X( _
    7. 设置需要训练的参数
    & [  v0 {& f$ [6 A+ F7. 训练与预测( v3 z% ^$ f& {# n' o0 e+ A& ^6 R# A
    7.1 优化器设置
    " c- A4 t6 L  h- h7.2 开始训练模型
    1 E1 ?5 X7 _3 T; c7.3 训练所有层% }9 A5 B8 a: K; \
    开始训练
    & D* w& @2 a( H8. 加载已经训练的模型1 N& J/ T9 P7 o8 p
    9. 推理  S4 d, @5 T# n! O1 ]
    9.1 计算得到最大概率
    : U, p, d) u; r" x  A9.2 展示预测结果7 i0 \) Q4 n' ~. S6 z0 i
    写在最后
    4 }' z4 N& b* o' `$ E; M3 q: F卷积网络实战 对花进行分类
    7 d& k2 b) X3 O3 ?" P5 @5 C: P- d0 ]7 s本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能
    5 t( ]. h6 g1 ?2 ^1 Y( i0 }! `, _- A0 W, S, U
    在文件夹中有102种花,我们主要要对这些花进行分类任务6 Q# f* M. J% x5 J; F. \0 J+ C# K5 j
    文件夹结构9 g; [7 A& N3 w! f! U, n9 \
    ) l% z2 ]$ K3 F
    flower_data( N: F- H* \8 g6 \+ e2 \# V2 f
    0 f0 a) u- w& O7 P
    train
    . ?+ F) B4 H0 D7 ?& p' V, J! v& K# A. F6 I) m: B
    1(类别)+ ]  j2 n1 O- P! W4 w
    2+ ^9 p: u! P; P0 {% M, R/ d
    xxx.png / xxx.jpg
    # s( l: W6 E8 e7 K. |valid2 l6 _$ {1 ~5 s; h& b
    - O, y0 @7 O( \1 Q' Y; N
    主要分为以下几个大模块9 `1 W9 g+ \1 a

    ) j- V: P% ]! U" [" t2 T0 m2 W! @% l数据预处理部分
    % M! D5 [; M+ z' D' [$ g( X数据增强
    9 x' F6 A1 }" ?! w数据预处理
    ' R+ i. o6 {7 j: |8 N/ a网络模块设置& f8 j5 v) M2 T- N* h
    加载预训练模型,直接调用torchVision的经典网络架构# A6 v) R! a0 G: E0 K! \
    因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务% Z6 |5 N' }3 Z. B% x3 j( t
    网络模型的保存与测试
    3 _+ f( I/ @. d" ]8 g: I& w模型保存可以带有选择性' g% U. U- d+ q% Q
    数据下载:
    ; q+ U8 D% q  Z5 ]& ]7 P5 [3 a% yhttps://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset
    + z! |  f# ], [1 P, D8 @1 K, E9 E/ H' ~/ x7 A" ?
    改一下文件名,然后将它放到同一根目录就可以了, x- ^, U' v3 e, L( h

    6 n9 h$ G. y3 n: \下面是我的数据根目录& [& _) W! W. O( d, A
    $ v; l% K. w6 e. b  ?% b6 X1 t

    & l+ ^- [# c9 h5 N6 f) |* g1. 导入工具包
    ! n! j' L6 W  a8 {: N; t6 himport os
    " @  J4 f, i& x( k8 V$ _import matplotlib.pyplot as plt
    , ?6 p! E/ b# V* N, {# 内嵌入绘图简去show的句柄
    ( a9 D/ ~$ f7 }& M2 y& ?) f%matplotlib inline
    ' C6 y& {+ U# [: `4 z4 y# ]4 Jimport numpy as np2 g/ }- Q5 ]: E- Y
    import torch6 x# B# n4 p  q' s+ @8 q
    from torch import nn
    % B8 G# L: G% ?( e+ G) f! U0 ^  B; C- r: L) {4 [  U7 `
    import torch.optim as optim
    $ u' P9 }3 S' D2 L4 G: Vimport torchvision
    ( X  }3 Q5 d9 Y6 ~% k' Ifrom torchvision import transforms, models, datasets! _. t5 e6 S7 s/ @4 O: p/ k* C
    : Z2 ]  u/ I9 l' y& E1 M
    import imageio
    ! O4 c0 a& {+ r* r/ Himport time5 D% K0 w7 Q3 P2 E) h- V6 W0 k
    import warnings! X2 `9 Q+ ?" ]3 ^* s: b$ I3 l
    import random$ S" f$ y& R3 n7 [# A3 Y5 w
    import sys: r" a' |9 l' ^" K* L3 p
    import copy
    - V0 }, R2 ]4 r0 v9 _9 y% J% ]* Jimport json
      W- f$ C% [+ E3 Q( |- tfrom PIL import Image; l+ s+ o. e. E
    # V- [. l7 @8 f3 c
    ( [- p; E/ C; H7 w
    1
    , p% x/ V5 K0 Y% d% f- n) W" R21 L( u  c' H& l* K+ ^
    3
    ! w  i. Z6 z* |0 o. @4
    # J, P0 H* b6 A' E5+ F8 W( n% C  n' G$ Y( `
    63 L: r& e1 i1 h6 q3 b3 y' E, Z
    79 ?  {( g" @: |& \% s" t
    82 W4 Q# n- @9 [( C
    9/ R. T) N; [$ u+ }
    10
    9 r; k3 r. R, N0 \0 F9 x. _' D117 Z/ n, E8 x* f1 e
    12
    ( _$ X8 k8 L1 t$ k( ]6 j5 @13
    / q3 W: G& \! `* N" b& X14
    8 ]  Y  X! g# T& ]153 K) D$ }3 n6 O& ^) ?  `
    16
    & @% z# w9 C! O: O4 ?2 m17
    - m9 i+ m5 f9 B1 F  H9 B18
    " {' L; H% S5 q19
    / Y# Y- {) }: |3 }/ x206 ^% K0 ^3 V* t4 t
    218 C( I' C1 ?# i& Y7 B5 \& V
    2. 数据预处理与操作
    & b, C) K0 [5 ?, l* R0 X* _#路径设置
    8 G0 i7 B' S6 m# W1 A1 Gdata_dir = './flower_data/' # 当前文件夹下的flowerdata目录
    " ]# R: _  N' L! q' P& d1 etrain_dir = data_dir + '/train'
    5 p6 e  [: K* Vvalid_dir = data_dir + '/valid'
    ) o- |# `) g, N' i. a" U1$ l7 r/ ^' w# l" G5 }  d
    2! }7 X& Z; f; B
    3% y' l9 c1 S, X+ @* v
    49 ?) ~% J9 k  @/ d8 q5 Z& H
    python目录点杠的组合与区别
    5 P5 q, G/ l: N* x$ U注: 里面注明了点杠和斜杠的操作9 u: E2 ^+ _/ r6 D- [

    7 s# Q9 Q* q* F- W: @% }/ b1 x3. 制作好数据源5 Y6 C1 ~) T$ o9 T5 g' A
    data_transforms中制定了所有图像预处理的操作
    ' K, S" {& }* I( v5 WImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片: ^9 s, K# s& i* V* E  h1 y
    data_transforms = {
    # s2 S! y, L+ I6 y; V    # 分成两部分,一部分是训练3 P% z# Y# L9 q4 @
        'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间! U7 x, u4 H/ S, d) d
                                     transforms.CenterCrop(224), # 从中心处开始裁剪% X0 {; i5 j7 g, e. f" c& y0 f4 R
                                     # 以某个随机的概率决定是否翻转 55开
      ^7 g: ?4 t" s* `3 i% z                                 transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转
    * M; b1 P! w) d& L3 F" [                                 transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转4 [3 Y9 p5 J4 w1 j5 n& v6 s( T
                                     # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相
    , o' Q9 I6 u. R1 w                                 transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),9 o; w- o$ G* ]( a9 g
                                     transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB
    6 e. |7 w# |% h$ Y                                 # 灰度图转换以后也是三个通道,但是只是RGB是一样的5 f2 p* K2 m7 g
                                     transforms.ToTensor(),) M# t# i7 O. i3 E, z7 g
                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差
    $ e7 R1 g/ G( \3 y                                ]),
    2 F6 b& s* G: r) m    # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
    $ c9 B; _. b4 B* h8 Z    'valid': transforms.Compose([transforms.Resize(256),8 F9 W0 Q" \4 v+ ~
                                     transforms.CenterCrop(224),
    ' X5 Y& n: W" p- t7 j& E                                 transforms.ToTensor(),
    , N6 M1 f/ p$ ]0 ~( o, B                                 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同1 f, ?; H/ L! N) j9 r" n8 r  [
                                    ]),) |9 W% u; a5 ~  g/ g
    }
    3 h9 {  t5 v6 l  {5 P
    - t" _/ |7 n( n7 M8 B' F/ O1) u$ S# W# K" I! ]) Y
    23 ~! S8 A, P/ a' C
    3
    2 _; B  T3 A2 O3 s4/ t" O0 S, o# g" E: s8 s) W% \
    5' B8 v% C  n# _4 \
    6
    ' p% E7 ~* `4 c% [7 F7
    " F0 B9 k; g6 K+ x5 K4 d8
    8 p, x  x* N, p/ M/ B# B( G9
    / ]1 @  a+ J, q( C4 u10
    : [0 M, r* L6 X+ d  }11
    - _6 A# k! ]! T' h- k12
    / o# b" t6 C9 U- T4 U13
    4 T) C! r1 f" T4 h: ~- v* X142 s9 h/ l6 V# R( u
    15
    7 e; v* t0 I# l5 L16
    ' s' H$ m& M6 h9 X17
    ! A# X: C9 J+ X) L; \% |" _188 K8 W9 V! V  H+ s7 s- f2 j& W
    19
    . c) t& J8 N0 Q, n5 g$ c1 c20
    ) }5 p5 g# A$ [$ }21* h: Y" u5 l; V: ^4 u3 R% U# V
    batch_size = 8
    & C) g9 a% ^& t6 e- `4 [9 ?image_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}! Q, @$ A4 I+ i$ Z- H
    dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}
    ! n: s! v1 Q; B1 vdataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']}
    . R4 s) L; r3 g! Cclass_names = image_datasets['train'].classes
    5 A5 W/ W/ ?3 b6 U8 x4 q% M
    5 r! g8 O2 }8 e" l. Z7 r* K#查看数据集合
    # B) ~- |0 R! U" a$ J4 J0 z4 yimage_datasets
    - K+ h! ?$ _, f% ~2 S. L( a3 K
    " W1 K1 \1 s" g7 ]4 @1) E1 C$ ^8 _) g8 C, f. Q, t/ k. o
    2" c! p4 V$ q0 q% u  V# `0 f
    3
    + @$ A( Q3 o. \1 {4/ w2 `+ E5 ?6 I% D0 X& {7 R7 ?; i
    5
    4 {6 V9 y0 W& k/ t% ]6
    8 s0 n; f8 z* i5 R/ _9 m6 ~7
    $ _4 ~  s' R! v+ T0 a8
    7 G( b+ n4 N" v6 y' v9
    8 ]+ a  ~5 R/ F2 `{'train': Dataset ImageFolder
    $ j) u8 ]  ^6 x5 v' f4 B     Number of datapoints: 6552" l7 R- {2 b( {8 }* E) O* r
         Root location: ./flower_data/train
    5 Q; x+ t0 z# o6 ?     StandardTransform3 f0 U+ Y# V) g1 q# g5 @
    Transform: Compose(
    2 s7 v4 Q/ g4 t  j$ e                RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)
    ; k2 R# _! }) T, n5 R6 ~                CenterCrop(size=(224, 224))9 Q7 @; P7 m8 j: X7 a2 v- Y
                    RandomHorizontalFlip(p=0.5)
    ! u6 ~+ ]: {0 N& ?/ F& m0 o                RandomVerticalFlip(p=0.5)( @$ r& [8 D. A2 h# c% R
                    ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
    # Q  ~5 ?" y( D% p8 I                RandomGrayscale(p=0.025)
    ; I* \9 b4 d" R. S- `0 G9 ?                ToTensor(). J- b  S) W: d* I; p0 u
                    Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])$ @9 F% d6 M; k5 }5 M. o8 S- n
                ),6 ]8 z# z; J& u4 _8 i5 O! J7 ^
    'valid': Dataset ImageFolder
    : o  x5 O# a' F- {) t* a' X     Number of datapoints: 818
    2 T4 Y3 S+ d9 L# b( w     Root location: ./flower_data/valid) \; ?+ Q$ X% A( O% L4 p6 A
         StandardTransform( H7 b9 _4 X" V8 ]3 l
    Transform: Compose(. }" Q/ L. M& ?7 n6 z. c
                    Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)) s! ]8 N2 C" m/ `+ b. h
                    CenterCrop(size=(224, 224))
    + f& u5 [* H) P% [                ToTensor()
    - y1 y$ m' u2 ]" f# G                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])( n7 r+ Q) B, E% ]
                )}
    7 Q; w, l" m+ z7 U
    9 q# e7 t' L+ L# y! d1# r# E0 ]# m! z' S2 }
    2
    ' F/ M1 T: W3 J1 S3( U0 y3 u1 x0 }, g# f
    48 b3 z5 M5 X5 Q& j; l
    5
    . R6 d9 P' w* O% @* t" P5 C, Z# n6% K; {) g6 P2 n6 e) \8 U0 B
    7% P7 C. F$ Q. P
    8
    * _- A; F/ J& S; ^; l3 n9) P7 u* i+ ?# f  m% U' ]
    10
    ' ]2 w- O7 @  x' E113 C7 u' M1 E; f3 r& G! n
    12" h, k; M; r! y$ z6 V7 Z
    137 O+ h8 q4 i% j% M7 D
    144 B; @% z- f) E5 T+ i, h# J
    15
    5 m5 Q1 s) q+ T3 y7 I16
      T0 B# e; Q4 ~4 g; {17
    4 ^% l2 E( y5 H18
    / V, _! e, u9 ~- O19
    , s5 j4 `7 ?% }6 U4 g20
    ! S" b9 I' I7 g* {+ I213 x" Y$ Q/ i) m$ ^  h
    22
    - r: ~6 ?4 C8 |; t' C* a0 ]23
    % l! v" f" T4 K9 Z5 @9 e' P3 c24
    + b" x7 U  L: R6 V/ W/ w# 验证一下数据是否已经被处理完毕
    - G/ t  Q8 ^9 ^* d5 Xdataloaders/ ]3 H: E: x. D: S
    1' q# `: U) K- t! F* o
    2
    ) T% f; N- `( X{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,; A% j: M9 m! \3 ~4 C: \
    'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}
    & Z4 K( r& U2 J6 S5 U6 o1
    " _+ P9 k5 Z8 u2
      B) w, s, R+ T8 u5 pdataset_sizes6 W$ \/ b  |3 V( j$ M
    1
    + r% @5 p* z; W" B# F" a( o{'train': 6552, 'valid': 818}
    0 g1 K, i& y5 o1
    , {! Q8 u* c' ?) u( t5 }读取标签对应的实际名字
    ; T% w; D! r/ w& Z. _, v5 m- p% P使用同一目录下的json文件,反向映射出花对应的名字
    + ?' X/ |3 K  q( f: o! `, \  K( A: U7 P& I+ p/ k: K& n% J
    with open('./flower_data/cat_to_name.json', 'r') as f:6 Z0 c. D  z8 [1 N# T. b
        cat_to_name = json.load(f)
    , S# |& _3 O* S: ^8 K1' t* u" M+ X. p( i% P) _
    2
    ; v0 _, r/ h3 G% z0 wcat_to_name
    # d1 b5 s' O: \1 n, ~1
    9 S) G/ B+ o. m; u2 |{'21': 'fire lily',! ]' x& ?: _4 W+ W& J2 U
    '3': 'canterbury bells',& \7 x, [/ M3 q  l6 s: v
    '45': 'bolero deep blue',
    : o& I7 M2 m: Q$ i0 b6 C '1': 'pink primrose',& K0 z' I: M! g0 a
    '34': 'mexican aster',& B/ C6 j2 ]7 b: j  [; P. @# J
    '27': 'prince of wales feathers',$ }: G8 }" q+ M
    '7': 'moon orchid',- G2 i" D+ T- n/ v3 i6 }  W
    '16': 'globe-flower',
    / b% f$ N# w( y" B' f) Z; r9 X4 \ '25': 'grape hyacinth',# h3 Q6 [- f+ m" K
    '26': 'corn poppy',
    6 Y  ~/ k6 O5 E' G '79': 'toad lily',4 Z3 ~- n! ^" M, N
    '39': 'siam tulip',3 g" ?3 [% y3 P( p. d0 F0 x
    '24': 'red ginger',
    ( G! f- i! D" ]! r  G' W. T '67': 'spring crocus',: J6 @; G6 k! u: {9 `+ b
    '35': 'alpine sea holly',
    / N9 b/ c- T$ ]/ F9 P( s! r '32': 'garden phlox',
    ; `2 J! p- p* i' k5 B '10': 'globe thistle',$ V8 a! o$ t! C  w# i4 X
    '6': 'tiger lily',
    % c7 _5 r% c& D( j '93': 'ball moss',
    , A$ O9 Q( p# s- J7 e% ~. K# d8 h '33': 'love in the mist',7 A5 g8 D( G3 X& N  {
    '9': 'monkshood',
    - b6 T4 v6 c% p3 R '102': 'blackberry lily',
    # u8 m& Y8 t8 e1 T+ U. \ '14': 'spear thistle',
    6 e, l0 O# H) C6 q' |# u+ Q' _ '19': 'balloon flower',
    , e1 L  a# J+ c5 h '100': 'blanket flower',* k' O* C$ \/ O
    '13': 'king protea',+ X! {3 G# Q6 f
    '49': 'oxeye daisy',9 ~, _. R% [) w
    '15': 'yellow iris',
    * P- v( h: R, o" _* M '61': 'cautleya spicata',
    7 D! z% C$ ]0 a$ ~" \; Q' g4 \3 W '31': 'carnation',, N) X) m  C& A3 R' B- ]
    '64': 'silverbush'," a. M# r0 b1 ~0 Y# d5 ~5 X/ {
    '68': 'bearded iris',8 F# d. w- [' g7 ]
    '63': 'black-eyed susan',
    ; Z) j3 g9 \3 _3 e+ a '69': 'windflower',' z4 F8 |$ G& l+ a; Z
    '62': 'japanese anemone',
    ' A- i6 p6 y( V* w1 l5 J: Q0 [! H# | '20': 'giant white arum lily',
    ! c: o8 c6 m0 s( \# R. a '38': 'great masterwort',! \  |, p# x7 w- [
    '4': 'sweet pea',% R1 z1 S( R1 Y" Z
    '86': 'tree mallow',
    - j! G4 o) @6 w  r# S& [ '101': 'trumpet creeper',; R  }0 B' i  W/ K  j& ~5 l3 q
    '42': 'daffodil',' J) L0 O! s7 U/ k+ v1 F! Y2 D
    '22': 'pincushion flower',* F5 s$ L+ `6 N# L
    '2': 'hard-leaved pocket orchid',
    7 Z& ^  B: S+ p0 z$ n. N '54': 'sunflower',
    , O# p( p) l% j5 S. ^ '66': 'osteospermum'," r9 i# Z8 w8 z# R( J1 z& `
    '70': 'tree poppy',
    4 f0 O# `" F# V; }2 H" U '85': 'desert-rose',8 ?8 b& S0 G! ?
    '99': 'bromelia',$ k, y' Z- }6 r0 S; e5 ^% u
    '87': 'magnolia',
    $ [* `1 q3 E* }2 i '5': 'english marigold',
    % b2 w3 ^. l0 B2 G: ?7 S8 _ '92': 'bee balm',
    ; }6 y# [/ J5 Q# [# A6 G  d '28': 'stemless gentian',( }$ J0 y) U, l# Y
    '97': 'mallow',3 X1 ~6 `2 z" m9 f3 j
    '57': 'gaura',
    5 z+ }: ~0 G9 _+ ]/ B '40': 'lenten rose',# N3 s9 Z4 Q- L
    '47': 'marigold',  H! n# ]; u& c; f8 q. S5 T
    '59': 'orange dahlia',
    7 A/ ]7 G1 n6 k. P4 T '48': 'buttercup',
    5 e" }( T9 g9 d! S6 E* c '55': 'pelargonium',
    0 F$ C4 z9 u7 E '36': 'ruby-lipped cattleya',
    ( @8 ?% t& P: F7 Q, `1 Z9 B4 K '91': 'hippeastrum',
    ) A+ x  h. q0 O) {4 l '29': 'artichoke',4 h6 p  @! X5 u  X* n. M% `( H
    '71': 'gazania',
    ! j+ |2 ]! d7 F5 r '90': 'canna lily',
    " [' b6 r) I( J, W2 \8 J '18': 'peruvian lily',
    ) U3 x. o% j6 r '98': 'mexican petunia',
    * K' h) d' y4 G- J9 B* I* q '8': 'bird of paradise',
    & }* D0 B: X' h. ?/ J: R7 {0 l( e '30': 'sweet william',6 l* D6 c3 `# C; I  O& ?6 [+ Q
    '17': 'purple coneflower',
    / r$ n" f9 {1 ]- A '52': 'wild pansy',' F3 |/ P9 g3 S+ Z- f8 z7 @: o3 K7 F
    '84': 'columbine',
    + Q/ J0 Y: b. z, \1 N- n% S '12': "colt's foot",
    # p) q0 N7 F! X: x  a '11': 'snapdragon',9 l' M, I( Q, K
    '96': 'camellia'," a4 H, B' z9 b2 G4 X0 }+ s5 ]8 T
    '23': 'fritillary',
      P9 d! k0 Q7 o '50': 'common dandelion',- |3 B! Z' c5 w- q- a
    '44': 'poinsettia',
    , k0 v, v: G5 u/ n3 [0 Q5 ~ '53': 'primula',$ o" Z! C3 I, s: W% u6 T) ^
    '72': 'azalea',- W$ {8 k3 p/ d
    '65': 'californian poppy',
    ( S( s5 W! f4 B. s '80': 'anthurium',0 S. m: \- z8 a
    '76': 'morning glory',
    2 Q" j0 z/ l2 m# y# w0 E) F '37': 'cape flower',# n4 W( z' m; @+ d  x
    '56': 'bishop of llandaff',; ^8 m, C* ]  i6 G
    '60': 'pink-yellow dahlia',
    + N! _' y/ N, G' B4 U5 a7 O '82': 'clematis',; r6 h- X" ]3 }) J  A+ p
    '58': 'geranium',
    ( l+ ^- c, W4 w# ^ '75': 'thorn apple',
    ! H" w- z8 v: r. ~- q '41': 'barbeton daisy',' G  Q# F# u$ a
    '95': 'bougainvillea',/ t8 \  {% @5 Z, e, J
    '43': 'sword lily',
    ' A) V0 B5 r* c8 ^7 _, v( f; Z" S2 J '83': 'hibiscus',
    * G- e0 B4 C% r4 K+ R '78': 'lotus lotus',# ?0 J% P  e4 r) r4 T& m4 }
    '88': 'cyclamen',
      P+ V: F% V# z, ~+ A1 b '94': 'foxglove',
    ' o5 }# n+ |8 v4 L( v& T '81': 'frangipani',8 ?8 E( C  C4 |5 E
    '74': 'rose',
    / ~- T5 G! w; F4 K2 C( X '89': 'watercress',
    " t* n) w5 F- _8 [4 F( q2 X '73': 'water lily',: f  n% l. L8 K% o& m* H$ J5 d2 s
    '46': 'wallflower',/ Y1 n6 O; w: K2 v
    '77': 'passion flower',' s; {( |/ e( T8 c7 z2 j
    '51': 'petunia'}
    ( ?% d7 ]7 {: T8 e8 y0 g# h0 k  s1 m$ ]/ [' \. U& j  |9 H1 n
    1
    $ g5 J$ a2 q4 E6 V6 H& J21 Z+ f: v) f3 \6 s+ d
    3- ^/ t' M: N/ j9 Z* r
    46 Z: G' z9 c2 q7 P7 D
    5$ n$ [. }8 f+ u" N2 d" O6 S
    62 F$ a) F1 e* T/ e: X0 }$ l
    7) w) F. G$ U4 C0 e+ @0 d
    8! i; e5 J2 `4 \- x
    9
    / @- l' Q( j; c6 M7 T10" y) a5 h5 E" v3 H
    11
    , b/ @* w6 l$ y/ M/ n6 M12
    ( q# A6 p" V5 s13
    # A, P' g7 ?/ w& {* |14
    % j) M" y8 b8 u) Q/ [15
    ! i! R# I  t: e+ b16* O, O+ {- ?; p' F) E( P, _
    17! ?: A. }8 s& M. V
    18$ k* @$ B, b/ c9 @) x9 O1 X
    19& D; E) q" [# u' }. ^/ w  A5 i$ x
    20
    2 V8 j! z  P# U4 S21
    ' ?4 y- c5 u! Y7 H& _& V: z$ [5 [% q; y22( f/ w+ S& i+ U; Y7 U
    23+ A1 J* k  K/ a' |. J' q
    24, m. `+ O5 a) `( ~( O
    25: r' A- B; l/ Y, V' v
    266 x. A% F  N( v1 p  e6 o
    27
    $ J2 e+ Z" e+ @28
    1 q) S- y+ V: g3 o* _) d29
    ! C6 Y% h8 h# k$ P30  i9 B/ o  P1 [+ j6 e. m
    31
    2 f" f/ k7 |" p! r5 Z+ n! P32
    : B' A7 O, _8 V! M( @* D% m339 |6 c' l) Q6 @5 [: X0 y
    34
    % D, N9 W8 w" p) D1 W% v6 f35
    6 j9 L8 R# q7 A8 w# l# [" g$ B36
    1 b! Q: @; _7 h/ A37/ t8 y( E" L7 [8 G% i) a
    388 j$ Z- g$ w0 D; I& k0 E
    39
    ' W- i$ Z8 {' V/ b; O4 y- K) @40
    & W( l. x. a! @/ z4 t: I41: q$ l( j4 J# a* s4 _
    42
    , R) J0 m$ O3 Y$ o0 @& l  ^# p43/ C9 J0 R3 \) r% ]
    445 s$ ?/ K7 r; A6 L
    45
    8 w' y# D, u, ~3 v. w/ W46
    0 l0 B" B3 Z9 }1 F% f+ l: o47
    3 i4 a% {1 U) ~8 l1 D$ [48$ I6 U% u$ Q8 p. I; W8 A& F0 ?
    49. R- l( ~* }4 J, a0 ]
    50# F, i+ f- I1 k, ~, M6 C! D' i
    51
    " M5 O7 w' v2 m8 s- E52
    ! R- X- ^* _/ D9 c; P536 Q* D, Q) B3 Y  U2 C- J
    54
    / i% l* n% A( ?, `+ O55; N) k8 Q5 e5 N
    56
    ; B2 l- a9 c: @) e4 d+ P57
    ! _9 l# m7 s) `58
    8 k8 c0 g1 |. t2 o" h: v59! j) ]7 [  p5 V1 g* m$ t
    60. l$ F0 m/ u0 p$ i
    61
    ! ]9 k4 {" w' o; b, s62
      Y) w) G8 f- Y+ z63
    8 ?" R* z9 ]( J2 ~8 L646 J, i3 y  [* n+ G& N: b
    65
    1 ~! ]* H+ J! A# ]5 {66
    6 \  H' S6 M$ b% g1 F67
    + F. M" D7 D9 U. }5 \/ R! u68
    1 O# F, ]- x) H69
    5 f6 \7 k! |* k/ z6 f- ]. `70* ?+ ^! l3 Q0 v8 U/ a. m
    718 r% k# v. {1 h5 s/ _4 A
    72
    ' }+ t' K# B& E% `2 \" r7 g73
    - L5 E% G6 s% X- G74' p9 e" ?" O* E1 _) x
    75( E4 @% ^5 ?- N# n' p
    76' m! [3 S# Q1 a) ]& ?$ Y: l
    77
    ( N. o7 X7 w) c0 l' x78
    4 f0 S; t' ^) d) u2 h5 C79. U( c% ]3 F8 V8 O
    80
    4 j& o" `" s9 r7 S' z81
    ) d$ Q0 X% b4 a/ J* a$ \82
    ' n, x8 Q3 U6 ~83
    * @( X1 G- o. _1 H$ u) w2 |84
    5 {6 ?  N% N% f/ E8 B858 |! o* G: N2 Q( ]4 o5 r4 t
    86
    ( q; I) C$ A/ y# y! @5 W- z3 p) G2 b87
    ) o$ {& u5 o( C; A" |- \" \88
    , p; d3 {; C0 O89
      t5 Q  S; g) c4 {  s7 w& w90
    # t* @0 p& L2 m1 z' b91
    ! k' |! b0 M9 n92
    : \1 j) @( F" u: |6 m93, U7 C& V; r; T' G- Z; w) |7 U1 X: a
    94
    - }. m4 e: @5 U7 ?95
    , n  B  A3 w0 `) x5 H$ x9 W96
    ( Z) Y$ R2 \. y97
    7 w+ P7 h2 E! x$ F; ]98
    ' e9 e' |# Y5 k1 p- b99
    1 B1 e8 g$ ?7 \( L- N; }100
      K% q3 R. R4 n2 R; `, O- i2 k6 v101' a9 k( b  s3 G
    102- d, K( L9 Q& P7 u& ~( B
    4.展示一下数据
    3 M) Q  F8 f: c: y% U  V, ddef im_convert(tensor):
    ; k- A# l3 q2 [( I0 x    """数据展示"""
    ; a) B- k/ g4 r8 q) F3 U4 _% K    image = tensor.to("cpu").clone().detach()% ]8 H' R4 g- d3 \
        image = image.numpy().squeeze()
    0 r$ S, T- x: i; l3 ]- G) K    # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图
    4 f7 h- \1 G7 w; q. ^7 H( D1 }    # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)
    ! a, I! Z/ G& h8 h% W# @! F. s, }    image = image.transpose(1, 2, 0)$ |: }3 Z/ u; Y0 R& X& A
        # 反正则化(反标准化)
    " J1 {) S3 u. z& O    image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))% n: X) }/ O3 a0 `, y. R" B9 K

    8 _  R7 O  V6 a% C    # 将图像中小于0 的都换成0,大于的都变成1. R: Z# b3 f% }
        image = image.clip(0, 1)* d: {' S2 U1 H" j7 I! o2 m

      T9 G- k1 s" T    return image* }) J; K- J; n' t& Q8 J
    1
    $ Z9 V2 w8 l7 M5 W1 x2
    : ~0 d. H2 O6 |$ Q36 o, X7 K, W3 p
    4+ J  R* K# ~$ S  R9 V  c1 i
    5- w5 V! `# B  F9 U+ j5 Q
    61 h, ^- ~, t/ b2 j$ |' B
    7
    ! N6 R: _1 r1 W5 U' L8; \& D+ j9 H7 n2 {9 P
    9' u) u  m! L! t* z, O
    10
    6 Q- M3 b1 L! q, ~3 D  c11" L: j9 i0 n. W. e
    12  l- b2 H) S! q3 g) l/ C0 b
    13
    ) ~5 r" z( ^! b) B2 X14
    ! f. I+ i+ s- s* w  H+ D1 b7 r# 使用上面定义好的类进行画图: C5 q: l) |% q/ Y
    fig = plt.figure(figsize = (20, 12))! X! k2 [/ w! b" X
    columns = 4
    $ Y! k8 T6 l% c' x( h; }7 krows = 2
    " I( F) W  G5 C; I5 _' g) W! n' c& Y: j
    # iter迭代器
    % p' e7 w8 j$ G/ m, S3 l# 随便找一个Batch数据进行展示
    1 y6 K: V4 M7 C  r( N: K; jdataiter = iter(dataloaders['valid'])
    + A0 ^2 S) {- Qinputs, classes = dataiter.next()3 p1 S5 B/ O% Q) X6 e, g5 C/ t

    - a7 U: X, Q; m# s% Jfor idx in range(columns * rows):
    ! U6 x' r4 z. b5 d    ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])
    ' R/ [7 ^9 c9 W3 d7 P    # 利用json文件将其对应花的类型打印在图片中
    5 I9 r, N/ d' Q. @  u1 F6 q    ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])
    " v' l0 g4 T4 a$ C    plt.imshow(im_convert(inputs[idx])). A6 K3 D. y- }, ^+ R/ A+ h
    plt.show()  s1 q1 I) D/ r( g! D1 q7 W% `
    " u! c  q! F8 v( G
    1
    0 H( M7 S4 d* l7 K% y0 s5 q25 q; e( c; a& \2 v- F3 M
    3
    0 e9 I) n) y% l' ~+ N. s+ s4
    7 Y. _! k: I! U# t0 c2 X: A. z6 X5
    & l3 B$ k6 N- f9 v" }& w/ j6! {6 y5 a2 C2 e8 \4 ]
    7
    % N8 ~- Z4 i3 H6 B) x! j* w83 @1 c$ a4 u* t
    9  i; u. ~/ W: [( i6 Y
    10
    " c% H, w( l: H/ p% H$ U/ G11
    5 f9 `  C6 x+ k3 G/ ?( E% ~( C12
    ! V1 l7 {6 x* T) D) \4 B13
    $ c( \3 z4 f: U' _" R14
    % s  `1 g) G9 O5 Q) ~0 I: i/ @$ |15
    % N7 g# A4 N: O$ W. c16/ V! ~: {$ P1 d7 z, y! ^% w" c

    % N+ n# ]$ h5 Q) @  {* ]$ z" j: R! ^( @6 |' q. o6 w4 Q. h
    5. 加载models提供的模型,并直接用训练好的权重做初始化参数& W& ]3 r0 o1 w* S. W, G
    model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
    # b" Y% U# v* n& g# h6 t# 主要的图像识别用resnet来做
    $ q8 Q& S; K+ _# N. z3 C# 是否用人家训练好的特征# Q# A$ D$ L9 G8 J$ |! p0 T
    feature_extract = True9 O4 `$ |: y% f# v
    1* H* E' T' p! |
    2
    : W4 V0 ^& ^6 \" _7 R* I1 d; k30 o/ ~/ B, m& f  z9 \
    4
    ' o* e  v9 b' `* D+ j# 是否用GPU进行训练) q0 Y) `* {# T2 V
    train_on_gpu = torch.cuda.is_available()
    & i$ W! S1 X& f6 m8 A* d; d% M, h6 v) U/ M8 H
    if not train_on_gpu:
    + b% z& V" p" o) w  T! f2 r5 M. _$ v    print('CUDA is not available.   Training on CPU ...')$ B* a' w% ~! N% ^* v7 H8 L
    else:+ z4 S  N! X: Z% b6 Z
        print('CUDA is available! Training on GPU ...')$ x& A' Y7 f0 `3 l% M& O

    $ ^0 b" h3 U! V1 N$ h- E0 qdevice = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')
    4 y  \' v: n4 @$ ~7 X- A1 T1/ V2 C9 G# ~( ]  l$ {
    2/ [4 ~% w7 N/ B* [
    3# z7 b, g0 V! c3 \1 [9 ?
    4
    9 E% W  [. S' @- {6 s& a& E" I5# c+ T3 l/ O# R, g- r) S
    6
    / L: @' p" j) B! [; i7
    ; y! v/ B1 W" K: v3 C81 R& V& C9 Y( [& I" Z6 m* J
    9' `3 {3 k8 W0 L
    CUDA is not available.   Training on CPU ...
    & N6 z/ V- w! F. S1
    * b5 ?1 n4 d7 ?& L; Z4 f( n' k8 x# 将一些层定义为false,使其不自动更新- C( }8 M+ x8 @8 l
    def set_parameter_requires_grad(model, feature_extracting):" C) ]8 l9 h( ~+ ]
        if feature_extracting:
    0 T  U9 O2 n# g3 U9 L% s        for param in model.parameters():0 Y, I* J8 A. Y* v
                param.requires_grad = False
    - G4 u, m% ]& V% m+ P: {1
    1 d3 @" `6 |1 g) ^8 m2" F* t7 I3 y% X+ T1 s8 @
    3
    6 e4 g8 u  w" X, _9 ~3 X& T% K4, N6 O/ z. C% S
    54 }0 Q( b* R% r) v4 i' L1 y
    # 打印模型架构告知是怎么一步一步去完成的, Z. U# [$ _1 Q, D& A
    # 主要是为我们提取特征的0 G7 @8 [) Y" x# e
      K% W2 x) j' v( p. r
    model_ft = models.resnet152()
    " {( W$ ~# w+ C1 g6 H4 imodel_ft( X+ L& k2 M  m+ G5 K6 j
    1% _/ I9 `" O! `* D# F
    2) n0 k& m6 Q, v; b
    3
    8 I- _" Z" G8 ?) i+ J* l  S: C4
    7 S* p; ?! q8 q# L8 j5# q2 `4 x' Z8 j( n, a0 o/ k
    ResNet(( f6 N& c2 B4 c" D' t% o7 T
      (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)7 _4 n; ]* ?/ O% F2 _# |
      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    . ]. g" U( b" L  ?4 Q- }0 ^  (relu): ReLU(inplace=True)+ q1 P* K" m2 G2 G$ n! L- e
      (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)" V: r: O/ S: M
      (layer1): Sequential(
    * i% W7 A: L5 r: [1 q    (0): Bottleneck(8 Y9 f( Y0 }( n7 j0 P
          (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)
      C) h0 n$ @1 Z8 q7 w      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    9 d$ K5 P, C$ u+ e0 P% y$ f      (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)  H& c. [$ t, k* t; w' r& _
          (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    3 }9 j9 M* k7 }3 b      (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False): j" H4 U7 a5 k& I5 k7 r3 s
          (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)/ V3 L2 `7 Q) o, V2 V
          (relu): ReLU(inplace=True)2 b& M( `( ]5 I! n
          (downsample): Sequential(& C$ |9 M, T) `
            (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
    ( N/ ?7 c  r( y        (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    , f8 `$ N* \& b: O' O      )
    ! I; b8 d8 O8 @" g( k    )
    - v' l3 L" \' Y& ?中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。7 ]. O7 B+ ]1 t" J& y
        (2): Bottleneck(& Y' u# t. ], q7 H0 F7 Z4 _& [
          (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)
    3 S: n  O8 p- @! Z9 g% r% F* I' @      (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)0 M8 q' P9 P6 J# B1 c$ k- P5 n
          (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)8 i4 l6 J/ E/ |% C( l' d9 i6 z& o
          (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)+ T! l7 G4 k1 t9 h: U# u
          (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)8 v, {9 ^- R+ X6 B7 V3 Q) p
          (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    7 F; @0 E; c: Z* S+ E9 J% Y& t) W      (relu): ReLU(inplace=True)8 d% S% j. q3 F
        )
    9 u1 c# ]1 B9 h* m' s  )# d3 X* ~) ^" J1 r
      (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))
    - D, f' i+ Z0 E5 i: p4 u" j3 o  (fc): Linear(in_features=2048, out_features=1000, bias=True)  h; u9 @5 h$ B7 R4 {& o$ o
    )
    / c& W8 F# }7 ?
    $ A& `  A2 n5 O3 b; w9 W1 A& S. O- C1
    ) [; k. Z" k/ K2! X( O$ R- l: F: `' |
    3
    + @/ m; a7 z( Q" Z0 C46 O% |7 x/ n! s, d
    5
    ; V& m6 [; _& X' Q& u& r6
    & o% a; t7 ]: A76 O- e  n/ c& k* |6 s" h
    8
    # f# v( p' A6 _" ?9* B/ |$ D+ `" U2 \5 H: @
    10
    6 t$ J7 Q7 g& Z) K9 ]11' Z8 z% q8 g0 Z% j% f; n! a3 M
    123 h" c8 [. o% N, P
    13
    7 l9 s4 f  P, Y: s7 B: E7 s14
    7 @7 l* R8 i  x- c% X5 s0 e15# i2 I/ }  P% l0 E3 s& i
    16
    # I* ?- v! ~& X, r. C0 W17  T+ |/ u. D9 h" c* z5 j
    18
    % [: C0 s% _9 q0 F2 B0 a  |* e: D19" N! S4 ^& `8 b  [( P! {- q
    20
    8 _$ @, G; t  \/ M: C21
    , K* U, R& a) R' r6 ~7 A3 p22- _) K: Y* |- ]( W+ G  C
    23
    ; E% V6 c5 ~, a! g# i0 r0 Z, I241 @$ j% ]- x+ e( K
    25
    6 j  H+ r$ B$ ]. q26% o' q- c- i" [
    27- Z& ]  t+ @5 A9 n# t0 a  D
    28
    7 z: i, N1 w) v. O, ^% V' j0 ]29
    " _5 _% s* \3 ^! K3 r+ O308 s0 u3 e$ h, l
    31% \  d# p& g9 P3 g) H
    32, T8 z0 ]7 V  [1 ?
    33$ Q* {$ X, W; I+ ]% B
    最后是1000分类,2048输入,分为1000个分类
    0 j( Z. z; v/ b. i4 H1 D9 H而我们需要将我们的任务进行调整,将1000分类改为102输出
    0 A( E7 s# s) ^( e
    ) X- F2 g) l3 m1 A6.初始化模型架构3 ?3 a. r" ^, a9 @8 z
    步骤如下:
    ) B6 E% q! n/ O6 j4 F4 L# g4 `
    $ B/ W8 E- x, p( K将训练好的模型拿过来,并pre_train = True 得到他人的权重参数1 o1 L0 d" n: O: v3 F
    可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)
    3 Z8 b" f3 j7 o; F% {无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数2 @7 ?9 [3 Q+ N) S9 ]
    官方文档链接
    7 p  m: V5 R6 r2 m5 ^7 ^& Z0 yhttps://pytorch.org/vision/stable/models.html
    7 u& _+ I7 T0 W  n2 g) W. B4 O7 E. a) o  h% g
    # 将他人的模型加载进来5 }6 a) u) R; u2 y1 U0 P
    def initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):) x! V# v, b& V# ^6 `/ }. E
        # 选择适合的模型,不同的模型初始化参数不同
    # l" C7 d+ }0 i8 s    model_ft = None
    / Y" f9 ~6 g& ]* r8 N. m    input_size = 01 |. {( p; F7 r# A# ?, Y7 E
    ; ~- }6 I! H2 O6 u* U
        if model_name == "resnet":: R5 ]- f" D1 G* T  I, a
            """; u( q9 Y+ z) B2 l, J) s- P
            Resnet152- Y, q/ u9 [) e0 n
            """
    0 e$ \( |+ E- g# V; {' b' s4 g& v' `6 Y4 @6 R
            # 1. 加载与训练网络
    . l/ H" v- [* X+ t$ x6 g5 E. O        model_ft = models.resnet152(pretrained = use_pretrained)% `" _. m4 ~4 f4 u: u
            # 2. 是否将提取特征的模块冻住,只训练FC层
    5 n. r/ C, A& _8 l) y3 Y! Y- o        set_parameter_requires_grad(model_ft, feature_extract)( O, q; e2 V7 T9 f
            # 3. 获得全连接层输入特征% Q% O& v2 G% x' |7 a0 G
            num_frts = model_ft.fc.in_features
    # p9 S! i3 _. q) g        # 4. 重新加载全连接层,设置输出102
    / @. Q; q% _7 A( }        model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),$ z* W( d9 }! v1 E
                                       nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1
    7 m8 X$ G; x, l/ ]% [6 w6 E        input_size = 224
    1 l8 u/ A( I  Z, Y" @/ E. ^: s# \8 k9 u+ l3 C: [2 n
        elif model_name == "alexnet":+ i) H4 j- d0 C! {6 I- _
            """3 P: R7 C) M3 u/ v! i4 r! R( A
            Alexnet
    " a- Q6 @/ a( [3 [9 V! |        """
    # n" q$ b% p0 ?        model_ft = models.alexnet(pretrained = use_pretrained)
    * B, c3 V5 {4 U3 U        set_parameter_requires_grad(model_ft, feature_extract)4 I' e0 b' y6 s; V1 [
    ! Y4 R, ~% O, P$ p) D
            # 将最后一个特征输出替换 序号为【6】的分类器) t$ Q( B9 M6 }/ u. `, M7 l
            num_frts = model_ft.classifier[6].in_features # 获得FC层输入
    8 q9 k, X% V$ u! e  n3 j( E$ Q        model_ft.classifier[6] = nn.Linear(num_frts, num_classes)% x1 i+ L6 F5 S$ J& v. d  c6 E: b" j
            input_size = 224
    0 m) s: d# F/ J; A" v
    5 H" @& F0 F. C/ ]$ Q    elif model_name == "vgg":# D( T. e! j# k+ ~# ]4 g$ y
            """
    , B3 G+ ~6 a  J0 J1 J$ N0 G! |; i8 L1 i        VGG11_bn  z* E2 b: L$ y/ f. [- n
            """' E' w7 b6 e( r
            model_ft = models.vgg16(pretrained = use_pretrained)) i: @( \% U& o" G+ T
            set_parameter_requires_grad(model_ft, feature_extract)3 k. e- y" m" s5 J: x3 z
            num_frts = model_ft.classifier[6].in_features6 @1 A3 X9 [9 p0 `$ B/ `6 v. ^
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes)* N0 `1 C" F; B
            input_size = 224& Q8 [9 v. R. x  G5 x* z

    8 }/ ~( N& J/ M- T    elif model_name == "squeezenet":
    " s% n4 i  R& T2 u' p- C* m8 ?        """* T; T% [( C( c+ k# O" p* h2 S; J
            Squeezenet. i6 S. W) F- q6 @1 B0 Y
            """& F, P1 q; p4 J0 q# F
            model_ft = models.squeezenet1_0(pretrained = use_pretrained)& d* b8 H% i( [6 v2 D
            set_parameter_requires_grad(model_ft, feature_extract)& n0 j8 L# \+ v2 p# _
            model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
    $ P2 S! t, m6 T, J" |        model_ft.num_classes = num_classes: J" ?* u! z# H  R1 J
            input_size = 224
    : j% ?* r+ A: _. J9 ?1 a; ?$ X% J' q) Z' h' }' O8 O$ a6 Z
        elif model_name == "densenet":. d# ?+ m2 E4 Q/ e
            """
    # @. {. i1 f: U4 P& |        Densenet
    9 Q# E+ `: B4 L: W* r% a, L        """
    " Z6 R7 [% X6 a) E! L: ^        model_ft = models.desenet121(pretrained = use_pretrained)
    ' g5 W  H3 S7 {2 T' \  e        set_parameter_requires_grad(model_ft, feature_extract)
    8 `$ Y  ~! H$ q0 T        num_frts = model_ft.classifier.in_features& G# j" E& J! i
            model_ft.classifier = nn.Linear(num_frts, num_classes)
    0 D: {. ]: i" v& n. G, r        input_size = 2245 G+ z) D" k$ P3 x
    * z2 M2 ^( O: {) d) w& G/ d, n! ^
        elif model_name == "inception":$ m" p1 X) W  a
            """
    * k3 F) r4 C" O6 ~        Inception V3
    ' e) @$ A& W) G# J) u        """" m# y, x) U( q3 e1 p. y1 k
            model_ft = models.inception_V(pretrained = use_pretrained)# H7 m3 ^2 v3 X& E) k
            set_parameter_requires_grad(model_ft, feature_extract)- H5 s0 a2 f5 r
    * P9 d0 L5 B4 O/ z! S  m
            num_frts = model_ft.AuxLogits.fc.in_features2 R4 Z1 Q2 A) s" n
            model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes); ^' ^' M. [, X" i5 b# D7 E
    / `: ]9 s" x1 C1 I) d# G! W
            num_frts = model_ft.fc.in_features
    3 e2 s( k: n1 b! C  G        model_ft.fc = nn.Linear(num_frts, num_classes); _) r" q! a* k/ F  i3 _" j. |
            input_size = 299+ P* W, u0 m' V
    0 e7 V! z* [4 }7 b
        else:, O( W: @3 H6 x/ A
            print("Invalid model name, exiting...")- X/ ]! `, i! ^7 h7 L) W
            exit()
    ! g  Z. z2 [' s& \! P7 Z3 c  r
    6 G3 d2 ~7 A% l. J    return model_ft, input_size0 R  k! A- Y1 k! C" ~

    " H& `- H8 v0 }7 S4 c1' J' R0 r! B3 ~0 N' o5 B3 |
    2, L! f. F4 H% m% L
    3% h2 t' c! g! n; E6 A
    4
    ; c  |9 w0 {2 K& l; m2 J5
    - {: ~1 i  M" J/ j( V& x6/ |- F0 y  n+ b, D6 D- c9 F4 Q, A5 ?- V
    7; H# c. o+ o7 S' d9 s% w- f; H3 X
    8+ v  M9 T9 |& v$ ?8 i' o
    9
    1 ^5 h) L) o9 |0 O: C10
    $ ]: A% d4 E. V  k' q* y11
    , U# x$ c! F) Z) d; C2 T12
    : s! r2 i+ h2 W9 O13. ?, x8 w  a3 n+ _4 |; M) ~& U$ N
    147 n+ a7 h0 R3 b! }* S" f
    15+ }- K! ?/ {7 W( [1 K
    165 f! E! Z. W5 T! S
    17) S/ {2 s4 d8 k' E
    18
    3 D  p' |/ u& [2 }193 w+ }1 y. O1 l7 ]6 v5 [8 N5 v
    20
    2 G# H6 n  z% N! c4 X" O21' z4 _. o/ n! [2 G, Q
    22/ y! R8 {; R& l& O1 w+ S- K1 f' [
    23# R  y" J5 P8 r  X8 @- N
    24- ?2 v) T% `7 x$ t8 y8 O
    25
    . j& U6 B" X! a$ }3 x3 w5 c26
    * w3 @7 A3 L! l272 E- g- a* l  s3 `# K# E4 X- O
    28) z1 `8 b' h% r5 D
    29  z7 t; \2 Y/ v
    30
    " ^6 R* I" I: D$ Z+ }31
      R, `8 X$ z1 x& E+ F* P32% r" n# \' V5 s( r9 ]! }7 z$ v
    33
    5 ^  P7 @6 d9 b% b+ J( L' b34
    & ?: _( k; z2 X/ t35
    9 ~, k8 z, A/ E1 ], q36' x- _1 u* V7 G% m8 X
    37  g% \) A6 a6 N8 p+ C) D
    38& k+ o+ ]# a+ |9 L4 k) G5 S' j; I
    39+ Q7 h9 B8 M7 H& A# o$ _4 H
    401 ^" K' F  }( q3 \
    416 O+ x5 b+ y' n5 e6 c
    42+ N6 s( i: m  x: x0 n- @. Z
    430 _8 P/ p) @: Q; ~' N8 K* \
    44
    0 }  {9 l; s7 v" w0 c3 D9 [  D45
    3 K" n8 Y& P/ N46; M. y$ p. G' \* f. C, s/ U3 @+ ~
    476 q0 k% |' w+ h6 l. A. o+ c
    48
    5 U: ]2 e1 {/ A4 m49
    8 ^# K: g9 ^5 ]/ G# m# \5 D50, o9 q2 z0 D- v5 ]# A6 W. O6 b
    51% r5 Z8 G; {8 v- b3 T+ `0 j$ h9 _
    525 t* n, g; G2 O- L: r* `
    532 S/ @& W% k- w7 Z- |! `
    54
    & K" A" |. [5 ?- {, b9 Y% U, Y55
    $ s2 h  k, B' w6 f" {, s% V9 a56( h' t) |# s; h9 L, L9 i9 c. A
    57: j0 {' t3 E9 S2 |6 ?
    585 y3 l: O0 g8 W9 ]) z0 X
    59" C7 m2 u+ `$ ]# j: }/ i
    60
    / s0 n! e" P  |) r3 ^612 l& v6 \( c! D5 M  e: T& B
    62
    4 X: g; B! D9 V; o; p63
    9 q& p$ o- P- F3 l64
    6 o  O$ A  U& H65
    7 D& \- u" U2 [% Y1 z/ G66
    0 ?5 \7 K! z1 @) V: \( U672 o, v$ I# H/ A$ P
    68
    2 J/ f2 N' P7 h: G69
    ! s2 w' P7 s4 D9 q% c# p70
    7 Z" J" n; C* |8 A1 X* W71
    - K8 b7 q: ?1 m  n& @$ V0 |9 O4 w72
    7 I( i' |% e5 ^' D- d. N, B736 S5 D$ R6 H3 B% i5 ~
    74
    + J2 X$ h* s/ \75) p8 a& u2 M  ^, v: F# t0 z1 j/ B6 ]. G8 P
    76
    - H( O/ i% [0 m" H. t) v% R77* ^1 _7 T& Q: P! e7 J- [
    78
      e& h7 \0 W8 T( B& p# l% M796 g* \4 c8 _% b
    80
    1 M; ^! P0 A0 I7 _8 l+ w+ ~81
    ' ]- @0 a3 t/ A2 O7 b5 s822 h% b+ H3 s3 F3 `8 R" S
    83
    5 v3 v3 j! n4 t0 N2 T7. 设置需要训练的参数
    3 U: a7 t! P* `9 J" Q  W# 设置模型名字、输出分类数
    , M' X$ b* Q& X8 o' \- ?. a# X% Smodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)* }: `1 d3 ]! `, ~9 c7 t5 N

    2 c6 e  A0 \+ P2 o# GPU 计算
    ( I" m. D7 G3 D$ B. emodel_ft = model_ft.to(device)
    6 N# p  d9 T/ P4 ?% B5 s  ~, t/ u4 ?8 ?4 S, ~) J& g
    # 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取
    ) s4 i, k1 \1 B0 Z# f& x! P+ m- K* gfilename = 'checkpoint.pth'+ R0 d0 l* V! `4 u$ J
    ; H6 e- X0 d  \8 I2 R7 N  Q& f
    # 是否训练所有层( d; k0 t( u- v" U. ^9 ~  @" q, x
    params_to_update = model_ft.parameters()
    ' Y7 ~7 G, E( R, d! T; D4 Y! j# 打印出需要训练的层
    + ?3 @1 M: g/ Z! i/ Vprint("Params to learn:")
    + x' m1 ^) M8 N5 W' ?- y! {4 r( J/ Mif feature_extract:
    9 @9 }$ B; l! i* D5 m5 y: n    params_to_update = []
    . J' `# x7 T- ~    for name, param in model_ft.named_parameters():+ {& @# O6 N4 ]7 [3 d
            if param.requires_grad == True:( z3 d/ K3 N; C. P
                params_to_update.append(param)8 c: f( N4 `3 G! o6 X8 ~- ?8 w
                print("\t", name)) d8 P; \9 R! W: J
    else:0 l0 j9 P+ v  \5 \
        for name, param in model_ft.named_parameters():
    & R* X. }* b. |- S        if param.requires_grad ==True:* @8 r; W/ t4 X
                print("\t", name)# E8 f  Y1 A. o+ w3 k) Y+ W/ z

    + q! t3 G, v2 H9 O" e1 s16 Y9 Y, R+ C3 {6 ~
    2. F9 W2 A) j6 A5 i& E9 N
    3. s  R) k% ]) F+ D, a* W
    4
    & A& r3 T$ E& E2 N1 ~' r5
    " P8 Q$ l- d3 P64 K) L, o5 G7 z; Q$ |- d" P
    75 T- T/ Y: t+ f# r9 @4 ^7 p
    8
    & R$ y$ d# r0 o  f8 K9; g9 L5 u. f# a1 k: @: b- @9 {1 k
    10
    ! h( K; @% {/ B6 P" ]7 e( k1 C5 a! E11  V, M6 E) L2 D2 Q3 w, N0 O
    12" R; x+ `% u7 F0 ?0 m) j
    13
    + h. G# J. v" @# p' ^: s, o14: U# a: g1 e& u5 f" d
    15
    7 n) }$ X! @( `16! J/ _1 Z  k4 s: N, I' S
    17
    8 T! N( C5 E# X6 D( Q18
    ! h0 W! c! s; K9 X- J9 e" y. D( M19% D* _9 i* P- k" }, A/ N7 x7 n
    20
    4 z0 `+ o& n$ C& O2 E21
    4 `; e" Z$ h1 D$ i+ c- [22
    . y/ H9 H6 W( E8 b/ h23
    5 z/ Y6 \. v2 e' m7 ^Params to learn:
    4 G8 c% j; `. F$ F) p. @         fc.0.weight1 E& H, B2 @3 r) m
             fc.0.bias) m) D. K  {' S# {' y6 |& |' }
    18 ?$ L, f/ a+ x1 ]+ S- K3 ?
    28 Q3 ]4 r& R: M9 G7 @
    3
    1 R1 y: C, r! g0 Y/ i7. 训练与预测
    ! E3 N  i  ~& }% F7.1 优化器设置2 D+ r2 ?/ B$ l+ ?
    # 优化器设置
    # u, S+ z! _1 _optimizer_ft  = optim.Adam(params_to_update, lr = 1e-2)
    * r% e' M1 c1 A' J# 学习率衰减策略: _4 _: S6 P$ o3 a) W
    scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)
    / h# {; c9 J8 d7 @# 学习率每7个epoch衰减为原来的1/10
    4 a4 K" m% ~1 u# L# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算  ~5 I2 ~4 ]3 x. C4 K

    0 c" }4 p$ ~+ K* r4 Lcriterion = nn.NLLLoss()
    3 N+ J+ J5 y* i3 M0 ^' Y& i8 I3 A1
    $ f. x* J* f4 n, s4 J5 d% K; t2
    3 D9 n6 U; |' q: b. D3
    7 v8 g* P2 o5 `% [9 v; h5 M8 X4
    2 o  ^% h  o& E2 O% E2 K* f5+ I( j& v0 i, t2 B
    6$ E  C( r% U8 j
    72 N0 W. @4 H4 q/ T' S
    88 ~* s+ M9 }+ G  r# e  {' D+ r2 P
    # 定义训练函数# U' T; W2 Z: i* V  k( c
    #is_inception:要不要用其他的网络
    , z0 d! ^) t: K' n. l5 Udef train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):, C3 k( p4 ^5 V8 |5 w6 `1 X$ m; u- Y9 O
        since = time.time()
    ) ]8 ~' |( K: ?( `- s    #保存最好的准确率
    * I# ~1 L. Y8 @  B    best_acc = 0
    : a: f3 K* A8 t2 _4 L! r    """
    # h* k3 S9 Y: C    checkpoint = torch.load(filename)1 z- m$ h0 i0 i# R
        best_acc = checkpoint['best_acc']
    $ |# m+ c0 F. C2 ?    model.load_state_dict(checkpoint['state_dict'])/ H  j4 O8 A+ @3 R- S
        optimizer.load_state_dict(checkpoint['optimizer'])
    9 P6 P4 c3 N3 U    model.class_to_idx = checkpoint['mapping']
    ( f" t' X0 e7 M  d- d    """
    ; i! _/ G& v4 H7 R( p    #指定用GPU还是CPU
    & B7 Z$ t: F/ l( m    model.to(device)
    ( U7 p, @/ O2 h3 ~: B    #下面是为展示做的
    , d2 ^& J6 X+ s8 F    val_acc_history = []3 R5 K" m/ K. d; f. r, u$ Q& m
        train_acc_history = []3 p4 M  y0 Q! Y$ Z' t0 U
        train_losses = []* F$ C9 }( f2 ~4 k) c9 O/ d% E
        valid_losses = []
    # L- `" C/ g( i- T: c" n6 e    LRs = [optimizer.param_groups[0]['lr']]2 B) M) H& ^9 h( ]3 `/ b! \
        #最好的一次存下来) C" t( w9 `5 T4 t
        best_model_wts = copy.deepcopy(model.state_dict())! v0 V/ M9 @* M" R
    7 ]( z3 x! L+ V# m! G- e
        for epoch in range(num_epochs):2 ~. I; [7 Z. Q/ b) l8 E5 O+ t. E4 F5 t
            print('Epoch {}/{}'.format(epoch, num_epochs - 1))
    0 X4 C, R9 s5 q+ F$ y3 W        print('-' * 10): E. k, w: |" C: D: Y5 q- r
    6 I# l, p3 m  W7 d& r" m$ Q: x
            # 训练和验证# U8 {* Y1 p) B9 z; j1 c: Q% T4 o; h4 D  Q
            for phase in ['train', 'valid']:6 o3 A, @! Z. _3 w$ C
                if phase == 'train':7 P( f" k- w; [# X9 f6 a3 x2 S
                    model.train()  # 训练
    2 F5 Q9 t0 A& r% z            else:
    6 W! H: x3 _  m/ Y                model.eval()   # 验证
    0 |3 o. H% n, b) e/ ?. y
    1 I+ V, k& m- P" [; X3 J9 j( |            running_loss = 0.0
    " w; X2 |/ S9 C            running_corrects = 0( d! G+ n; b) |  q. Q) v
    # H, A$ Q& h  s! u  X
                # 把数据都取个遍
    . {8 q, N+ ^6 x            for inputs, labels in dataloaders[phase]:
    . M; p3 A% W! S7 b                #下面是将inputs,labels传到GPU
    # T6 n0 ~. T# F2 j0 D                inputs = inputs.to(device)
    $ G' ^* L  {, I1 P6 k                labels = labels.to(device)
    6 p2 d* M, A1 z/ j8 N- |* _9 ]! y8 M; D" ?
                    # 清零
    & a9 j6 X: u: ?* M! a9 o                optimizer.zero_grad(), \& V7 N. U5 {: X4 B
                    # 只有训练的时候计算和更新梯度
    . Z& _- p+ g8 Z7 q: \                with torch.set_grad_enabled(phase == 'train'):2 M6 ?# s6 ?- d* f- H( I
                        #if这面不需要计算,可忽略
    / \. \. e5 t  @" G' z& E7 o' C, H' ^                    if is_inception and phase == 'train':
    % K' `( a5 t5 b4 p4 a$ J! P                        outputs, aux_outputs = model(inputs)+ B) o( Z- G5 E( b3 z6 r- R
                            loss1 = criterion(outputs, labels)$ K' V9 J0 ^% r# F, `: x
                            loss2 = criterion(aux_outputs, labels)5 w& _0 G; d1 Y8 L& c
                            loss = loss1 + 0.4*loss28 Y8 L  e' [  j; v- [! g
                        else:#resnet执行的是这里
    7 n7 e; F- w4 u3 v0 F                        outputs = model(inputs)) k# j3 q" l; G' W$ M
                            loss = criterion(outputs, labels)
    7 s5 }7 d# V/ E% h1 y7 }) N/ w2 A! L
    , H( @# N: R  Y  Z# A$ L) w0 B! m                        #概率最大的返回preds
    / ]5 h: `7 O( @5 U3 }( o' D% H                    _, preds = torch.max(outputs, 1)9 w: k* W" |4 x% {: f- e+ b' N0 R
    . N  V0 y- w7 r6 N
                        # 训练阶段更新权重2 c  Z: v( h- s1 r
                        if phase == 'train':" ]1 P, T$ q7 z
                            loss.backward()" ]7 b' B0 w, x" c
                            optimizer.step()7 p8 n8 e' R* {' J; {% _
    # A5 T8 F6 r8 U0 U* R
                    # 计算损失
      j6 P' B2 q: q5 Y9 A) A; B4 t                running_loss += loss.item() * inputs.size(0)
    4 O0 N8 L* r; H* K% K' K. o                running_corrects += torch.sum(preds == labels.data)& n2 e5 V  F* v2 q( w

    ; e% e8 b$ O/ d. ^# O# }            #打印操作
    8 ?0 a) T& M! Z  I* I            epoch_loss = running_loss / len(dataloaders[phase].dataset)$ w& b7 g# C/ d7 d
                epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)
    & X  ]) |4 i! y3 P$ ^6 }7 R1 S% B4 `' \, Q

    ) A' t8 }4 U, S; z            time_elapsed = time.time() - since
      f3 z! Y% a6 ]5 f( _5 _            print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
    & @+ i  T) J( m& Q6 D; H            print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))+ R' P6 A% F8 M. w4 \  C
    # U$ l, |: k  B% A/ o

    8 {# c- ]2 j+ @# D) i% M            # 得到最好那次的模型
    4 d( N" w9 E- e: t- ~) l0 Q            if phase == 'valid' and epoch_acc > best_acc:2 r, |4 `1 W5 s0 g/ B" U% ]* ^: w
                    best_acc = epoch_acc+ L, i* M  h: l7 f
                    #模型保存
    & m& i# `) j8 M  I: l% a                best_model_wts = copy.deepcopy(model.state_dict())
    ) b/ w5 c3 L8 D* g" M/ q                state = {
    & [7 ~0 U0 ]: i  d$ T& ^                    #tate_dict变量存放训练过程中需要学习的权重和偏执系数  ]! K6 C) r6 j; R  y9 C1 E' G
                      'state_dict': model.state_dict(),* ?+ u8 H9 l! \
                      'best_acc': best_acc,
    # D6 Q* k2 @& d( T; x                  'optimizer' : optimizer.state_dict()," Q" I5 `% P* [  Z( o1 ^. k3 l
                    }3 P8 S3 _6 I* w4 l) j
                    torch.save(state, filename)
    : K( C( ?8 W( v2 k. F5 h3 w            if phase == 'valid':
    4 N5 k* ~5 Y3 e4 t& z8 U                val_acc_history.append(epoch_acc)
    4 L+ ]5 t: w7 f1 J+ M                valid_losses.append(epoch_loss)
    1 K; w5 u8 v5 \                scheduler.step(epoch_loss)
    . c2 q* |( r  {            if phase == 'train':. s& @4 }4 v8 X# q9 J
                    train_acc_history.append(epoch_acc), G& o3 C$ @) Z! X) d. o
                    train_losses.append(epoch_loss)
    $ [, X5 V- M4 R  v; f
    2 Z* s9 V* u: T        print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))# b% w. f' N' e6 G, n
            LRs.append(optimizer.param_groups[0]['lr'])
    ( ~/ T" ?5 @- P% ^5 Q9 J% g9 j        print()
    $ W% H, T- }, j3 F+ E2 ^: ^& O# N/ F+ l" k" J* G0 v
        time_elapsed = time.time() - since1 v$ D. R# Q/ J4 Y6 p
        print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
    6 H% j. G" m, X/ @2 W* F$ u# V    print('Best val Acc: {:4f}'.format(best_acc))7 R" n% r# V3 B; A0 k

    $ B4 y# Y/ M+ \. J% A    # 保存训练完后用最好的一次当做模型最终的结果
    / c5 Q' g/ T- I# l7 A    model.load_state_dict(best_model_wts), N5 {# E+ @) G5 F: ]% e
        return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs 2 b1 [" d/ Y  }) F2 Y
    # L9 _" [7 U3 f' J

    2 v6 ]' J  L% s1
    8 ?, d- K1 i+ e! a! B! p# H2/ A' e; k$ F2 P; L) L
    3
    $ p  w' T7 A  P4 e% k; V4- S: ~( C, G& s
    5) T% A) D; ^6 `- f' q
    6
    1 u/ `' \7 T, G7
    ) W4 e# |2 U& X+ S9 ^7 @- f& W8
    & |- A2 ?, \" E$ p91 u* O, ?: L' r3 x- X
    10
    4 B4 ~/ n* W* l* c2 o$ h11
    8 P0 H6 ^) q! A  S% }12
    6 u2 J4 F. `9 h" G13
      Y* N1 X( g+ F9 X+ G0 X6 B# L14
    - G0 O! T# ^2 a* y( l  O9 M15
    3 h$ ~9 g. t; O7 n* i+ W2 B160 G0 J3 j( F3 ?) G! V
    17
    ( R( i2 v& l: D3 c6 V188 i9 g( T2 ?5 I$ M2 x4 ^" d! h" ?6 F
    19
    4 x9 n' S( l- E1 m20. O% G+ v. Q: Q) B* `3 ?4 Z8 X
    21, E) Z- |3 [8 Z  q6 R2 M6 r
    223 ~: x* _/ M1 u/ G
    23; k& W, n8 N/ X" q0 ?
    24
    6 k$ ?  s2 H/ d/ a25) u; E6 A+ Y3 q0 m/ R0 G! I
    26
    ! }/ ?! H* L$ i2 N27
    / p. u4 y: z( X# X" G/ X5 {6 E28
    * |. Z5 f7 X0 K5 k  r- d' G292 I6 ^" E/ A  f& o9 S1 Y3 Q
    30+ M2 R, F4 X+ A' G7 O' Y! r
    31
    5 `+ M  _! N1 L3 I32
    ( C! k* M5 m0 A% ~; t33
    1 w7 z9 E$ K; B! x% W34; G! Q& t1 B; G% T8 Q' M2 M
    35
    2 ]' F$ ^5 E( b36
    7 ?8 d; w6 @, L) f1 F37
    7 N8 p) @0 t$ w( y' @. z! J38. K& j' ?# _) O
    392 n, t* t, h9 e) b7 \
    40# S( I/ t( ]  ]! N
    41
    8 K% N& A7 P8 ^3 N" y. m5 b42' M2 r  B' B  r  t
    43
    ! a, D* [- y" u  i- w9 Y3 u$ G44' S* q3 v; x+ C$ e& p) q' k; s
    45- ^) w5 B9 l9 T! D: @) Y3 M# h
    467 p" \) p& e3 H5 k( y* a, g, l5 n
    47
    , A: Y' n8 a3 c, X/ o48: M8 _5 S9 s& y
    49  v) U. l! v/ a0 z* f
    50
      W2 H$ i  r. e3 S5 l8 v51
    4 m5 j3 Y# O: B526 t% h  I( Q+ v/ d- N7 r
    53
    8 F: e; \7 {) F4 {0 X54* \- z8 w4 ^6 ~
    55
    + }' C" C6 q! k3 T56
    : m1 y# j# e$ n6 ^3 D3 p  K& q57
    8 [6 d& E8 f' y58
    . M1 B& g- {# n" {8 S59
    ' J8 C9 H8 q, z$ Z$ j60/ {5 S; q( J/ \
    610 R2 q1 `2 O" K0 J3 ^" {
    62
    . K2 T6 ]/ a% q6 ?63$ g  ?% g* P) \- p* f
    64
    # N- d# Z  y2 x" `% l/ g! Z65
    / ]' b3 X- d3 V7 V; K66% T) J! H! q5 [( j, ^% H
    67
    % P5 a$ l  ^1 L- F68( }, w$ \# o8 Q4 c& B. V
    69
    8 T7 y! R% d( r) t4 |# f70
    " p& U& b2 X# |9 i71
    # H, g8 M* ^" e. V) C72
    $ Z7 N- ?1 ?5 }2 q8 S/ t" G73
    $ {, N, b3 P7 f1 c74
    " A  s% E4 C3 q75
    , V7 f/ y: j/ u1 Y* P* U- z8 w76
    4 r8 s. L( w2 a; k# j$ q77; h. J: I. [# t+ f
    783 X! u& C. G6 Q# I9 Y
    79. P$ S; j, t+ W( ~% S
    804 b7 o' w- I% T6 u# f$ `' Q; D0 `
    81. d+ B. X- l0 Y- ~
    82* m  R0 P0 }4 o
    83$ _) c# b7 Q: k2 Y2 C
    84
    0 f  r, w8 q4 y( w9 p! _/ @: q0 D0 Y850 E) t, F5 U) B, Y9 W
    86) A0 Y' m$ w, |& j6 {6 ^
    87
    # s0 o' _' H  d0 z  ^" D886 |8 a, a, i% E% R+ a  b: @
    894 @3 t% G, z- w- s5 {
    902 Y4 c, M4 D" X, [8 G/ B- D: |
    91
    7 L9 D/ ]3 z5 ?* t& L) h92/ N. p- B0 m1 R0 |7 ?; ^
    93
    5 n+ Q) ?2 K+ O3 g6 t# x( I9 \3 a94
      Y  Y  R: c( U% S( _# n+ X4 P; L95" @2 R# N* K. t4 M* ~# ^
    96! B) ?* f& |6 M2 b% W1 b* I
    97
    " q! Z, J3 s7 k" o# o98
    $ N& Z/ G# \# u+ a( n# ]99
    ' C" h0 k, }9 n100
    $ z" c/ L, J! _& J3 c101# F3 s1 i3 P6 f( ^+ z
    102$ k5 t* K: [$ x, x& J; V* @% ~
    103
      C5 K5 B0 C4 V- N7 P, z104
    " m) w! F% K! d4 i% P5 e2 N3 D105
    * e$ e- e- N! J7 y9 [. T& C106
    7 c3 \! c" @* F# ?) w107/ c, g$ Y9 j8 E0 i
    108+ ]' v. T, }0 i) z2 r4 U
    1091 v# C( r1 J) z$ O' L( b; j
    110
    8 V9 @! W: H  L# |1 A( c+ p111
    , u: }6 q9 t9 a. u112$ p0 {% K# }( p) ?3 v( `0 ^# x
    7.2 开始训练模型
    1 e, o* y  p+ P6 c' B: j  [我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次- {( ]  Z4 L3 v( U9 P( Q/ N

    4 _3 R% V+ I5 U1 n8 D: a# O: t#若太慢,把epoch调低,迭代50次可能好些: P3 z  W$ f7 n, O5 n* Q
    #训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合, G* c: W' c3 l: G* o- 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 o7 F+ C- W" N7 s
    ! Z! v; Z4 y3 Q# ~# D4 F
    1
    ( |& `+ J( h, \) C2
    " u9 D0 }/ i) t( _$ ~# }. @3
    6 z7 ]) E0 N8 Z2 d  n4
    , J# S) E# T% ?" D" \Epoch 0/43 V4 K! _2 z6 E, P" u6 {
    ----------0 a: p, v! b; h# W; y# H" B
    Time elapsed 29m 41s
    6 v- K% r* x- }% m  v: g7 Ntrain Loss: 10.4774 Acc: 0.3147
    3 A# n  |7 \5 V+ P- E4 pTime elapsed 32m 54s5 e5 {+ Y% t, g) y+ F
    valid Loss: 8.2902 Acc: 0.4719
    - h" R" n- S0 m. f! JOptimizer learning rate : 0.0010000
      Y" L- W: }. S3 G: m: v: r7 d% b& E+ D  v# V- F% o' K
    Epoch 1/4# p0 |* L" ^* [8 I4 n# J
    ----------
    ! [7 E2 K0 @( ?6 zTime elapsed 60m 11s2 i) _( H. o7 d  W
    train Loss: 2.3126 Acc: 0.70537 W9 }& L1 d3 u* k: N" r( W
    Time elapsed 63m 16s
    8 J  o/ W+ _; I7 O4 b8 p% vvalid Loss: 3.2325 Acc: 0.6626
    6 c# q" ?- t7 z6 NOptimizer learning rate : 0.0100000
    ) o; b8 Y5 H; ^1 [; ]4 p, C* T# @$ B+ r/ `5 @. j2 C6 s5 U
    Epoch 2/4. A1 }- C9 g  D& S
    ----------
    & U! B; z! c$ v% D* K5 D) [+ M5 yTime elapsed 90m 58s/ x* v1 [/ [3 b. r% t& B. [
    train Loss: 9.9720 Acc: 0.4734
    $ _, V* t7 y& u$ n4 t- K9 p2 H; R- hTime elapsed 94m 4s
    / W) O; r& i6 i% l, ?valid Loss: 14.0426 Acc: 0.4413" N( }1 m: N4 Q1 a1 K8 {# X
    Optimizer learning rate : 0.0001000/ J. c: O* b9 t+ {) O/ A+ A) G$ Y9 c

    3 Y) ~' j( C" b0 y- |, _Epoch 3/4
    - S5 r2 L0 m) A8 z+ t; b----------
      J- R; E$ N, o5 ~1 _8 y! K$ ]Time elapsed 132m 49s* ]$ V& S  U) W1 j& n: c" y# d
    train Loss: 5.4290 Acc: 0.6548# A" o4 P- K: i* o, E& y& K
    Time elapsed 138m 49s
    ) m/ @2 R) l4 w) n7 G) r  _valid Loss: 6.4208 Acc: 0.60270 h/ {2 D9 L. I
    Optimizer learning rate : 0.0100000, B& e/ y3 j- k, E& K) g

    ) H1 K, g: I: P$ [Epoch 4/4
    1 ?0 H+ d2 z2 G/ O: @& T4 v- c4 }9 o----------1 k4 ^4 i% @5 Y! o8 _* p& b
    Time elapsed 195m 56s
    5 A. H, G- \- |" Ntrain Loss: 8.8911 Acc: 0.5519
    . Y9 H4 ^# K  R4 S4 C- A: \- iTime elapsed 199m 16s8 G% C' c/ R4 U2 `7 V
    valid Loss: 13.2221 Acc: 0.4914( F# J& @* t5 e/ Q
    Optimizer learning rate : 0.0010000# ~1 K8 s9 G% H/ w6 ?% \; d+ h

    , q: x1 @* F$ s3 ~Training complete in 199m 16s
    9 w/ A" R. Q, W! C7 E. h! tBest val Acc: 0.662592
    , ]4 W! @  l7 u; X# ?# V# j2 H( N; {1 H; o$ s# F! D7 O
    15 J/ p: W8 H9 ?' S( E
    2
    9 H% ^- A$ x  T3
    - t6 q' L; B2 \3 q4
    1 p. ~: d  e( j. ?8 H, g5
    : p9 u+ @& t( C1 B- N9 b5 A7 o6, [9 u2 Z3 D' z% ]+ D
    79 T! t3 C) H7 Z0 [! H, r" s5 {
    8/ H3 t- p  ~! U0 l/ j& W! d) E# G
    9
    3 e: L* ~8 H% K+ A102 r9 s5 k  H. `" H
    11
    2 C7 ]. [) s% e0 p12
    + i. F2 I- `5 b$ R138 Y/ y. y( x8 a; S) ]$ T& [
    14
    ' h- r: u4 O5 |' j  i15
    & l: @' W$ g* m. E16
    3 P0 b& {5 t$ T17" ?; a" `0 D# @
    18
    . I2 v8 J0 k+ X* U$ m. Z! Q1 s19
    " n3 H% @4 Q$ C! J20
    ; N: Y# _# _) R# k' g  r( D2 i" Z) Q21& L& y# B1 g; J5 O- Q
    220 P; o* h# I- r* I" @
    233 k4 W2 e: T) N7 S3 I
    24/ C- F8 |' C7 n& _4 t
    25
    ' I' |4 {0 u) v  v26$ P8 ]# |, ~) o8 L/ A0 ^" K! e
    279 Q+ O! H; ]; h/ r8 W+ q5 V
    28) ]7 Y' x, r5 m4 e
    29
    ! F7 q8 ]3 S6 I& d/ B4 a30% r  G+ u5 b& Y, u
    31
    / w2 S) R. n$ f$ S& X32% ^) y& ^, \: x  B# R
    33
    3 I& A: u+ r& ]34$ q( K$ {1 ~9 W( U0 k8 x; V+ U5 i
    35, `# g6 ^- |1 T2 ^+ ^- [8 z
    36
    * ~6 |: R- o6 U9 H37
    # S/ \: z1 i; O$ }38
    ) Y* a3 M2 C, \& C) w" u4 ^39" z# J& o+ x9 E5 \# d8 M
    40
    5 R; m9 e' @5 o1 z3 I, v41
    8 c1 V" D" V) G42. W; @2 k8 k, r, X2 H+ s
    7.3 训练所有层
      b" K7 p  r/ D! q# 将全部网络解锁进行训练  B5 R" w$ G5 h% i/ n
    for param in model_ft.parameters():* h* p; n% D( k) ]: z
        param.requires_grad = True! O7 i) j' r# B; ^
    ! i, N# @  T5 t- P* R, q
    # 再继续训练所有的参数,学习率调小一点\& O% o( U9 y& [+ W+ b5 Y
    optimizer = optim.Adam(params_to_update, lr = 1e-4)
    ' }% T, A( q) I; Hscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)
    4 x* s: B9 z# b: V7 {2 a- _$ i% n$ m6 R# D) R+ E: K$ H  Q5 L6 x
    # 损失函数
    $ F8 ]0 G0 ]& b: s1 ~& a3 [criterion = nn.NLLLoss()
    8 D: ^4 z- M9 K13 {$ w. q4 @. {) a9 h( Y
    2
    7 ~3 i5 A7 d! z; B3
      U/ }* \, B8 T4 m& X" H' H. d4
    ; z5 Q0 y% ~8 _56 ]2 O, C  u. I+ ]8 X: Y# Q' A/ K& v
    6
    $ A- b% L6 ^' J# a72 j0 k1 v) d6 g
    8
    5 h$ i& G: y. T) J4 N$ {9( V4 @. b" M5 \& ^' D
    10
    - \3 M. h" W, q+ Q9 y; J$ O# 加载保存的参数* B' R  n! Z. h
    # 并在原有的模型基础上继续训练
    " H6 w; A' V" i" f6 z3 z# 下面保存的是刚刚训练效果较好的路径
    5 |- M1 f& e) B  O  l1 Y! o% J, L7 V$ hcheckpoint = torch.load(filename); I: z+ I( l! J. D
    best_acc = checkpoint['best_acc']8 }7 N2 d7 W! x$ X$ [3 {
    model_ft.load_state_dict(checkpoint['state_dict'])/ }; X* ]; P- o1 F
    optimizer.load_state_dict(checkpoint['optimizer'])! T: `2 S8 `, b
    1
    7 a$ N9 m6 w$ t22 W/ h" d4 w6 ^6 |" o
    3
    ' a7 U' {+ \& B41 J# I3 E5 w0 J! E4 B
    59 q3 P% ~+ V2 L! A$ B' l- s: U
    6: A. V( Q% k0 N1 m+ Q- M
    71 p: Q$ w3 s% f
    开始训练( G) M- I$ q5 ?3 |, @, H% w
    注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考
    7 t4 \( X; R  F, a: S7 L  x5 l* X  i/ Y, s2 g' x6 z
    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"))# m) F+ d& j5 T: J3 w
    1" b4 u  D# Y  ]6 P2 z7 C
    Epoch 0/1- m* I$ o# b' Q# s
    ----------5 {  g( I# F( H4 _( K' g7 @
    Time elapsed 35m 22s
    6 }, q' ~/ m# mtrain Loss: 1.7636 Acc: 0.73463 A/ A2 i5 W' K9 z) u0 ~7 X1 U( s# e
    Time elapsed 38m 42s7 a9 o1 v: T3 C1 y: q
    valid Loss: 3.6377 Acc: 0.64556 Z' c+ g* k  B
    Optimizer learning rate : 0.00100002 D# d" e  w  Y: s: R7 P
    , Q. v. U6 a6 o
    Epoch 1/1
    2 P+ Y( ^; v5 H2 w4 s----------8 E& m9 m, A$ @
    Time elapsed 82m 59s3 J) T& N7 J, D2 Z' G: V8 y0 b1 L
    train Loss: 1.7543 Acc: 0.7340
      q$ c! |; ^+ Q7 q  v0 o1 qTime elapsed 86m 11s
    0 x! z- ]; l2 w3 avalid Loss: 3.8275 Acc: 0.6137. \  w7 ], |0 C3 h. ^5 T7 Y
    Optimizer learning rate : 0.0010000
    9 |( _$ g  D" g( [- m4 b9 U0 A6 {
    : k9 C  E* F3 ATraining complete in 86m 11s
    0 A6 b+ q2 C" C1 M$ |# ~+ }Best val Acc: 0.645477
    9 \8 E) o! v" x2 l9 a2 C9 L8 n! w) \# W" q( B5 A
    1
    7 k6 J- W1 V% V9 u6 b2! J0 N1 j6 k) e' [8 `! J- }
    32 d# b! Q- u" I" g+ ?$ o; Z) |; G
    4
    : c% p8 P4 k  Y; O3 Q. r8 a; Z5
    4 M4 K$ ~% d, q8 k. L0 W4 v1 {6
    ' T% _( D2 O! |: R7
    ; \7 }0 V+ B& \* I- ^2 ?$ |4 k9 e5 c87 l) w- S5 D; c
    9
    % {" ~0 _  i* J1 u) a2 h5 j/ V108 D4 Q/ a: x$ t% w, n( l/ L
    11
    ( G0 y8 j- s( W1 Q) P) D8 u. q12
    7 `0 {1 i$ o: N: E. P1 G& o# Y2 x13
    1 U( x. ^: U) y! T7 R& B+ c14
    ) w$ U! V, d( D15+ x- q9 i9 p$ m* M
    16
    ! j, b2 m! ^! A5 g  C, y! V17
    ; L6 F9 v* ~  x3 {18
    ; B0 L# E6 z4 J+ @* B+ a, ~6 z8. 加载已经训练的模型
    & e- j! K5 \8 y6 ]2 B! L3 S# z- i相当于做一次简单的前向传播(逻辑推理),不用更新参数
    / s0 s! ]% D+ {: _' ?% E( V1 b! G. M) [0 ~% ]+ e7 m8 T6 n$ \
    model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
    % j- M8 O- E0 M" I) |+ K( x: b2 K; G; H+ i( ?" e  t: {
    # GPU 模式5 M6 q0 ^# V/ c$ N
    model_ft = model_ft.to(device) # 扔到GPU中: [' \( q; E! k. \) f1 ?
    " Y! z; R! @: h7 ?0 ?+ X' ~3 X
    # 保存文件的名字
    / }' @5 Z0 P7 Z5 t2 Q, f# [; I: d% u) @4 Vfilename='checkpoint.pth'
    - J5 \8 E/ S* U2 P+ C0 N( ^) b1 p! |1 O( o
    # 加载模型) w* e, i, s# {0 d3 _0 ?0 B
    checkpoint = torch.load(filename)9 b' I( _1 a/ P8 o+ @9 S* p
    best_acc = checkpoint['best_acc']
    9 E. i* U: c2 q) D& \9 O/ Fmodel_ft.load_state_dict(checkpoint['state_dict'])$ L$ ^2 |: H% a
    1
    / P. }4 g  c) Q$ l; S2* ]$ t, ~1 r+ o  [6 d& O! b1 U
    3
    0 k4 _% Z* n- V: |& [4; p' }  F# x: Y, d* k: ]7 p
    5. b8 V% k- b. H! m
    6& j3 C7 w2 u; D7 C* r/ v
    70 [  x% O& B+ M) `$ A6 C! z4 b
    8! J0 O( h) \3 g0 N1 [
    9  F, x+ G6 N9 \% V2 d
    10
    ! w- e0 |# n: V8 w" d11
    5 ^6 Y2 x0 r* |5 z12
    3 e1 x9 k' G+ }8 M7 v' E- h5 @( ?<All keys matched successfully>
    : ]9 R" g" `( T1  v2 j& K4 G1 w. ?$ r- N
    def process_image(image_path):# e; i; a8 |1 T3 e9 x
        # 读取测试集数据+ S9 M2 q1 X! z6 ?9 M( \* y# P2 ]
        img = Image.open(image_path)
    ) w: k6 ~; B" D    # Resize, thumbnail方法只能进行比例缩小,所以进行判断. _$ r4 P. t* w: D
        # 与Resize不同9 u! ]% K6 ?7 x# v0 o' Y5 Z
        # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小  S1 `8 N. Z2 t$ M- T) J$ w
        # 而且对象调用方法会直接改变其大小,返回None
    6 u/ m- J$ W' w2 Z7 [1 I7 V3 W    if img.size[0] > img.size[1]:
    ' k. E. z/ l$ O        img.thumbnail((10000, 256))/ I5 {) o4 k4 M  l/ `, {8 h
        else:9 W6 |& u$ g; F4 g# ?# t
            img.thumbnail((256, 10000))  m1 x6 G% N# [3 F$ ]9 \. J$ j; X

    . l/ U4 h# T  @* ~+ x9 @$ f8 ^; p    # crop操作, 将图像再次裁剪为 224 * 2248 c7 T4 G$ E. |
        left_margin = (img.width - 224) / 2 # 取中间的部分
    + }. o' h# d9 M  m    bottom_margin = (img.height - 224) / 2   I1 X+ p. b8 \: n
        right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度
    2 o6 ~. E/ u* ]; y    top_margin = bottom_margin + 224
    $ R; v5 L1 s4 w& D* S7 p4 {( r, F' b7 p5 k$ {1 Q! J
        img = img.crop((left_margin, bottom_margin, right_margin, top_margin))5 v* y: |8 W4 @7 _0 p6 a
    - Y- |: H$ J0 u$ ^
        # 相同预处理的方法
    ' O3 r3 R1 o4 q, b    # 归一化
    0 J" B* O8 d- j% O, l    img = np.array(img) / 255. y! }) P$ L  I6 Y- P
        mean = np.array([0.485, 0.456, 0.406])
    : j; c* G& q. A+ j2 u    std = np.array([0.229, 0.224, 0.225])
    8 r3 T% T! w% c  k    img = (img - mean) / std
    7 C' A# M. d% X9 ~& v. ^
    - \6 I& {5 U5 J, D/ c: r    # 注意颜色通道和位置3 L- H  W7 }. q4 v$ U& `3 q5 o
        img = img.transpose((2, 0, 1)); V. v6 G8 e; M  A; [0 }5 q( j6 v

    8 Y% C2 u& I5 X1 \    return img
    * I+ `0 p; i1 F1 i$ {3 H0 \( l, j/ S: g6 g4 \6 }  F) c. F
    def imshow(image, ax = None, title = None):
    # M3 Z& f7 y* I6 ^7 r! R3 G    """展示数据"""
    5 W; p" n1 R, N& D/ v4 g6 d( P    if ax is None:
    ! x. o$ y& E8 i        fig, ax = plt.subplots()
    + i3 w7 z$ A8 T, J  P- g' h7 Q# d6 @# Q7 e& N
        # 颜色通道进行还原
    , _6 |) Q1 N% p" z1 r0 J% M: t; D    image = np.array(image).transpose((1, 2, 0))0 t& U9 b  y. p* C& a4 v
    * P! N# A6 d$ c2 F, C, m
        # 预处理还原
    : w; l: A& r! |$ j9 A    mean = np.array([0.485, 0.456, 0.406])% y- ^( |; k- S! E& `
        std = np.array([0.229, 0.224, 0.225]). {$ l& G9 _5 A0 p2 i
        image = std * image + mean: U' d% O* h$ C# c( L2 a+ [' `* c$ A
        image = np.clip(image, 0, 1)" U" ^9 g# g  K/ Q

    / a/ X+ W. ?4 ~9 R0 T( i    ax.imshow(image)
    # Z, L, @! t2 u) ^9 t% L# Z    ax.set_title(title)
    ) y! E- |9 \5 e5 J' \: ?/ ], E* V1 f) o' d0 u6 i. o
        return ax
    6 Y# a6 m" t0 f, |, y+ u) a% g7 c- y4 _# w( p2 Y$ f
    image_path = r'./flower_data/valid/3/image_06621.jpg'$ M0 d4 m& ~+ ^
    img = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理
    / Q7 \, V# r3 \4 p0 Y- Gimshow(img)/ U4 ?. g' o$ x( l
    9 I2 |8 \" J* G) T, j( ]1 a1 s4 j
    1. u; W; d0 n+ C
    2
    & I6 p7 l# Z, h$ Y3 q3
    3 a) C. l6 |7 u$ \4 x7 ?4' D& J# T. B. M: ~) N  ?+ C. c  ?
    58 V: X2 k+ q2 s6 R. u: H1 i( \6 Z$ _
    6
    ; u& E+ W% T" M( z" W7
    " T/ n1 L7 Q" Z5 l% S1 o8
    * Y: y7 S0 m" r2 v& u, X) p2 D% v( ~9
    ( ~( c  X! T0 {' M+ x10
      |4 C- U  D( {# v* L9 D- W3 L% C7 X+ l11
    , Q8 T. l. i2 i. D: y) ]& \! X1 c8 u12
    0 o& i* a8 \& J" C# ~7 ~, k13
    % t8 d8 c! O3 K2 }146 V# U( E/ M" n1 z; _
    156 k3 M1 Y9 n5 W9 b  h0 j
    16# x3 s" q6 j+ i/ q4 E5 i
    17
    3 }  e; Y; Q  o) G7 b18
    9 j# [- S4 q- w9 C19
    7 {/ Z: d! ]$ Y8 S  g+ g2 p20/ M# h0 t9 O( e3 _+ X4 v5 C
    21
    1 W& }- `& v$ |! _3 S% N+ t22* X, D: z) k- s: ~6 u. b9 x, K
    23
    % i% a% O( a* E24. L, U* i$ U2 _' K( g' P2 n
    25
    + C6 |( O; v/ N6 P' I4 p+ Z' @% g261 _. _, h7 q7 b7 D! U. e
    27
    # S* c; |- D: |2 U28
    - [. |- g4 J: N! A7 P& x# p29
    - j7 I8 v% s( J30
      X5 L; N9 a" t- }7 e' X31" Z) ]* R+ f( ^) z6 A( _% T9 Z  i1 R; R
    32( Q7 O7 C* D8 z9 W
    331 b! r7 @5 q0 H7 w- C, L
    34) ~4 {/ S) r. a6 J! W1 x6 X, U
    35$ c- D' ^. v8 J1 e* G! M
    36* R: O& G8 F$ @6 L5 m
    37
    * I7 [: m2 I% _. s0 `38
    4 A" T$ h# J% [: D& S4 Z39: V# n8 h. \9 r4 F6 I" W
    40/ {5 l0 D9 w8 f3 F3 t6 k5 r! j' e
    41
    % e& C) Y1 K2 a, K2 x; Q5 N$ h42/ C' N' C8 k6 |8 p3 s
    43( g7 T/ ]! K! |* k: v, g
    44# x: H2 q8 u9 E  M3 ~1 R
    45
    2 y/ h, }5 h8 Q46& {5 J; c( g9 R' M/ X. B, z
    47
    9 R% j4 R( V8 ~$ V48
    4 H, [8 L' H8 Q8 o0 c49) T9 ]- h1 X% C, J# ]
    508 |/ b9 f$ n. `4 x  K) ?9 A
    51
    ! v, p5 l5 B6 C+ t, P52
    5 G& @3 g4 Z& `) }, M53
    0 {/ \& N# X2 t% E* a54
    " @% t5 s+ S* C7 ~8 `6 U+ o0 @2 ~<AxesSubplot:>
      U  a7 I. }4 @* Z, I. R17 c4 \- F, r8 f

    2 N6 R* K* U$ D7 }  k2 s0 R8 V上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确& f3 R' [1 _1 Q5 [# y3 U9 A3 k
    ; p1 s6 ~- I8 ^! C. P# i
    img.shape8 q4 z# \& e1 P  P
    10 ?* @) m3 s. V% ^: C
    (3, 224, 224)
    1 \, M& C* I  j6 l2 }# h11 \6 O% g- ]4 c' l7 _9 t
    证明了通道提前了,而且大小没改变
    3 y# s+ f5 U, `& `
    $ [* {5 _/ l/ W8 a/ \) e; v2 u9. 推理
    9 h' R/ W* t4 o' }& G" ~img.shape1 p& r6 x+ w/ b
    * k/ O. F9 L& K) u6 M
    # 得到一个batch的测试数据& {" m5 e) `& t, Y0 ]2 N
    dataiter = iter(dataloaders['valid'])
    , Q4 a2 t/ w$ e$ r- e- O: b. l5 Kimages, labels = dataiter.next()
    " F8 H/ F2 f: B. m8 h# g
    . H& o# s5 |$ A4 ~model_ft.eval()
    ( u; G9 @3 l& }/ o( u0 M8 f
    ; b1 J- W2 f  p0 S' U, [3 k. @0 {* Jif train_on_gpu:
    4 Z, h# s5 e# }% Q, b3 k4 X    # 前向传播跑一次会得到output
    ; O6 M0 S" h. q7 a1 c7 w: J; `    output = model_ft(images.cuda())
    9 T) z5 g* [8 Z: x/ ]else:
    $ H6 v1 k+ V# m5 N. [) X! G6 E    output = model_ft(images)) L% \& X: B) G/ ?0 _
    ) Q- F0 y# s/ H8 L
    # batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值8 ?1 j& [" x- Q0 W1 \
    output.shape
    : ]! x. ?. u/ u, Q# v; f0 l! M4 \
    1
    " q9 T$ q# l8 v$ f* Q" F4 q2* c2 Q" @- I' t% B" N& @; m7 Y" o
    3
    : @9 L5 ]9 l6 `) V4
    ( M3 q1 J, ^7 {) f! e( ^5
    ; @: \/ b4 Z# ~6 H, ~) C- i/ C6
    5 P! \" B, o) b# C. `( O+ C7% Y$ c$ H% o3 p) S
    8+ F6 D/ ?/ ^3 L: A7 {, \7 H5 `. p
    9! b9 Y; [- _% ~! @+ ~
    10
    3 B7 j( X) r1 e6 ?11
    - E1 j5 O: n0 c12
    & d! h9 S' E7 ]2 e& ^# {130 B/ h! }, t5 r! s7 u
    14
    3 E, g( R7 T! z9 x- o) A15. z! H. n: ]0 T5 ?: A: w
    16
    . y3 y8 x% @( J4 [7 Itorch.Size([8, 102])
    ( W+ P2 z5 L" H' ~1
    - h, ], N4 `# p; l4 G. K9.1 计算得到最大概率7 }4 ^0 d" z8 Z* i
    _, preds_tensor = torch.max(output, 1)
    1 k8 G% b" \% E" h9 o0 o  {( S
    : v; }3 o3 l5 |1 [8 V( {7 R$ Jpreds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量
    ) V( W1 K  R4 G# k6 [1  A7 I' E; a8 z
    24 X8 d1 _/ t1 I' T
    3
    3 Z8 k7 j2 }8 A7 L  a5 a9.2 展示预测结果
      \4 f+ z: ~% ?: [5 t. l' Efig = plt.figure(figsize = (20, 20))' a: e" W4 y2 F& I( q
    columns = 4, l: J! `+ t% S
    rows = 2
    6 ?" l, h3 ~( d+ i
    $ b: d7 U% k! G: ~3 Kfor idx in range(columns * rows):
    ' z0 A: c) Y9 r1 e2 H0 F/ y0 p    ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])4 }* S5 x- X! W! s' P
        plt.imshow(im_convert(images[idx]))" K5 l7 I  T; a
        ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), 5 K5 s5 {* m% ^: |* B" m% t
                    color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red")), M/ _0 ~5 ]# D3 n' g
    plt.show()  H; K2 i1 @. y: ^- V0 P
    # 绿色的表示预测是对的,红色表示预测错了6 {0 }( i2 Z0 Z; a/ D
    1  J4 Q% D  e1 I
    2
    % a+ d5 u3 u3 g# E4 Y3
    1 R! O: S' d0 U) Y9 H4
    5 i% J1 ]" j7 A) R" ^5
    ( ?6 t! T$ N/ U# z8 ?" y7 M6/ j0 \* R: O4 r
    7
    / G1 r  u; N9 t* U* D: ]8
    : X  y  U2 h* t8 E" q' W9: q7 z4 q3 z, e* O; l3 l8 a
    10
    / _: L. G# Y2 U: v11! D. f% H3 l# A# e6 o+ u: O7 x& _1 b

    6 V3 o0 {: ?  [& N& N* H/ O- {
    , R" i( v$ S& D9 g. M% |' S
    + }( m( u8 o- U) {————————————————
    " x% _, u/ S: f版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。4 k5 g# g2 W) S
    原文链接:https://blog.csdn.net/LeungSr/article/details/126747940
    9 |1 ?; x! P% ?6 f. `) f& t1 Z7 M9 ^2 }7 v" Y6 ]0 N8 M! R3 I

    9 m! h4 i( O0 Q6 s* L
    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-4 04:04 , Processed in 0.460545 second(s), 50 queries .

    回顶部