QQ登录

只需要一步,快速开始

 注册地址  找回密码
查看: 2779|回复: 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)实战案例
    ' F4 Q* Z: S1 p, ?' J$ |1 u$ z4 B& |  k$ G8 s
    文章目录
    1 F  N/ s. d$ u5 d$ i1 P卷积网络实战 对花进行分类
      t  P4 g2 o! B' L4 z数据预处理部分
    ! p8 U. r2 ]" Y5 A网络模块设置
    4 J1 }0 I" Y/ K* R6 H网络模型的保存与测试
    ! n4 ]( |  q- w5 ?. }8 f数据下载:
    4 O* a3 r% L- R4 x$ g9 E9 F" Q) R5 o1. 导入工具包
    : `2 K9 O' J- d; r1 _- x2. 数据预处理与操作
    + T* M, [7 i; y; D& H5 K# k3. 制作好数据源9 i* _+ F/ V0 o5 Y. X
    读取标签对应的实际名字
    " Z& X* U+ Y6 }: L3 I) y/ O4.展示一下数据  j0 |* F9 k* C# o; a
    5. 加载models提供的模型,并直接用训练好的权重做初始化参数- ^& @5 M& K: F
    6.初始化模型架构
    . w0 c  @! U3 s7. 设置需要训练的参数% a, g" M7 n+ t: t1 F
    7. 训练与预测
    7 d, k. v, S9 _% t) \: O- ?' c7.1 优化器设置, a: W9 _& D& t4 r
    7.2 开始训练模型
    " z) u' j% r- v2 z4 l) ~6 x, b2 m$ W7.3 训练所有层
    3 w8 c7 _* w" _/ Y开始训练: D# {# t5 v" j; G8 a+ U
    8. 加载已经训练的模型
    : c2 J  S/ o- N9. 推理
    2 X8 I" I; R6 u: z9.1 计算得到最大概率  `# B3 ]5 G3 d. }
    9.2 展示预测结果
    . i( V" S# [& f2 O写在最后
    - q. C# R' c7 K2 _6 H' H卷积网络实战 对花进行分类
    4 w* \) R6 Y/ @本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能  }! ^: ?. D. v5 j: K) j

    * E* j. a1 U! v8 j在文件夹中有102种花,我们主要要对这些花进行分类任务8 X. V3 r/ V2 }8 T/ h- e
    文件夹结构. L: m1 k8 B- ]
    3 t0 [7 d" b3 Z, ?, D* g0 ]
    flower_data+ f$ x1 ~1 L/ _+ S( e

    + p: |2 @7 [  _% Ctrain
    7 c; {+ c8 o( d1 g, I% m. M' n4 |/ ^! r. \, E
    1(类别)" P; M$ i) y; k# \- v3 ?, N
    2
    ' x( m: u) K' exxx.png / xxx.jpg
    " I! |" q+ C0 z, x, d( R2 W$ Jvalid; z5 S2 s, |* w" K8 A; e0 P4 y
    0 e9 P. c% M1 w7 I6 b! x- j2 J2 B9 n$ }
    主要分为以下几个大模块
    9 @$ z/ X: k* _# a/ F( @+ S8 U
    数据预处理部分/ s3 B  R0 v8 z3 S$ H
    数据增强' q6 u% S; b( g; m
    数据预处理
    3 D; _. h& ?8 b( J  v( N网络模块设置8 I( P3 Q( k1 ~6 V
    加载预训练模型,直接调用torchVision的经典网络架构" l- ~: F9 h! b2 I+ ~8 g/ ~# r( y
    因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务
    - D  W  v& J: f9 I* l网络模型的保存与测试
    $ L' C  H+ r  {+ [2 |: R- U模型保存可以带有选择性& O/ f" B8 U8 o) M% m8 @
    数据下载:4 i/ |8 `( H  Y& V& {1 C. V2 P! u. Z
    https://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset
    2 w; h/ [, H) Z% D9 O/ \. S" H7 @
    5 \5 \. S) Z3 `, c% Z$ l9 o" t+ b改一下文件名,然后将它放到同一根目录就可以了5 o: q$ [5 M$ R

    - p# w( `( p/ S9 p# `( ~% `9 L4 _  C下面是我的数据根目录
    , L7 V+ {5 a/ p$ V( O' ~, t( n' Z8 k8 P5 z! W: e  s8 F! o: K. t
    / q  d# e4 J% `, ^* j/ v: d
    1. 导入工具包% v( \: d' z8 {1 ?5 j) q' C
    import os# C# M. ^: _- ?
    import matplotlib.pyplot as plt
    6 H( C" H1 l' n# q: W7 C# 内嵌入绘图简去show的句柄" c% g2 b, J2 s  O
    %matplotlib inline - f& A+ O4 `4 |& _* R; g4 B
    import numpy as np
    * ~; J" O8 K3 ^import torch$ s) z4 d: N' d- ?6 v: s1 b
    from torch import nn
    8 F$ [1 I" B# T' A4 {0 G6 q0 I
    , z- l7 w- |) M- L  ~9 C+ yimport torch.optim as optim
    " W, _% ]  v; f6 J! l% U+ Z; Rimport torchvision4 z9 j) v. I2 W
    from torchvision import transforms, models, datasets
    / V1 D/ ^, ~+ P) b9 X' r- I8 _" a5 \4 J/ e) Y
    import imageio  V, d. k" q2 C
    import time( u' Z8 m* a9 ~+ \) b; A2 L) Y
    import warnings! j* O( J  \0 F( F, I
    import random
    7 ?; q8 x* O% Dimport sys
    & l$ p/ i- S( N! q, Gimport copy
    ! a# k9 W+ B) d4 Dimport json% M7 L7 q1 ~& j
    from PIL import Image
    : T% J, h1 G7 E9 _
    8 X9 l6 C- m2 R8 ^, b' Z! V& v
    ; U2 T% q" p7 ?" ^' `13 c$ [3 q5 X) O
    21 I; R0 E7 G' u2 r
    3
    % R. w& J0 [3 Q' o4
    + F: i4 X; h. \57 g4 E/ s8 H7 h" K& X3 D# Y4 Z' W
    6
    9 r3 H* {. T8 m* H4 M7
    & l2 G" {7 W8 R8  a4 N8 ~1 W6 q- v9 q/ R4 d
    9) G: s* `- ~+ V2 D' S& k
    10
    ) S# ?4 v. E9 C. L' u4 F11
    " l: S, J9 {3 F" h12+ N& L- Y! q$ B5 N; _. D
    13  {0 u$ V3 X" o8 \# \- W1 j
    14
    / x* k: w; Y  \8 }15
    : g& r; c  H. ]/ Z16
    3 f9 H" G1 Q& T6 X8 i177 F3 l3 H, S: S( f7 z
    18
    8 o( L3 N: M* j* d19) L8 [1 z( X/ `: I
    20: y4 i9 U& X; u2 j
    212 S" S$ y0 T, d! y! Y% C2 o' R
    2. 数据预处理与操作! Z4 u! ]+ N; y& r3 f; x3 B3 m
    #路径设置: w% x+ M6 _2 ]& `! P, w8 z
    data_dir = './flower_data/' # 当前文件夹下的flowerdata目录) V" _1 J, v5 n; p
    train_dir = data_dir + '/train'! a* r( i" B, z; O  j
    valid_dir = data_dir + '/valid'( v9 \, W/ M' O0 ~
    1
    % ^  r. m2 ]& L2' O6 R. t1 S& L$ Q! J! ~% T2 x
    3, h! l) ?3 W* y* [
    4
    ' \" j! O0 Y. z+ Epython目录点杠的组合与区别
    3 I, O! x9 D* J6 s注: 里面注明了点杠和斜杠的操作3 ~# }8 @" y: J, F

    , _) @5 R& x- k3 ]' f, ]: i3. 制作好数据源
    $ B" e) A% ?; ^$ A( H; Mdata_transforms中制定了所有图像预处理的操作
      [/ s+ G5 w3 O* B" pImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片
    % \0 f. a, O, U0 j. Z7 Idata_transforms = {9 G5 J5 g9 V# d3 g5 I9 y" p! J/ @  `
        # 分成两部分,一部分是训练; b. T/ J; l. n' P
        'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间! |, G# x. X9 g$ P1 m. f
                                     transforms.CenterCrop(224), # 从中心处开始裁剪
    ' n) o6 _3 t) s' j0 ]                                 # 以某个随机的概率决定是否翻转 55开2 D/ M1 A7 w$ L5 r
                                     transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转
    ! B0 L* F' j; [+ ?                                 transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转
    # p4 S  D5 n3 f/ O9 w& [6 T4 K                                 # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相' B- \6 {$ {4 [: e2 J, W0 J2 K  {
                                     transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),0 b+ h9 U% z( A3 Y( ^/ V, x# @
                                     transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB
    1 g) K3 x8 @3 i  X$ |. y                                 # 灰度图转换以后也是三个通道,但是只是RGB是一样的
    / H9 u, p7 a* T# m. K, s                                 transforms.ToTensor(),9 R6 L$ }" D, }: |' C6 l) M
                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差
    ( e  r, A8 f: _  p0 U4 `                                ]),5 f9 z5 n+ {+ K" D; h0 T1 P
        # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
    & y8 ^: ^/ o+ ~    'valid': transforms.Compose([transforms.Resize(256),* P0 [* I! u& V
                                     transforms.CenterCrop(224),
    ( U0 Y7 z; V- A  u% i4 C8 _1 M; f                                 transforms.ToTensor(),; t. @, S$ j# D( U+ m! s
                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同
    : D5 V/ I/ e4 ]3 z                                ]),( ~" b$ F. Z* {8 s1 B6 e; A
    }
    7 L/ R+ d8 P5 ^3 b; T: c5 \5 e% j7 d, }0 u5 o: V- l( {6 z
    1* `3 P& F/ b( S8 t3 s2 t
    2; U& J+ i* `9 D
    3- f1 O/ H% d" k$ R& `( s
    4/ w$ F1 T" v+ S" G( x& U7 [" B
    5
    : ]6 ]& Q( y6 F* h6
    : K+ }) _* H& b0 @- J  e% _1 N8 F7
    * a' _$ Q6 X- X& _( h. }0 `5 R; I) c8
    # o2 a, z1 `  r3 G' ^9
    7 e6 \# E& l+ b& t! n0 }10
    + U+ a4 x: [; d) h: O1 W11
    ) D3 o4 m$ a3 A12
    ) X" v9 q7 F8 e' }- v" n9 }& S136 s. U7 z0 L8 k
    14; ^! v! d  }, p& m
    15
    . r; j0 b* n  |- _0 A7 p16# w- W1 A! V- U& `
    172 {% N: k- b0 e$ H5 M$ {
    18* _4 t/ b7 Y0 N
    19" ~6 I% h4 {5 F' r# V+ r  z5 a
    20
    . Q- n5 L1 t( ^8 l# D21
    0 c$ i; ^" b) X) T+ X: ?batch_size = 8
    5 p- K3 k  _- a3 I1 j9 Timage_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}/ C$ ]" j. h. n
    dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}
      B' Y! }: d6 t; E) B" i8 j! {1 ~dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']}
    " ~$ P$ M" V) bclass_names = image_datasets['train'].classes, v( P" s# d) V: P9 O' ~
    ; U* A) ~8 c" J% n, B
    #查看数据集合
    & T2 d* E0 a7 {# j% x# uimage_datasets' L" M1 m3 {! p- @; S
    2 E5 f8 t5 N+ r$ l9 c2 z
    1
    ) u& e/ a* O5 j* |; j2+ U/ Z. }5 D1 U" b& q5 I
    3
    ) _, j, V: H' l8 g& a; o- M44 R* \9 H  @$ a/ U- m- ~
    5* ]# b, f/ N! I: j% h+ K
    6
    9 w, g3 ^9 p) n3 {- S! @7
    ! O8 u3 ?  s8 j85 z# D5 b5 W/ N( F' x
    9# T# T, Y4 X! x" q
    {'train': Dataset ImageFolder
    " X7 D! S& |$ ?     Number of datapoints: 6552
    3 ]. @& T$ \! N0 M( x2 ?  [* S$ r     Root location: ./flower_data/train' B. o3 V0 D/ v: X: ]( c6 N
         StandardTransform
    $ I4 ]+ I$ I( [2 a; D% k- R Transform: Compose() P7 U8 U& B% i" j2 R
                    RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)
    & g' T. n6 u! z' v) c                CenterCrop(size=(224, 224))/ j3 C/ k* o" L/ l
                    RandomHorizontalFlip(p=0.5)
    , ], s# L5 C6 }+ |# ?3 G* H/ V                RandomVerticalFlip(p=0.5)1 U2 p, z- Y! @2 b- Z
                    ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])/ v: @3 {. S9 \% R3 p  ~
                    RandomGrayscale(p=0.025)
    2 I. }! d: u7 _                ToTensor()9 [- }( \" k' Q- T4 t" ~! c; t* d
                    Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]); A4 m! _0 H7 ~" W. j8 `
                ),  d  P) i5 {: J9 x
    'valid': Dataset ImageFolder
    - `3 D+ M, l" o     Number of datapoints: 818: `: [) @7 Z$ f0 Q, a) y
         Root location: ./flower_data/valid6 K" e8 C. y0 S
         StandardTransform2 h/ b4 |4 N3 ], N4 J' ~* f
    Transform: Compose(0 G) j# I( W* A& C
                    Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)3 w/ S4 p5 Q: K" i
                    CenterCrop(size=(224, 224))3 y( s- s. L* P2 D
                    ToTensor(), Q+ ^, Y; E, M: s& z$ Y3 F. O
                    Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    , g1 {, K' P# f            )}8 c) k4 [8 `- P3 U& K
    ; O0 c" g! |2 G
    1
    # F' o5 e) O" z! R! s8 r2
    1 ^8 o5 k( S# V" Z# P! F, c: s3) A( u' ], q; y- i$ z
    4
    , ]1 X+ s! U8 Q( P53 x+ \6 u+ B3 z- ^4 g# E0 C1 G0 f
    6& N! S' y# E2 |3 s; e/ K0 i
    7
    1 u0 D- d- ~! `) Q8
    % f1 _2 H! s. Y) H; C. U) W1 g# Y91 r. w8 C$ f9 j+ G( r' i0 G" R
    10! b& H- M; l0 v! {% i6 u1 S; |
    11
    : z# `# R% X9 d& B: }: I' }5 O122 d% f# M2 X" f$ j. I
    13
    6 i% \: F, j( p3 b: Z8 p9 H$ U# B14' P8 p' s5 v7 R4 {: a1 i( I
    15
    6 ]% J5 Q% q3 x' K! B" a: @16
    8 k3 e$ a6 I$ r/ [17
    / d* c+ c& u9 V; w( L; ?# `18/ ?7 B9 v  p$ z/ H. z. ~
    19( ^9 c% n5 i8 K" b
    20+ {& i; @% O1 K) \& x  w( O
    21
    $ V' `# W" A& C4 t6 L$ Y8 e2 [22' |4 I) {7 ]. h& ?
    23
    5 t$ p/ F+ j' H% n' E/ [% ?24- O0 `4 f( Z% l
    # 验证一下数据是否已经被处理完毕
    3 u0 v  L4 s* b! A  C9 @& _8 Udataloaders0 _* {0 u$ |. d& d
    1
    4 N9 A" s4 {% c7 u3 g+ K2
    ( V9 G, a/ V. N2 k' l3 n3 @. ^{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,
    * u: C1 g4 C" w5 D" a( H+ P 'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}
    9 T6 o" `  B  D$ n1
    2 X, ^4 E8 G: y1 |' L2
      \, P* R" W# k. D( m. @) wdataset_sizes- r8 Z" h' q% s9 @/ ]
    1& e+ T: |) d+ G5 K; {
    {'train': 6552, 'valid': 818}2 p5 ?! B/ b* }5 @. g
    1. H! F$ {5 e% Y
    读取标签对应的实际名字* P8 D2 @4 U+ U- F: G
    使用同一目录下的json文件,反向映射出花对应的名字3 c4 W6 _) l! {3 i& t/ f

    0 f# Z' H3 L# `5 R2 `0 \, i4 p. qwith open('./flower_data/cat_to_name.json', 'r') as f:
    ' X8 Y! |" W( }& u! Y    cat_to_name = json.load(f)
    " ~  P) r5 H. c3 b( c1
    ) }! P) ~# W$ F: e; d: R2' C- A# L  y$ [. L, |% E* M9 U
    cat_to_name
    - W) @0 d* E6 F6 ~, U% r$ T1
    , u  d. z# Q$ D. o% W& I{'21': 'fire lily',
    ' l' {" Q" T; ^8 Q '3': 'canterbury bells',. y( R/ x) w  U, E+ ?
    '45': 'bolero deep blue',
    ! Z5 \2 K$ m: O '1': 'pink primrose',
    * p; L# q% k0 b; f4 j8 `  d2 z2 o '34': 'mexican aster',4 M1 m! l" ^/ J* a. @
    '27': 'prince of wales feathers',
    9 Y2 m! o6 A( B1 E9 @ '7': 'moon orchid',8 u# _. ?. N9 Z, P5 ~
    '16': 'globe-flower',% T* V6 U  N: Q! j; b- N* a
    '25': 'grape hyacinth',# b) f0 D7 A. l- a# x
    '26': 'corn poppy',+ P/ R0 O4 g% A# s8 r6 m
    '79': 'toad lily',  l# N. i: C- g% m4 Z$ U5 O( U9 l
    '39': 'siam tulip',9 U5 @; `  j3 U. `6 j2 r3 w. E& @) y
    '24': 'red ginger',
    $ Q/ \0 l9 G, x4 e% a '67': 'spring crocus',0 Y* f& I: p3 p2 X- l
    '35': 'alpine sea holly',
    / M) J$ N- n: ~$ ~# v '32': 'garden phlox',
    7 P+ E" Y, P5 |) M: ?3 F '10': 'globe thistle',) \: {4 y- q, Q9 @( ], d
    '6': 'tiger lily',
    " a0 O9 u( ~) b( o6 @ '93': 'ball moss',
    5 U9 Q- L# k9 {- [4 M '33': 'love in the mist',
    3 V: `# F2 m+ ]7 _ '9': 'monkshood',
    . h6 @. X' ]7 L4 s# c '102': 'blackberry lily',& {+ ?6 Z5 e7 F  p
    '14': 'spear thistle',! J* Q  @6 \) M5 O
    '19': 'balloon flower',
    3 a0 V2 J: _0 x4 W '100': 'blanket flower',
    # N/ {' n7 ?9 A1 D( o0 A5 X; H, F# ] '13': 'king protea',
    . B: g  d5 N7 G '49': 'oxeye daisy',
    & p" J) L8 v: q: f5 ~# o# r7 J# _ '15': 'yellow iris',
    ' X. c( b* {  r5 }  m, n '61': 'cautleya spicata',' O& h7 ~/ W! Z% _  D
    '31': 'carnation',
    4 r0 @2 `% t7 G! x' l '64': 'silverbush',
    8 r9 X) i: c8 H. k  `& \5 g& o '68': 'bearded iris',
    , B/ Z& N7 C) \9 x; |* k4 J4 k '63': 'black-eyed susan',
    * X" m# \- V/ m; @ '69': 'windflower',  \5 z2 n; c. x) c, _
    '62': 'japanese anemone',! W5 G0 Y0 s4 ~( f
    '20': 'giant white arum lily',
      u- ~% B1 s; |$ H '38': 'great masterwort',
    # d; P8 |" M7 g6 p '4': 'sweet pea',
    * W- J. T% ~1 Y& ] '86': 'tree mallow',
    " j+ l5 S+ x2 Q/ J5 b9 t8 R" r; R '101': 'trumpet creeper'," \0 m* T( v# u( I- w& Y
    '42': 'daffodil',
    7 p9 P; r9 F  @3 W '22': 'pincushion flower',/ F+ ?" C: R4 r# ]: q5 h
    '2': 'hard-leaved pocket orchid',
    $ z! W0 `( f/ k9 l* O '54': 'sunflower',$ A4 a. J2 \1 P5 j  i1 U- g# g
    '66': 'osteospermum',
    : n! e+ V3 R& L9 g6 ] '70': 'tree poppy',
      ~' f" c+ O2 M0 [ '85': 'desert-rose',( \# O1 R- t7 P& f% Z& Y
    '99': 'bromelia',
    5 a5 i) ~1 f/ P$ f4 _ '87': 'magnolia',% T1 U8 ?1 ~, `3 }. |
    '5': 'english marigold',' b, ]% S. U1 s' a& \. o
    '92': 'bee balm',7 x1 Z# v. v* m6 q
    '28': 'stemless gentian',# d& P# Z0 Y, c9 ~9 [1 j. w& Q* m
    '97': 'mallow',
    9 c* w, I2 t3 i3 | '57': 'gaura',& P: d4 ?0 n6 S5 n0 j+ y5 H
    '40': 'lenten rose',
    ! P! }; E" ]: C) t, i* X3 A: p '47': 'marigold',
    5 j' o: ^# ?2 ^- G; ]! r0 _ '59': 'orange dahlia',& b. s1 `% Y( ]0 u
    '48': 'buttercup',
    ' j8 Z/ L, s( A! E+ L2 Y( O8 I '55': 'pelargonium',, i2 P: g' D, l0 a+ n6 o0 \
    '36': 'ruby-lipped cattleya',  ~) r& l) j. s0 r, z
    '91': 'hippeastrum',
    + C! K1 i; E/ u' o. H9 Q '29': 'artichoke',: K" F1 Z/ B+ ]% B- n: C9 Y
    '71': 'gazania',5 y/ D7 ~+ d1 L8 R0 _0 o
    '90': 'canna lily',
    . Q) L+ l, O* @0 @& P8 R- v. L '18': 'peruvian lily',
    6 c- `6 U: }" F4 U" l '98': 'mexican petunia',) L8 |& ?7 |( ]8 z9 t% i& q( N, A
    '8': 'bird of paradise',7 B7 H6 L( e0 s$ M$ a5 E
    '30': 'sweet william',
    , `$ {* f" `" z. w4 x '17': 'purple coneflower',: w8 @# W+ [0 ^; G+ W( {
    '52': 'wild pansy',( ~6 m. C* f" H3 x' Q$ m
    '84': 'columbine',
    + u( q! |4 S% w, b% V '12': "colt's foot",
    & N. p) W/ H9 v" `/ l$ p4 B '11': 'snapdragon',
    # v9 M! y- ?* B  d8 `0 c% O '96': 'camellia',
    ! ~7 @! j  U5 W '23': 'fritillary',
    4 H9 w( f! v2 d! v9 S# F8 W '50': 'common dandelion',
    / S. |+ _. K' i! C: Y0 q '44': 'poinsettia',
    8 h: M3 d5 w, o! n$ r+ O* ] '53': 'primula',
    ' q) N9 Z, ]/ z6 z+ [ '72': 'azalea',
      ~7 y# c1 Z; I7 Y( M+ N '65': 'californian poppy',; c6 k, `% Z: p2 F$ R9 O& V7 }
    '80': 'anthurium',
    " s8 x5 }- A$ ^6 { '76': 'morning glory',# o. @8 W8 m0 _+ k
    '37': 'cape flower',: N# ^: @* |: z1 W* h8 N
    '56': 'bishop of llandaff',- B8 K6 e# Y3 R
    '60': 'pink-yellow dahlia',; k, s" H2 y% P. S" F0 r8 w* p
    '82': 'clematis',+ o& u- d& |& I$ C2 _/ @
    '58': 'geranium',' t# J, o$ t( R. K/ J) V5 j$ n
    '75': 'thorn apple',3 Y3 x- M4 [# M8 d' W
    '41': 'barbeton daisy',; s8 b  n+ R, ~
    '95': 'bougainvillea',
    5 ^  m4 |0 F' L' }% g '43': 'sword lily',3 m! B! G: T2 h; B
    '83': 'hibiscus',
    6 U, V' ^) m% H4 ^3 d9 g. z '78': 'lotus lotus',
    ' c  w3 F6 I! N9 Q% Q '88': 'cyclamen',
    - ^7 ~4 Z4 q9 h' _3 e' m- m '94': 'foxglove',
    5 x  I; L; S! m- O, X, i4 R& d '81': 'frangipani',; j+ O$ y$ q# |9 T: y
    '74': 'rose',2 K) d) L5 }; r! K" M$ }  F
    '89': 'watercress',
    % A3 }5 Y5 H+ Z  b '73': 'water lily',
    $ U8 @! E3 [% z '46': 'wallflower',
    6 n: m- |+ s. b) X7 o- D '77': 'passion flower',
    3 N7 F( _2 {+ w9 f" p( | '51': 'petunia'}2 _" Y% l) s6 B5 q( ~1 n4 Q- b  E$ _5 d

    + A& Q5 f" z9 ~8 q6 z. ?1
      {* l; e% X3 X2  d- I, {2 v1 n) a5 J/ I  l
    3
    3 Y% k. @' Y( k# a5 d; W- G; A# O4! a& S4 p! d$ T
    57 Y, {# z- M, U8 g
    6
    $ J1 F" T: y& ?4 P4 V7; T8 n6 a4 O! ]# g+ ^2 w: O
    8
    $ A- Z" a! R1 y; `7 c. d( C99 F' c; [5 W; H; W: H, g# \
    10; L$ J  w  Z* e, @1 z2 _" i+ l
    11; i! Z+ z: D/ n# q3 G2 f- q% d
    12
    , _" B3 g& v, c3 M, d13
    $ D9 M8 Y- X* }! x0 Y: o# S14
    . D3 }8 {8 |8 S1 S' m. n15
    ' V9 K) }" n+ c2 \" y) S163 G; _. Z: V) _, ~
    171 ~0 G  J/ j- @) k8 f0 A0 v
    18
    3 N) \/ D9 N1 }' ?, {19
    , ]% |2 k9 e! h) h20, e8 g* H; C# a8 P- B
    21
    + L4 L2 ]  g, w2 O22
    * `2 x6 |' \: J+ L+ [. E' x23' t, `  J' K) R1 m) h
    24
    & G; k" ^& S, ~# W! Y25
    & W9 I& U8 p* y+ q0 o26$ o* Y" u$ Z& o' }- q
    27+ A$ }" t/ ]5 e$ F
    28
    4 Y8 f4 T0 l5 T4 M, p29  ?( q; ?) U5 X  I; C) C; V$ S
    30
    ( l5 \! ]# S! g+ b31
    ( A: j; [9 W/ o" N: f: \2 g( ~9 b32
    " I* D& Z- I6 v33
    $ h0 o1 j+ i0 f$ @0 e: u/ \34
    ( F" h* q- @9 O: f: F- S2 Y35
    4 G$ p- J9 D& _- a6 z365 y3 Y+ f' H% c- \  p2 A
    37
    ! \. b4 }6 q2 ~38
    ( q9 f8 K$ T/ u1 O39! {& n3 B$ S1 r7 r0 M
    40+ r0 C. @, F# F. M
    41/ K5 c/ G$ l6 _6 C( o; P; T
    42
    & x  K/ y9 W# u, L& ]# F" a* o43
    3 |8 o6 }3 ?' O! a2 P441 ?' e. @& i, v8 u* Z6 c
    457 v) h- A+ K2 r7 q' m
    467 ^4 n- T, q, j1 R
    47
    6 V! b5 ~7 V5 q+ \& l48
    , J& N: \, L7 l/ V! {496 }* m5 z/ L; ]4 ^' L0 k0 q
    50% j- ]/ Y0 B) w% s6 n2 Z$ n
    51, Y; A- [( H* A0 E0 @8 e0 q+ n
    52. u( g* x' s3 m8 a4 M
    53
    % {$ V8 u+ q: b3 Y( O6 g543 f) `/ [# B. j! i$ `
    55
    ' k1 F8 n8 P* q9 y. W% {2 J56
    6 d% S: R, S* L' S57
    ( H2 N* H7 U$ F58/ l! x( ^: O" a) V/ j
    59
    + \7 e+ U3 H1 u60
    : d+ f8 ~$ W) d$ `' `$ J61" d. b' \$ s! W5 z6 r
    62
    ; ?' z  Z' q. f63
    9 T: S# q: @/ m: h8 z* _& u- _7 B64; N0 |3 ^7 Y" w& ~2 |. |" N
    65" M  j2 j" k. ^  m" g
    66. G2 @- i6 ?) S: Y  l7 T
    67
    2 J/ s! S7 ]. _/ w$ V  r; r1 R68
    0 D' S# y8 p7 o69' h" P& r5 S, D- b+ r$ A) x
    70
    9 h5 R8 D! U3 K! x' r8 C71: _1 H6 u5 q2 f# F5 o$ ?
    721 G) l! q, u1 t4 K$ D9 d2 \- u
    73
    . o; I/ M) q* x" c& }$ z74+ U) K1 E* u4 ~/ E0 ]7 f8 `3 }  w0 ]
    75
    ) _0 M. z5 ^* i4 n, u( N765 h& @0 [, ]& d' Q" z
    77
    + D5 J6 S" b% _9 J78# ~7 k2 }% a- _9 S0 B5 }
    79$ q3 f9 q& W+ y) J4 z( n) l
    80
    : g5 {7 U4 ~# C2 @( O6 ^7 z813 Z9 h' N0 L' o7 s  Q
    82
    9 ~$ H' h9 F  j! g0 z83
    8 a5 S  r/ t7 P1 m7 ^  G" Y5 @84% f# D/ ^% G& }9 H! n4 d
    85
    ; {9 G, Y5 |3 `9 ?" j. [/ W86) `( Q+ x9 k  ], W5 ]6 n0 f8 B
    87
    % A# {" w) s/ p3 A88/ C: N/ W6 W/ l8 m  a
    89) y  j6 n" R5 v' a+ z1 v; ]
    907 y- ^, h% F5 i& P8 B! B. i
    91
    + r. |# @9 O0 h$ C/ E- l2 ^8 ~/ t2 k92/ C( U4 b6 Z& @8 M/ S
    93
    " V3 ]5 z( r8 ~8 a94- j  u( \7 W+ w" l; M# f7 f; E
    95# L' x# j/ R2 ]$ c1 g
    96/ ]4 Q# D( ?! R8 G) r
    97
    $ Q' K# ^5 Q  e: C8 L; \( Z% E98
    - H; I" ]  p9 R7 ~99
      W4 G9 ?& L0 }: F# l( u0 E100& b% l( e  p5 F; G0 }3 U: I8 j! M/ U
    101
    # _5 P; i# R, c9 G. K102
    - C1 X* W  @, I4 U- \4.展示一下数据8 k8 B; H5 A5 a' E: V
    def im_convert(tensor):1 F; m. `' Z- W) r6 z, o
        """数据展示"""
    ! v; D2 x& Q1 Q' g# x: ]* E    image = tensor.to("cpu").clone().detach(): n" F# ?& \' ]* K1 f$ F
        image = image.numpy().squeeze()
    9 A, I& w# d# |$ S2 p! e; m, n    # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图
    5 U0 c. _) B: `6 I1 G    # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)8 ]$ r$ r6 B  [; X7 w
        image = image.transpose(1, 2, 0)
    + ]# b. F  S' @+ S    # 反正则化(反标准化)
    1 L4 H$ l) x$ g/ d    image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))  W7 I: G) M6 k+ O  V! X4 o$ V; r
    ( F& h2 h3 E4 R, v7 a. t
        # 将图像中小于0 的都换成0,大于的都变成1
    / h2 h+ U/ u" B/ C9 h5 v- P: `    image = image.clip(0, 1)  a! M1 _* H  C! N

    4 l+ ]- e! U( z* X1 ^8 D: O: [    return image
    ) _+ i1 Y# U8 Y1
    8 x6 m( }# N' E& O0 o" m2* W1 f# {- Y( k/ t0 ~8 R
    3! `9 i) ~# T5 t2 N  [' _# n
    4
    / M* P4 C/ [0 S' L' }5
    7 E# C  E5 d. ^: Z) c6 e1 L8 Q7 j6
    0 b# j% x! A# g" V8 s% ^72 q  ^, }3 k& W
    8' x+ ?( E0 J5 _' c. K% g- g0 r
    9( x2 M  g7 S/ o% i' y+ z2 N+ e" _3 c
    10) h) |4 {; r8 E# T$ q
    11
    * ^! N) @, f+ [' n6 f; T3 Z" c12
    % O' H! K* {/ H* M' C13
    ; ~7 `$ x2 M5 r! W14: O6 N3 z$ J1 l2 k: m/ J$ X% Y
    # 使用上面定义好的类进行画图! \4 ]1 K! K) O9 w+ a$ e
    fig = plt.figure(figsize = (20, 12))
    4 s: O7 H8 ~; Xcolumns = 48 s! {0 ]$ Y) Q: t( f
    rows = 2
    ( |- e, Q- P) e* `. U  ?7 E
    7 d3 o7 _9 C/ E: l" K) c# iter迭代器) y4 I( ^6 w7 c' m# V# A+ l" U  C5 f
    # 随便找一个Batch数据进行展示" {  U  m* C3 |5 K9 ^5 Y& l. y
    dataiter = iter(dataloaders['valid'])
    1 R/ L) a0 F4 i- I0 n# i" [inputs, classes = dataiter.next()$ z! @8 W6 L7 F+ v; A; F0 i. o
    : V( a2 K3 Y  A: F- K
    for idx in range(columns * rows):
    / j) t1 l- L( J! b$ F    ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])
    8 r, G9 u3 P" s    # 利用json文件将其对应花的类型打印在图片中) N1 D' F/ d$ w4 x! I1 X& M$ k% @0 S* B
        ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))]). c/ d2 X) T8 ?
        plt.imshow(im_convert(inputs[idx]))' t) ^) g! I' L5 m' {0 O3 u+ J$ i
    plt.show()1 d" h1 X. u6 a. E
    2 w0 A/ m' p0 c; Q1 d2 i! K2 t
    14 M8 l/ q' ~/ u! X* B+ }$ b, z
    2, d# g+ E- G6 `7 n
    3
    9 d) P+ i. v" X2 g! P# p0 K4, O" Q. z# B+ F4 ]; @$ S% t8 B' ?
    56 j$ \5 o# i+ B' u
    6
    : {/ H* z9 D! ?. E. A: f4 ]6 n, q7
    : B$ ], J! C* z6 D' O  h9 h8& p' j. |' M3 l. K
    9
    1 q& u9 }# X2 r% V7 _- Z10
    ; d' _! |7 C, r; s3 a8 I11* A8 @' P( {" R* x; O+ a/ `: m
    12
    # d. ]& N! X( M& N0 }13
    ) o' u* h0 m- R; a( s2 p148 n4 ~4 Y9 o# C7 z6 d
    15, a1 A; x; M$ J, R
    16
      M( s7 }' D8 z5 B9 m9 J8 a- g$ F. y  _0 S
    . E5 [/ O# g( m" m
    5. 加载models提供的模型,并直接用训练好的权重做初始化参数6 D  e0 t0 ^0 t# N  v2 S# l
    model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
    $ e9 J# h# V7 O+ B+ [9 q5 ]; Z# 主要的图像识别用resnet来做
    % z7 |" i9 s% s% A% K7 D# 是否用人家训练好的特征
    , H  \# k6 b( Y# gfeature_extract = True
    - R  a3 S! |2 B) o2 p0 \, m( F1  L" ^# _3 [& E! a4 z
    2
    6 R/ @% H& w' z* F. P3
    * F1 O7 u3 O7 r4/ M4 N8 u9 m* S/ y
    # 是否用GPU进行训练" j5 Y" S9 T8 H  a0 V1 R+ ?3 R
    train_on_gpu = torch.cuda.is_available()
    - z* S3 l' J4 r, I
    7 P/ g; f/ s4 e- C5 f: J7 }* Z$ @( Yif not train_on_gpu:* w2 e& h- z8 c
        print('CUDA is not available.   Training on CPU ...')
    3 J$ h6 l/ b5 [7 X1 Xelse:
    : s8 y+ i# L% \8 H: \. x+ q    print('CUDA is available! Training on GPU ...')
    3 d. w/ f4 O; `4 P+ ~, P* p6 _! r. ~3 z# M+ i
    device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')- e( m+ V8 \4 c; ]& u/ l
    1; ]  J0 u; S: i& P
    2% ^8 O5 l: Z) B; I
    3
    + u8 n0 j4 |) A/ ^. e& U; J46 R$ M7 u* {/ H. Z6 n6 O
    5
    + `0 D2 E. c( C$ x* Q9 J4 ^6$ a8 ]- F1 |' f  \1 u) B6 i! j
    7/ N+ M$ a9 `1 y" _: Q- G( O5 {+ A
    8
    " E6 q! E, Y* ^& [9
    $ D, ]( c1 W% R# RCUDA is not available.   Training on CPU ...
    " H* w# {# J: ?' O. Q: B1
    3 R  u. s! \5 h! n. D# 将一些层定义为false,使其不自动更新
    0 z3 k8 r' r0 c/ P2 @. d( s* adef set_parameter_requires_grad(model, feature_extracting):. }- g. E; p! V0 f$ Y
        if feature_extracting:' u  l# m' H3 [, g0 B' e
            for param in model.parameters():. m% T9 H' ~0 u  G' p
                param.requires_grad = False
    " r& i5 `5 t, b3 V9 @. }  I19 N! D. s& H" O* D" P( x0 N
    2
    0 z6 J) A  d, t$ W: b  _7 K  J34 ?9 ?2 w: G! l& [
    4( z( F: c3 Y/ Z9 y! q
    50 R6 i0 W+ c0 M
    # 打印模型架构告知是怎么一步一步去完成的. |' U5 q6 D+ q3 U% f
    # 主要是为我们提取特征的+ F" G& R: z! K; f
    " K- k% n( @/ [( W2 @
    model_ft = models.resnet152()" u" \$ n+ _$ e9 B
    model_ft
    $ K5 d  V4 i# {, v% p' G! B1# l+ K' Q5 q8 ^* W% w+ q' H$ k
    2
    6 a( p. K% t/ g+ |$ J) y- z4 a( e3: j' h6 e. {. I' `0 V
    4
    8 Q2 X3 k! m3 t: R1 o5
    6 G+ {, Y: y  o) d( l1 hResNet(% L$ y" n7 `( `9 |/ S  X
      (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
    ; h0 T0 _$ ~' A* n6 Z0 }  (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True); m9 Y; C. u) r/ N) b# u5 ?5 w
      (relu): ReLU(inplace=True)
    : f' ~. M1 x2 C) Y9 p  (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)3 K+ T( v4 h1 l9 k0 n# K
      (layer1): Sequential(
    ! f. H) t) J" b5 P7 p    (0): Bottleneck(3 P$ x  _1 D% {$ b$ y
          (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)! h$ j( A9 `& U5 z4 m1 _* S
          (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    / P+ ?/ I4 @: h9 X. `1 A3 \      (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
    & I( C! o# a1 |6 ~* C6 b/ b      (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    4 P/ S. H2 g4 T8 e( ]% c      (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
    - V6 ~3 P4 }! H- x9 c$ h      (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    / _8 S- @* y8 a* R: s) ~7 T0 r      (relu): ReLU(inplace=True)2 k. N( _) Y, ^" d4 K2 |
          (downsample): Sequential(9 o& n- U* R  F, j# Q; Q/ q
            (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)( y* O! p1 F" P) B% v; a; ^& V# b' l
            (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)1 c% V, Z1 D1 W  L' o4 u& F; K3 m
          )' x, r8 `  C2 H+ \" G7 z" ^
        )
    ; Q& d3 l4 ?7 D1 ~; @( T6 x中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。3 P9 `. H" _) q" W
        (2): Bottleneck(
    $ @/ u5 b: e6 D$ }      (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)  s* ]7 a3 v* \5 F
          (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)+ C0 _! b9 P5 s
          (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
    4 [# n- E9 e4 c/ B( }: y9 d- }7 D      (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)& L9 n: B- X4 B' j, Q& Y* I
          (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)) t2 o, v2 u. d3 y% R! D% U
          (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)* u- @$ I, [3 H% x8 ?, O9 S) `7 v
          (relu): ReLU(inplace=True)9 @6 c7 W6 U6 I% Z( I; d- a( [
        )
    0 i% W+ }: R* l$ ~% ]  U4 B  )
    & c" _' O1 p$ a2 w( b' O  (avgpool): AdaptiveAvgPool2d(output_size=(1, 1)). g' a) S! h/ Z5 Y
      (fc): Linear(in_features=2048, out_features=1000, bias=True); }, U( l5 Q& e$ J
    )
    , n) @- \6 u; J( h
    0 N4 M8 A% ?/ e9 ^8 c1
    , @3 M# l& T' M% x) ?2
    ) Z# X4 M% `' w3) |8 ?1 t- W' P7 D" U2 k% \% b: V
    4. F8 r5 B4 R. J5 O' C5 n! p9 x' }
    5, d9 w- l- T9 x( _
    6
    2 B4 X* r/ C" t+ V7 D& s3 e78 F( ^0 J% Y4 ]  R3 c' E; B
    8
    8 q1 c8 r9 u, o2 q- F: H) h9% Q4 u# J% V( d7 ~6 n+ u
    10
    ' c' q, `0 A  O( [# _11
      Z& z! ?& Q! X2 L) f8 `12/ ]" q2 r/ J5 [0 c7 G
    130 w/ e! F8 M+ [8 v4 _
    147 i! G9 P) K" x/ B3 ~) g
    15. ~7 v2 T' ?8 Y! E9 P% e
    16( l1 }2 Q# a3 R5 Q* O; \
    17$ R' K7 h; |  k( @8 W+ M: ?
    18  A* q) a8 o) b) h
    19
      z9 L/ n3 ^8 S) ~& r3 a* X20
    / Q$ B# D9 w9 c" U& f21
    5 K- X! C! Y' ~6 R1 F% G2 L22
    # ~0 Z0 V+ p  K/ u6 U2 u23
    & n8 B( W/ I$ R; }+ ]/ W& i  p24
    9 O3 l+ B% \- P6 U& j. M0 |25
    / g. }1 M+ R! p% x$ L) ~$ @265 u' h1 S0 F4 ^7 a- [' s+ l7 m
    27
    - S: q1 R/ m3 i' q+ G/ g. c. \286 t5 m& [6 v8 [$ }# X7 ^
    29" L$ _! p" h: a+ p7 M
    30
    ! E8 G8 P& e8 T31
    4 e+ Y( a! P/ D) s1 b32
    ' I$ j8 c% g) p7 J. x( Z% r' G33* w3 _' j* Y) |0 G( u$ j
    最后是1000分类,2048输入,分为1000个分类  C4 l. h  B" x" M3 P( J
    而我们需要将我们的任务进行调整,将1000分类改为102输出
    " V8 {8 f: E+ ~: N9 J8 ~% `. q5 w8 \  O% r' v0 G; X" g% V
    6.初始化模型架构2 H0 n7 H6 i& M- {! h. f7 f/ a3 J
    步骤如下:  P9 [$ C+ Z5 R9 x' E6 x' H

    0 X) J9 H. j7 f% J' l. i将训练好的模型拿过来,并pre_train = True 得到他人的权重参数
    + h0 K7 m+ L# V# K% {可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)
    ! r" ]0 I( G3 O1 P1 r: Z* Y: [无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数
    8 x' |' b0 v; t& Y官方文档链接$ S$ Y$ V  |" V9 X0 z/ s# }
    https://pytorch.org/vision/stable/models.html' o5 u" s' X, d7 R" r0 b8 {9 z

    * q' M" N+ z- j4 Q4 c# 将他人的模型加载进来: P8 h5 x) f* [/ R
    def initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):& n% H: C$ z4 E7 |6 i, a8 g
        # 选择适合的模型,不同的模型初始化参数不同
    3 r1 A0 N9 Y2 I8 ^    model_ft = None: y& D: x7 R# r) m4 U. [7 ]
        input_size = 0
    ! E; E# I1 ]& a3 y! p' M3 @' O2 ~
    ( E% R  @1 ]. }: h    if model_name == "resnet":
    6 B+ F7 z; i5 Y: S( U. ~% ?  U        """& B$ r0 ?* `6 \# u
            Resnet152( i- D8 C9 q2 h0 q7 j  K' l
            """0 G2 _; x' B8 U7 ~" @$ @( B' p

      R3 P- ]& t& h. b8 L) U        # 1. 加载与训练网络
    3 n2 O: Y" X4 S  _        model_ft = models.resnet152(pretrained = use_pretrained)
    * v+ Z6 W$ m" F) D* c, g1 \4 ]/ W        # 2. 是否将提取特征的模块冻住,只训练FC层" u# e6 e4 t+ T5 l% g2 t
            set_parameter_requires_grad(model_ft, feature_extract)7 i3 r  G) r2 f6 y* `6 D+ y
            # 3. 获得全连接层输入特征) `) g* Z/ I9 y4 t, K" b) w
            num_frts = model_ft.fc.in_features
    9 f& m) T" a# S, o! N        # 4. 重新加载全连接层,设置输出102
    # H5 L. o' J# n' \: b% L1 K' j( x" d        model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),
    ) x+ o  J% v  @2 P) s  V2 n                                   nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1
    1 D- p6 n) F8 O. Y        input_size = 2244 V9 [; T9 w1 x; k. i  S4 h5 L7 @
    ! C7 u, _5 q7 w
        elif model_name == "alexnet":
    3 C! f- g. v% ^# }3 }# t        """
      C! f" j: U3 a4 n7 B# C        Alexnet
      B" ^9 `' m/ I0 A        """
    6 L2 }! T4 F/ ]2 l: u! _% _! ~& }        model_ft = models.alexnet(pretrained = use_pretrained)* U3 V; h& S/ h# t" v& r
            set_parameter_requires_grad(model_ft, feature_extract)
    5 T: c% k; X! k8 Z) k+ C/ O3 N2 _! X+ H& b9 B- b' o
            # 将最后一个特征输出替换 序号为【6】的分类器
    4 I- i. d: [7 @/ P5 |        num_frts = model_ft.classifier[6].in_features # 获得FC层输入
    , m* g7 N0 G' E: C( V( K* C        model_ft.classifier[6] = nn.Linear(num_frts, num_classes)( Z1 U/ z8 M3 ?. l( a) B2 q9 w' U
            input_size = 224, q2 l: A$ v: m. d, h5 Y

    , P- `7 r# }8 h! R    elif model_name == "vgg":7 D3 h% ]0 v' \% Q$ P6 k
            """4 G0 I- p4 C3 L9 u* y
            VGG11_bn, _9 Q# E, _3 D* B: A! |
            """
    # l4 p2 M! R: G  Z1 B        model_ft = models.vgg16(pretrained = use_pretrained)
    ! o4 F. R2 Q( e  f% f        set_parameter_requires_grad(model_ft, feature_extract)1 ^# I; {( h6 ~4 {- c1 ~: f
            num_frts = model_ft.classifier[6].in_features
    # q3 g6 B  ~4 N" D  E. v5 \        model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
    * p, m  P7 _+ B/ C        input_size = 2247 Z# M& ~2 ~7 [2 b2 o( b+ a. a
    , S, ^: `7 u! T5 m2 ~. L9 |# j
        elif model_name == "squeezenet":
    7 K) v0 V7 p3 ?4 j7 x6 M        """4 Z3 V" ]. C$ O" P
            Squeezenet
    4 I" d+ I# C- U  P        """% B0 B  U$ Z/ @% @2 \
            model_ft = models.squeezenet1_0(pretrained = use_pretrained)
    5 B. ]  w  U5 u- s1 a        set_parameter_requires_grad(model_ft, feature_extract)6 c6 ~/ |& Q+ O; t
            model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
    ' R' M: F( I4 ~1 O$ l1 U+ }" N# I+ }: d        model_ft.num_classes = num_classes
    # j  h% a9 I! W, w' B) {) D        input_size = 224
    # H6 }' B& X6 b; h& b; O
    % h/ w8 G% T; k+ ~: k7 ^    elif model_name == "densenet":
    2 l  s1 R6 |" Z7 x& F- f        """( f" x: Q$ d* Y" k& ^6 O
            Densenet9 ?8 S! ~" g! A+ B& {4 y
            """/ _; F4 Z$ w: G' m  x
            model_ft = models.desenet121(pretrained = use_pretrained)# T& w' H  c$ ?; n7 \  s# L
            set_parameter_requires_grad(model_ft, feature_extract)
    8 d" |3 e4 L5 ^& w, j7 G        num_frts = model_ft.classifier.in_features/ `) @. U5 ]  g8 w' t
            model_ft.classifier = nn.Linear(num_frts, num_classes)1 J* `3 }- T, |3 Z
            input_size = 224; [# V% j2 ?! _! K0 d. H0 M! L- L
    % s3 {' s& O! C4 O# B6 i
        elif model_name == "inception":
    7 P4 r- Z6 v. s; R        """# }! e2 o! E1 P5 u$ s
            Inception V3) o2 R, Q6 i0 x
            """
    ) J$ f; r. ]8 y% V4 I        model_ft = models.inception_V(pretrained = use_pretrained)
    : c: S" ]6 Y* N- \* t+ {( M; H        set_parameter_requires_grad(model_ft, feature_extract)
      b  i# x$ r% }& O) j  s4 I" Q) v5 r( @6 `0 m
            num_frts = model_ft.AuxLogits.fc.in_features4 r# z1 [) r9 \
            model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)
    6 _0 H+ k) Q4 {# I3 p, I7 m& l; \! F$ m. }5 B; ?: w2 ?* I
            num_frts = model_ft.fc.in_features, F* P1 {) j9 `! x7 |* ]' S" g
            model_ft.fc = nn.Linear(num_frts, num_classes)& y4 }* t! ^9 M8 g
            input_size = 299- L% \3 y) ]% j$ Z6 p
    7 K6 u  Z/ X5 C, D" t4 A
        else:% v$ X0 b* _: k
            print("Invalid model name, exiting...")3 r" g( N/ k+ D
            exit()
    ; C4 V6 N; x  P* Y8 O5 U2 w  ~- q, Y  h" D! t/ ~" O, X7 D' r
        return model_ft, input_size5 X, P( m5 X5 c8 m, ^$ s0 z
    ; m: b* T! C" r5 z/ c1 Z; w
    1
    6 R' A) C- t( E9 q1 X2% V# I7 a" Q! q3 S5 Q# p
    3* w; T. x8 C6 u7 B& H
    4
    5 u+ `" E4 E5 t* d- k5. N& n( {6 ^0 \
    61 k% J7 _$ A/ w/ ^8 z
    7$ e/ P3 V) k* ]. p
    85 H% M$ u# u$ U9 {6 y3 M- M
    94 |* L0 @( ?( a) \
    10
    . A5 w) Z# h9 ~. h( ^" e11  x  M" G' Z( s! }# ~
    124 X) V% g  L% W3 s
    13
    % K4 j/ A7 j3 w( e" E4 N14
    : z& @  m8 I5 K5 K+ P2 \15
    0 d2 v! X3 @% _0 N- [* _16
    $ u; [, O+ J' r* B8 S8 {8 s6 V17( B* y! E: }2 m7 u, Z, X; f
    18+ o5 f3 D* f  U$ ^8 w) z
    19. e# T) E4 r- q
    20! ?: P. i: a/ P) J9 g+ n& R
    21: C- {: g' C; P
    22
    + |1 J2 g! j+ x* g23
    % f2 s, b. f3 F0 P: x, G0 }! w24( h! i' v$ P6 n* S, K. O& F
    25
    : v: T9 F' `$ W8 V! C. Z2 |( j26
    1 K* f. v6 G' C( z+ n- J$ |& Y( H277 h5 k7 f( o; x; T% b- ?' M
    28
    ; B+ X8 ]& T9 V) [) P! @29, V3 P8 [4 W0 C  w+ w- t1 j8 z# n
    30
    $ p( F. C5 e! U% |3 y$ F31
    ! _* v3 B3 B9 k  _; {! ^329 a& M: F8 T. E) A, x9 h: A
    33
    7 U0 v! H& i; f  ^/ J/ q; X3 V( C+ ]349 X/ ?, V) X/ m! F! }/ r
    35
    + J4 k) L3 P9 S8 V36' g( Y: l# c! M" i& Z; t7 `
    37
    1 X# w& A% E, @& V5 u38
    " Y/ f7 l$ S8 ?1 E39, h/ R3 R" A% U8 ]
    40
    6 Z5 L3 ~: s$ E& Z; e41, z+ O% ]" F2 p1 S' B5 Z1 b
    42
    0 u! X7 |# K/ a. r( G/ y43
    ! d7 o# D3 h7 F- b+ s44
    " g! q( C  ?& \0 W: x45; y6 L. n$ E3 t. o: k4 a& c
    46
    . F1 F& P" F2 O" \+ U47" w9 c' L8 ]& w
    480 T; K; ^& f$ c2 i
    490 H9 c8 W* p7 y
    50
    ! Y- J( ?) {: K' h" W: `511 _* i; I; n- R3 m, C/ S- R* q
    52( K0 R6 }4 G3 y8 F
    53" h& @4 Y8 v4 E9 l$ e
    54! V5 b8 [8 T+ A. y( c! _9 {
    55
    % ?" o, F% i' f" h# N3 k563 \% ?9 `! a, A/ k& T6 _
    571 f- U: {( F$ J! A. f
    58+ y5 ^8 E, B! w8 j6 Q, ^
    595 m0 ]0 L5 _3 h+ f6 N
    60: w8 `- b1 w% t* a, Y7 Q
    615 d* c; O( }2 e5 N  [
    62
    + V1 I4 M# ^% b/ G! [* g63
    ! |% V) z5 l& L" m$ g& l2 V, V7 V640 u4 I: i$ j# }, P/ |
    65
    : w) K/ j: D4 x. ?66
    . a( C9 p0 j" G" @  ^67# Y# }; t2 A* t% R( @+ z/ f
    68
    & c. F' e4 q# `! F4 z- A- p699 i$ x. ^1 [6 ~/ `0 ?
    70, @% E; e( Q1 Q& s
    71+ p& z- E1 p$ J
    72
    % r& b# z$ f7 t- C: k& z4 v$ u73( t, h1 h% D4 c! d, j# g
    74: f- a( h# Y: u# f
    754 R, V% s7 L) J% g# _& V: e7 Y# d
    76
    5 |9 y' s, C, Y  v/ z& e77! r1 x, @' a$ c9 ]
    78
    . b; }5 L, L, f791 `+ y% B! H* a/ B0 M
    80
    % }4 _" {' z8 f81- E: @. j+ ~7 N: w) ?& d% V: n
    82
    , d6 ^- {. |/ Q+ `83/ u$ L$ x2 B* T- t
    7. 设置需要训练的参数
    $ R, K4 `# r; R* I8 |# 设置模型名字、输出分类数
      d7 l! u6 S7 }" f; V, A- }% Fmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)
    % H- r/ u  a7 m4 q1 R4 C) `( y7 s6 y; K5 b* a# E
    # GPU 计算
    9 B& Y& z6 X: d1 W7 \9 J$ wmodel_ft = model_ft.to(device)
    2 a& i6 T3 g% k: c7 F* n6 ?3 h/ y( V7 o( W$ N7 V
    # 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取4 J( V. u" k6 }( `
    filename = 'checkpoint.pth'; m# h9 Y( I7 R5 d

    $ W4 x6 G4 B* j$ D  Q* G# 是否训练所有层# k6 W, a  p* e, H5 u
    params_to_update = model_ft.parameters()
    4 `: a  x* T8 |, ~# 打印出需要训练的层4 U0 Q1 _, H6 [$ D0 ?# m0 ?
    print("Params to learn:")
    ' s5 ^4 h/ F" B' z- s) Eif feature_extract:" C! v7 o; ^  P; `+ u. J
        params_to_update = []/ B: J6 A- }& |3 T/ j3 l' o3 r
        for name, param in model_ft.named_parameters():
    + Y. @9 L7 e9 G0 u  k        if param.requires_grad == True:- g* s/ y( z( W& ?
                params_to_update.append(param)
    8 H" ]6 v8 F6 E" ]  T3 B1 P8 i            print("\t", name); B6 x/ z! s9 j6 `7 I9 v- ?
    else:& l# w6 l* h: j! z  Z8 Y% ?7 u. U
        for name, param in model_ft.named_parameters():
    ( l# z2 T! e3 w3 k        if param.requires_grad ==True:6 [( a) O0 @# C; g
                print("\t", name)+ W% y) Y( i2 `8 T, z+ W
    " q. v. |8 E, O+ R4 y
    1: s. }+ Y8 Z* E  U1 @
    2( B4 n8 a# B% h/ {! S8 F
    3, p6 r2 q$ B) W( }
    47 J7 }9 X; ]$ |: o) P# Q# a" ]2 P8 S
    5- E# Q9 r! Z( K% \2 A9 P
    60 k: R7 j, I( }9 \
    71 j2 V$ Z7 k* K
    8
    # Y( R: ^: G7 N8 ?3 D. p9 b9* P' W  {% B$ I
    10
    * [# `( V' }. h5 w4 f2 B11
      p  |; n$ F: x# c9 _8 \12; U, D) \9 S6 k
    13
    " w: p) G  {1 M* r% b4 ^146 `+ ^4 A! x: p" C! k
    15! @7 b2 q/ i2 \, t" I' ^
    16  ^: Y6 t4 r& V! @* c
    17' T- U! }. Y" V8 {5 Y
    18
    7 }' {5 u& I9 e% a8 U192 K( I. q( h& y8 o8 H: ]
    20" J5 Q* O( T8 c5 l/ y
    21- f- J$ T1 Z) ~) V2 b0 u
    22% p. Q, t; J" p' @9 p$ u& f
    230 M: S. t$ I! r
    Params to learn:
    ) [( @3 P# f9 t  c% W/ E" M         fc.0.weight
    9 q) b" v/ ~; l; L         fc.0.bias* N$ f& t% H! ]2 o& X6 j& Z
    1
    3 m' c* P$ `. S- J* Y2
    / ?# [7 g+ o9 A4 ]1 }39 C1 q  o7 Y% u2 M7 _7 I- v
    7. 训练与预测
    ( c* h0 D( \% p8 J7.1 优化器设置
    4 F! E; j1 Q1 A$ [" s6 G) h# 优化器设置
    & F) w  U' t, Z+ x$ ]+ `+ [optimizer_ft  = optim.Adam(params_to_update, lr = 1e-2)
    4 l- o$ N9 a9 j. B! t  a- z# 学习率衰减策略
    6 c' m& K* |+ }! I: T4 C# y+ nscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)# b1 a4 E! w* p: y6 @
    # 学习率每7个epoch衰减为原来的1/101 c8 b1 U8 @  l( ]: u* T- \* ^
    # 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算
    . W; P- j& \6 Z4 h4 X- M/ \1 W
    + }5 b5 K4 ]7 {8 d8 xcriterion = nn.NLLLoss()- [) f9 s" E. E9 @3 ^0 q0 K
    13 V% \9 B/ s+ i- ~( f) L- x) q2 s) f
    23 p5 Z" e+ F9 F$ H( t+ m' E. H
    31 W8 e  c4 R: n) A2 r+ C
    4' M# i6 n# p& K. U
    5" A: G9 a$ D, f  {# d7 \
    6: @5 _0 X* y( T4 i5 b7 S
    7
    $ A$ ?/ X# _$ m$ y4 Q8
    ! V+ a) K! M! T) s& a3 G  R% S# 定义训练函数, m. H) z! N3 b9 k0 ]
    #is_inception:要不要用其他的网络2 L. \8 D0 e  z$ C
    def train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):6 m4 D' q1 S% \& h2 n
        since = time.time()
    ; {1 x/ S. X8 }9 D; ^    #保存最好的准确率
    % \+ b8 }. F4 e    best_acc = 0
    ; c" P4 e* m- L    """; f$ |: X8 W; \. |
        checkpoint = torch.load(filename)
    , D1 M# M( K' q7 m- Q- L' [, D    best_acc = checkpoint['best_acc']; g' r3 t( W, r! g' l8 r
        model.load_state_dict(checkpoint['state_dict'])
    7 Q% s  H( ]" r. l0 [    optimizer.load_state_dict(checkpoint['optimizer'])
    , Z( e, b8 Y$ J- l! V    model.class_to_idx = checkpoint['mapping']) m3 Q/ |$ }4 Y
        """9 ^, |- C" w1 `6 s- ]# y
        #指定用GPU还是CPU0 D' [1 D9 T9 K, i; t# b
        model.to(device)
    / K1 ~7 S+ N) Y    #下面是为展示做的
    ( x. x9 y* ?5 Y" [7 J: t/ O* T) f    val_acc_history = []
    7 ]: G! X9 S5 o    train_acc_history = []* C% k$ l1 D9 S' i1 {- m
        train_losses = []5 }/ @1 Q+ z. Q& E8 Y
        valid_losses = []/ N" n: Y5 N+ H* q5 C4 [
        LRs = [optimizer.param_groups[0]['lr']]! b' a: d; t& R  c  y: P- a
        #最好的一次存下来0 s' |, Z' @: Z
        best_model_wts = copy.deepcopy(model.state_dict())
    ' l6 B) q) V( L5 ?& [# ]$ c$ @2 W% r6 t+ n
        for epoch in range(num_epochs):
    6 z) N( ^* f; J7 Z' w        print('Epoch {}/{}'.format(epoch, num_epochs - 1))/ ^5 d4 @+ n* }8 y  t" M
            print('-' * 10)3 Y' o% q1 c( e1 `+ A( Z

    & v1 \9 H7 j8 j1 q        # 训练和验证+ Y8 z6 e4 e; b
            for phase in ['train', 'valid']:
    3 d1 f4 l" }0 r* ]9 Y$ i2 e) b8 S            if phase == 'train':# D* r) M* h5 F9 Q2 q  A
                    model.train()  # 训练
    4 j. E5 `: w4 G1 ~            else:
    " R! _: i3 [, c8 [: H3 m# g                model.eval()   # 验证
    & j' \3 y/ G1 R4 H, r. d: W7 \
    2 u, y7 J) L3 t1 Q1 d* w            running_loss = 0.0% L* X8 M% E9 Z: a
                running_corrects = 07 E4 ?+ k, \! V- q4 f+ @

    : y8 t4 e4 V! @2 |% Z" ]3 Z3 B  n            # 把数据都取个遍
    " A0 f% j# k# Z7 F  t$ A4 n            for inputs, labels in dataloaders[phase]:
    4 v6 ~1 w6 L: b, i  g* @6 ~0 Z& f+ s                #下面是将inputs,labels传到GPU
    * p' T) u3 m7 Q) y, D# z                inputs = inputs.to(device)
    , R" S4 D+ {: o9 Z                labels = labels.to(device)5 v* q- J) k, G' x% j% C  S- Z
    + J* ^: d0 [: `# F3 A. w
                    # 清零
    : ~7 k# r8 W$ Y3 n( t, R: [( K. q3 j                optimizer.zero_grad()
    6 l4 ]0 Y! c8 c2 G* A' I# _9 n                # 只有训练的时候计算和更新梯度
    0 e6 o( z! }+ @                with torch.set_grad_enabled(phase == 'train'):
    ' |& F8 Y6 v+ @+ K9 [                    #if这面不需要计算,可忽略3 b2 e5 d( R5 {$ s
                        if is_inception and phase == 'train':
    ) ^$ F0 N8 Y3 p- p6 A& I                        outputs, aux_outputs = model(inputs)' c1 M3 x0 M+ V! n3 G# a4 b2 B
                            loss1 = criterion(outputs, labels)
    : _; v8 \+ m; [0 t1 O: x! O+ z6 c! t                        loss2 = criterion(aux_outputs, labels)2 `6 }) H! |' e; Y! N
                            loss = loss1 + 0.4*loss2/ ^0 e# c+ n2 x# K( K* O! p
                        else:#resnet执行的是这里  j' [8 a0 e1 b4 F
                            outputs = model(inputs)
    * L3 K& E4 v0 G: X4 }                        loss = criterion(outputs, labels)
    + O' |2 K, g3 y1 N3 s0 R5 P2 I
    + G+ B8 ~+ }" I( d                        #概率最大的返回preds
    " j& y+ @/ A  H2 _$ t: Q( n  y/ x& I                    _, preds = torch.max(outputs, 1)
    5 G, X6 @# C8 v1 v9 _2 x" i% N! E1 ~2 |
                        # 训练阶段更新权重
    5 y* ~, D8 c# @0 o% i! T5 k                    if phase == 'train':
    # A$ r0 ~! R& \                        loss.backward()9 W( R, y# K7 Z# }; }9 M' r
                            optimizer.step()
    0 |! v2 }; l( u9 s9 E8 @, L( P7 W* v. g4 B
                    # 计算损失
    . ^( a. d' Z- B& Y                running_loss += loss.item() * inputs.size(0)( b5 i( v6 O# i& L& E) K2 q
                    running_corrects += torch.sum(preds == labels.data)/ `3 U% G0 a/ k- Q. T5 \1 `" g  Q6 p

    , T, I1 p; F' i( d            #打印操作
    / f0 C; O' @2 h" F; Z            epoch_loss = running_loss / len(dataloaders[phase].dataset)* ]+ Q3 X7 C4 N( r5 S
                epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)( v0 E( |* s2 z7 A+ Y

    , \5 P8 V' Y; @9 D. j/ ~2 Z
    - K0 M: N/ h7 d. n7 b            time_elapsed = time.time() - since
    4 q$ O$ f3 N1 k7 B) ]# d            print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
    - M" y: p' f4 P% l            print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))3 e% l/ q+ [/ @/ n1 C* c

    + q8 R! ~* q* m
    * F9 D+ @7 d7 `1 a8 @. f            # 得到最好那次的模型
    9 X1 i  Q: Z$ M: w* T            if phase == 'valid' and epoch_acc > best_acc:
    , i/ b: m; l  t, B" e                best_acc = epoch_acc
    0 E1 i0 h- f0 K  x                #模型保存
    " a; |/ [% L6 @) [2 B                best_model_wts = copy.deepcopy(model.state_dict())
    ) O9 N! |8 C! }                state = {- {+ {% V: ^, T" b
                        #tate_dict变量存放训练过程中需要学习的权重和偏执系数
    " C- ?+ U  X  F2 @. _                  'state_dict': model.state_dict(),# s6 @) A9 ^8 P( {$ W
                      'best_acc': best_acc,8 p: o4 ^& ]+ `' b. u6 d9 t8 K
                      'optimizer' : optimizer.state_dict(),+ F+ p. ?$ x( T, p! c$ O
                    }- j  [; @7 Y! ]4 C9 T! j
                    torch.save(state, filename)
    , \, h+ r1 Y3 E            if phase == 'valid':
    * `9 ]5 T6 q+ C$ ~# Y                val_acc_history.append(epoch_acc)# i) w- ^- ^0 h
                    valid_losses.append(epoch_loss)
      @- S! {9 K+ _                scheduler.step(epoch_loss): ~( v. b4 i! L/ ]
                if phase == 'train':
    - _* K% y. x+ w: Q; I0 D$ Z6 F                train_acc_history.append(epoch_acc)
    . m' |8 f1 H1 [1 f1 H. ?                train_losses.append(epoch_loss)' x* D2 l. Y+ F3 G8 j
    ' a0 T$ g) j' W& V. \
            print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))" X9 \! E9 K. q7 y1 V; d0 N
            LRs.append(optimizer.param_groups[0]['lr'])0 m7 d0 ^; ^# g. i/ L2 K/ u
            print()
    8 ^  g; `1 o, g/ }% D) M, P' P4 p8 ?
        time_elapsed = time.time() - since
    * e7 K4 I/ o# ~3 W9 Q" ~" m* Y    print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
    : Z: }9 V* G" ?9 C/ h2 j    print('Best val Acc: {:4f}'.format(best_acc))' y. w# L% E* |+ F

    " k2 |& u/ \  T$ [) m# ]# T    # 保存训练完后用最好的一次当做模型最终的结果
    # H* Y! o/ K9 w# F2 ^; R4 A    model.load_state_dict(best_model_wts)6 i' f$ a' T- H  Z
        return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs
      J, d; t/ j) a8 C) ?( v
    ) G, e2 V8 O0 q, E; m2 T+ b$ b  I+ G9 Q
    1. f( P9 K5 ]& D" w
    2
    $ f, ?" h$ i$ i5 ]/ j; j$ g: d3
    , m0 T2 w' b: S& J0 I$ @& {4
    ) g0 H- x' ^& d- T( r( m' r5
    1 X* E: u. ~$ v- X6) c2 \/ D# D1 D3 K5 K
    7' j8 u( Z6 k+ X8 \
    8
    ! D$ E$ T8 j' V9: \7 i& H* ~% r* x3 s/ t. m
    10
    # B& c( A$ [) U  W) J( R& a  p119 h1 K: U/ P. l( ^: o% E
    12
    2 |6 S# g2 q6 e. ?1 ?. D3 X13
    2 P/ Q9 r. K8 Q; [, d' A14$ x! X; ~; q7 _6 Y( r* n
    15* g) x. X$ ?& o/ j& W
    16
    ! [4 ~/ L( S: k2 x( {174 B  o1 ?# v: Z2 n7 o% z
    18
    + k) ]% X# l3 G0 W, j; R  \% O5 Y19
    ; S6 s, _$ a+ N# o20
    " }! z" M  O5 J. {4 o7 Y21
    + z" j9 w, H7 X7 H# h1 i. W- O22
    ' U/ s6 S0 |; f$ T23# d" k9 A- v6 A- z1 q. W3 z! i# ?2 O
    24) M6 @/ n0 b0 c7 o$ Y& C" g
    255 q# F  p$ s0 v# c, m+ z4 F: B
    26, N5 p; B* i+ a9 V: j9 g
    27+ f' k0 {* K; d, V5 w9 l9 l6 g
    285 V# {& {* ]' w$ \# T7 V" o
    29
    7 o/ @  W4 p+ h; X1 z30% @4 w6 `. O, y. F. y' _/ J
    31- M4 t& U6 a+ R% I) ?
    32
    0 {4 U) T) Y6 f8 O* Y6 c0 ^* t33
    8 D( z# h# t) v( _! t; a  ~34
    ) n5 R0 N" [5 Z35
    ' g( y4 N* W8 D$ P9 v36
    ; m, Y( O1 x: r. m* o$ L37
    ' P4 c; t% D+ S/ R0 ^6 p38
    % S4 d4 q5 S) M: Z, n& {39
    5 T) J4 P& _( I0 e% |# P8 {( k% {409 o3 r6 _2 _4 z! d% z
    418 M  Q5 X/ z( W; M
    42* f- n4 v  v( c% n, m$ ~$ Q
    43' N% Q9 S5 G7 w
    44
    : c; X1 h: ~! q# N; g45' q$ N9 V2 ~$ i1 |* p
    46. k; f4 {: o2 C4 r5 ^9 `
    47
    ; J. J5 J  H, K( ?* I- v1 D( j; C" f48
    " @" d4 s- s, b, f3 R7 C; E$ _49
    8 \# G9 v2 X" ^3 g! {4 d3 \50
    . ]2 }0 `% L/ n- E; e/ F# \2 W! i518 {8 H  R  J) v( B/ N( A
    522 Q9 G# g8 t6 k: P' {8 e! A5 b8 ?9 C' C1 [
    53
    3 A, K# Y7 u# w" h4 R54# {: U  v8 J, m, v! x' B
    55
    4 i4 S! X9 c( U+ b56
    9 @, ~5 ~( [: u7 H9 W# G# p6 ~% a5 i, }57
    6 o0 E& P9 C" F58+ x) f* K, r8 k+ W; Q! @& I6 I
    59
    - K+ A8 ~, J" {; x60. |* O! Y7 @0 ~# G; `  `& Z
    61/ X: x3 c, m# E, u
    62) K  F- s% K% q; [- c# r
    63
    ' ~- X) q" I( I3 v' g- h2 ?0 `648 F6 [4 \1 J' l6 u7 M/ C2 `
    65% ?" z3 H- y6 {  W
    663 f' g4 D, n4 Q' v' m. N
    67, n- m. o9 T2 y- T% i% C
    68/ T- Z; f$ w. F( [
    69
    2 N0 u3 Y, K# Z700 M3 o+ t0 ]7 l4 X. ~6 i# E" I+ x( B
    71: J8 M& m6 w" H# y  d
    72" t' Q; d6 v" j9 }" l
    73
    2 j7 m" S+ U5 ~6 |74
    " |$ J- l% s  C: x# Q* v! Y2 F% ~2 y9 ~75
    : i! j8 V' K  i. W% l+ J6 o% Z76
    " J2 B& t& b- N: |6 O77
    5 L+ f/ {, i9 M$ f78
    ( g% B- ^; a( _3 @* _0 X79
    7 x9 G: M7 D" C7 \9 _6 K/ M802 Q" N: I0 p# h( ]% T, v
    811 b2 f( `7 ~$ }5 }& w
    82
    , P; ^! X7 ^6 S7 w" |8 @. G83
    : W! e2 ~, N0 K) E' a5 _84# }& B) d7 A9 A
    85
    * n! C. Z- }/ n4 U" K' x- |( S86
    1 U: G3 c% V# z7 g9 R$ J5 l87
    3 x- t6 v  w! D8 q1 v. c9 V* C88
    . Q0 x, |. y' w1 u$ `7 v4 Q: K892 L5 Y. \) |  V( B* N( }0 ~) Y0 l& S, _
    90
    * J2 i& h* z& t! r* P8 p* u) P3 }. o8 E917 \' Q) s& }& Q6 m- ^& X( {
    92
    ) J% L$ Y4 R: G" i0 v0 l9 q/ O93
    $ n& L) U2 ~0 U2 }# s) P, j8 u* U94
    + @' _% L- l* G+ m& R6 ]. ^5 B) J' M95
    : D/ o# I0 Q8 k. ?# T; C" y96
    ' Z  T* L$ E1 p9 z% u" W: @97& X( Z" O" ^6 g' l- a1 X. \
    98+ p: E9 L- }9 F4 @1 T
    99
    0 m* [7 Z& z( a100
    6 m2 e: G0 z& t: K0 M101
    - {2 {6 d5 q. N" e1021 @& R* q/ u- Q, m
    103
    7 U1 ~! g. K) f& ?$ p104
    % Z6 \7 E3 Y) k: B; C2 J1 j6 h105
    ! j$ j, ]6 x9 i; E106
    & q3 s# x  ~2 u/ D! {' n107
    6 S; N" p9 U: S( n, X- m% W9 N108; g5 m; g$ X6 d' ^
    109( v) E* v, \0 r" j
    110
    " @, t; l$ z; p. F111) d( i- H  \9 u: H
    1128 d5 a2 d# Y+ Y( h; K6 I/ j
    7.2 开始训练模型
    7 j9 f" p$ l# d! K8 o" T我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次) D* B/ }8 E' n) E% [& P
    ) l6 Q3 t% _1 B2 q- W
    #若太慢,把epoch调低,迭代50次可能好些2 T# M) |7 u6 y( G7 t
    #训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合
      `$ d1 I# t: ^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"))$ _1 F) C  Z$ ~4 Z, B2 S7 G! V
    8 o6 ^+ B, U9 h' U0 k6 X
    14 B( f6 |$ o8 I& X& |& s6 K" s2 A7 c
    2
    , I" N7 E/ `  b5 O3" P3 f! f( ^+ ~* a* I
    4
    & \8 A) @) E+ h1 \Epoch 0/4$ O! J# A6 ?. A: t7 ^; U
    ----------& o% G) E+ X9 O5 I" D
    Time elapsed 29m 41s
    ) A. @% Y. K% e' T! [, I7 D1 J7 w: J% L0 Gtrain Loss: 10.4774 Acc: 0.3147
      ~$ ^4 h+ C% M3 z! T; UTime elapsed 32m 54s5 ]8 W) K* w6 y2 w
    valid Loss: 8.2902 Acc: 0.47194 K2 U+ [4 l% N$ Q  Z
    Optimizer learning rate : 0.0010000
    9 b) t( }, a" \3 b% b7 L) D) ~" G
    " N' P) j3 n6 l: XEpoch 1/4
    1 Z, b: Q: U# w( j0 i- w----------- f0 R: g1 o9 ?, k
    Time elapsed 60m 11s* E! g2 }1 M& q, u5 r" P
    train Loss: 2.3126 Acc: 0.7053
    . e- Z; N( R! W2 ^$ qTime elapsed 63m 16s. c! h% a$ l2 Z, I/ g4 \
    valid Loss: 3.2325 Acc: 0.6626
    : c! x9 k) _0 h0 h. XOptimizer learning rate : 0.0100000
    . I& m3 W4 F( o2 G$ S2 u- ?7 S) m- w; z6 M( Y+ Z9 r
    Epoch 2/4
    5 m+ i. E* B$ V----------( g7 T/ ?7 K, D5 i' \' x
    Time elapsed 90m 58s+ c: N' w6 u" N0 `* V. k$ X
    train Loss: 9.9720 Acc: 0.4734) I; s  o! B4 |; T. ?. Q9 @  Y6 X
    Time elapsed 94m 4s
    ( D5 m: S8 j7 S3 {valid Loss: 14.0426 Acc: 0.4413
    ; g7 d/ [" {7 }" W+ iOptimizer learning rate : 0.0001000
    ! N: A, T6 q. Q
    . A  j1 i7 g: yEpoch 3/4
    8 G  N' {/ }7 H. ?& D6 J- o$ g----------+ X* e4 |) O4 d. y( c
    Time elapsed 132m 49s# j: z+ }  g9 L; N9 b4 q" ^
    train Loss: 5.4290 Acc: 0.6548" ~6 |: Q4 v3 T( ^; B) b; L! j8 n2 `# _
    Time elapsed 138m 49s; V0 H1 i5 M  j* Q% x+ J; S
    valid Loss: 6.4208 Acc: 0.6027
    . }4 @& [/ D7 ?Optimizer learning rate : 0.0100000
    + n2 d# {$ w) b- `1 @* F2 T3 c' ^- x7 ~7 `/ v& x% q
    Epoch 4/4
    1 q% j' ~$ [, a# b7 x& o----------. {4 `& W" B( L" Y, j- E7 c
    Time elapsed 195m 56s
    - d8 e/ h1 G* |' N' G& z, Dtrain Loss: 8.8911 Acc: 0.5519
    & D. a; N4 t( \! {" PTime elapsed 199m 16s
    8 n6 M& r' }7 [3 s/ g1 M6 ?7 Ivalid Loss: 13.2221 Acc: 0.4914; x2 I  L7 W8 N7 k! w
    Optimizer learning rate : 0.0010000
    5 G" ]: i0 m8 h$ A) S4 i
    * ~* M- m( q$ W0 q5 LTraining complete in 199m 16s0 x1 c1 n4 w& J, N" v% q
    Best val Acc: 0.6625923 p! r  m- Y/ B& U7 u

    2 O3 V/ [1 `4 u: W$ b( W1
    5 y, S* S' [  L, [: h3 O2 P2
    " F5 r' ?0 Y- M' I, t3
    / x( \( ^8 Y. ]) E* C4+ T' Y3 _* x) N9 N! _9 s
    5' L) `  O" f# ?# f- F7 w) B
    6
    6 q4 O% g9 J% E8 K7
    ) J# z) o+ Q& t0 a$ M$ e& u# |8+ o# L' b7 C) [9 w# k
    9
    . a' Q& t. `  k  x5 S1 H3 }1 d* m6 K10
    0 y2 n" F5 x2 y$ E11# {7 ^' r6 f0 ~! ^
    12
    % U& b. K1 o" w( p13* O3 s. B6 o  [$ ~8 D1 ]
    14
    # p3 \# }! r, o8 ~% N151 \, H! @0 d# q0 W5 A
    16" O6 U5 q# n3 v& d; x
    17
    ( U$ A3 c& }2 W, f* U6 T  G187 {% f; P) p  b* L
    196 {: C( e$ ~) T7 Z9 b
    20' L. H4 K0 @+ a' ^
    21
    7 S' [1 m' Y# U0 q4 z  N/ y! C0 T22, Q; X3 L8 {3 k5 l$ ~
    23
    8 H1 p1 w( a, r5 C7 Q# s24* A% }& S" e  L- n" W
    25
    8 E/ X% z& I& u% a; r/ s266 x$ n$ U5 O! o) ]& q; B
    275 y# O/ a4 D" f: Q
    28* \2 A9 c3 T; l' S" j2 Q' R( @
    29
    / U, M% r" ]8 N9 F7 e301 w$ F  O: c* z( k% ?! p9 r3 x
    31, w) ]2 b& X4 D' z
    322 r4 r/ n7 G' e$ h. Q0 Q
    33* n! D) a9 _% \
    343 _  ~+ c4 I* ^; |+ N
    352 g" N! E5 A2 T2 [- |
    36
    - l1 i$ w) A1 q3 e& @37
    & T; n4 Y, |* Y7 n. i4 [4 ?3 X+ t38. s4 ]2 x2 `, c- @1 ?* T
    39
    % y; ?& f* k/ w40
    6 B8 ^  j9 [: l' {" _" f41
    1 E2 R+ n( y6 h4 ?9 w$ }9 E/ e42  r  F3 r! {. T7 q& [. {
    7.3 训练所有层; `1 b2 A. J# i7 z. O
    # 将全部网络解锁进行训练
    * v6 r3 c' F  Y& Mfor param in model_ft.parameters():  `. \, j% e- c9 j6 p, j
        param.requires_grad = True
    + W& p7 o  N5 @  T' G4 u# \: ^2 G8 u1 J0 P6 ^! a4 [
    # 再继续训练所有的参数,学习率调小一点\0 w& H0 B  |7 e+ p& _; b1 p
    optimizer = optim.Adam(params_to_update, lr = 1e-4): @' Z/ Y. C! ?% n$ S
    scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)2 B! H, {7 I+ k+ p( U: _0 x) W

    ( a9 d/ h) ?- i' [- M# 损失函数0 ?/ A* C' W, z% I& _
    criterion = nn.NLLLoss()
    7 P$ A# x) p0 X- T& Y1/ K2 U/ n6 v  Q
    2/ a* {; A# C$ p& G( w8 P& M9 T5 ?
    3. o5 X# l' Z- r/ n1 p) a
    45 y2 f! v) {$ I- \( V
    5
    # V: u6 {& B2 {& e" Q6
    ( `6 G+ w; g" |& R" ]" S; }9 P: @( p7: N/ `* y, t! }7 e% K
    8/ l! i2 L9 I# D; X/ b
    9
    6 J6 u8 Y; J, F3 R: o5 @; {7 Y: x10
    + a2 ?) s' ?; y9 N) [8 d) y9 U: Y" s) Z# 加载保存的参数
    4 q- l2 E& t- R# A/ W0 ?2 u: I# 并在原有的模型基础上继续训练
    8 z  [7 V. h( L: S' ]& R; x3 i# 下面保存的是刚刚训练效果较好的路径- Y$ k) \! }3 Z4 k) q& U
    checkpoint = torch.load(filename)
    ( k6 I2 H5 {' R2 ]best_acc = checkpoint['best_acc']
    % X) J. L1 \1 c& fmodel_ft.load_state_dict(checkpoint['state_dict'])
    & F" g  ?& Q( m8 V0 v% E( Y6 Roptimizer.load_state_dict(checkpoint['optimizer'])
    & I& q6 _) j( ]& V; p10 P* \  @! Q8 q0 e
    2
    " o- q" z' V; @3
    2 Q5 o" _+ W! N4
    # I5 R0 x8 Z% V6 D/ q, m! P6 q5
    ! d8 F4 g( }1 y7 i+ k" {6) N5 ^. q5 A) b8 H
    7+ V) T" a, g* V% E
    开始训练- [" `" V2 ], ?$ E" C
    注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考8 i& A7 J& j. d9 F

    2 P: q5 Q' P8 Kmodel_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"))/ v6 f) N7 c3 k7 C$ p# F
    15 X4 z$ r( [) B% J+ H
    Epoch 0/1& R, n. V' }9 M/ v# C
    ----------
    , O; D) U; K3 G7 O" b4 m6 y/ b4 [Time elapsed 35m 22s
    0 O; O) U3 ^! F+ @* h$ k: \9 G: e7 htrain Loss: 1.7636 Acc: 0.73461 ], j" e3 U/ {1 t
    Time elapsed 38m 42s8 \9 p4 L8 b6 l9 h
    valid Loss: 3.6377 Acc: 0.6455
    + P6 P- i; [4 w' j( J, h  nOptimizer learning rate : 0.00100006 t5 n+ q, a" f  Q0 _; r

    , H( w7 K6 ^, M8 c+ h; q; wEpoch 1/15 K9 A2 n" \! o0 w) u! Q4 s8 G
    ----------8 {' d% C% l+ i9 n( |
    Time elapsed 82m 59s
    , ?8 y( k* G+ d3 |) _' Gtrain Loss: 1.7543 Acc: 0.7340
    : S5 X" W" d/ d5 ^1 J3 nTime elapsed 86m 11s! a; w* W% I8 Y6 z2 }. z
    valid Loss: 3.8275 Acc: 0.6137' f9 z5 Q+ m2 ]# r
    Optimizer learning rate : 0.0010000
    : U+ P% F9 e" I( `
    - n; c; X' Z8 f4 }9 f# RTraining complete in 86m 11s, F  T7 U, W* R! w
    Best val Acc: 0.645477, {7 Z2 p8 V: [) E

    2 G6 n# H$ O) y6 T) t1
    ) s' O$ K: u& ~8 S2 D9 n1 @2: v% \. I0 @/ Z% R% m3 p
    3- @7 {" L) B& }4 D) f! `* I
    4
    8 J4 K9 x9 N4 C5; }& s4 [; n, G2 a
    6
    % v: S! Q) T0 e6 I7
    2 w5 e" x/ f6 F( ^$ e6 {  @8
    / n. p/ i: h, G3 O9 A1 p; Y; j2 m90 p) e& w5 Y- v- ~
    10& P+ v5 c- n  I% s& I1 P9 L
    11
    . c7 B* p* Y/ T8 v12: X9 y5 w( @  x' h3 s7 q$ D
    13
    4 X8 m  P8 U8 K) ~" T. `, D14( ]7 U, d9 {' c: c+ \
    15
    + U6 o! d' R7 [( l16
    * v; P% k3 ^! [1 S7 r17. A9 J) D) E% p" G  W
    182 o. U1 d3 L' W  d; Z: T
    8. 加载已经训练的模型
    ) o- H+ S. e  w4 ]5 P( H6 K相当于做一次简单的前向传播(逻辑推理),不用更新参数
    % @& d. ]& M1 o( }( b+ Z) T- F5 y7 z# a, p8 {( f
    model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
    - V9 V% w; u! w+ U# \" A7 @
    $ C, ]# P% W4 i- p+ R6 [# GPU 模式9 c  E, k  S; f2 ~
    model_ft = model_ft.to(device) # 扔到GPU中+ u6 q2 u! u" [( I) v' O
    / G1 l! l5 T8 E  B/ E6 f- V3 H
    # 保存文件的名字8 V( x  ~+ U' V+ S
    filename='checkpoint.pth': G- X' ~3 U$ u" j
    4 C3 T3 \7 c* s. y4 o  d
    # 加载模型1 d1 J$ A: U4 S  R# D
    checkpoint = torch.load(filename)( _' z$ C; d. J. o
    best_acc = checkpoint['best_acc']& D5 W: J: E  i6 [" r& ?
    model_ft.load_state_dict(checkpoint['state_dict'])# u+ z8 N- L7 b8 i9 X
    1
    % D' k" ^5 [' n2 i0 L, L28 w# T# N, P/ h! W
    3
    . T" U/ T, Z& q$ L4, [3 n/ H) O2 x4 I
    5
    $ J& |9 @8 [- D6; P+ ]1 m3 Z& |+ I
    7
    9 I6 \/ l3 f0 N% z! t' x8
    ' I" Y( N% k9 W' ^0 b: M/ ~91 v4 O% b3 {6 M- {/ q
    10+ K/ c& ]) x$ Y9 K* L
    11) u& ?3 E; t: q# f5 e& \
    12* p0 B  B# G+ Y6 _( }
    <All keys matched successfully>
    1 L8 s. e4 L+ M" j1
    : v/ I" n' g3 X$ n7 B  ^0 @def process_image(image_path):
    / A3 o/ m. L8 k. \; h    # 读取测试集数据9 K, N: j1 k/ F$ @5 Y3 c9 N: Q
        img = Image.open(image_path)
    1 B: p8 S  \' {  x4 C    # Resize, thumbnail方法只能进行比例缩小,所以进行判断; p% ^# L; K: v' J# O+ p# h
        # 与Resize不同$ @* V( V0 V# M* }) h# U) W1 i9 N; r
        # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小# g# S$ n7 n( @4 C
        # 而且对象调用方法会直接改变其大小,返回None2 |# V  t" F3 j
        if img.size[0] > img.size[1]:
    " |; t! z: C  L$ ^        img.thumbnail((10000, 256))
    3 l& Z1 H; _& S  S6 t    else:1 U2 A9 h1 b2 n( |2 X1 ]9 S
            img.thumbnail((256, 10000))- Q7 k' c8 w9 d5 B, C
    ' J; P- y) F, [8 U+ P/ `# Y& N
        # crop操作, 将图像再次裁剪为 224 * 224  K9 g2 Y( S! x% {1 P" s" Q
        left_margin = (img.width - 224) / 2 # 取中间的部分% Z% w9 t" K- `8 s/ I/ ]% x
        bottom_margin = (img.height - 224) / 2 1 v4 J: T9 @2 M
        right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度! l. z. P, w. ], w
        top_margin = bottom_margin + 224$ I' U. H6 g, [; u3 y
    3 R9 o7 J9 a8 B! f
        img = img.crop((left_margin, bottom_margin, right_margin, top_margin))
    + p- P5 Z4 v. @" K, k! F- y: N( f' G& \9 K' T5 j
        # 相同预处理的方法8 l3 Z% l4 y# `( D& a% z5 Y
        # 归一化
    4 I, P1 f6 @4 h' ?! z1 z+ p    img = np.array(img) / 255# g- B7 O  i6 X( s  R+ z* D% Z
        mean = np.array([0.485, 0.456, 0.406])
    2 v+ v' |4 {: s+ \8 D9 L1 h2 ]1 |    std = np.array([0.229, 0.224, 0.225])
    ( Y( Q, e4 x0 k) m    img = (img - mean) / std* Q+ F5 ~- n/ c; [5 [0 k" Z

    7 V$ s" k7 R+ B8 h, |    # 注意颜色通道和位置: x5 D" v& J7 O+ j0 C4 `
        img = img.transpose((2, 0, 1))
    $ `. N, A$ l. s
    + q; r. W: _0 o. J2 D    return img0 V3 x2 \3 q6 S/ }* l' S
    ! j* Q) {- A- S) V
    def imshow(image, ax = None, title = None):
    8 @$ s9 n; j" b- F    """展示数据"""
    * i- I1 H/ e/ N- z, Y" Z0 e    if ax is None:# {2 G7 s! Z/ v6 A
            fig, ax = plt.subplots()" Z, Y( b. y4 s+ i
    : N0 U0 |! L* h
        # 颜色通道进行还原
    2 l0 r, L& |2 s    image = np.array(image).transpose((1, 2, 0)), w* Y. t1 P7 d
    , r: V& b: n' l
        # 预处理还原& G/ w  ~& t( N5 N
        mean = np.array([0.485, 0.456, 0.406])
    2 J+ e) Y9 A2 d; @% S    std = np.array([0.229, 0.224, 0.225])
    9 D+ W; F* g4 G, \% h& |    image = std * image + mean
      c: J# i( @5 b9 f8 l+ F    image = np.clip(image, 0, 1)
    : u8 M2 v2 [* ?6 e. n$ H" m, i
    $ R+ F# Z+ ~" A6 `8 m/ \    ax.imshow(image)
    3 Z+ s& i9 G& g1 x    ax.set_title(title)
    2 C  K2 e" s/ u2 q' L1 w5 p% ~6 R) r( d1 v
        return ax
    ) }( a  M1 |' |' N8 F% Z) f4 d' ^; i9 U4 ]4 T, o* \1 ^! b
    image_path = r'./flower_data/valid/3/image_06621.jpg'' @/ i' d% p/ T# t; M
    img = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理
    ; x& U' k0 Q+ @% m4 f1 E# s4 nimshow(img)# S' x$ D( y. G0 j& E+ i
    6 e, ^# j# G) R: P) r& r8 Z
    1
      Z4 ]7 ?/ L  l5 q8 S( y, Q/ S2" e+ y1 H& W8 @$ \& ^
    3
    % X8 A* a6 b7 T5 V. I- m4
    9 M! X" p$ y; e9 |# E5; q" t! f% m- {; E2 h0 k
    6
    % c; Z$ X8 W" F  x7
    . v$ {5 [) A4 J0 d. ?8 w8- G# \7 s9 M! A; V; e( @
    9: P6 W. `& k( _! Y, V
    10% X0 x% w% @5 E( Q: t
    112 {" B+ ?( F9 M9 b" F
    129 m5 P/ t. Q. S2 V1 q
    13
    " r  [# k" r, J; u14
    ( |8 w) w% l1 F. }" ^. K15
      v* ]5 T8 R. L8 ^, f* Z16
    0 ]4 F3 d* n; X$ N/ [$ `173 B+ `" V& y+ @! o* j
    18
      N  ?3 Z% Y0 p- v/ V7 N8 h19: u* _0 P3 V; K1 y: D
    20
    & t# y4 d! D9 N+ ^$ d21
    ; N% o6 k/ p3 m6 }22# M: r0 g# j' p
    23
    / q$ K$ X  g& ?240 v) y) k" ?, x' U* U( Z
    25, g( O) M( ?5 R
    26
    6 V, F4 E3 g/ v6 ^- U3 Z27
    ! F0 M0 q/ d, H% u28/ a. w, _3 `. q; W. z: U- D* \
    294 O+ i" l8 K  x: p  @- u! u
    30
    ; }4 G, O6 ^! A$ A! R" |% D31
    : d/ `  X. p3 E9 |32( K7 i! ]8 `: }# p: g. F5 J  m, K
    33- r$ f  b* f9 Z8 A4 b. t; \
    34
    ( j. u' Z2 }6 C6 v( r! k35
    4 r* n; `4 g3 k. d36, J( L0 v) r# {' j0 f5 S
    37
    + i6 r: ~: _  @: M1 g7 q38
    ' [& i# ?  h- Q9 g6 G1 }39
    2 A  r- Q* P. u+ j40
    . W) {3 y# F. L  V. J41' x) ]* K) ~0 x: U
    42
    & [. `/ Z# }. E: E% V+ A; B2 E8 u43
    * P: y; Q: p+ t' u1 F44/ l, w1 c/ s' m# I/ L
    45
    1 l: F. `$ [6 d$ X( T( J1 n46
    ) v# d+ F. \/ L5 f8 q. Y4 K3 K47! t& @3 c1 n+ l" i6 n
    48& ?" V3 U2 [' r+ w! \6 L! K
    49
    ; H, \! w: w1 A, [7 c9 z9 i# S' q50
    ' b7 M; O! M4 C51$ W: l5 y& f7 K: e4 }" V4 X
    52- K$ U; g5 k  T" X' ?
    53
    ! b3 s$ E% j- H2 Z549 @; Z+ Q" F7 N# W; d
    <AxesSubplot:>3 h3 Z# M1 n9 `. |9 U
    1& t0 H! f2 @# c6 d: y' ^

    * U3 T( W. H3 c: ~& Z+ g  W上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确; |; a/ U7 R, f8 \8 E" q
      h3 \; d8 t0 L" X4 L
    img.shape
    7 a8 A. }; j: z* F% S# H! ~! H) N1/ E$ M5 y' }. {% q
    (3, 224, 224)5 E% D. W; f$ k) L6 E2 y! A
    1
    1 _4 O+ X( L& P- W% z证明了通道提前了,而且大小没改变( f; g# K+ L9 d) S+ E- t; v$ O

    6 p/ \5 S3 n4 b3 G. x1 r0 [9. 推理% J' _1 u# O; L( O) }7 m6 V- M% @
    img.shape( c5 v6 J2 H: O/ i2 G) U6 T" u
    & @* [, G: z+ e- O
    # 得到一个batch的测试数据0 z4 v9 `3 G$ e1 \) y; ^! f9 h* {
    dataiter = iter(dataloaders['valid'])/ X* v* v4 M! @
    images, labels = dataiter.next()) ]0 z6 q# o" ]1 m  h

    9 E* X* r! ]1 \& l+ M( Tmodel_ft.eval()& u7 A+ E9 Q4 V- l/ M
    7 }* ~" G9 V3 Z2 q2 x
    if train_on_gpu:2 Z$ h' C' n( g% l
        # 前向传播跑一次会得到output
    ' U1 M2 s1 A/ N, s6 n% |8 a- O    output = model_ft(images.cuda())
    ' ]! x; H7 J3 T7 I- ?else:- i. s! o7 r4 K2 @' Q
        output = model_ft(images)' Q  A* W$ T7 S0 a8 b( @4 q2 i0 ]

    ' ~; v* L" a8 C0 C# H# batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值
      q& |9 ?  p$ S; x! _9 Aoutput.shape
    0 o; D3 b3 i9 ?. s6 Q+ A7 U( c  n$ |1 R2 H& D
    19 Q5 `! L) i/ D$ ]
    2
    # q$ t4 q0 t- ^  `3# r/ X6 a; [. e7 h; C1 Y
    48 ^& q0 G; M2 `! T, i. N
    5
    7 I$ b* w+ t7 k. s% d+ s& K5 d6
    5 l) }: v# T+ ^  m0 I7
    ) b3 q' G6 e. f8
    6 b- }' v+ w- R- U  h7 h7 y9* i" b5 T* ]- z
    10
    # T; D: \1 ~6 t* g$ e9 k118 p1 A0 O: e6 n( D
    12" H7 M9 W% Z* ^% L, T* `! D0 ^
    135 i/ ^' Q' k& U5 S! |8 v
    14  }$ z% j8 Q# W8 R5 i  a7 S% R4 e$ p/ N
    15
    1 y6 E, t/ r& N) e16
    / r, I- U+ \+ S4 |) Ytorch.Size([8, 102])
    ! y& ^' A# p2 k% T: s1
    ( T& V) ^  z, b: U6 b+ R! B% P9.1 计算得到最大概率' s8 B& E+ }  T- z; }" p$ d0 X4 c
    _, preds_tensor = torch.max(output, 1)
    3 |- H4 e# e- G' w0 d5 Q' ]! {8 W1 c5 C6 }5 M7 ]' O# h
    preds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量8 U7 j9 c: x6 F
    1
    ; G# b0 Z+ o$ U! w5 m( a2$ G& w4 E  \: p; X" n
    3
    0 W4 ]6 Y& [7 \$ p% o9.2 展示预测结果
    ( t2 `( Z9 C' H1 q' rfig = plt.figure(figsize = (20, 20))
    6 ]' X" }, M) {* b2 X* |1 ~1 jcolumns = 4
    ) {/ V+ I8 |8 [6 X0 j" Mrows = 2
    " r* S6 C1 |9 X; H# Z& ]0 t' B7 o0 G& T% W
    for idx in range(columns * rows):/ L5 q! ^9 {8 M1 c$ C7 d0 F" \6 F5 u
        ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])
    ) i; l" J/ b/ k, b    plt.imshow(im_convert(images[idx]))
    2 J$ B9 L1 T9 H! @& ]1 ^7 A    ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]),
    . j$ V8 r  A. L' y, E                color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))9 n& A( S7 p4 e9 R
    plt.show()) w" y# ^* Y8 J  S
    # 绿色的表示预测是对的,红色表示预测错了
    ; H; E) O/ F3 L. y+ v( ?+ G  U1% U, ^+ I( Z" f0 ]3 d+ I
    25 y/ m& [2 w+ l" y
    3
    * r0 |/ S9 V' F  ?8 o4
    ) A0 D5 |; ^8 t- C8 L  ~- i$ |) f: E5; k& ]# a8 m" z( ~
    6% z7 N0 ?  l4 O+ e
    7$ o, b- j& W0 C. S: T* X  q& {
    8
    : u* Z. ]5 I* N. m9
    " Z1 y" q, C+ B, W+ f10
    ) E$ ~4 H, j* c2 q! R11- b& E5 g& ]% S. ?/ J% A& D
    , W2 _& ^! R# |# E% W
    * n1 I5 i5 k8 Z% D
    4 J' F4 v! v& `% V1 X
    ————————————————
      D6 I, _6 s5 @& |+ W9 m9 M版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。
    ! p1 @  Z- V' l% d  z& V原文链接:https://blog.csdn.net/LeungSr/article/details/126747940
    6 t; y" h- U+ \7 k  T& [/ ~9 l' {: f- a- b
    " W0 h( W2 A2 @) P( S% [' b
    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 07:50 , Processed in 0.448604 second(s), 50 queries .

    回顶部