QQ登录

只需要一步,快速开始

 注册地址  找回密码
查看: 2825|回复: 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)实战案例
    0 \# M( }9 k& c  Y8 ]2 N3 G
    ) h9 q, s6 F3 v文章目录* x. ]3 Q5 _: j- A6 r5 |
    卷积网络实战 对花进行分类6 k* v, s! \$ h6 E6 d
    数据预处理部分; W+ A8 P8 F3 h& W. i3 _
    网络模块设置
    8 ~- E) N3 i! `; o: G9 a4 S# ~/ a网络模型的保存与测试
    5 `1 D4 b* H8 |0 F* i数据下载:9 I1 J9 Y, S6 s7 Z" _' r
    1. 导入工具包$ Z( c, I: }! N8 ?" q2 }
    2. 数据预处理与操作5 B7 a% y( ^/ ^; }. H8 G
    3. 制作好数据源! |' W' e5 E( @6 u
    读取标签对应的实际名字
    # D, A6 F' ]4 a# \  U4.展示一下数据
    ' G0 N: c; n) ^/ [9 `$ e! Q' `( s5. 加载models提供的模型,并直接用训练好的权重做初始化参数
    / Z2 V* h: D* B' L8 a' a; h' @! U1 z6.初始化模型架构
    + |5 |" ^: O& L* i$ t7. 设置需要训练的参数8 m5 y% q5 b1 s
    7. 训练与预测0 t7 E8 |& @% E( x  O" x
    7.1 优化器设置2 b2 y7 `4 e, R
    7.2 开始训练模型
    ! T! E2 x& I# V/ j* Z- S' @! ^8 j7.3 训练所有层1 b/ i4 G1 O; D6 u$ ^3 N
    开始训练
    . n' q- x3 f3 j* I* s  K+ H8. 加载已经训练的模型
    / v! J3 O# G4 {. c9 j9. 推理' e: }7 Y. G' A+ V% l2 t7 K
    9.1 计算得到最大概率3 [: y4 a/ }9 G* Y3 g/ t4 y
    9.2 展示预测结果
    ' b  j& ?  x1 e  M9 a写在最后( i* l; i- y0 q' U6 N
    卷积网络实战 对花进行分类
    % D7 w  f7 Q$ x. j( Q& K, t2 a本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能: ]% b0 j: ]1 E2 d1 _

    6 a. o8 H- N$ ?7 r在文件夹中有102种花,我们主要要对这些花进行分类任务+ @! R% ~. h: [, @& i
    文件夹结构
    6 c( E* D' P* d% Z' G1 R) o; v  D9 n) v: o7 `# P7 `. ]
    flower_data
    $ V- l' I) R0 `8 _$ W
    9 A5 X  a2 X. }: W1 v* P4 D2 Strain  Q/ h6 I, D1 N! }
    4 G' O' T& n9 \8 N
    1(类别)
    8 }, ?  o, ~! D+ V6 c9 {21 Z: l  X3 @& H6 x" _: @9 l
    xxx.png / xxx.jpg
    - Z) O) [& Y, F  M2 Ovalid! }$ u$ k9 U9 X

    6 e! S1 T/ u: n主要分为以下几个大模块  f+ N8 V7 R' I

    , r, ]- P5 I! o% f/ G! M2 F数据预处理部分
    2 l! f4 C0 Z* b) p+ h" z4 O数据增强
    & d! q0 b7 o5 k数据预处理, f/ K4 l* c+ I1 Z8 m9 j4 B0 H: E
    网络模块设置
    - P2 F$ a0 U% G9 A2 K8 W- q( R. S加载预训练模型,直接调用torchVision的经典网络架构
    0 h; y( w, r* Y/ f因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务
    3 w. [# l1 Q; r网络模型的保存与测试& Z3 T( J5 B, f, P. B
    模型保存可以带有选择性2 @' ]# f& |" Z$ e6 m$ u
    数据下载:
    " [! G+ }( B3 y9 V1 M* khttps://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset
    ; W3 g' ^' u$ p+ C2 {2 \  J* B3 Q1 `& i
    改一下文件名,然后将它放到同一根目录就可以了8 e$ I/ {  Z) h0 D6 c' Z

      @# H: B' L  ]2 {3 }下面是我的数据根目录  H" @$ X+ p8 x/ l

    - z/ i/ A( Q) P4 `
    ; u, f) B$ ]( s, \5 p1. 导入工具包1 }2 P: I7 ?' l: i: B
    import os
    # N' W7 j: x& G0 y8 {2 v! bimport matplotlib.pyplot as plt9 k$ z6 ^2 D% \9 S$ B5 p: b( K, [0 e
    # 内嵌入绘图简去show的句柄
    5 @! _9 `  y8 K: O' L: ^% b%matplotlib inline
    ! o/ ~+ x2 Y) u7 ?import numpy as np  W# U4 @+ o+ C4 X; ?- M
    import torch
    4 |' H! v+ h1 o4 A6 c9 \  Pfrom torch import nn
    " P8 |0 \3 H0 G0 V" l) \5 t' k
    2 B; @0 I+ @6 w0 s! L4 jimport torch.optim as optim' L  T5 C' k, T( o
    import torchvision
    " E, u+ ]6 _: [7 N: zfrom torchvision import transforms, models, datasets: [+ j/ Y* z/ i5 n* w* @
    5 \1 _7 c  G/ I
    import imageio4 ?& M- l3 w, U0 k6 s
    import time6 v9 G( ?( s1 W$ T0 K' J& ^
    import warnings
      h' z7 e0 h! t3 H9 himport random
    1 Y: G2 P7 ~; x( Eimport sys
    9 ?- w2 Y; Y8 ~: [import copy
    6 N" e5 A7 }6 t1 G; ?9 vimport json
    9 W. l4 f  O1 l; t: Z7 C, [$ dfrom PIL import Image
    ; B( W4 C/ X+ u+ K  m8 G* \+ u1 p
    . @1 W6 |: h/ {# O1 V7 r
    1
    6 x0 R* U) D. j3 Z8 r& k2- z2 y! M. r% K6 b
    3* V2 s7 q: P( T4 p9 G  f
    4
    ' N$ w& |) G' F3 I+ V. ?% f, M" u& ]5
    + Z4 }9 m9 _- Y5 Q5 g! E( |3 q6
    7 W1 K# r2 W' V9 n7
    ' n+ D' d# t4 a8 f( I8 Y9 h86 t* }/ e, B# O; E& t# o1 `
    9
    ' M$ j, C. _/ m10
    2 ]; f; w) n* `11; \" C0 j* D( |7 m9 c
    12
    4 N. ], h2 W6 X7 @; z6 Q6 z: N13
    4 E) d8 R4 ]( g) u5 E145 p) B% G- B  _: P0 |/ u9 o8 i7 X$ J
    154 Z3 H7 f7 k. r. |3 R4 N
    16
    - k/ C4 m1 L7 R; J17  ]: Q( _% R- v- G: H* D! L: m
    18' ]: @) M/ H0 y
    19. m% H8 z8 r& B- E- y
    20- l5 A6 w+ s- n
    21
    9 q/ k6 s2 e6 o$ n: A1 W2. 数据预处理与操作
    ) f. F( C3 q& H) L) C  ^3 x" c( e#路径设置
    % c8 u, o1 t2 {3 }( ?9 vdata_dir = './flower_data/' # 当前文件夹下的flowerdata目录8 a- A% \! E4 ]: V# N8 c
    train_dir = data_dir + '/train') ]& b& C" Q( [
    valid_dir = data_dir + '/valid'1 q/ N# @8 B: Y7 |# q% Y) C
    1
    9 G3 P: ?  U2 |4 T5 A& Z* [2
    : a* q3 x( k1 Q31 ?1 z4 o/ v" h( Q! [
    4
    ; _( W( q  J$ c) `python目录点杠的组合与区别
    . a4 d' l9 ~  _1 u- |! D注: 里面注明了点杠和斜杠的操作
      T7 k- z) A% u- G# V
    . @" H1 ^" w1 k, s6 V$ O3. 制作好数据源
    ( {0 t6 v( K9 q8 C# ]# ~3 Xdata_transforms中制定了所有图像预处理的操作
    9 ?' n  h) r; b3 d+ x. wImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片/ n% B( x) e& [0 j* e. ?4 |( e
    data_transforms = {
    7 Y, |4 K3 Y, q4 w' ~$ |0 o' ?$ ?1 F    # 分成两部分,一部分是训练
    + v. s) G1 U. Z    'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间
    6 r6 D3 f1 U4 c/ c* P                                 transforms.CenterCrop(224), # 从中心处开始裁剪: o6 k( l1 O1 U$ ~5 v: U" Q
                                     # 以某个随机的概率决定是否翻转 55开
    . m/ {# R: C6 |4 t& @                                 transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转
    : q5 s$ P/ {/ [6 n$ r                                 transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转
    + I- i) i/ j- ?. u                                 # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相
    7 [; X1 W+ U. d; q  X- M8 X                                 transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),
    / \& t5 G" o* v1 u" H9 e                                 transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB% f% M9 R% t/ P- n
                                     # 灰度图转换以后也是三个通道,但是只是RGB是一样的
    8 x& L4 q, j+ F, j" x* O. L& I" X                                 transforms.ToTensor(),5 R6 c* Q! k* F: o6 n- v
                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差: O5 R8 h! T6 X. B; j
                                    ]),
    8 f& V$ j8 `4 C    # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
    % r* U( ^: Y4 k, R. x    'valid': transforms.Compose([transforms.Resize(256)," l* ?( `0 B" z1 {. U) B5 D; V! z& d* z
                                     transforms.CenterCrop(224),  P, o3 u0 F. G7 s$ W. n4 ?# b
                                     transforms.ToTensor(),* l- t$ B; @' _) }+ Z4 H# z
                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同
    * n6 l" h! r' \5 \) K! `                                ]),
    , X) v7 {) ^1 J9 t* Z' T4 z( c# L}; Q- A7 T( Z  U, r) `

    ( F+ V' N4 K) z! \1
    ) ?7 P8 g  o" @/ r. W/ v  q5 s2
    5 P( T, ~* Y, V3
    " G( Y4 F: b* p' E41 d5 ]3 l# @& u1 l) P  A
    5
    + c' d! ?9 l$ p+ e2 S9 G! p6
    * F9 D$ Z  H  l4 `7 _* D7- h, x9 C) A5 m9 }. |" T- o
    8
    8 X4 ~7 T; Y2 y4 Q96 s, E) p2 @- P3 a0 A5 m
    10
    / B- q- L2 g& O4 P( P4 o' D. c" O11$ ]$ e9 m+ ~1 B9 {
    12* Y( S0 e; s% s8 v2 c: ?, N: O
    137 ]) [9 ?, x9 T0 p( x
    14( ?' O9 R/ e5 V: [) u
    15
    ; v4 X2 t( D" v9 q* T/ |, L3 G$ v16
    ! ^# J$ k/ ?( b2 F3 B4 j3 |17) G! o* T) D% ]7 U) C5 z: u$ _, U
    18
    8 O8 B# b  S, d. n& v; }6 H19
    # P. Q- G0 P: M; M6 j" m8 n20, {  L- J( ~, j6 b: C
    21
    1 h& O, y7 v" \# T7 E! R! nbatch_size = 8* ?5 x* p9 f7 M  u
    image_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}2 i. l- N5 ?# N) ?
    dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}
    6 h. M& f0 J1 V+ f& C6 n5 [dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']} " ?$ U2 R  Q2 I+ M  Q
    class_names = image_datasets['train'].classes# _% d( ^$ {6 Y6 D5 f9 H( s
    4 g0 V3 X* C+ l5 Q! I' q! N1 M
    #查看数据集合
    . b! u6 f; R; b' \' Z6 C! s2 J% Qimage_datasets
    9 _% s* m& u  ~% N  \3 E
    4 k( ?0 M2 H! y: y12 n7 E2 _" g. `& }
    23 m. J4 ^' G2 T2 N2 ^
    37 ?! w0 m( `' @( E8 j, ]* O
    4
    $ K" c/ E  O3 @; h+ q$ w5) g# M  Y# k. ?
    6& x* N) P2 Y8 e0 c* o- L/ a' q
    7
    5 A: q9 f" \6 m3 H) I8( [! O: o3 u. s2 p6 N
    9
    / M# c" `1 f# i; X' B{'train': Dataset ImageFolder  s7 p0 a* v7 P: ?% Q
         Number of datapoints: 6552
    6 l) s5 _1 E1 D1 {( F     Root location: ./flower_data/train; P' W( G( W9 J" I7 \& r3 _# w" B
         StandardTransform
    3 P- t  l( s+ A Transform: Compose(: T; X$ D0 r: Z. F
                    RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)( A7 [' }: t2 g
                    CenterCrop(size=(224, 224))
    ; P( g; R" X8 c7 }7 L                RandomHorizontalFlip(p=0.5)$ L; ~9 ?0 W9 }. u/ z2 C
                    RandomVerticalFlip(p=0.5)3 T8 V$ _6 C7 D2 M" D, S3 ~- Z
                    ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])+ ^% @3 q" t$ c' K
                    RandomGrayscale(p=0.025)
    4 C3 n0 x5 Y$ R! h                ToTensor()
    / I" E, F2 A9 g; s% S- H) m                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])* O) w8 v0 W. ?* B2 v
                ),$ b, R! d7 E9 Y2 [* V0 C1 w( U
    'valid': Dataset ImageFolder
    ) w2 x& g, j/ j2 J0 `# M     Number of datapoints: 818: U- R# @6 ~. R! w( {
         Root location: ./flower_data/valid7 K. U4 m& ^6 _( e" e
         StandardTransform
    + O- `4 b7 L9 S: @ Transform: Compose(
    9 d5 V4 Q; p5 m/ a                Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)" Z0 c; P9 {9 x; L: G
                    CenterCrop(size=(224, 224))
    5 f. c! Y3 g7 K, t1 _                ToTensor()
    0 A; z" d: N9 r" o7 P4 l7 p                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])( N* p) z  I' m! N9 @" g/ s
                )}- S% [6 R1 Y  x3 r
    $ N1 h. ~4 \$ A+ w
    1
    , h* c2 b; C2 o2
    / Z1 e  F+ X2 R( ]% q! T: U39 ~" r: W# Q: J
    4
    % W* O  Z* f# G. F5, _$ a# V! S% ]/ B
    6
    % p* F+ h( v  ~$ Y6 W( B7
    # k# j6 p5 _$ j; J: S8
    2 N% S5 v7 ~2 e- T& N9& J9 V6 x2 U3 s& Y& S3 ?6 J
    10
    ) J, X4 r: y% s# n) ?! S5 E% @11
    . p  ~  H6 S, z12; s) a1 |9 \$ D( z8 D6 E' ^
    133 H; e: {& c+ k+ \% a- d
    14
    . d  r& S  m9 i8 h15! [# M3 v+ @' @# A3 X
    16  F7 q* Z# J) x7 N0 [& u. G
    17- Q- S  w( n1 A! g' W8 f: H: }! y. U
    18
    8 k3 l  n$ V1 v& W19
    , K) |( t) G* {20
    $ D$ b. \( h3 ]5 A$ z$ _$ J21
    - b# ~" V; {% T# |$ h, o22
    % ~- l9 B" _  X  }23
    9 C' v* e) G, Z& b6 k. A& S( T6 M249 l$ b$ ~, h4 z" {
    # 验证一下数据是否已经被处理完毕
    7 t! u  U- V& y! O! edataloaders" l, }  a" B; n; ]0 N/ T
    1! L5 n6 [( V% O( t$ @; L) D
    21 b0 ~, `$ b2 y2 H, k' q5 d2 O/ V
    {'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,. q7 k+ O; k" o/ }/ Q
    'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}3 w% i3 A" F7 G1 Q, Q0 l
    1
    4 d. ^) h( E) K, M( Z6 N2- {) [( N5 B  T7 x+ Q/ b" Q) [2 I
    dataset_sizes
    + v7 S- U& |6 S3 G4 ^& H7 w16 m* l2 p0 u% n; E
    {'train': 6552, 'valid': 818}; |0 I  @8 z7 k- _" H
    1
    4 L" B" x1 y# M% \4 G读取标签对应的实际名字( A6 N  Q- q  E2 s* j  g1 B0 Q
    使用同一目录下的json文件,反向映射出花对应的名字: s3 F8 \3 r7 L6 V% p7 r0 I% @

    0 E" A7 c7 i+ R. v. Nwith open('./flower_data/cat_to_name.json', 'r') as f:
    # ~5 i0 f* H* S& Q  C    cat_to_name = json.load(f)
    . z" Z& a4 z& i  H' C+ {" \1
    8 @; a4 e; A; {* |4 Y21 X% x) w6 F' K9 {. M, w, S4 ?
    cat_to_name" Z7 a" a8 n1 Q: l
    11 D* e* `. S3 O" Z
    {'21': 'fire lily',
    ; q, g8 w8 u8 T7 y '3': 'canterbury bells',' S! e+ N  g, q& w& y# b: L9 Q" I
    '45': 'bolero deep blue',7 Q9 P  D- d: V) ?$ w: B, j% Z: e" \
    '1': 'pink primrose',. E# o9 p; s0 W2 t- t
    '34': 'mexican aster',
    3 z/ [9 `+ F: a1 x0 f6 V+ z% t. \ '27': 'prince of wales feathers',
    , \. p/ g" J) N '7': 'moon orchid',+ x5 L3 x" Z  s* W3 h
    '16': 'globe-flower',
    0 C1 C. T- \( @. g  A: T '25': 'grape hyacinth',8 V8 G& D" u4 \
    '26': 'corn poppy',5 }+ v- Y; M, c( ~
    '79': 'toad lily',. i. s2 M! f! l* E' q
    '39': 'siam tulip',- @7 m  X9 ]* h( e  `
    '24': 'red ginger',
    8 L, R0 D% x6 D$ [6 y/ ~ '67': 'spring crocus',0 [0 I% c8 E' n+ Z8 n; c5 K3 B, u
    '35': 'alpine sea holly',
    : N3 |8 o: E  s0 x '32': 'garden phlox',6 v* k% p  L, A) X3 D2 o* B2 ~' `
    '10': 'globe thistle',
    - k! G# ]# j' S1 e  _9 n '6': 'tiger lily',7 I0 @5 Y) r, U0 r  {
    '93': 'ball moss',
    & Q' }* G, n! P* ]$ O/ w- ?- o '33': 'love in the mist',
    2 l5 h9 }! Q& Q# I2 k '9': 'monkshood',- G3 x; P/ R* N
    '102': 'blackberry lily',1 F9 f5 q- W) w* h& n1 E" [$ H
    '14': 'spear thistle',+ B  e* B& _# Z
    '19': 'balloon flower',
    1 I! L, P- e, ^$ J" r" Y% V '100': 'blanket flower',
    ! {( U- c0 H( X) f. f '13': 'king protea',
    , k" i: |, i* ~3 y8 {6 t& | '49': 'oxeye daisy',, |6 l4 c8 R4 ]) J% u
    '15': 'yellow iris',
    0 ]0 }, N% j9 A3 }  V( W '61': 'cautleya spicata',1 @- m8 w4 S" ~5 X+ B4 {* X
    '31': 'carnation',
    # o; d7 V8 z/ y: X '64': 'silverbush',
    ' P  |% [& B8 \ '68': 'bearded iris',$ w3 N; e/ Q+ P  T0 l
    '63': 'black-eyed susan',; ]% S' |: s# K+ z
    '69': 'windflower',7 }1 y+ c7 h) D7 c- a" \) t! @# P
    '62': 'japanese anemone',
    9 i+ Q$ M$ R3 \- j) h7 ~5 p' \ '20': 'giant white arum lily',
    2 E- c% b  F" @4 z( ? '38': 'great masterwort'," c8 }9 S9 i. r
    '4': 'sweet pea',$ @! n8 j7 z" N& S5 P0 F
    '86': 'tree mallow',8 s( w2 m$ Y. C1 X5 D" x) Z, a
    '101': 'trumpet creeper',, P0 L& A+ O/ q* q) C+ W' Z
    '42': 'daffodil',7 }7 b5 V8 d$ Y5 E$ d+ R
    '22': 'pincushion flower',& N( E* F  J7 @- ?
    '2': 'hard-leaved pocket orchid',
    2 _% H; v; m# L) L$ C( z: B8 d: D3 } '54': 'sunflower',. ?! N! q% Y' _$ U( ]
    '66': 'osteospermum',
    ' L. V; C: a" ^$ J '70': 'tree poppy',+ ?- j( g( }" Y% E1 i" _
    '85': 'desert-rose',
    5 p$ g5 `) T! W: V$ _4 T '99': 'bromelia',6 |( H; c' y  M* l: L& t$ G6 U
    '87': 'magnolia',
    ( b. M* G& u0 n) @- h" S4 F- L. d8 U '5': 'english marigold'," E( @7 i" Y, H) D. c0 b( l8 w+ k
    '92': 'bee balm',. z, i8 G: j2 K' a" `3 J- p
    '28': 'stemless gentian',
    6 R) b& p) ^* F+ b '97': 'mallow',
      i: i) y$ O% Y- w. Y. j' Q '57': 'gaura',  G$ F9 @9 j$ a/ J4 D+ X
    '40': 'lenten rose',* F2 S# i# u( D0 @2 S, ]
    '47': 'marigold',
    7 u, `7 T& o8 t- A& f& w0 v '59': 'orange dahlia',' U1 R# q9 O4 I& D
    '48': 'buttercup',
    # ]% B% k7 Y0 U+ A8 |. f9 d3 { '55': 'pelargonium',; G6 n( L" {; s4 L1 e( u9 u
    '36': 'ruby-lipped cattleya',
    ) \- E  K( q9 _; g/ Q '91': 'hippeastrum',9 E( J7 e7 x2 L, A$ n. Y4 u0 Q% {
    '29': 'artichoke',' O2 d- o$ ^3 ?& T( {9 E+ K0 Z
    '71': 'gazania',: b. O" D2 K, N9 g' D
    '90': 'canna lily',* A! b" X) e- |, N+ s- I
    '18': 'peruvian lily',6 V" {+ x+ d- e( D
    '98': 'mexican petunia',
    5 R' ?) J# i! E4 p/ |, F '8': 'bird of paradise',8 w. G# i% K: W+ W8 q+ K
    '30': 'sweet william',
    " ~4 {: m1 i' ^9 m$ u '17': 'purple coneflower',. L$ n/ |3 E- @* K/ G
    '52': 'wild pansy',: a5 z# c( g9 r- _
    '84': 'columbine',
    $ [  Q+ Y# b3 ` '12': "colt's foot",
    - N3 N8 S) [  Q- S& T '11': 'snapdragon',
      U% P0 _- u8 l  ^( U8 ^ '96': 'camellia',
    3 N9 r+ S! r# u '23': 'fritillary',
    - c$ S2 Y' S' V$ z( e$ r) m '50': 'common dandelion',
    6 {; N- w; k2 }, `1 G* ` '44': 'poinsettia',0 n, n" T+ S$ v
    '53': 'primula',
    5 @- N  ]+ z9 r) u '72': 'azalea',
    1 v$ b& e$ l, i* F* @8 \ '65': 'californian poppy',
    / P: ^7 C" }. F '80': 'anthurium',! O' {5 j6 U$ [: l2 I+ |
    '76': 'morning glory',; o5 N% H- K. y
    '37': 'cape flower',/ I3 K& u8 t: m$ M9 K
    '56': 'bishop of llandaff',
    ( z' J& ?! V0 X% Q- [+ C '60': 'pink-yellow dahlia',1 ~, F8 Z! _2 J) t. g, x4 D
    '82': 'clematis',: d7 _( w5 a. Z9 i; D2 e
    '58': 'geranium',
    ! t3 }. }4 L% s7 r '75': 'thorn apple',
    + n# x/ r7 y0 S# W6 p '41': 'barbeton daisy',
    1 o2 P" b% u+ ]7 V! s '95': 'bougainvillea',' p6 j; a+ A3 z1 F9 c( c+ C  E
    '43': 'sword lily',
    ! ?. f! ~6 z% Z* U% x/ X5 @ '83': 'hibiscus',
    ( a! a3 L1 q" l/ S6 n' _ '78': 'lotus lotus',
    " s: E3 N6 |2 B' ~# [1 f '88': 'cyclamen',
    ; e: ^; p* d% s6 ]+ f8 L0 O3 m% t '94': 'foxglove'," G3 P2 P" {! n! P
    '81': 'frangipani',3 }/ W; i! \$ w3 S. f7 i7 L' m
    '74': 'rose',3 A7 m: d9 O: w: s# s# w8 _
    '89': 'watercress',
    # i6 h5 _* K. A% L '73': 'water lily',
    # M8 |$ A* f& _; A& O/ q '46': 'wallflower',
    1 `, N( _0 f+ v/ E# _ '77': 'passion flower',
    1 `/ l, p  f! t6 O '51': 'petunia'}
    $ ?6 s. E7 l& P7 v0 w; u: ?* C( y7 d+ L3 ?8 ~" ^3 s  H
    1
      A, \3 r1 m! q6 h2- S: L* ?9 ^/ v, C8 J, S4 b& f2 }
    3
    $ \: f1 b3 J) \" Y; `# A2 p& g48 s' P6 @, Q2 `/ C4 D8 E4 ?
    5- \" {& B- v- t7 W. }
    6
    : v+ T) k: ^1 ]7) c! ]% M) O. h2 z; y9 A
    8: ^% O) g9 G+ \7 A
    9" g8 h9 P' p0 B/ @- O# w0 {/ n, Y
    10. ~% ]; T: ^  J
    11$ j; Z8 e+ q& a$ r- e; A
    12* E* D9 m- r3 _
    13
    1 A$ G* i' ^( x9 v& Q14
    ) _# p/ k+ k1 a: n; k) X; W9 Y15
    / P# P. Q9 N: Z( Y, N5 v- r16; B( e& c0 l/ C0 L& e4 I4 t
    17
    * A* [4 g3 O0 {$ D" `% \# `' I18
    2 J5 A+ f+ a. J+ _( N( y19
    * m( L' {# ^% q20
    ! v" ]  a, O# U7 K7 j: B' O1 p21
    2 }7 h9 e0 t2 g# \. D. E22/ q0 i0 {" y5 ?9 t- B  u9 c
    23( E, [% [3 W3 b6 I
    24
    ' s+ r. M0 N: T# [25& z+ x. }, A2 b2 p  ]1 L
    26* [7 j6 y: b8 ?, r" x; E
    27
    6 e! E' M0 C2 T0 h3 e- _28
    . l/ c; O. ^0 N  ]) k5 ~4 T29: O& q. J# [0 m7 D0 I5 E; |
    302 Q8 R* A4 b1 }% a; n0 r
    31$ X" U& e) j1 H8 Q( n
    327 ]8 M( b& F2 G7 T) W
    33* n6 x: I' {0 M6 [4 ?: ^7 }$ E
    341 }6 U" t( z( W+ ~' \3 ?
    35: F+ D6 L; Z: y5 q* |4 n
    36
    ! N/ I! V- j, S5 W' Y0 k' Y5 a1 O+ p37
    1 Z! Z/ r6 m9 S2 Y5 m38
      X( \& q' h; @39
    8 M, W  M: J" h40
    6 n; O( t- z: _- d, E  u4 x" g41
    ( U, e" Y# }% V+ l8 F42
    0 V+ f' |9 U' a+ x9 X) r4 g43/ r, b# u7 m9 H. ^+ e
    44" Z8 u) ?2 W  K, \7 j" |" U
    45% t$ e1 \0 ?; H
    46
      H7 q- x/ ?# B" q47+ K% k/ B5 p% L# w" s& x  n
    48
    + o0 L3 v" z6 K+ O* o' {; [49& ^4 O3 s: A1 ]2 S/ |0 P
    506 B6 D: h1 \+ n+ a5 n; A
    51
    6 x9 l/ z7 n9 n- }6 `520 Q& W! t' S* t
    53- W& D" o7 I5 p1 H1 T
    54, r" }& e3 M: \) w/ ^
    55
    " }! p; ^8 d# f! R  m) G; k56
    / S. P( v0 e) ~  O57* Q* b2 ~8 z7 l+ ]3 E/ b
    58
    - M9 t/ T' b4 h6 C59
    ) i/ c* f3 a5 ?6 ^, i' X60
    4 l5 I5 Z) g; A! m; G! E8 L+ E- h& N' `61
    $ y5 X9 @0 Q( ?( h; }5 y62
    0 W, D: c; r1 @5 ?63. N3 ~" i# Q' r. ?5 z4 a
    64
    0 B) O' E. f7 T/ ^) A652 y7 f8 ~" g- o8 d8 O1 h
    66& s2 h2 G  |# z, E1 S( G4 I
    67. d1 |8 K' E7 U, W2 ]4 r5 \$ W& y
    68
    6 P2 U0 X! o- D( f8 F9 a69) j/ e' ?' `* |" e8 |
    70
    1 }4 `2 t9 [# I: X  m3 h! x71
    8 R6 s( y  e4 z' n* \2 P. l72
    1 a0 v& o! u2 ^! F% C# ?73
    $ i0 U; `  A6 j9 }% C74
    & D2 l, t' {) L75
    " k8 K: l. d) @5 o  q9 U* c76
    # J  V( V7 O3 e* S+ z77
    7 W9 l5 `7 B7 \" t. l, d% f$ G4 D+ V8 b78: _5 E  Z9 y( [) r7 [
    79; ^; x3 g3 o( I5 j$ D# S
    80- I. L8 \( M8 \% N* j$ S) U
    81
    ' f; X# {4 P8 Q: m% v$ l* Q, t' F# t82
    ( j/ b& @5 P; \7 D83
    , t  r( i& U2 U" o84+ Y" b9 B' X9 b! Z% n" ]( l
    85! V- K8 }: ~1 V3 t
    86$ m6 D' @" ~# x% L1 I! ^6 I
    87
    3 d( ?* c& w* u* a  a( m88- ?* r, x0 c, H  h6 k
    89
    7 c0 v, W% m# A8 {3 D+ p0 M90
    ( Y" F1 H- w& i) n: y91
    / P1 |3 e1 `1 O) d92
    6 c6 c8 }! X  q/ U* U93
    * H% Q2 D( N% d2 _; k- l$ n! j94
    ! r! R) M2 _) l7 e7 ~" J95
    9 W' D# u7 N8 N# L96
    ! v. B2 w2 L1 x97
    , N' f" \; V* T. i3 K3 G98
    * H( i- D2 j3 ~1 Y' G/ F99! i/ p/ C4 C  D: s$ o! x& j
    100
    + N4 m# j  T4 w7 h1 G& k7 Q, {$ `101; l% r! e6 l  k
    102
    1 t5 [" N2 t3 @4.展示一下数据
    5 \; p4 }0 D- M, y) G  U1 y; ldef im_convert(tensor):4 q' x9 ~! V; `" j1 L
        """数据展示""": [% P1 Q& y% W( M* p
        image = tensor.to("cpu").clone().detach()/ |6 ]8 {4 S! [9 z% }% l
        image = image.numpy().squeeze()
    ' f6 d7 \) G- f( s- U    # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图
    + t) F7 R9 Z3 n1 o    # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)
    5 C! W$ J4 b" p3 V: A    image = image.transpose(1, 2, 0)9 o/ n2 M0 J  B* V
        # 反正则化(反标准化). O  I3 h" u$ i% @# f  S/ u
        image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))3 a0 ^& G! |3 X; L# u; o: a
    2 T! K- V& {' j6 [( ?
        # 将图像中小于0 的都换成0,大于的都变成1; g/ e) ^' b8 Z0 n: l7 d( x
        image = image.clip(0, 1)$ a: ^  O( P0 q) m) r

    ! n5 l9 a) U  _8 Q3 K    return image. q6 g. O9 W5 d# l5 `
    13 K, @1 w+ G4 d
    2% d  ^5 \6 J9 p5 Y( H5 k
    3
    9 k1 o6 h3 P; P& A, l) i4
    0 Y% q' w- e/ F. C5) l. W# I% P( ~- z+ a- b% w$ ]
    6" u' l1 @  U$ K) V+ r3 q
    7
    * s! O2 m: C9 Z  E; G7 h89 i% X' J5 ]6 }! B9 V
    9" w! H/ }. P' D6 \3 a$ K
    10
    0 [( ~7 j8 J( r# c& b# b8 x111 p5 p) z  z8 K
    12
    - c3 U* H& ~9 T7 s% W* ~- D& [13( p: O2 R/ n9 L; k& k# c
    14
    2 Z% F! X; p) q1 I( l# 使用上面定义好的类进行画图
    5 F+ ~* _" E. q  ifig = plt.figure(figsize = (20, 12))
    # B( h+ o' p( n4 K5 l( t$ i  ecolumns = 4
    : m. i0 e) x3 x: F9 k! }rows = 2
    9 n+ j% t* p8 H; S# K7 G3 f/ N( W
    # iter迭代器
    5 A% T- ^3 R# `  B" W# 随便找一个Batch数据进行展示' w% P. |# P3 X- ~/ D( G" a
    dataiter = iter(dataloaders['valid'])
    4 w0 x/ ^( S) `. yinputs, classes = dataiter.next()6 R! _$ v6 `  S) v
    7 N3 H" w$ h) y' Z6 b$ Z# h6 m
    for idx in range(columns * rows):
    $ k7 u* c% j. R  Y8 V    ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])" X: X/ B, t% s# J% h8 p4 c
        # 利用json文件将其对应花的类型打印在图片中- J$ J* [  z! g" ]
        ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])7 c' Z- M6 p+ {& H  M
        plt.imshow(im_convert(inputs[idx]))7 E* [  |0 B1 t/ H% B0 a
    plt.show()
    3 s3 l6 S1 @! E2 H: d9 A/ j; g  M5 S1 m5 Z1 f) L& U
    1
    7 `' @, p# W6 ?$ L+ |2 }. K2" `! Y+ g7 Q4 w0 e
    3
    1 F- k' s& c( c6 x& R9 ~+ M47 b# q, ], k5 d" g, O8 a9 I, E+ ]
    5
    3 d8 O+ v# L4 ^0 S5 x2 F6
    , X8 h$ E' s" @( ]$ h7
    - C* ]- C3 {# q  y( p5 c8% V( W7 l& V" x8 k8 h
    9
      |3 t: V$ i$ Y  D4 o# x10' B. I, ~% M4 G
    11
    - r) K+ k% c) F% m2 W1 G123 j) p4 w7 p; e3 v( L2 @# V
    13
    : B3 Q; C8 M3 Q  i0 ?$ @0 m1 v14. E1 y. v; A% a  O
    15+ V5 J4 O6 q1 {. u5 Z/ B
    168 b  o. o4 p9 S9 f, Y
    # t4 w. W+ T4 P5 i& ?; m# t2 B

    5 X7 a/ I# ?4 Z5 d9 U% m% @+ j+ f( t5. 加载models提供的模型,并直接用训练好的权重做初始化参数) S& K% k9 I$ T- J0 P
    model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
    5 H) t! a! N' j1 I3 B1 y7 v+ L  k5 j- K# 主要的图像识别用resnet来做. p: d$ G& }1 a! j: H
    # 是否用人家训练好的特征
    * n& ?$ u) Z% k- Ifeature_extract = True
    $ ]  s4 P2 r# |: m. [8 O1$ T- f, w: A+ s- h" k' r
    2! Y0 }; Q% I: }% ^  J$ ?( y9 j: A6 l
    3" B. k: I8 f5 n) k
    41 q3 Y  R5 J: m) D6 h6 y
    # 是否用GPU进行训练
    # r3 |/ \( Q+ x3 s/ p' ?train_on_gpu = torch.cuda.is_available()8 d& m, I  w, M; W" V: @" y0 r+ P( L

    1 i  |' g/ M; M6 A# Y  y0 eif not train_on_gpu:0 z$ y, L" ~4 r, l  \
        print('CUDA is not available.   Training on CPU ...')
    ; v: W* Z9 l( d* F4 V0 pelse:
    9 y7 k: N: [  v2 _2 @5 z! r8 Z, l- h    print('CUDA is available! Training on GPU ...')$ R2 w" v7 s4 R
    0 R% y+ }+ b* B. ?3 F7 _4 `" W
    device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')
    - B8 f* j4 Y# q$ {+ q: y* K; b: n1. a" p& u, ~6 A( ?2 C' y+ e
    2
    ) F5 `" ]+ G( q" q1 H3% |& L* Z* n4 ?+ e: q( }6 H
    4
    $ b1 [1 P5 E: S3 Z6 H5
    3 [. Y' |4 ^+ l6
    1 W# \% P& M9 t7
    ; ~( V8 r- W- b8 i4 @( E! j8
    7 V6 u' j$ K; \7 K" X9, q, x$ t5 Y) j( a& F/ M. R; e& n8 O
    CUDA is not available.   Training on CPU ...9 W: Z" z/ ^, o. ?/ M
    10 W* g0 @8 q2 E( N. x" T8 A
    # 将一些层定义为false,使其不自动更新
    * [- _3 x( J0 L2 r- V% r+ ~3 udef set_parameter_requires_grad(model, feature_extracting):. y% O2 H) r9 B$ E% }1 D
        if feature_extracting:
    # w" l3 G& i' J7 I        for param in model.parameters():1 K: x0 ^: r9 L$ ~1 y4 X. Y
                param.requires_grad = False# s% Y5 w: a2 G
    1
    / ?5 u4 ^" A% s  m$ W20 b) ~; M8 g; z, Y, ^
    30 R* s3 o  s2 O3 X
    4
    # R$ ]6 A: ~: i" Q55 B. L# E% c" A% W
    # 打印模型架构告知是怎么一步一步去完成的
    1 O1 N7 @- y( R8 }, N# 主要是为我们提取特征的
    9 a$ \* b: f8 ?5 T$ B& r% k5 [9 g/ u0 \' h8 j5 d: J
    model_ft = models.resnet152(), F6 @5 l' O: u) ?3 D1 U: ~0 D
    model_ft1 r4 @- q6 ?# Q6 @% G7 |  z
    1
    + B4 X9 N2 |2 j; p28 z1 C7 D! z! R
    3
    # f2 j5 b+ k. J& [& F4
    0 N& }5 g7 ~% g5 \5
    . K. r. t: U2 {4 K1 rResNet(" `. h0 L. T# V1 `5 ?9 F
      (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
    " t4 F& v8 e" Q7 \! Z' H9 n% E  (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)0 x+ y# Y) j- ^) H; ^
      (relu): ReLU(inplace=True): W! F/ @5 j0 A
      (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)- W9 J. h* u6 G0 r, I) d, K) E
      (layer1): Sequential(
    & o; s: ~' T, W1 `) r2 b    (0): Bottleneck(* B0 c. ?( s8 n' P1 v$ t- }
          (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)
    ; Q2 b* u0 @! G- s& c/ A: [4 A2 |      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)+ m" R# r) S4 y4 L4 E4 `
          (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)3 `! ^2 l' `$ E3 F
          (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    ; {. i* x' d8 g: O      (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)% o6 P! c: D$ O9 \6 D- g* u; |6 A
          (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    # ?5 \! [& r/ D! Z8 R" @" e      (relu): ReLU(inplace=True)
    ' @7 B2 l, l! S  K4 n6 o      (downsample): Sequential(
    3 j) E) O* y8 r; P        (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)8 j$ X. D& e9 h8 g6 _! j5 V$ G+ V
            (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)2 B. v6 g$ ]) k& h8 X
          )  l9 R+ H) v. C* F" j
        )
    9 ?2 W) p) B2 J5 u2 O中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。
    % H" `0 S) B+ j$ |+ o    (2): Bottleneck(
    % c: s: r$ o* t! a5 x& L      (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False). s9 K- `. r5 I9 W* Q0 {
          (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    9 f+ F, p/ ~4 j1 l0 I& c      (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False): d# ]0 N7 i  E+ T5 ^7 v( O# Z
          (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    , B$ r' L8 s, J: }      (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
    8 u9 h1 ^3 A' U1 E  x6 k      (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    ( d3 i. @( G# t. n& V- F      (relu): ReLU(inplace=True)# B5 J% [! `# U% B  D6 Q
        )& i/ @* L  j! h
      )  [- e; A0 p1 A2 W
      (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))
    . [0 f. @- b0 P  (fc): Linear(in_features=2048, out_features=1000, bias=True)" a5 O+ z, x6 V1 X! k
    )2 M, h$ L. t: r5 c2 a

    , f0 j! f, B, K+ h1, g2 T  h/ k7 N/ s2 R
    2
    4 _: k, {( \' {, O- k34 Y: y8 i  ~( T  O( D7 h- f; x
    4
    8 [6 ]# m- y" q5
    $ u3 s2 [  d: f4 D* G* I; L1 \6  t5 T, N7 w$ f, c
    7
    / S! t7 V& o4 L- j2 U; r4 i86 B5 x8 V# E9 O
    9* c' v$ Y2 S, P' J+ v. a; z* P
    107 ~# Q4 _5 B' f8 E! ]( d
    11
    ; x2 \7 k2 E/ q* Z+ G  R, N12
    * d9 r/ x# K/ H' h4 \131 J' `" E6 H- G( Y
    14
    6 ^. j% W) n4 b15
    ' [1 d  t1 h" o, N- c4 W164 k1 {8 P: Z* _, |- r3 n: ]
    17
    8 c5 m! ?1 T, G18. s4 j; r( z8 j3 k2 s
    19
    % l# q. I' d# o( l20( z7 j% g5 C( i: I/ Z- I
    21" `  x. i% u/ ?
    22
    ' t. j  J8 x; }0 B- O4 t& P23
    4 b' f9 a/ b! X! c& n3 J24
    % T5 }2 a* }0 y+ r) X25
    : M0 w# A7 m( h3 D5 h( e+ g8 ~26/ [3 @3 G' P" f, h
    27% e- K5 w3 l9 `
    28
    , n  j7 [% S/ G. F0 Y293 P' {, j5 A0 Z& p/ Z
    309 x/ T* y1 U- h5 {2 D0 B$ m% W
    31
    ' [" X5 f, x! i+ U, s4 w& u32
    . c+ O  f- r. [33% h4 Z4 l- ]7 L# b$ d9 s
    最后是1000分类,2048输入,分为1000个分类' L4 V" P/ V( G4 G9 F  z; ?/ u6 _
    而我们需要将我们的任务进行调整,将1000分类改为102输出) o: C& R' f7 q) W6 U1 n

    1 v& w/ ^- y) K5 D/ w# A; ?- L5 d6.初始化模型架构
    - E! Z6 w: \3 P3 n/ u8 E步骤如下:
    9 J% y3 @, c* }0 j' P- L* O6 h0 f" |* ~6 g
    将训练好的模型拿过来,并pre_train = True 得到他人的权重参数$ t- }) [4 E  m; J
    可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False): V# ~" C7 ^* V5 R7 _
    无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数
    9 _( J7 f2 P; w官方文档链接$ i2 _! s0 T& T8 A
    https://pytorch.org/vision/stable/models.html, V2 q0 C$ q+ W
    4 ^# A; `- V6 ?) J6 G0 w& O/ N( s# W
    # 将他人的模型加载进来, `# i6 ?. C* p% v1 H3 k
    def initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):
    , J+ X5 V' A- o- ?. p/ A0 I    # 选择适合的模型,不同的模型初始化参数不同( R! X  x5 O% P1 W8 B5 T: M
        model_ft = None, `( x- b  P! O. `# r
        input_size = 0
    % C6 l0 A" U. P- t1 c5 o1 N. Y- s1 f7 b
        if model_name == "resnet":
    ( o/ H% p2 M) p# K6 f        """$ V& y5 e5 [/ `2 J/ e
            Resnet152
    # Z. w- C" U( `$ K: F        """
    / s1 c3 u' ?* f/ E
    " |) v$ h$ L1 i+ P: y: d/ u+ }        # 1. 加载与训练网络1 I0 p4 N$ Z: W& B
            model_ft = models.resnet152(pretrained = use_pretrained)
    ( a2 b! B# ~5 Y4 @% |+ A        # 2. 是否将提取特征的模块冻住,只训练FC层$ G: \% K& t: R' u* l
            set_parameter_requires_grad(model_ft, feature_extract)
    % I; G. q" j$ p7 ?- t9 s        # 3. 获得全连接层输入特征% a6 h# @* O. L  L. e7 x  u
            num_frts = model_ft.fc.in_features
    " h% w. x4 E# d) C. V% p3 }, V* H3 V        # 4. 重新加载全连接层,设置输出102
    9 W. V: y1 |$ B8 b  M  ^        model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),
    * k$ i, N2 e) j3 V- b0 E6 m                                   nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1
    ' M, _+ p2 O* _8 C1 p# p  E        input_size = 224( i7 ]; D& U  u1 a' c

    2 |/ k& b' X/ c4 ~; @5 o+ w' p& e    elif model_name == "alexnet":! T) |1 [; u5 q
            """
    3 V9 b* ?* B: q- d        Alexnet
    7 W* i* m$ _; R' K        """6 r( t' ], m2 K: R6 E' U
            model_ft = models.alexnet(pretrained = use_pretrained)
    & N3 M& ^' b8 v6 I0 o        set_parameter_requires_grad(model_ft, feature_extract)
    4 K. j9 L' N2 S3 d
    5 Y+ w/ o, o% R& z6 ?( T        # 将最后一个特征输出替换 序号为【6】的分类器
    - C9 S3 \% l  F0 Q6 i6 |        num_frts = model_ft.classifier[6].in_features # 获得FC层输入
    ( g0 U7 l; m) E1 J        model_ft.classifier[6] = nn.Linear(num_frts, num_classes)2 M( r9 U0 B7 W/ u/ C8 a
            input_size = 224% {+ Q, q4 L: R/ g% n3 x8 o7 G! q
    . K) a6 e( c5 I! K
        elif model_name == "vgg":
    " n3 B% r) s; \' C6 x) b7 i        """
    # l: i% S2 r! |: Y0 \7 F        VGG11_bn9 f. V8 A, @  l* ]' ]2 n+ T
            """
    6 W; R& n- i; `% D        model_ft = models.vgg16(pretrained = use_pretrained)3 U# ]7 e7 o( z/ v$ \, b
            set_parameter_requires_grad(model_ft, feature_extract)
    # C( V+ s, z( r& v8 S        num_frts = model_ft.classifier[6].in_features; M! U0 B# u8 }
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
    7 O  E: A9 h* K- Q5 K% I3 z: ]        input_size = 2242 r$ m! r& p  v: F6 I+ |
    % j$ N2 z0 Y4 ?* Q1 ]( a8 Y* b5 `1 f. [% P
        elif model_name == "squeezenet":
    9 n- N' l; s$ l        """" O- y  F% a. o" F3 t3 R- C* F
            Squeezenet3 t: k& v* B7 F5 D
            """9 F  P" ^4 D) E, V& N1 m7 v
            model_ft = models.squeezenet1_0(pretrained = use_pretrained)
    ) ^  o6 J, e. \3 r$ d) a; \: J        set_parameter_requires_grad(model_ft, feature_extract)+ g1 z: G+ p5 g. Q
            model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
    2 x* T/ R. R) N% x( _  r: V5 @        model_ft.num_classes = num_classes% |9 y/ f, {7 B* L: Y
            input_size = 224
    9 m8 G) |& b1 p. Q0 v) X  l
    ( b  n+ c7 _+ p7 L5 q; M/ B    elif model_name == "densenet":3 g  k/ [( E5 L! b& L$ W% b
            """
    # T7 {* [2 L6 Q6 Y0 @  `        Densenet
    5 a: f; s9 @2 r        """
    2 q4 B! l8 ~; I" j        model_ft = models.desenet121(pretrained = use_pretrained)
    3 h! j2 N3 |2 f7 [6 Z7 ]        set_parameter_requires_grad(model_ft, feature_extract)
    - `3 `8 {" r( l        num_frts = model_ft.classifier.in_features+ K# S" z  c/ l1 @
            model_ft.classifier = nn.Linear(num_frts, num_classes)- R; ?! F& Q( d" x
            input_size = 2240 ?" R) Q& @( ~; O& X" |: \5 r
    , A( H) y6 @" x& y- U% P5 e% o# Y
        elif model_name == "inception":
    / }: d# e2 y9 d. `4 l, q. n! _% o        """
    : n# f- n" e% m) N) C7 `7 s, z) O. B        Inception V3
    9 t5 p; P# Q% ?' [. L# B9 Z        """7 R( x" P) h6 E! c
            model_ft = models.inception_V(pretrained = use_pretrained)
    & H+ A( R- z/ S2 t4 n        set_parameter_requires_grad(model_ft, feature_extract)
    / E& q. `& G* _8 t3 R
    - p: ?! C3 F5 h( y        num_frts = model_ft.AuxLogits.fc.in_features
    6 w8 o# G* i" M' j3 C- h6 b        model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)
    ! C- i. [( ]; \. E: e0 x' Y( \  _+ f, L2 N% Z
            num_frts = model_ft.fc.in_features
    " n: Z2 ]) v: b5 i6 E8 `* B9 S        model_ft.fc = nn.Linear(num_frts, num_classes)
    " a6 E' {4 R, _4 z& x0 L        input_size = 299, j5 L  U2 p3 O5 h6 _$ N( k, ^
    9 L$ H: _4 v2 U) X; w5 F
        else:/ a1 E" e7 G4 m, E( h
            print("Invalid model name, exiting...")4 K# _. E6 C) ?
            exit()
    # ^. i8 A& U% ?# {2 t9 o& m2 R+ m! Z9 u* I1 R& I
        return model_ft, input_size
    3 q& W6 y& R3 P& T$ u, Y- Y1 l+ i# @" A9 ^" G% f# x. V+ l7 l
    16 M4 w- M6 u7 W/ ]4 l5 q
    2
    + w7 k. K/ ^6 W' Z7 }3
    & X* y' t& _6 r& w: U: S; Y4$ f  ^' S. u! @
    5
      E' l$ z9 V+ p6 q7 ~' p' U) j. [6
    4 S4 F4 \+ p) d7
    7 [% x* y$ ]: J' {8 q* V0 v% B8
    2 u, @. M# }5 m3 B8 g9) b! ^' h% \0 D( U1 S3 P
    104 T! m, j  K; H' F, }
    11* }  Z, i4 ^) d/ K. ^  {, N8 y
    12; s( |( Z( B. p8 I9 j7 e+ \
    13
    ( X, o' o( }! I% [2 D2 Z/ }; v0 ]14- r; t) F# g2 m
    15* I0 |2 |; P8 T) P: a- I
    16/ M- p; R/ w. o  I% Q
    17
    ; w3 w0 C$ h: S' G0 D18% u8 P& U+ W  @
    19
    $ Y; Z& P0 M: V/ B7 i% L( O20
    " u/ _: ?( Q0 L) U  ]21) q4 C' a. g5 M! U8 j2 o* c
    22. B7 K7 n( }8 I
    23, U& b$ p$ D$ o1 d
    242 |0 k7 J* |% `
    25/ U/ S3 S5 T# \) h
    26, A% c+ l5 m7 U9 E0 f5 A
    27
    1 r7 J- O6 u2 R8 t( i28
    6 |. g& ^2 i# s! T) g29
    1 c7 I( ]* y0 j& m30
    - J- T3 H) s) |4 R" F2 f3 g319 @) p4 v5 D' e) n
    32
    ( V( q* f  H( E5 f8 y" L8 k5 x8 Q333 F5 ]8 K4 p1 y0 v
    34
    ( ^- V; ~& [$ A5 u( e35: |( b% a: ~! u+ g
    36
    ' D5 N0 M/ R  @  e37
    : F6 ]" p. S) G. a* Z4 F38
    3 ~4 j$ W) i4 H7 c) `, o39' w; q6 n; y- }& t
    408 C' P$ u8 |: ^
    41! q* B# }! `# D# @
    42
    , l3 r( N- N% O3 A" e# S% G43( `0 ~- {, r1 V  ], C3 Q# C
    44
    9 g0 D! Q4 n7 @. H% ]3 ^6 d45& i, ^6 ^/ i+ `5 `+ n! i
    46: X9 `2 L2 W- r1 \
    472 d& H9 D2 ?4 s0 P
    48- ^6 K& k: E2 C- S% K5 l
    491 {( Y6 N2 ^; j. ~: u
    50
    3 z6 z& y* F5 f+ B51
    / P- l: y, n1 r! |3 R2 a, Y52
    ' L8 ^- W- U# W) Y2 S53" }/ U* X' T7 A4 N" J7 u- V
    54
    ( ]7 W0 o) \& I& l* x+ k1 [+ _55
    4 q9 B, o& ?" {8 k- a3 J562 d, g  Q2 I" J
    57
    6 I' u! y( k8 ?5 J, l# _" H2 y& I58
    . o4 c# Z. K* V9 N' L59
    , m1 N1 O1 g: c3 {  ^$ T5 G60$ x2 {) \/ v3 W% o/ ^4 n; J
    61& f5 K( s% Z' b1 z9 O, S
    62
    5 h  o8 j6 e% c2 \5 ]63
    # G! a9 H+ @- ^( q1 x64
    , p; z# y* _2 ?( B9 L$ x653 J/ S- u4 h: W9 ~
    66: k' m9 P5 J; j9 ~" R8 p& Z/ p
    679 Y, @2 t$ G0 R) W# {/ f: d
    68. ]8 y. J/ y; n& J" x. D2 o$ B
    69/ h( J: k, K' G1 Y  r9 F
    70) ~: t0 @0 Q. D2 j$ {9 k3 C: ]3 }( \9 h
    71
      H* B# C4 A: _/ G* \8 j72
    6 X( u/ B7 r- s73
    . R: u! p6 t5 r, c8 _/ y( Y: a74: S: J7 Z; I5 j; i( Q  f% Z5 P
    75
      S/ U9 q# F/ O2 B' c! ]76' U( [  _: X+ }, R
    773 O  k" O2 U) B3 o: Q1 Y
    78
    : d* m9 s, a3 \5 i9 r; [/ H79
      d% C+ k5 m+ ]" m) r/ N( A80
    2 }) ^( w: j% K8 j- h0 ^, D81) h( L% g$ ]# m
    82
    9 ~1 `  K6 G# r) L83
    4 H' P) u" |5 P$ ~7. 设置需要训练的参数+ S5 ^3 v0 o! X) g5 z
    # 设置模型名字、输出分类数
    % S' p( u" W$ v& B" Amodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)
    7 N- w0 k7 |0 M& k3 B, w
    * {7 P9 X, J3 @7 G% {# GPU 计算
    ; w. k$ n" S& G: e% ]$ f$ R$ `model_ft = model_ft.to(device)
    0 J" P- Z9 p  |9 D2 ~4 A
    8 u; b1 T* U8 v" C% V1 B4 n* I# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取
    ; M6 r: n' J) T6 w* e4 w: Gfilename = 'checkpoint.pth'
    3 O) G# b" R! |. d% v$ M- Q
    1 b4 C8 U/ E9 R1 h- }! W4 k# 是否训练所有层  x: m( J  a! D  k2 n: x$ o  x
    params_to_update = model_ft.parameters()
    ! o/ f" v+ F+ U8 s# B# 打印出需要训练的层6 s4 w$ T, Z3 h  L. j& ~
    print("Params to learn:")
    7 b6 \+ A' _& s/ _: q$ z" Sif feature_extract:
    2 L6 W7 O; o: I, V! r- U; H1 X    params_to_update = []
    , P/ e7 {4 M5 I    for name, param in model_ft.named_parameters():; c! L; `; `) w
            if param.requires_grad == True:6 {6 r( m) p+ n: T
                params_to_update.append(param)5 ?& Z( |3 ]. r
                print("\t", name)2 R/ P5 ~: S4 l8 t
    else:
    - C# k5 R1 P" p& y/ _7 N  D7 X    for name, param in model_ft.named_parameters():
    : Z* Y& T: O$ _7 Q& O3 k' n        if param.requires_grad ==True:+ T; x& }; }: q1 w7 x
                print("\t", name)
    ( E9 Q( X/ ~' Y  I' _# ^& `- e- p" ?! m* M
    17 S: O7 ?  C3 s0 e
    2
    % j: p8 h; D' T2 S2 y3
    1 u, r, C$ t8 m4 J1 N& y- H! B4
    ( M- g$ N; G9 y' p; e; B0 W. ]2 i51 B5 u3 b/ W; A2 D. P2 ?
    61 {$ P$ y8 }0 P# J2 d) k
    7
    + L& Z/ E' n0 [5 a% F( O0 b% q8
    ; L9 J) b1 a; g7 p; e9
    " X* z$ m9 U# L10/ g! A1 G8 W+ H2 k8 O3 I
    11
    ( p2 p6 Z' U& ^& }+ E% U12
    . |' K5 [. G" [% p13
    $ i. }5 U% I& Q& s4 q* f14
    - F6 }4 [) [# g2 q6 H2 m15) o# t( s* z- O- H' X. a
    16
      v5 ^( ?" W- h172 S) Z4 E4 {3 u: }! j6 Q$ `/ t; m( R% M
    18+ h, D/ ?# ^( I0 N- N6 z8 k
    19
    7 h3 j" U. ~% _( {8 A20
    : d+ G$ o% X" X' F& \- V( y) q21. `+ B. G. B) g8 T) f* X
    221 c9 p' r% ]+ _; \0 B/ [( z
    23
    3 l, B6 h1 o1 ?/ vParams to learn:
    2 L3 e8 I. {  D% x/ C. o/ ]         fc.0.weight
    4 c" Z" c) j) e% X         fc.0.bias
    ( E5 u% |2 w- k; O' i1
    3 ~& p* H/ u# @0 {+ a2 J9 s( `2& b& Q/ k: d4 K" \8 x2 |- b
    3) C- H, t4 X) C; N
    7. 训练与预测2 C* ~+ P, s$ ?, ^2 M6 x: L; M; c
    7.1 优化器设置  W  ?" r- O4 g! ^0 f  ]0 M; o2 {, X
    # 优化器设置
    ; r3 Q" j7 `; D9 @5 ~optimizer_ft  = optim.Adam(params_to_update, lr = 1e-2)* {% N) G1 g7 F. M
    # 学习率衰减策略
    & p' i4 S+ K. m0 F- {, n+ O9 ~scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)/ K# P0 C% `8 Z+ l0 d/ A! p
    # 学习率每7个epoch衰减为原来的1/10
    % L2 i# B7 Y3 {* s# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算
    - L5 a6 [8 f& N& |' t6 d- j5 O% O6 r9 F0 M3 V
    criterion = nn.NLLLoss(). ^5 U0 T  \. g1 ]( h
    10 z# k/ g3 o1 W5 v1 H, T& V) h1 M
    2( m( I9 K# h5 P3 s
    3
    6 {& H" f' T* g# L" u  `7 d7 e4
    1 d9 o. r" p" u. C! ~5
    " e( _* g/ x+ T6 V& s6
    , \3 C* C+ @. s5 K- B/ V- u  L7/ ?- l3 S. b2 ?
    8
    + G$ v8 }! N, b$ i" M2 E/ S7 Y5 ]# 定义训练函数& R. W7 {7 h$ I
    #is_inception:要不要用其他的网络5 l- R; ?7 e1 u3 b) o' K: q. N$ c6 j& n
    def train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):5 ~9 J: H: w8 y$ U1 |
        since = time.time()
    1 M1 ~; `9 @. G    #保存最好的准确率& {% o) F0 @1 N1 ?" i& o6 U; `, D9 |
        best_acc = 0/ C2 q8 s- {+ X: s/ H7 v
        """3 u. m) E; u8 r, [7 e7 |
        checkpoint = torch.load(filename)
    . q0 @' u. \. K    best_acc = checkpoint['best_acc']  ~3 @( ?( d8 q0 l- Z3 k
        model.load_state_dict(checkpoint['state_dict'])7 I$ l# v5 c% y' i! O) l+ X
        optimizer.load_state_dict(checkpoint['optimizer'])1 @, o( O$ h" V6 I( {4 a0 Y
        model.class_to_idx = checkpoint['mapping']
    # r/ }: N  y$ Y+ E* i. ?% F    """
    : O" {; W; J+ w0 S+ z    #指定用GPU还是CPU6 a8 W; g4 H: m# E+ g
        model.to(device)
    3 R# o% M/ k9 P$ G% g4 W* W    #下面是为展示做的' I1 q1 t: N" t9 D: g( z( h- i% O/ H: q3 [
        val_acc_history = []/ d  q4 }5 I, v+ q7 [+ ]6 x" D! o
        train_acc_history = []1 r9 Y& K* e4 b7 }* {  u$ R1 C5 B
        train_losses = [], h& r! r8 v$ J& n
        valid_losses = []
    " b$ F( [9 ]; n& Y  b! C! U, c% U    LRs = [optimizer.param_groups[0]['lr']]
    7 a4 L% g! d1 I    #最好的一次存下来% a. Z' T! `. ]. ~( |) O% p. `( `
        best_model_wts = copy.deepcopy(model.state_dict())3 o, {* ^& a2 M
    . B$ w- U" a6 k% y3 S8 L& E
        for epoch in range(num_epochs):
    , c: I, i% V' F& o) l: @# j        print('Epoch {}/{}'.format(epoch, num_epochs - 1))
    8 T: k* S4 X/ s9 N1 c. ?) q        print('-' * 10)
    . ?) D5 s% {1 q0 c1 ^" x( q
    " w7 Y" p( B. E        # 训练和验证- j9 o) `' Y# V+ R- f: Y! A& ~) _
            for phase in ['train', 'valid']:
    + r$ @. Q+ T  X3 `            if phase == 'train':
    7 q; s* L. A' }/ S2 x; Z                model.train()  # 训练8 x2 Q& c" R: H+ c9 ^% B8 T
                else:
    ( U: }& Y, U2 Y% N4 D2 V                model.eval()   # 验证5 P# H2 h. B4 E+ b, z* g
    9 b* S) t" \5 P# |
                running_loss = 0.0
    & O5 k1 c8 D) ]6 p* `            running_corrects = 08 X0 J# |/ M& D& `

    + g3 X3 G# M" n: P: @. u            # 把数据都取个遍4 H/ t0 Z- ?9 [# D4 r, o& d' D
                for inputs, labels in dataloaders[phase]:
    , j5 w, q( J  D" C; s  }6 f- c                #下面是将inputs,labels传到GPU
    + @/ k, Z/ {; N                inputs = inputs.to(device)8 {, W6 V, r( O% P; ]- V
                    labels = labels.to(device)
    ) F4 `7 @0 G  p0 a- Q% R$ I9 Z+ e. B# @8 A5 O
                    # 清零
    / {/ m3 a! g) I# e                optimizer.zero_grad()
    - X5 L# [( |. Y3 |# r' n                # 只有训练的时候计算和更新梯度, I. G1 P/ ]. n7 N* h+ Y
                    with torch.set_grad_enabled(phase == 'train'):
    ! c; x, u/ N7 f! m6 R1 O                    #if这面不需要计算,可忽略, r  U: o. L9 D9 y& z8 `0 o
                        if is_inception and phase == 'train':/ F+ ~( H/ X4 M# j
                            outputs, aux_outputs = model(inputs)  n1 }$ O7 U1 Q% i3 F7 R: n
                            loss1 = criterion(outputs, labels): i2 Y) x* _9 V' ?- y
                            loss2 = criterion(aux_outputs, labels)
    . d. t  l5 D* u) R                        loss = loss1 + 0.4*loss2
    * T6 M9 E% W6 N0 O. }8 F# n' u+ O                    else:#resnet执行的是这里, H8 r1 H% Z- s  \3 [
                            outputs = model(inputs)
    3 S( M8 W: a/ h+ Z# E7 o) c                        loss = criterion(outputs, labels)' s6 R+ S/ L) Z" D6 v  p8 t$ A9 X

    4 y/ U( a# P6 }. h                        #概率最大的返回preds" c* S3 C" P0 ~. Z3 K1 l+ @
                        _, preds = torch.max(outputs, 1); T5 g+ p" y3 z7 R, x/ [* l

      G/ F, c6 G, l) N2 U0 `2 \* {                    # 训练阶段更新权重
    5 t" y  m9 L/ k+ N2 r+ a" _1 q                    if phase == 'train':  W4 w& ]- y0 d2 }- j9 R6 P
                            loss.backward()* C( K9 G  k7 y, u2 n
                            optimizer.step()
      A1 j8 {) k( L# b, L8 c9 T- v+ h+ }  l! C" k0 b) ?9 E
                    # 计算损失8 }, p6 p2 U7 T% p" ^& |
                    running_loss += loss.item() * inputs.size(0)
    ' L. R! N( u: Z                running_corrects += torch.sum(preds == labels.data)
    8 P% ~: V3 E7 ^; c4 M) P8 ^0 g) q& n
                #打印操作
    " Q( ^5 [8 l# a4 Q            epoch_loss = running_loss / len(dataloaders[phase].dataset)
    / Y1 _1 C5 g" i7 q( e; B            epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)
    ( C) A4 R+ Q/ p, k! C, V* I  c9 b& Y: D9 y9 O& r
    / N2 r  r) V8 d
                time_elapsed = time.time() - since0 \4 H- z* B# a; ~6 G6 q
                print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))4 _' b$ E( C# i; b7 E: j
                print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))
    " q: s( l9 S, R5 [9 v
    * e/ P6 M" p( {% X, w1 I* ^7 \; n  {0 z
                # 得到最好那次的模型9 K4 o% q. o4 l2 V7 T
                if phase == 'valid' and epoch_acc > best_acc:
    1 f3 x& y8 _1 i% n7 G- [                best_acc = epoch_acc
    % _( z0 X1 N( r9 j' O                #模型保存
    % e" i0 w* W6 x                best_model_wts = copy.deepcopy(model.state_dict())7 P: F* ?3 G0 y8 n* J2 d# H
                    state = {
    3 A  A6 u4 A6 R9 j9 `$ z, r; h                    #tate_dict变量存放训练过程中需要学习的权重和偏执系数
    4 p% u2 L6 P, w$ k' r3 w- B& W                  'state_dict': model.state_dict(),3 \# k: T7 U6 g
                      'best_acc': best_acc,/ d: S4 c3 u' n- P' @. v
                      'optimizer' : optimizer.state_dict(),
    ' _, d' O- F8 Q+ o" B0 t+ K. C                }, Q7 p) v( f6 [
                    torch.save(state, filename)) j+ Z# W  x7 P, D9 ^
                if phase == 'valid':8 \9 E) w& b4 i& Z8 h
                    val_acc_history.append(epoch_acc)
    5 T' d) B  D/ v                valid_losses.append(epoch_loss)
    & H9 N9 a/ x  B& P6 J" L9 l                scheduler.step(epoch_loss)
    * Q4 [1 Z' R$ R( o0 C" A            if phase == 'train':
    , m0 C4 ^  c0 \/ U/ Y* T/ U                train_acc_history.append(epoch_acc)  r2 F# {. J, Y8 X& Z4 N6 q
                    train_losses.append(epoch_loss)
    ; o! T6 q, t5 U# `5 Z# A1 y* e+ E& [
            print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))
    3 M+ G, T8 A- q" D        LRs.append(optimizer.param_groups[0]['lr']); X0 K: k/ n! W) ?" i$ P
            print()9 }. P9 ~! w2 J; h

    3 R( Q( X) ^! m0 E/ U; `0 |/ }    time_elapsed = time.time() - since
    0 P% b4 y8 I8 ~6 u( k: G    print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
    9 V) R- k2 ?4 P    print('Best val Acc: {:4f}'.format(best_acc))
    6 x1 i" p0 M: ^2 k* ?' p) x/ R) N1 L3 L0 e
        # 保存训练完后用最好的一次当做模型最终的结果" O: [7 M( R" Y  s7 c
        model.load_state_dict(best_model_wts)
    2 v" J7 R( l8 k    return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs   F6 X3 b& g  W9 c. K0 X

    3 I0 w% h* w8 G2 E; k( K0 }: z3 m, c) n7 W/ z9 F) S0 W$ T0 r( r7 K
    1
    0 O# X- [2 Z* ^  u( z2
    ) A1 Z5 X* Y4 F8 ]3
    ! `. F1 X+ |; c& f) N4+ n6 k0 y( t5 Z1 R+ M
    50 z- z; w6 X! }4 [5 P% w8 C
    60 C  V: O( P' N' Y2 \
    7, R* c1 H- T( ]5 W
    8
    , ~) y# ]- b1 k97 ^) n- u' [' t& w- ~8 s9 e
    10. W$ q0 k7 M' c; g. O( R
    11
    4 P/ s* o, A, d  B8 t! N12, c! d% q5 c9 ^$ x- n% c
    135 @, t9 ]9 v! _" ^5 D# T! h" O
    143 B) G9 q, r0 g# `
    15
    9 D9 h+ ?! r3 M! _167 D+ [, n- ]0 {; Y0 O( V: A" N- E
    17, G3 g; ]3 l# T  L" a
    18
    ' U& P% s' Z8 X$ X- m3 b/ f192 I5 @' B# M$ w/ X& K$ f. C$ Y% h/ u
    20
    , _) d+ U2 q8 R21' e/ e8 B0 v- ]/ N( Q
    22
    6 P& R! {% Q2 c- f- [' i8 @23
    / s+ u, w" _0 |8 ?* O24
    ( f( C: @' l$ i7 j' g255 }4 B1 X! p% Q6 y: U" A9 V8 c
    26
    0 C0 {; U. h* y% w) T2 S1 N  k; a27
    & @9 E( F2 [: u/ o# B- X( }28' o9 Y3 f. ~/ V' U
    29
    ! v$ u9 ~$ c+ W309 |9 W$ w- [' y& O8 ~0 X/ S
    314 y' }  }( _4 B5 F& g* |
    32
    3 E8 T6 |( l# h" ^4 X33
    7 G. Z7 m" q, u- z; y: `3 [4 w/ A" z347 M" m( q3 S8 b' m3 ?
    35
    - n6 |* G+ i9 S4 F36. {2 s9 _; M; h7 G- r
    37
    * t' e4 U  ^) ^+ U38, q+ i* [/ \" J5 J2 b
    39; }. a0 l1 t( w9 r+ ?' a/ l
    405 O: f+ ~7 k, F4 t
    41
    2 ?9 L, X" R5 p" ^% \42
    % r2 v4 O' P+ d; N43! l3 [& y. X" m( V" ]  v
    44' o) A2 D2 x; X8 c5 c: s' `. h
    459 }, ]2 z2 R4 o/ L
    460 G. u& V, W0 ~8 ]+ {; g
    47; ~& n+ Z- `, V
    48  L1 @$ i  h5 D- U! d$ u
    49
    ( A4 N; |) {7 B  j; `50
    % o/ _, `; X; O: I* P; z7 C51
    , Z1 b: Q+ i7 K, h: F52
    4 r& H/ q6 H8 T; }2 t  F% J" f; `53
    - {8 V9 _6 g0 f54
    8 e$ X: o7 S% e0 h9 A" I  |556 v$ l5 p# N" \4 Y. p# P4 v
    56
    9 L" e2 F1 U9 v2 p5 p8 c57
    . `  o" C+ L- @* H58
    ; k2 _% N" e/ L59
      I1 u$ E5 `# E/ s# e" j/ [600 E, R  p# H4 z9 p. F( b, |! K
    616 X6 Z" x; r6 L' d' z
    62- L! Z; G/ s6 f! p0 s6 O! a
    639 X% ?( P) n/ [1 ^; e7 J
    64- P) _* r, I( K5 ?
    650 k7 h" Y  F4 r& a+ P$ o7 O+ q6 I* p
    66
    8 K0 Q. v2 b9 J! f6 ?4 C676 E+ x7 M1 o2 R
    68- T! g! E. X6 o% R2 ]
    69
    ' Z; z& k% G0 O; I( G70( O2 W/ G% O1 m) x, D+ Q/ s$ m
    712 n! b, P& k- t8 [* I5 d9 |
    72
    7 o' Q% {8 e. u5 c, s73" F! D5 p2 L3 F) t$ s6 C1 o( b
    74: ~0 }, _5 \* z) F
    75* x; }. P2 W7 o/ h
    76
    # B- L* x9 b* K2 v77% C# F, O0 P! ^0 P+ f4 ?
    78
    3 U1 r. O" m& a: p9 ?$ I( A4 L( t79
    $ o$ B( C: Q6 g% ?80
    : E, i& P4 c% t6 {0 @$ E: V2 c81  _; X, T6 d$ S' I( C* q
    82  b2 |1 [) @1 Q
    83
    % J; \5 C0 ^. |+ |& }" `& Y84: ~! b3 K* C8 Q- i; D0 [% \
    85
    ' {; o% Y9 {8 V" A, q  H86& i2 H" p% ]' o  X" \. b
    87
    ! Y8 F( j+ B7 s88
    8 O" D* B9 O; p2 H+ B89
    - U( x" R4 |$ Y8 C! b; M' s/ m90. v) j( Q, d+ U$ g4 V" V6 y
    916 u4 d2 U. S4 j5 ^0 D( r1 k% K) Q1 T
    92) @( d  i) _) g* u; h+ e
    93+ s9 _  Y' d3 V7 A" |- f: L  c
    94; |5 u" q4 A! G: C5 M2 }3 Q7 s) ^
    95
    & F4 l# B) y$ Q1 f: K5 j0 D96, z4 m- u) I8 F7 f+ z* H5 }
    97
    $ ]* n! M! r. t. q& r4 k982 d, L5 {% \9 t; n. N( l; Z' c
    99
    $ I8 o+ }, ^: B: i# a1 E100
    4 L( {( g/ Q# f101
    4 ~- H5 A: d2 V) N$ q8 \$ f102
    ' _- Y) b  U# Y8 n+ |* M9 S( Z3 W# p103
    * t' {6 M% k  k+ G# y& A104
    ; A7 p5 E# K' o; Y* |9 U3 p105
    ; }) j, W7 r, t106+ `( e+ H  b1 C! f# s4 }2 N
    1077 X% Y( `/ A3 K5 \6 V
    108
    " \0 G- w: R% I109* w4 H* N. h& d1 I6 b0 B
    1100 @2 m: G( L/ d+ N! z
    111
    1 W' T' h" l5 I( ^% A9 W7 `112* j" R) ~; H1 Z
    7.2 开始训练模型6 K9 P; O  ?, }  ]. Q5 X: w
    我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次
    ' J* j- J: w6 d
    / }8 A  a& e- m% h5 O8 M4 V- [#若太慢,把epoch调低,迭代50次可能好些
    . Y. {) s. Y2 b0 ]* ?8 ]#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合
    # }! d' q% p3 M% ~9 M! Zmodel_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 g% F  A3 `. y" }. H9 E* X: W1 e9 A4 J2 I) C5 x0 c9 L" q
    17 `9 ^- H- w/ \% Z  ~9 Y2 x
    25 T9 }/ S8 ~( O) Q. i
    3" |2 @5 w  ?  c) r5 ~
    4
    2 C( k" Q8 ^  z6 `4 k* ~Epoch 0/45 L# p4 ]3 z. u$ i: r
    ----------8 i2 x8 g1 |# f
    Time elapsed 29m 41s' b3 Y; w- o5 |* q. X
    train Loss: 10.4774 Acc: 0.31478 O+ \" e& o6 e! t% F2 Q7 j( q( X
    Time elapsed 32m 54s
    $ f4 J7 \- y5 Y3 F0 A( Dvalid Loss: 8.2902 Acc: 0.47190 o$ P3 K$ [' C7 }  {4 _9 c
    Optimizer learning rate : 0.0010000
    4 _) A# F& m9 j2 f4 z$ g; Z
    0 `3 F- k) c0 GEpoch 1/4) t2 n: D( L( ]7 D2 H$ p1 R7 o
    ----------
    0 _  Y; X" X" S* V7 D4 H7 S% hTime elapsed 60m 11s1 y5 Q3 `' }3 y$ N# h& k. d) f3 q! s, u
    train Loss: 2.3126 Acc: 0.7053
    4 @' P. _3 F/ S" E. GTime elapsed 63m 16s
    9 L; V. r" C9 zvalid Loss: 3.2325 Acc: 0.6626
      \( R/ o% M. p- P( m! NOptimizer learning rate : 0.0100000
    9 R: p2 `, N; J
    7 _+ [" b/ J- i7 YEpoch 2/4
    / m) A3 X- q) i  y, K0 W----------
    % c3 I( i0 ^# \) YTime elapsed 90m 58s' |1 b: |7 [8 T6 v+ S, A8 Y9 R# L/ E3 E
    train Loss: 9.9720 Acc: 0.4734# l: ?) e9 s, _- u/ d! L8 f
    Time elapsed 94m 4s7 J; t6 H/ B( F: h2 F+ H
    valid Loss: 14.0426 Acc: 0.4413; h+ S2 C; N/ x3 h& k
    Optimizer learning rate : 0.0001000. z, a2 H% t% T/ B' S

    5 v& J& L5 {# _: d3 REpoch 3/4+ {; q5 D6 w2 l/ A) W
    ----------4 G& `1 y6 t7 r, V' g0 A, ^( T
    Time elapsed 132m 49s
    ' j- n6 l8 _; Xtrain Loss: 5.4290 Acc: 0.6548
    ; c2 k* j. R0 i  r( L. C) NTime elapsed 138m 49s8 P: y" z" D# E( |7 @, l+ n
    valid Loss: 6.4208 Acc: 0.6027: A5 m+ j9 y" f! f6 L* H% w7 y* b) d
    Optimizer learning rate : 0.01000003 f: D3 ~' s: S' R2 t- r( ~
    5 s% v( D+ N0 H) s# K) E( b& C7 \
    Epoch 4/4
    ) r, t: p4 T& e5 O& E1 |----------
    / h7 g! P' m, v2 ], {* C3 g6 LTime elapsed 195m 56s
    $ n2 r) l% `- p4 A( m8 q9 Ltrain Loss: 8.8911 Acc: 0.5519
    1 J; V. ^' ~( `0 q* E% A/ hTime elapsed 199m 16s% I$ O3 D" t! K; b: U) [
    valid Loss: 13.2221 Acc: 0.4914
    4 k1 T6 L6 Z- v% V: G9 o" `% _Optimizer learning rate : 0.0010000+ t2 {0 W# U0 W' y  m
    ) e; `8 p  P7 A/ `
    Training complete in 199m 16s
    1 C& _6 o- z  nBest val Acc: 0.662592* @% o3 G4 y+ F. a0 ~: `: l& q
    ! V7 N& V; B2 B) L
    1( V# u$ j- [9 s6 a
    2
    ' h: u: u& o; f3- ~' s9 H2 D& x- B
    42 I1 v6 T* t' T9 k' p; m* O
    5
    7 L# s3 c% {' ^7 t69 U6 X# N+ ]( m8 n; T1 B/ p3 `
    70 s* v# I% E" G3 m
    8
    # a& }7 J5 D9 d- Q& d' n9
    3 J: z; V6 C) G  N' ~  @, |10
    + H* P' R- W2 H" a% q: S" v11
    1 I2 ~- A2 U3 a. Q129 C) [1 ]$ d  X, v3 W3 L" G
    131 e9 E0 {- j+ j+ H& R- J
    143 J* X. O/ e  n5 s% `: ]
    15
    . k. F( j/ T' P& B# c) N163 J  y' s! R" Q2 c! G) y! k( u
    17* _$ I) ~3 f4 w% t0 ^+ F  }/ I
    18
    & ^; w) F& N' t7 D19
    5 s( {: ]6 |0 n4 Y! i20
    0 J2 a# d. C& `4 o) a7 \- g21
    * D$ u( v& l9 T8 r8 C9 L' l% y7 Q22
    $ i0 T. p# C- w1 K- {9 z23
    5 [; F5 e0 t" q245 V6 q& K; `7 y
    25
    0 q; w' L, C/ W& i26) f- q. Y2 F- z
    27: E9 u3 A, O, h' a% @2 f
    28
    ' D9 Y: k: |4 p29
    2 b5 g0 U5 z  g/ r: }/ O6 y30
    : o* i% O9 `: p3 O! d" O- I31& C5 H8 ]- H$ g
    32
    8 H# S* O$ R# k2 u' x9 X( K& {33
    2 f$ Y8 F) J+ ^+ D. X349 k! \; K9 [" ^" X
    35$ F/ w# y* n; h3 k2 w
    36& I* d) b# @8 W8 p8 i/ Z  E: P
    37& r$ C7 n9 P/ `" L% ?0 S: m; p+ j' a
    38
    3 C( e) G& w" r4 y! j+ W39
    1 w' g- _7 T2 _! m, g+ s40/ g8 l% m1 j$ v" d) s
    41
    & o& k; \+ E, e423 V8 z9 S) i6 e2 ?; P- H
    7.3 训练所有层
    6 T5 S& ]% Z4 j# a; S4 Y/ T! A+ z2 L# 将全部网络解锁进行训练
    5 V% J) P, {! w7 d2 o* |for param in model_ft.parameters():& u; J5 @5 o; U( B) S/ h& t: j
        param.requires_grad = True
    7 S5 a8 B( K% _3 J2 R# k# f9 `1 n5 G+ W( r
    # 再继续训练所有的参数,学习率调小一点\7 Y0 G& p. ~- W7 M- y! `
    optimizer = optim.Adam(params_to_update, lr = 1e-4)
    6 ~8 O) H) G$ Wscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)
    + e! Q+ p. D5 d& i2 |' z; k1 Z1 M. Z8 W  h* {
    # 损失函数
    ; I$ c/ K& M1 o- P6 ~criterion = nn.NLLLoss()
    4 Q" D/ f+ Q0 W# b# i+ _( D1  K( P3 |% ?  j7 B
    29 r& i# g$ x2 z0 K; T0 h
    3
    " h" q( N4 s- ?$ C48 b, k- |( \+ G
    5
    9 u# L1 v8 N* k# y/ k9 ]6
    + z' ~2 K7 z: o1 |* K' H* d9 _4 ]: E/ M7) ?. S' W9 t- e! F  s
    8
    8 B3 P/ k0 q3 f! o  `9
    5 _& W0 M' X2 F9 T; M: O. \10- T, @3 g. s7 F6 I: N7 L/ R
    # 加载保存的参数4 b$ y# K4 U' ~4 q7 @; x! Q- t
    # 并在原有的模型基础上继续训练, P* x8 t5 X+ K1 S$ l- X7 p
    # 下面保存的是刚刚训练效果较好的路径
    - n' d" _2 n4 _9 k; q, P2 U. r1 Ycheckpoint = torch.load(filename)
    7 u5 |) z$ ?$ U2 Ubest_acc = checkpoint['best_acc']
    6 T9 {. t6 f! Z1 qmodel_ft.load_state_dict(checkpoint['state_dict'])! R, O  g6 y: U
    optimizer.load_state_dict(checkpoint['optimizer'])7 Y, E8 f# M2 i5 L# p3 o& m
    1
    ; ?  R. Q0 y3 Y9 K2
    : h: y0 Q4 L; z9 s% C3
    6 V. r! @# W; v6 X2 i! R9 P3 `! {4
      [/ p$ Q/ @$ b5 m8 p- A1 e5
    9 ]# K# j/ |8 B" J( Y+ E6# }2 n% O! ~" @6 y" y
    7; }4 k7 o$ v$ u$ ~" t/ ?
    开始训练
      B; y2 Y3 `) c注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考
    $ z4 B# O+ B3 d# X8 i5 C9 h3 U
    1 {+ n9 `5 X( Nmodel_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"))
    : K" T* x% t# {! r* d: @% a. |2 b1# o( \+ t, D+ O0 A
    Epoch 0/16 t# g% l# [$ S4 I# O
    ----------
    * y6 N+ P6 H7 S: a4 pTime elapsed 35m 22s- ~4 U, q- G! |% i# Z1 R1 v& o
    train Loss: 1.7636 Acc: 0.7346) N/ |9 I3 j9 m& [
    Time elapsed 38m 42s
    1 G( P+ u- H. M* S2 _, N' d) Fvalid Loss: 3.6377 Acc: 0.6455
    9 R9 }- k$ J0 Y, E( YOptimizer learning rate : 0.0010000- |+ v$ g+ q. X7 p* [% ^
    ; Z3 \/ f9 n3 w# ~# k
    Epoch 1/1
    ( M0 R' S+ P# `+ j----------' q- c. N% }1 p7 Q
    Time elapsed 82m 59s: i" ~, b; b. d9 K, h
    train Loss: 1.7543 Acc: 0.7340
    * r5 ^( }. E1 M' E2 b3 ?1 tTime elapsed 86m 11s
    ) ?$ d. g5 p$ L# ?valid Loss: 3.8275 Acc: 0.6137# P! D) y6 X! U3 G) J8 j
    Optimizer learning rate : 0.00100005 D0 U! e, c" S7 v6 V$ c( z0 X
    8 `5 H9 l. L! ?+ _) z7 S( f
    Training complete in 86m 11s
    ) M$ }* O! t; n. d/ b7 QBest val Acc: 0.645477' w+ y8 y6 `& d6 k, a

    $ }9 J* g  @( [8 _1 n4 ?/ a1: E! T# L6 |$ i: c
    2
    + C+ |$ j# R  _7 B7 x: o1 o0 P3
    7 Y" K  O+ J+ [/ B3 F4: Y& B0 C& z7 Y3 a' r4 |" ]) o
    5
    ' h% z. {+ R* Z% \6
    7 \) Q" R# Q9 N; A6 D0 V5 Z7; P2 r  M7 t2 n) R* J: p
    83 I3 ^" {) d: e8 H6 \
    9
    7 ^4 a  O( C% g% }7 H. u108 D8 f: Y5 [4 [; w8 }
    11
    ; Y. E6 ?/ D. C+ l  O; j12, F& J% G4 N$ T9 s
    13
    - t% M( b1 H. [3 ?5 n; ?14" {; f& O+ V3 I* i. A" h+ P& _: H& o
    15
    2 }( `! v6 O2 o% t$ O5 z1 S16
    4 q( g% {6 I8 U" n1 L- T+ D170 \9 d) t5 G. Y4 t; B! T9 T
    18. Y1 q! e; K5 C+ s
    8. 加载已经训练的模型! Q# z) R$ l3 H8 P: c
    相当于做一次简单的前向传播(逻辑推理),不用更新参数7 Q6 S! z5 \/ s& H0 \4 ?

    5 L+ s: Y9 {+ V5 }* o! Nmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
    8 z* c: C. P  l% V
    . z$ X4 E( @- K# GPU 模式) x' Z% d$ _- `$ P+ R
    model_ft = model_ft.to(device) # 扔到GPU中6 K7 f% C4 r8 N2 C" N( k0 R( K

    " Q6 L7 `( I5 u$ C# 保存文件的名字  C; R- F6 E) x
    filename='checkpoint.pth'
    ) X5 h+ q' u, w+ l: [
    # T+ O' E+ I- u. Y5 E4 }: E# 加载模型; r' m" ~# |* e; |7 S
    checkpoint = torch.load(filename)
    / |2 Q5 B3 x9 Y( k+ k$ T; `best_acc = checkpoint['best_acc']$ j4 P# E9 B  H, G
    model_ft.load_state_dict(checkpoint['state_dict'])
    : w8 V! J) N: |7 O: p1
    & B+ q+ e$ p" ^: ~2" o/ X& e5 ?% M- k$ }2 U6 S3 V
    3! Z# Z0 v4 B7 _1 C; T
    4
    ) J- j6 e6 g2 @' K- h1 v5  G' r5 o" y! }; o' T- l6 c( [) Q
    6
    ) V4 u- q5 @) d( Y, K7; U: n# w& ]! O7 M; @
    8
    & x% E& Y, T7 J+ b9* \3 z2 T1 I. C, q; ?$ c
    10$ q6 g' X4 k2 [; c
    11
    6 q( k2 y  z% u12
      e$ r8 c+ ]5 m<All keys matched successfully>$ i! J; h8 N5 m# o+ s5 T
    17 C; G: p1 r0 k3 d$ H# S2 m' y
    def process_image(image_path):8 h# w7 s0 Y$ R8 Y2 C
        # 读取测试集数据3 d! a5 {8 x5 C) s
        img = Image.open(image_path)
    8 B; ?, `- Z) k9 C    # Resize, thumbnail方法只能进行比例缩小,所以进行判断. `8 B. n% t, `, Y5 ^
        # 与Resize不同
    2 ^. w( C7 e- x    # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小" B' ^; R: [! U7 R
        # 而且对象调用方法会直接改变其大小,返回None2 ^3 }! d. n' B- H) V5 @7 u' a  {
        if img.size[0] > img.size[1]:8 I* l. }1 @" d; f3 l
            img.thumbnail((10000, 256))' R2 `0 r! `% `8 M
        else:" A7 I+ s7 D& v4 e( x4 K9 a
            img.thumbnail((256, 10000))
    5 `3 P/ x' u4 i" m( q! G5 }4 t* b% ~, k2 J' Y( C3 M  W5 r
        # crop操作, 将图像再次裁剪为 224 * 224
    ; \# ~, s1 P- i  b* e8 b    left_margin = (img.width - 224) / 2 # 取中间的部分2 y. T8 C" {; A8 q* O/ b8 [
        bottom_margin = (img.height - 224) / 2
    , ?5 G) t0 Z0 u/ H2 x    right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度
    7 D& A# t4 g2 q- d2 J% {    top_margin = bottom_margin + 224& x; u/ X- {( t/ \* [! M
    9 @" [$ {# ~/ n# q" a
        img = img.crop((left_margin, bottom_margin, right_margin, top_margin))- t0 U0 a; E' O: I5 W; e

    / l/ e/ L8 F, h9 X0 l    # 相同预处理的方法9 @2 I; E0 j8 p! ~3 L! ~
        # 归一化. m, J9 I9 ^( d9 i) {# o
        img = np.array(img) / 2556 H- \( \/ P$ H5 U+ o; z
        mean = np.array([0.485, 0.456, 0.406])
    + r& r! B$ \5 Z- m    std = np.array([0.229, 0.224, 0.225])
    " U8 v  E0 M* J8 B    img = (img - mean) / std% b4 H- _  `1 K
    2 w* ^5 |' L1 y% Y; ]2 Q0 Q, `
        # 注意颜色通道和位置
    + }! x3 F( K7 o) H    img = img.transpose((2, 0, 1)), S$ t9 A& f+ M( x7 v
    ! V5 ?, A# w# p' @% F: }% U2 c# X' V
        return img
    9 k3 m; k- V# ~: Y4 j( g: f7 [# M
    8 ^: s: g& Q/ }: |$ xdef imshow(image, ax = None, title = None):
    3 T6 ?/ O6 w8 t( F    """展示数据"""- u, D7 R' p) O8 D
        if ax is None:5 U6 ]( L6 z( V9 d0 N9 y
            fig, ax = plt.subplots()
    9 |' s3 n# L; \! ^& [4 B# ]& a0 h
        # 颜色通道进行还原
    / b1 }- l0 C) p/ r1 }- c1 o    image = np.array(image).transpose((1, 2, 0))" F+ D( H( j6 Z( V( Z( m

    " M' r& t" ?+ U: p: c    # 预处理还原+ G/ s# q$ s' h
        mean = np.array([0.485, 0.456, 0.406])
    $ n. `/ a1 G# q0 `" P, G    std = np.array([0.229, 0.224, 0.225])
    ' ~$ w3 B1 b' o6 M" f/ `: O! N    image = std * image + mean* J2 ^9 D( e& Q
        image = np.clip(image, 0, 1): |7 C; m: `% Q

    + @" I2 a) P1 t, B0 J    ax.imshow(image)
    1 v. I- w4 e+ @* p5 s! g" C! l    ax.set_title(title)
    5 U# I% v* _" x" r' m- E) b
    9 k$ L# \- B4 _2 X# Q2 j    return ax. u; C7 f# ]  U( @7 @

    , K# w% X+ l7 X2 }/ X, Yimage_path = r'./flower_data/valid/3/image_06621.jpg'
    * Y4 u: _8 W8 l1 O; Gimg = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理2 M4 K- {% ~8 [- @: w1 t' U
    imshow(img)
    ( P' X1 k: Z9 d6 g, S% ?& [& P% M3 o
    1
    8 d: T% _) {1 j) s0 m2
    6 s6 J3 n7 A8 Q3
    " l* p, q: S/ y* a9 W% T( H4
    ! S9 f6 A* y& z3 [% }1 f4 e$ j# N5
    & A. z1 N3 y7 h' f$ V7 i$ h. B6
    ! u' z# b! \3 I7
    # @) @, {7 r3 X, p# x* w8
    6 S+ V7 I/ q, ?- |/ Z7 |  |9( k) j$ O3 T( \+ v5 [) Y5 s
    10" n& W4 w2 r, N6 k9 k6 J4 _  f/ x1 f/ t6 _
    11
    1 j. a; o2 M1 Q3 ~12
    7 `6 u3 J) Y* a7 B6 F( c13
    " Z* K: Z( n2 \* T. H( ^14! \0 |: \& }. v9 |  O
    15
    " n. n) i" |% V; l2 q2 J; y, I16* X% b: e! I4 v4 L; V' ~
    175 T! m/ v. o! z5 u+ Y* P$ N
    18
    6 F1 d5 m: F" E7 y) ~+ E" h, f19
    + e, L; h$ J: C( X6 M20
    % U8 I. R' ^2 r. p- T21" j$ A" d- l+ i) @+ B
    22  N; O2 w: h3 u( ~
    234 t  d3 g$ i8 x! Q: Z' ^
    24
    # M8 f. [1 Q5 H8 X$ t25
    # n$ L2 ?1 u* z3 ^26
    + G  x/ X* n! Q* p) H1 |27. K) M% e  f) U& H
    282 \$ a+ J2 h7 M3 C
    29
      z1 v0 z+ l$ N! b7 [30# I9 \' w( K! l7 {, _
    31
    4 B1 C6 N( w# g# j3 g# W2 d32* F7 _3 m: K/ W% J+ I
    33' J9 R, B) f6 n% ?9 G4 q  p' C
    34
      g" u- ]* i$ _+ K35
    3 W9 |. F! a5 M" Z, L$ _3 n36) c6 d! u' Y2 y+ ^/ U) z: G
    376 Q8 z. G2 `. s3 q0 `: G- O
    38  F3 O8 I4 S( Q, {) g" ?
    39  {+ a) _: O9 n3 J% u6 @* u; I" \+ f
    40
    / N% [( _# i% _$ _* t% O" p9 D41
    : a. M) W- I0 B42
    ) q( z8 R1 d! ~9 |43- ~% K4 g$ ~# e, P. \1 i* u
    44$ h1 e9 M8 i. c5 }
    45/ A7 {3 v3 `0 {1 o. J: o1 }9 \
    46
    5 Q! h! _7 W' X0 q6 `4 Y8 Z% B+ J47
    % A' U6 @6 a! Z4 ~8 n7 y! D# s2 }48
    ( L" Q% h$ O. S! u$ t5 l49
    $ l) M* m1 d( p* d! C50. J2 Q: w5 U' I# W9 |/ a$ E
    51/ w) n3 p% E% h( {) S3 I
    52  F# h6 w2 e6 }" j
    53
    , Z+ j6 u& B% z54; b9 c- }; @3 d: ?
    <AxesSubplot:>
    6 |  Q: h- P' [7 p/ G1- X: q" O$ V  X
    7 o4 u1 Y* J. y; X' U' [' A
    上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确
    0 K& u  [  f. v3 F- U1 h
    " M2 g, i- I1 ^9 ]& p) V7 K! Y; Qimg.shape
    ! j; p) u$ o8 W) k9 W19 A& B: o8 Z4 m& o! [& ~
    (3, 224, 224)4 @# K" G% n/ |+ q; {( Q) F8 {1 v! a8 O4 t
    1
    ' C% [: C( |& R; V5 x/ `3 K4 l3 |证明了通道提前了,而且大小没改变7 a! r& c: ~( {( m
      E5 Y) s. t( K( L
    9. 推理
    . M: |. f/ ?1 c$ p8 N& C1 o3 u; C1 ximg.shape# A8 n5 a. U* Z/ _% n1 M

    % q- r/ `# Q; n% Z% ~# 得到一个batch的测试数据
    0 s: t0 e9 b( R( ]8 L( udataiter = iter(dataloaders['valid'])
    ' b/ B$ Q! F0 [4 Pimages, labels = dataiter.next()
      m; M! u0 K$ I7 U) x. K1 N, E1 x3 ]% v# I0 `: o) _
    model_ft.eval()
    $ a- H4 Y# S/ `' S# _
    0 H( }$ X7 E1 Y& w2 ]- t& J/ [if train_on_gpu:
    , n4 L% _% B1 m2 y) V* O$ j    # 前向传播跑一次会得到output% z0 y& k  l% e$ [/ x# j* q( D
        output = model_ft(images.cuda())
    # K: ~8 O, G* ?- G# U- Oelse:! k, P# W' J# n+ g' L
        output = model_ft(images); u) m4 \+ G9 P
    . A& ^- i$ J& s2 v0 U) s
    # batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值2 x) u+ Q7 x" D) V/ Z( H: F& W
    output.shape% f; u2 i  v8 l' [* D0 K
    3 y- X0 F! H) Q( p3 ?
    1, n! Z8 a  @  h* c, X3 d
    23 G* Y  B0 t1 l' E8 f  _
    36 h# y2 p* M6 n
    48 V2 \" I# ]  i6 C# o9 Z+ c: M
    5  I6 V& K/ X' J' B
    61 G! B- t$ l9 n+ L
    7
    2 W0 v- y4 Y4 W* e$ d8
    8 F1 X8 U/ L7 X0 j: a0 j9
    3 _. J1 ]+ C8 m10
    ) I$ {; g; j# L: W0 I3 j* o11
    2 _$ Z! b+ I- D$ Y( A$ D+ H127 d) k1 h+ P6 o
    13
    ( ]( z/ o3 C6 ^1 X2 w% g3 E& V2 f145 x  y; t: D( [- V5 p
    15
    1 R  U) H: Z4 Q6 D1 O6 |16
    / y. E/ a! n+ `. I7 vtorch.Size([8, 102])$ P4 R& v( ~. @  }5 Q: m8 h9 K
    1
      B* |1 c0 T, ^9 t! x9.1 计算得到最大概率
    . v2 K- T9 D2 O, I# u_, preds_tensor = torch.max(output, 1)
    1 x( M# R% g+ f* y2 [4 _/ K. p$ M; }( E( L. h3 S
    preds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量+ W% }. z' f+ ]. x# \& `& S* N/ ]
    1
    " w: i7 ^" [1 f2) L" q3 G; Q9 {# ]- Q3 A5 \4 ~" K( E% P
    3
    # ~- j1 q- e0 V% }  [9 K3 d/ A9.2 展示预测结果5 _) M8 d' d+ o4 m9 F% y. T
    fig = plt.figure(figsize = (20, 20))6 m( c. Y! ?6 z. D  ?, P
    columns = 4
    5 ?5 Y; [( e0 D6 \  K2 \rows = 28 o! n; S# c# a. d
    8 f* @8 N/ @" F4 G4 R4 `$ m7 M
    for idx in range(columns * rows):
    - b$ y8 x& D8 B3 n" O# U    ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])1 g4 S4 p1 d3 ?$ h+ z
        plt.imshow(im_convert(images[idx]))
    ( R. f6 Z$ \" n! h    ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), & I, R+ t* a$ G2 j5 ?, S
                    color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))# o) y6 l4 t6 U
    plt.show()' a; X5 a( E* _) v* z
    # 绿色的表示预测是对的,红色表示预测错了  J3 v* n- W$ y! p
    1# p4 [' Y! w& m$ A
    27 h& \% E8 A/ f) V- F0 e/ a- r
    3% w8 ?' a% I" }5 R& _( B9 N$ W0 |/ i
    4
    ; m$ J3 M7 R' m- Z3 q5 C0 N5( r9 w9 t1 T: |' K$ W9 k( h$ v( s
    6
    : ]- G, z0 ]9 s5 \. a7
    - }; l+ r. o  n) N  I/ U$ v0 V( P! \8
    : H; j& ]  ]% d8 r9 C' h* `/ W. A4 J9" ]/ z4 A% |* v+ [% h0 V
    10
    4 V  _/ W( D+ A/ \4 O11
    0 H" _0 D! u- S5 I% |6 C7 l8 k8 a+ ?' b; o
    # B1 c! P* v" [

    : D3 V# G; v, G5 J! X————————————————. |* ~% p  x2 u4 T( G
    版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。
    1 w  A! R! J8 ^+ m& w原文链接:https://blog.csdn.net/LeungSr/article/details/126747940/ p  l4 b/ L: c1 n# R5 x6 C1 \* `# J

    & {1 L" k$ R4 F: I7 w7 C7 A0 q* e8 A
    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-9-22 20:58 , Processed in 0.382466 second(s), 50 queries .

    回顶部