QQ登录

只需要一步,快速开始

 注册地址  找回密码
查看: 2824|回复: 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)实战案例1 k* a/ \' k$ O5 P" U
    6 S' P0 {6 E+ Z! v1 ~3 Z! f
    文章目录' N9 }1 C$ t% ^( r& B2 o7 K9 n
    卷积网络实战 对花进行分类
    . N; E& I0 B& w$ ]# ]数据预处理部分8 G' D) p5 P2 m1 O
    网络模块设置
    2 b+ r$ h+ c0 G5 G# Q1 x网络模型的保存与测试
    1 f9 n: _4 |7 Z% P. D; V2 d数据下载:' m4 C$ {5 I  Z4 j4 R
    1. 导入工具包
    - P7 s, o1 i1 j) X/ h% ^/ o2. 数据预处理与操作
    2 R' r2 P0 l+ m3. 制作好数据源
    / M" @$ ]' B# @* s" R+ N* |读取标签对应的实际名字
    7 S3 P* f3 J# Y5 [/ T0 A3 r% K7 b4.展示一下数据
    ' ?% e/ ~, g5 |# I% c! m5. 加载models提供的模型,并直接用训练好的权重做初始化参数7 M( T5 a- P8 c1 `# s5 V
    6.初始化模型架构+ D4 J# R9 @4 {* ^2 A0 ?! R
    7. 设置需要训练的参数
    $ _; B* m4 T, D+ t! R  r" R9 B7. 训练与预测8 w( |/ H" @/ g: x
    7.1 优化器设置/ D3 j& q( v' k" k2 T# y
    7.2 开始训练模型# r. d: a/ k' s1 o
    7.3 训练所有层
    - G) z9 @2 z; d- t! R3 L开始训练6 s: h! m7 l6 J1 a. x  T3 H  J
    8. 加载已经训练的模型
    " O( n+ X3 T" t$ K& x9. 推理* b( ]% v  S; }( c: m3 Y
    9.1 计算得到最大概率, J% z4 `  ^3 ~) A& n
    9.2 展示预测结果
    : B, d4 g- H4 {& Y' f% k5 c" {: D写在最后
    ; i" u. c  D/ Z) l1 r7 o卷积网络实战 对花进行分类" k. C$ n  T5 J6 h0 M- f/ }
    本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能( X6 a9 v; D9 c4 \# _

    + O+ Y9 a  X1 W: O, _. s在文件夹中有102种花,我们主要要对这些花进行分类任务
    9 n% K, h4 c1 y5 V1 ?( T4 F6 ~文件夹结构
    - C" a6 m! w  b/ V4 ?8 z  @
    2 ]( [2 O9 K. P' K* Hflower_data) V/ T; D& c1 @
    $ i1 V7 |- Y8 m. f5 x
    train
    1 [" d: ~6 Y; E  w  g+ T- ^, M
    ) x6 N7 N; s* @% @1(类别)" l+ v% W4 K- M- _/ W
    2
    # e, G. |8 E! u- `% g5 ~xxx.png / xxx.jpg5 ^2 l' ^; f' I; n; ]' g4 Q
    valid4 W; Z" \3 e! ]
    4 G6 j2 R& x' U3 \: h: ^4 j2 p
    主要分为以下几个大模块; d, d' @8 e) o7 ^3 Y+ X
    1 I( o% `' w' V  u" L0 F
    数据预处理部分
    . u" ~6 [1 \2 }, S; z数据增强6 t3 G4 F$ s$ M, P% U0 u8 i
    数据预处理2 g, `7 [  H1 d( ?
    网络模块设置( W2 z: @$ b3 ~6 ~- X4 p4 S, L
    加载预训练模型,直接调用torchVision的经典网络架构/ t) G' }% m: |
    因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务
    8 f; G" I* |) D网络模型的保存与测试0 _' N3 x9 _1 `2 v- m/ S9 W
    模型保存可以带有选择性( ?# @4 `" X1 f9 I, J& S
    数据下载:+ b, \9 [3 c  f5 A
    https://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset. d8 C- D7 e/ n# j

    ' t* I4 N" M4 ?: P/ C: z改一下文件名,然后将它放到同一根目录就可以了9 _' R, W5 X! z/ k& J! x$ q9 i1 c! o5 r
    . p& R$ ]4 w7 y' N# @+ W, ~2 z
    下面是我的数据根目录6 g! q3 a6 F* L) G

    + ^' J9 C5 a) e( L9 d
    # A# K0 R& h2 t/ H( Q1. 导入工具包
    ' N9 n  K* l) Z% G+ |. R) @import os
    # ]5 r& F) Q' X4 J: u3 e! ]1 Eimport matplotlib.pyplot as plt( M% ?9 k7 x+ Z) K  L
    # 内嵌入绘图简去show的句柄' b8 `, G$ v" j* x1 }7 ^
    %matplotlib inline
    / u6 X& [2 ~. X5 U' I/ L2 G0 S* _import numpy as np; X3 g- d8 g0 `( e
    import torch
    6 s/ v& s; f7 u0 ?& x5 \' l4 N0 pfrom torch import nn1 m0 u7 m; b5 S- s/ j1 ^4 c# I

    . D! W& H: i: B5 W: x5 ^- Aimport torch.optim as optim3 m" F5 B9 c; x+ S% R: t
    import torchvision& v, h% Z9 L9 X  @9 Q# D
    from torchvision import transforms, models, datasets5 m( i' N) H% w

    . ?, ]( {! p' D+ Q0 O0 L$ wimport imageio
    , M8 ]! c1 p0 |2 K, Yimport time, J0 Z6 b: U; h0 D; I5 }" r) m5 g
    import warnings" O9 p' _9 s3 C
    import random
    6 B; g$ k0 K6 s8 m4 Jimport sys' e2 Y- Q4 H" O
    import copy/ B8 n$ l, h- G5 X2 T: C
    import json9 C$ X8 S. G, P$ u  v1 j# {
    from PIL import Image: j) y/ U: M  r/ d, U

    * I1 _7 v% O' d1 v& v2 k! g$ V' H% T$ h3 W3 y4 L, X& g$ d
    15 d% C7 _+ x6 }9 ^5 F; J, Y
    2
    * f; L) x& p. @1 R4 W7 D3
      c* k; \# r& Y" d2 p4
    + A) a8 \2 p; x7 C& Q& b4 Z* \5, H% q! y3 l- w
    60 A, G: h- q( g" i
    7  q' m, }- G0 @3 \
    84 q( j' V6 H4 C+ F; c/ }: Z
    9/ p  Q  r+ x6 B5 z# K/ r& x
    10
    4 `' N# T- \8 _9 f11
    2 u* S, e: x7 `2 q1 T: K4 q2 l' q/ d12* ^1 O" V, W3 J: f( w1 Q6 p" K
    13
    4 k7 t' t( Z5 b; A9 c0 h14$ W( y- L/ A1 [( W- L: Q
    15/ Z7 `6 {4 E4 S: P+ M' j% v
    16
      e: D2 a* ~9 b; l$ ?/ \0 [17
    ' s! s$ y! k, @6 [6 b% R- T18
    ! V! E: m' U9 d) B! ~19
    - o' D$ [, z- d: l- [( U20( d4 M  Y( y9 S6 ?
    21, `$ h) u, F, W% k$ p
    2. 数据预处理与操作9 g/ ?8 h0 ^0 ^6 `# \1 \0 z
    #路径设置
    5 S& x1 X+ Q# A- ~- Cdata_dir = './flower_data/' # 当前文件夹下的flowerdata目录( j; s* ~' @4 I. O* x! x
    train_dir = data_dir + '/train'
    - i- B' u$ y& g3 D5 ]/ z+ ]6 g9 a/ yvalid_dir = data_dir + '/valid'
    7 W4 x3 Z: E& C! |: |1 C" h7 \  p1/ b. f9 ]8 e8 H
    2
    , I1 V" d/ j* t+ Q# T/ O- u* ]9 S3
    . ]% U3 P- M6 E& H2 o# s2 H$ `4
    ! _* Q7 V5 o( mpython目录点杠的组合与区别
    ; Z: \! m1 @* Z6 U0 W8 [8 o4 Z注: 里面注明了点杠和斜杠的操作
    ; @/ {' b* J" f) p, L% W" Z9 S7 L8 ?) t6 y, @
    3. 制作好数据源8 I6 o! B' r8 ^
    data_transforms中制定了所有图像预处理的操作
    5 N( T6 L, s6 f. ]8 t1 Q4 e: QImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片3 l" J, z0 }6 ]( t% E
    data_transforms = {% g1 P+ k4 S( H+ Y. u* O( W+ J
        # 分成两部分,一部分是训练
    1 a5 m$ a( @: W* C: f3 j6 b0 T) G    'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间, W/ ?- ?5 a4 ^; J( S5 g
                                     transforms.CenterCrop(224), # 从中心处开始裁剪
    7 ~) D" t& d5 O  S% q8 D- P                                 # 以某个随机的概率决定是否翻转 55开. J0 R$ K  D4 ~6 P1 o* K& n# w
                                     transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转$ P2 k3 ~& a" B' f4 b
                                     transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转
    3 F& D" R5 J- M2 l% O& _                                 # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相
    8 R1 `# L0 K; f8 \4 P2 l! z                                 transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),6 J* t+ z+ J% U" |1 c" |
                                     transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB5 `* Z/ q3 C+ B: w
                                     # 灰度图转换以后也是三个通道,但是只是RGB是一样的2 W" S' ]6 a7 L# Q- s9 e
                                     transforms.ToTensor(),# u$ L! b* I1 N! b5 E
                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差
    ! p4 s  O1 D" ^% S( `7 ]                                ]),. @7 _1 T' Z5 |* j% K9 }( e! A
        # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
    7 i: i1 J2 e; p. D+ [/ v    'valid': transforms.Compose([transforms.Resize(256),
    # x% z3 D$ Y1 J& P$ b  C                                 transforms.CenterCrop(224),$ z1 k- S& D$ A  |/ l: u, A9 ~
                                     transforms.ToTensor(),
    * m( T/ @# {/ ~8 W5 \3 f+ w* k: H                                 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同( H. w9 p( @; b; M$ v# {# Z! P# A
                                    ]),
    $ z% o/ R* `% t& ~7 X! `}
    , A) @% _1 u( X) L+ d" ]9 X! F
    . i- ^: M2 u! C8 t8 F1
    # p0 p' v5 a" ^( j& _2
    0 O2 z4 {  E7 s6 _30 L! C. ]7 v  v0 [# M( q
    48 o* e3 F) Q6 e$ \4 |8 j) }/ c
    52 C  E- D/ o: r7 Z2 J
    6& M$ a+ o) s" S" _. R- l4 o3 }
    7" B" @/ i. u1 I! H# o+ w6 h! ~- p' A
    81 {$ f. I' U( d2 J
    96 P% j  c, D; x5 }  A, y! V  s
    10
    4 Q0 m( U, Y, }. Q& }! V111 w6 S8 p8 h; A4 B% Q! x# k# L
    12
    - u: V+ y8 [# {5 I( L+ J, d13
    9 q" s5 t- q4 q( C14
    : l) V! Z9 r( D! H9 C' x" D3 _& M15- Y  ^7 r; _4 v: S5 q) [3 {3 ]2 c
    16  p' B* R. R/ S3 q: P
    17
    ) ^$ U% L" w' ~9 z4 L7 K3 l$ V18# }4 q( M. ?0 k( u# s
    19
    5 a3 s  c6 Y8 c203 q; ]( H( N$ @. x
    21, p( K. e3 S7 ^$ A2 E$ {
    batch_size = 8, F$ w9 k. `6 H6 R/ H  [2 \  j
    image_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}
    0 A2 \$ n) d( P% S5 |+ \6 A5 j% y2 \dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}
    : u( j; F8 G! ]0 z- L: A! idataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']} * n  I  P: P: \* w
    class_names = image_datasets['train'].classes
    + F. n' z) u, T# h3 C  [8 C" ?9 F2 q( N# |" t
    #查看数据集合$ m5 x" B) p- U7 ^5 y1 n5 p! G
    image_datasets
      o/ j4 U3 \# ]: w( v, T& I
    , G4 h" H7 A) r% ]* Q  Q1 \1 k1
    5 l2 m7 S4 A& Q' S' k2 R  r' e2
    3 v7 @5 Z; r6 A) C5 t3
    1 T; h- V8 ?  X4
      G0 N7 h" B9 |58 f2 D6 }! i, f
    61 f  O# A! E9 E2 ]; n
    7
    + |+ o8 Y: [- l8' `4 F; ]7 N  `7 I
    9
    " X% V  ], P% t  l{'train': Dataset ImageFolder
    2 c8 S) ^$ F9 v     Number of datapoints: 6552) K- a! P/ l" X* G+ H
         Root location: ./flower_data/train4 n; c) o2 D4 z, X
         StandardTransform
    ; M4 S5 T; |2 }7 U( E# M6 l4 X$ J  ? Transform: Compose(
    * L# j4 w' M& L& |  Y# K- a                RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0); L/ `; l3 Q  F# u: r5 l* L/ [
                    CenterCrop(size=(224, 224))
    ; e: @5 n. _* ^' o. }" o5 V  \4 ?1 O                RandomHorizontalFlip(p=0.5)& z' V+ s: n  G! c5 p3 V
                    RandomVerticalFlip(p=0.5)2 w4 S; }2 {' ~; U. _) _
                    ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
    9 T% f4 a( }$ ^/ z: B5 Q$ @! K                RandomGrayscale(p=0.025)9 a4 @! c1 {6 P, T% P% a: y& S
                    ToTensor()" I* Q2 J: P" S/ v
                    Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    # n5 D2 e/ B4 x7 z3 D% {            ),1 a& _% G' l! c& F1 K
    'valid': Dataset ImageFolder
    0 d" Y/ U+ N2 R: v1 a: ^     Number of datapoints: 8182 a1 k4 Z3 ~% |% p+ E
         Root location: ./flower_data/valid
    3 W; T0 r1 H/ D6 z. u0 x% w4 M, `8 {     StandardTransform3 W# q: w7 v) J8 p# k
    Transform: Compose(
    9 E. l1 r1 P* [$ ^6 R- w                Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)# |) Y9 c% X; p" e- ?" J; I
                    CenterCrop(size=(224, 224))' U5 h8 Y3 V$ ?& T+ [* K( ]$ K1 n
                    ToTensor()  f, M- T* K4 z
                    Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    4 |  M8 u4 _- c5 `2 J# y1 v            )}
    4 i7 d' W- k! W$ E0 ]7 |! t3 P, E; \
    1% w" a4 _6 c" \+ M
    2- [2 u! g/ u" e$ u2 P0 f
    3
      r8 d0 |( n7 n8 `" {  |; P) Q4
    4 y2 }% B: J! F  I4 B5! r- T* _6 R. }. J
    6
    ( p% X5 n. X: c  X2 N/ \79 b7 z" M6 l4 J' E+ A3 S
    81 }7 U- o6 T- S+ J# m% \: y
    9
    ! b7 M; D' m3 X6 M& p, ?2 Y102 j" W# ~- C9 H) ], i, m
    11! i2 m9 W& W9 X. a4 u3 w* s
    125 {+ E$ [( |" a. F6 `, x3 F
    13
    2 D. z, H; W* r9 x9 E2 J. g( T14
    " f; ~, L6 W, ^2 k% {156 A9 M% j# j5 u! ]1 U3 S& q6 `( ]6 p
    16
    3 y9 ]" K$ o9 e! j17
    5 Z+ A5 v7 j6 }2 i/ q18" z- f9 |- a- `3 u% \
    19
    % F( p3 c* }) ^$ [8 l3 `( L  p20
    # ^6 Z& V  C% Z3 @  r213 f: `' ^3 f! x+ M
    22! H; l! U; ?5 x* {7 o5 M0 F
    237 X/ a3 N/ d0 Y7 q, o
    24
    " H1 C1 t( s2 m5 k# 验证一下数据是否已经被处理完毕
    6 t* C. k, e0 T4 `' Pdataloaders
    1 u: r7 N( u8 e9 U0 c' U/ T9 |9 M  q1
    * O0 C+ e& `# U20 [: Y( z; X- I, [8 a& |
    {'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,: P: M0 @- X6 J( v9 D" n" {$ I6 E
    'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}
    3 t6 [' W  t- Y/ o& [  }1
    ! p  h9 O) e0 j- |3 o+ C: S2
    , J* @8 O" I3 L  p6 K3 B- ydataset_sizes9 W6 `8 d* z( C: X
    1& B& O* G2 [5 M: C) Z1 y2 i
    {'train': 6552, 'valid': 818}
    ; s( E9 X% J  i. K- x13 ]7 b  q" o& ]; u" N" d
    读取标签对应的实际名字
    & J1 q/ E% u( C) G$ y/ c" f使用同一目录下的json文件,反向映射出花对应的名字
    % C) ]# t+ G$ K, u0 N' }2 Z
    & J2 j, o* [8 i! _1 k# Xwith open('./flower_data/cat_to_name.json', 'r') as f:* m: m8 L- k9 ]( O* Y1 H: `
        cat_to_name = json.load(f)
    . K+ w2 z0 J3 N8 F14 |) L; ^5 p2 Q6 q
    2. O; k6 ^) l) d
    cat_to_name2 E3 {8 L' |; O1 h, @* a$ G% Q
    1
    $ q9 g( E% b- D! D! @$ v{'21': 'fire lily',
    : z# j5 d7 B6 a9 }6 L% Z '3': 'canterbury bells',$ O  E5 I; I6 a- D3 N+ K
    '45': 'bolero deep blue',$ K) w$ k8 n" e& {: y
    '1': 'pink primrose',& s5 v- R* C( i8 ~& s' t4 Z
    '34': 'mexican aster',5 R9 Z+ e' r! p) U, X" d6 y; z
    '27': 'prince of wales feathers',2 A: H' s$ H- e+ J
    '7': 'moon orchid',
    6 c# c$ c4 D6 ~! w9 q- p3 P2 N '16': 'globe-flower',
    + Q/ w0 L1 O/ O  |8 z& q9 q '25': 'grape hyacinth',+ i5 ^" F! K! ?3 s- X/ _8 |( j
    '26': 'corn poppy',
    $ g8 s8 l, y) D& s '79': 'toad lily',, @( j7 p8 D# f4 B
    '39': 'siam tulip',
    2 r$ [0 \9 S  f' |  W  H# G/ h '24': 'red ginger',
    % h6 {! P, U! H: E) j7 j' U  j '67': 'spring crocus',8 K$ v" k& W  B
    '35': 'alpine sea holly',+ @' r& [* j) t8 T
    '32': 'garden phlox',
    * I( i6 h7 B. d4 i$ G- y7 \ '10': 'globe thistle',; S' P+ A9 Y1 j  P, T, c6 _; }% z' R
    '6': 'tiger lily',! A) p+ M0 K1 f. w1 C6 x$ ^
    '93': 'ball moss',3 o- I8 C( g, ?( A6 t9 v2 e
    '33': 'love in the mist',& v3 Q" h9 ^  c9 |/ c  s
    '9': 'monkshood',
    : Q+ c7 L! R) c  _5 g '102': 'blackberry lily',. I* |( Y7 H+ e: ]8 I
    '14': 'spear thistle',
    3 Q9 `) q/ d' @6 e( G( p6 Q/ e '19': 'balloon flower',
    ; ]2 W9 C. e" s( S1 e, @% M. j '100': 'blanket flower'," R$ D6 F0 p$ U) a! p* d# a
    '13': 'king protea',( O& B1 v) f  `9 U& v
    '49': 'oxeye daisy',
    ! I  V5 b9 r) {3 E$ | '15': 'yellow iris',
    ' i2 T0 n% u! J* g: r' x '61': 'cautleya spicata',
    ( U4 b; g% ^2 p7 [  Z '31': 'carnation',
    5 x4 x2 A' `, g  Q; J; o  g6 q '64': 'silverbush',
    # ~* |* P1 g7 l  F# f9 i '68': 'bearded iris',
    * T4 z: j$ d; f5 A- d3 z* ?7 \) ? '63': 'black-eyed susan',
    5 C  h* }1 h; ? '69': 'windflower',6 H- ?7 E, }& Z) {  A
    '62': 'japanese anemone',
    ) N$ ]! H6 z' _% g" e5 { '20': 'giant white arum lily',
    ' u- y( w& P* E9 X '38': 'great masterwort',
    5 R, X" a' U1 F/ z( s '4': 'sweet pea',7 Z- H* T+ ?$ l. g7 \! Q* f
    '86': 'tree mallow'," m0 y8 i* x) a; d2 k
    '101': 'trumpet creeper',
    : y) X9 d( T1 f# V" i '42': 'daffodil',7 V# J$ B0 a  x4 P
    '22': 'pincushion flower',2 @* N6 ?$ |) O& ?
    '2': 'hard-leaved pocket orchid',
    3 n. e5 G2 G& f. V '54': 'sunflower',% f, g6 ?  u$ N7 g, b/ T" c; m
    '66': 'osteospermum',
    6 G+ E7 ~2 Q: P6 | '70': 'tree poppy',+ Y- t1 C; |4 g: V
    '85': 'desert-rose',% T, ~- F! I4 u. x7 o, l- [' V
    '99': 'bromelia',
    3 }( F* o! S$ V2 l$ E! Z$ U1 V0 q '87': 'magnolia',
    ( m6 C( w6 v- d$ m/ a- I '5': 'english marigold',
    9 U4 n8 W2 F' e: |7 `! r '92': 'bee balm',9 v/ A& X; _, X
    '28': 'stemless gentian',0 `+ o. {) V& l
    '97': 'mallow',
    $ b6 ^1 B& p" U '57': 'gaura',, X5 B' u& ~( t1 c* X
    '40': 'lenten rose',
    1 t0 R" a- o5 Y/ h( _9 ~ '47': 'marigold',: ~* u7 \2 y) }, w0 D& W; a; `
    '59': 'orange dahlia',
    6 ?& N- f/ y6 `7 i8 Q4 T) E! F( X '48': 'buttercup',; G, b( N3 G/ h, }6 O' v: J
    '55': 'pelargonium',) p4 `; Y6 m# u" H
    '36': 'ruby-lipped cattleya',
    . c' n( b7 l: A, {& n& z* x '91': 'hippeastrum',
    8 ]# e# g3 @: @, J- s" a0 F$ p '29': 'artichoke',
    , x2 W3 q. Z! M1 `1 j4 ?1 k '71': 'gazania',' U0 C6 ^& C# T" m0 {
    '90': 'canna lily',% w: ^$ @; r  H. z2 B
    '18': 'peruvian lily',  h) `, X; p' L! S
    '98': 'mexican petunia',
    & J0 R& @7 \- _1 E8 O; y* i' [" T '8': 'bird of paradise',
    ! E! T. r# |2 v2 o. v2 i '30': 'sweet william',
    , z5 q$ E6 p1 C! g; a6 B( f: k '17': 'purple coneflower',- ~: I: B9 @1 ^
    '52': 'wild pansy',+ g& c9 ?; x3 Z1 U  [0 Y: k2 }, \
    '84': 'columbine',- g2 @( I' f0 Y; D* R, \8 W
    '12': "colt's foot",
    4 c6 y* y' K, c4 Z1 e# ^ '11': 'snapdragon',( Q! T+ d- C; X, b8 H
    '96': 'camellia',3 g" H: R7 W- j+ I% p/ K
    '23': 'fritillary',
    " e. u2 v1 g! O '50': 'common dandelion',
    2 t" q1 }& o- K1 d2 g '44': 'poinsettia',( D1 t6 J6 @3 V
    '53': 'primula',
    2 q  @8 P$ B5 }  X& o; W6 Z; j/ p% ~ '72': 'azalea',0 L# K, Z+ y. e' Q0 ~, c
    '65': 'californian poppy',
    " s6 K3 s0 [1 [ '80': 'anthurium',
    ! s0 j: S  B) L. R9 O& O '76': 'morning glory',1 h0 S& D' C2 \
    '37': 'cape flower',
    : W! T9 s7 m, w4 c2 M6 D '56': 'bishop of llandaff',
    6 t$ ~# K9 [! ~9 {9 g '60': 'pink-yellow dahlia',2 u1 N9 C! J: j+ r* ]& j
    '82': 'clematis',
    / n, x1 Y6 L' S$ M4 M '58': 'geranium',; \# q0 `- X1 \, W7 s; f/ l
    '75': 'thorn apple',
    ! B9 F$ [/ C: G, \$ j '41': 'barbeton daisy',
    2 M& T4 w2 P: V5 t% C5 Y/ c5 `; K '95': 'bougainvillea',
    - X: P3 ]1 V; y: ~  b, U '43': 'sword lily',
    4 A: A0 G+ P' u& t% m  X3 K& Q( w '83': 'hibiscus',+ u6 R" {" f* [1 n
    '78': 'lotus lotus',( d0 _- R# a4 E' \- S
    '88': 'cyclamen',
    * o% K6 G+ y; j '94': 'foxglove',
    & L8 y5 k- H/ o8 j) a '81': 'frangipani',
    * c! `: ^# g; H; H/ O" G- P, E+ ~: ^ '74': 'rose'," l/ p% y4 v3 F- ^" S7 n# g6 B
    '89': 'watercress',
      e0 P7 j: }$ z2 u- h& R '73': 'water lily',
    ' v+ E+ S0 X! B( k5 | '46': 'wallflower',
    ) d/ @2 A( r2 `, O1 N; k. r '77': 'passion flower',0 U* j3 ]9 a1 i( X4 q. R
    '51': 'petunia'}
    ) ^/ l. i( V/ ^
    2 H$ a3 x! N0 |7 v4 W13 g5 m6 t/ m  ]/ s! ~' z5 t
    2+ V5 S( D) P& w6 u: y: e
    3) e* q& N% f9 x2 ]& W2 F0 [! X& D2 l
    4
    6 H9 O; c$ l# _: b: x) b5
    ' A4 V9 c* K7 H7 c0 [5 Z3 T3 \6
    3 L) @% m$ M4 Q7 m7
    " {7 G/ a& Y- i% i8# L, T9 S" _% j4 l. Z5 S1 |& N
    9
    + P7 @( K2 X* u" o10
    9 T9 S* D( \2 p) e11
    ' K) z; L9 @8 T3 `12
    ) r; G$ T5 M0 V' T- J130 J. g( @! }8 p
    14
    : d1 X, {4 K$ |2 I, H( B9 X' ?15
    8 Z+ Q3 }  U: f, _1 _% z0 j16
    1 Q" O+ d9 J5 a, W: {17
    ; x4 N6 o3 W# m- n6 C18
    4 H- \8 L5 ]7 g1 S- f19( y) N) ?1 v7 U7 K8 G
    20
    9 l5 L- t4 _! x* d1 E& x21
    $ J4 v1 i* M- |# k" z! F) m2 T7 u22
    5 n9 |. W: m( h5 t2 Z) M23* _, v7 R, {2 Q# g8 L. J
    24
    8 I5 L- N! \  k1 V25% K# N$ e2 O. |5 [6 Y9 Z2 o/ G0 h0 Z* ]
    26
    ) u# ^& O7 M8 i+ T) y27
    ! ?% A3 `4 C0 l# y1 p1 {28* Y9 ~% P$ a9 t7 v* ?& `+ L
    29
    / F4 q+ l7 ]7 M: O% }309 }$ \, X6 W7 p# F; G" t% r
    312 f  ?# D- ]/ |" f; `1 ~* T
    322 m0 v/ j- s, _/ I0 g6 Q- h
    33
    ! x# H( i: K  P# I4 ~& [34
    , M1 Y6 Y* S& A* V) {, K* k35( s6 Y% U$ w3 b) @% X: T2 [
    36( h* E2 P* k9 b5 z8 i( c" Z; }) Q
    37
      A) T+ a; ^/ b4 n/ f% q38# z! s% J, ?* A8 V
    39, x- ]& N. M/ d8 l
    40
    / f2 V, C8 P: b  I41
    5 L% ~; q. a, d; W2 h42+ e, h8 Y, N' }7 p3 t; f# q4 n# Q
    43
    4 I5 s/ @4 {( |" P$ B% k; B44
    ; E4 N4 c( K6 H  Z* C9 D0 f# N45, t6 B& P! K' `
    46
    7 y6 p" w3 }: q+ r+ h0 h# k470 y! }  {' @! g3 _- c, b% w! O: k
    484 o. H3 Y! |: ?
    49
    / l# Y8 y$ _, Y6 l50( |) p: P, e& T( v
    51  z- D; p' z1 [, e2 m
    52" ~; n2 _* J1 n
    53. E2 W+ B/ K4 l# N5 H" }
    546 ]1 i7 I+ K1 H9 J
    55" i* J: Q+ V: W
    56
    7 q: n; h4 z9 f) I) [576 V6 g9 L" V; c7 W. k0 ^0 _; v
    58
    ' a; Q: s/ r- Z5 ~59$ L: ?* S0 u" {5 T
    606 K2 w4 O* ~) q0 _. g8 b- o5 u0 @& |$ G
    61! j- a9 r: H6 N; @" Z
    62
    5 Z1 d* j$ ]) V% B2 |9 Y7 m+ A6 b63
    % a: o$ [/ t  W. W- U64, C. p2 r! Q, y" O, L
    65
    6 F) s" S% w; m! G! ~+ U661 a4 _1 K2 q- e& T. R* [+ G
    67/ g/ e( ]4 a4 D1 F6 M1 M0 v
    68
    1 f% d4 k) A# j) i+ Z' S69
    ; r2 k/ g9 p% j70
    3 U/ n7 C- o& d  z+ [9 J71  @  G9 e& C' z  p2 |; l% M
    72
    . r- J5 w7 [3 ^. E2 d* C% g73; y' Y  ~4 W0 p; T+ ]; `
    74
    & G2 ~+ x' H7 S- ?75
    5 `( Y: q2 s! D0 q76
    ; ~( |7 z' d4 ?1 v+ g9 Z2 [4 g771 w5 N; Q# Z0 B5 n6 a. w
    78
    7 \2 N  h: x5 f/ b5 D; b; A79
    ) I; o3 O* {/ T/ p803 i9 k( S& a% @
    812 C3 A9 s4 Q# y5 t2 m
    82  H; m  G6 F$ q7 A  I
    83
    $ o$ h  h$ @9 F: Y0 {; q84
    1 u2 X( _, i/ o& }+ C85( w4 D) e- i6 j& ?$ R* [+ j
    86* r8 u* V+ s+ Y7 t2 r4 d/ y9 U
    87% `$ m0 N" S6 x5 v; R* u
    88
    6 `4 e% g% Q( P6 q' c( `89
    ' c6 S" w/ T1 |7 c* w90* u! y) y5 H3 L, L  i, C2 d! ^6 i
    91' A+ z# c" R0 C9 B0 s" r3 s+ t& X, A
    92
    # X- d+ |0 @3 F# R93; F, i) ?  `( l5 ~+ I
    94) i! C- k( p# w9 S& f# F- h
    95
    " X: \: ^& |! k4 m96
      W5 v- N1 B  C! Q/ o97
    ; D0 m  J9 ]4 n4 [3 `5 q9 R3 c5 D98+ m. B) s$ k( Y- f
    99% @: Z# [2 M) I9 e
    100% \; G" l& Z0 J& U2 f
    101: t- s  z/ G' T
    102
    , L. X- u+ M7 k# n% E2 N' j6 E- d4.展示一下数据
    + S' W( W- x4 W7 E. @: g$ ~def im_convert(tensor):% r" P& Q2 u; k* }/ L% k) m
        """数据展示"""
    1 e, E$ \$ {+ {/ G4 ?) m    image = tensor.to("cpu").clone().detach()
    # Q( t% z7 z* S2 i9 I2 ~$ n) E+ [    image = image.numpy().squeeze()
    ( M& k+ n9 c  V3 X, }! \1 t    # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图) A  i+ l( t0 O" F) A
        # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)* ^- G4 g0 e) j* e# h. }$ S0 I
        image = image.transpose(1, 2, 0)
    : B4 ~+ H5 V% t6 t& \: X    # 反正则化(反标准化)
    1 R9 ~- x+ l# U1 G    image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))
    ! t/ ?2 Q0 I! O) K3 @9 L; b+ u' I
        # 将图像中小于0 的都换成0,大于的都变成1
    3 ?9 |3 O$ d! S    image = image.clip(0, 1)" X5 {; ?% D$ ]' z& B/ G

    $ Z# o; f, J  F$ D+ ~$ j1 }    return image* J' t$ h# }$ C. L& x" ?
    11 v2 `& `( R# s: T- W& O
    2
    * I5 k# }1 c6 L: O5 E( y8 C3# j% O7 y, [0 F+ \6 s4 Z
    4
    & j1 X9 F  A) b4 z6 K8 X5
    1 C# |9 P7 B& H0 j: g0 v7 g" s6# H, w4 x9 P* O0 f2 a, J# {
    7, Y5 R& X% A# u/ |5 g! L
    8! r; A: F) d5 m+ t
    9
    $ h3 y, J' I4 l8 ]0 m: S8 ?10
    7 A" W) Y$ f# J- m9 }/ O8 G11
    ! }5 K) `, h' h* G3 x( i12* V4 M2 J6 ]  n- l
    137 j9 d' q6 B( k' o
    14
    # ^' P/ a4 R" E4 N. a5 n- |# 使用上面定义好的类进行画图. K: C: r4 w/ s: `
    fig = plt.figure(figsize = (20, 12))
    # O5 n  T& n; ?2 M, |0 X0 Bcolumns = 4
    - O; V# \$ U' H  ]rows = 2: n8 {& A  H" u& ^) X3 ^

    8 q6 Z  b( P% N) O# iter迭代器
    ! g3 R3 I: [! Y# 随便找一个Batch数据进行展示6 S4 a/ Z  i% u' G: y4 v
    dataiter = iter(dataloaders['valid'])
    % N. x/ ~/ r7 b1 X6 Rinputs, classes = dataiter.next()
    . M8 d4 M1 \; j' l) ]5 }4 h5 Q3 o! Q7 g! U3 i/ S- P# o6 o  t
    for idx in range(columns * rows):6 |2 b4 @1 h3 G3 b7 P; A1 M
        ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])
    * L# @, }% T1 w# v    # 利用json文件将其对应花的类型打印在图片中
    ' V  Q* q6 |/ z, a7 G- ]& [9 c1 `    ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])/ q4 r; w' M5 f( }1 L3 q5 u
        plt.imshow(im_convert(inputs[idx]))
    9 d3 Q2 @8 ^/ W; Uplt.show()
    " H2 [. s4 ]* [4 \4 r: |2 D
    ( ^/ L' l8 ^" O7 o# r6 c1
    * D3 D5 a) o; T- D2
      M# K+ v5 b! D, O9 E3: i" d4 Q' W8 w; N) }8 E. `5 ?
    44 x" l, t5 s* o0 e. p
    5
    4 Z' J; b9 V! T, u  c6
    + X3 `& \% ~+ H; M" b7
    # x; R; C# y+ S- U2 O2 z/ i85 x$ |& a; h; a( w" E8 C
    9: }; N5 z5 `; x- A! ?
    10, v% ^, D8 p$ D8 p* o
    11
    ; q% T4 l. R8 R9 h0 F! C123 ~8 l1 m) D( J
    13
    3 I5 O% j$ a8 a9 A14
    % u8 y! |% o0 b) V6 J7 p) g( P15
    2 e1 J! z$ g/ Q16! ^2 ]' V, H5 p: ]0 L
    1 I- c0 V/ y: O$ x# g
    / D; T! v8 l$ s% U: y) x6 m6 d
    5. 加载models提供的模型,并直接用训练好的权重做初始化参数' k+ b- W: T0 `: g5 n7 p
    model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
    ! V  Q# m& {. ]) y# 主要的图像识别用resnet来做% p: O4 G) M5 _2 _& m  k
    # 是否用人家训练好的特征
    * o' G! x$ p; a5 `feature_extract = True1 Z! j  |/ B% I0 o- a8 I
    1
    " r1 C! |, A: g4 n! R2* M; s; ~% G: k: e2 L
    3
    * b# W. O& D- U! r' V4
    * M" j! x* l6 O6 a# R# 是否用GPU进行训练
    ; i  V' ^& e# o& m9 x; B- R1 Vtrain_on_gpu = torch.cuda.is_available()  [9 m: A( a, j3 Y6 L  q' |! O5 y

    ( u0 g, r2 `( \if not train_on_gpu:
    4 a- [! l, D( M1 g' f    print('CUDA is not available.   Training on CPU ...')
    2 z/ U. B$ p% S& L5 {$ F: i+ felse:
    9 i' J" z2 L3 x1 N1 _8 B  w  {6 A2 f    print('CUDA is available! Training on GPU ...')5 e1 H+ w+ o' E( U$ d0 A

    & e# f/ t9 X$ ]! Y* Z. Gdevice = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')
    " [" y' M: R: |. v7 {1
    * c& p, x' B5 a7 J9 W) B5 G22 r& V( I! _( |. p
    3
    2 [" C+ i: O9 ^  V4
    ( i% _! ^$ y' u; s; j3 k, D5- i: C* x* ^7 w  l* D
    6
    + h/ g0 `+ X$ C( g6 q7
    4 Z0 D; F5 D" A0 O) z6 e% [" L3 W& E1 P8
    2 d0 q4 H- L; V# C& M5 V/ X9" [0 K& g! @3 K, C" E
    CUDA is not available.   Training on CPU ...
    " d, N; ]! q" Y8 _. x, [' q1
    ' @1 ~1 {# f6 f7 V) U% O* ]# 将一些层定义为false,使其不自动更新' f1 n7 k& l& a9 a
    def set_parameter_requires_grad(model, feature_extracting):
    ' N+ o- F" A$ i/ Q2 O& A( _    if feature_extracting:0 h; }* c# k7 ~/ p) ^. l7 y
            for param in model.parameters():
    3 k& r' {& Z- J# ~            param.requires_grad = False  I3 H& V% ?. {- h1 N& j) E
    1' l+ T" g3 B+ }* M+ j2 T: @
    2# G" |* W. w; i' x
    3
    * e6 [& }6 T, J( J, X' X4" u1 u5 W& a4 r2 k4 D( p* R
    5
    ' f4 G. N3 C& a# 打印模型架构告知是怎么一步一步去完成的6 [% \0 P' I  _7 D: i
    # 主要是为我们提取特征的
    . u. V; g: @! U# X4 w+ W' ?, ?# m9 K! G! Y* v' f* [! F6 @
    model_ft = models.resnet152()$ }+ V% `7 s. U* I/ t
    model_ft
    + D, g; ^6 ]$ l. Q7 x4 X: l. `1
    . L3 X( J" \6 T- _2
    ; G8 C5 U4 T, b3
    7 M2 G7 w% }. W8 j5 s5 P41 c% ?. ]5 ~# P. ?" J
    5
    , T# |9 V/ m* e. |ResNet(
    - y# L4 u1 [0 c' U% b- T* N' a* N" v" s  (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False), A1 M! A! P& j( g1 c5 g. D  \
      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    , o" B; v4 G7 }( Z- d" ~# Q) S  (relu): ReLU(inplace=True)1 O2 d5 \0 o: v+ t; t
      (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
    ' d% h1 q2 U" N( x- K  (layer1): Sequential(6 P* L$ t- \( i$ J6 G! y6 s
        (0): Bottleneck(
    # f: F. r: @5 i4 N      (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)% }4 D- `+ Y/ @1 M. {
          (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    5 e5 ~9 a; _$ \9 ]4 K8 W. _      (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)) R1 m& C2 T% Y, p& Z
          (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)$ u0 F+ t/ Y4 D" }( M
          (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)4 N' Z- k$ i* b0 @. G$ g
          (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    * s) E7 J. z& R/ M      (relu): ReLU(inplace=True)
    3 B4 A' ?7 f9 v( C' o9 \; z      (downsample): Sequential($ |/ v0 l$ h5 b/ V4 s
            (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
    7 F6 s6 A2 P" ~- ^8 Q% u' K/ r- |        (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    : y  C, U. w/ b& G2 q; n      )- Q- r. R, X- K& I
        )
    $ V9 A' M" D+ K% P$ M0 P/ ?; A8 a中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。3 m4 h4 h6 q6 ~6 P$ P; f
        (2): Bottleneck(
    4 n% Q- X# g* h3 |# M      (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)& w" M' D6 B9 u
          (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    & @# [# b. s/ `. T+ b& V      (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
    3 A+ [8 ]; ~2 u3 \+ x      (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    # h* A( s0 o1 Z- d5 |      (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)0 a! ~( E% g2 m+ K: h
          (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    7 h4 [4 x1 P5 M% H2 p3 r6 {      (relu): ReLU(inplace=True)1 x9 ]/ U) k& p8 K
        )
      h1 o2 X, @1 g1 r' b: R+ Z- }  )
    ' ]5 t, S0 |9 R# p! x* [# M" j  (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))
    1 i7 W: C& s6 {6 n# h6 {  (fc): Linear(in_features=2048, out_features=1000, bias=True): S& j3 r( E. z
    )
    & L. e3 e( P4 e9 I. T/ _0 F% U% K4 U  c, D( F+ x$ d. S
    1
    $ {( K, T3 f! A0 M27 L2 S) E) g8 K, k( \
    37 T) D% a  r( V1 N8 u
    4
    / M3 m5 o, C. W: ~' j8 d5( O! U! k0 G6 q; D% p6 m+ {
    6
    ; t. t% D# S. t' a9 |/ l7- b6 Q. ]7 v+ M4 Q, V5 B1 B4 ]
    8* v2 x  Z" v4 A1 H# Z
    93 b" G9 e* U; I3 l6 A. A
    107 u% s3 w8 G- S4 d( {/ ]  D
    11, o) Y4 _$ _$ Q5 ?3 o
    124 m# w$ e4 W3 c6 i$ d) o
    13
    8 {4 o% B. E8 g+ L  l14( U& I+ b" F- m1 T  O' ?
    15
    0 ~3 c  z* H, [5 c$ F- g16
    # H: m+ m1 K' `  t0 N; i9 x$ Y17
    ' I8 p% c5 ?$ z  A' _6 c7 ~18
    3 t8 Q2 F# ^0 U193 w7 j3 I" _, U1 H4 Q/ A) n; s
    20
    7 E* H. j. U. X& M21
    " @: W: G% B2 @# F; w, z* w22$ j) l  G' b1 `: V
    23% M8 P, X; ~; D9 X
    24
    9 M: N6 ]' M9 y& d' X, f259 z$ ]% T; @9 b. Q5 K- O8 m& u+ j
    265 d4 x9 G; V% `8 o$ b
    27
    3 n* B2 U3 [" a28
    9 v, v* z8 w( a% W- E; N6 {3 U29
    6 K. G$ g3 j4 u$ R308 P; a! @4 A5 V- h2 `3 ^. c
    31
    # N  \7 i' ?% x' u* U) v# B$ P$ X0 |32
    ) K1 [; l3 W) l1 Z: B  W" v8 B8 u5 T1 ~33
    2 s: J4 d; s7 Z3 E. F最后是1000分类,2048输入,分为1000个分类
      ?7 d, @, b2 Z, j而我们需要将我们的任务进行调整,将1000分类改为102输出
    # P1 ]5 |$ s+ I
    ' `$ q3 Z6 y" x0 p. G4 U6.初始化模型架构
    ' f) D/ c+ @- ]步骤如下:5 S: o7 |6 |$ H/ q6 V' ~
    & t4 v' B3 s9 O$ B8 X$ ^5 ]
    将训练好的模型拿过来,并pre_train = True 得到他人的权重参数
    6 G! T/ T4 U/ y, q8 Y可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)0 j1 [2 m. \& T5 _* P( e* c
    无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数0 G. N" a* G! C( K2 J6 q: O
    官方文档链接: ?3 A; h9 v! D- f
    https://pytorch.org/vision/stable/models.html. O3 x0 H4 P0 V3 L& G
    , d: T0 v' m7 y7 C) N8 |
    # 将他人的模型加载进来  n  S' s- k+ W/ G
    def initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):7 J0 }  _+ ?4 P- g* L$ B
        # 选择适合的模型,不同的模型初始化参数不同
    : c" V: G1 M/ }* P    model_ft = None
    * t- n# U1 z, v0 n7 K    input_size = 0
    ) O$ b# H9 a+ n! F/ l0 h1 D5 U# E% N
        if model_name == "resnet":
    & y# L0 i. M) F) P5 L, @" c$ D8 g        """
    9 `- V/ `8 p$ o$ ?        Resnet152
    % y6 H. O/ f  S        """
    " s, R5 J3 g  [+ ]8 s; U- t
    ' c) D7 J8 P8 Z# f1 `$ x6 X, t        # 1. 加载与训练网络
    # y" Z! w4 K2 k* i" M  ?6 K        model_ft = models.resnet152(pretrained = use_pretrained)6 ]0 m4 E% r+ x: H$ _$ L- M" u
            # 2. 是否将提取特征的模块冻住,只训练FC层
    ( _  V/ _3 S& n9 e        set_parameter_requires_grad(model_ft, feature_extract)" E6 F, Z7 }; t! T  k
            # 3. 获得全连接层输入特征
    4 z1 K4 `8 R) `# w        num_frts = model_ft.fc.in_features: x7 F! {5 `3 h! u
            # 4. 重新加载全连接层,设置输出102
    ; e7 n6 N% r! R# s/ _0 w        model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),
    1 U% _  Q8 ?) G8 P0 D7 b. ^                                   nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1
    & {" T& |+ D- f/ N# A        input_size = 2243 k  X4 C, I! \4 ?. p" l8 I
    ! ], r; m$ I9 g' \) }/ V% M
        elif model_name == "alexnet":
    ) e5 N% A. s  N5 ?! t; Y        """
    ! v- X$ ]8 [. g+ L$ Z2 T6 A        Alexnet
    * r8 T' k- c4 w$ ^8 M  e, b        """
    ' N' p" \* O: E" B) X: h        model_ft = models.alexnet(pretrained = use_pretrained): \5 K' f; I) B8 x  `
            set_parameter_requires_grad(model_ft, feature_extract)
    9 c& D1 l( E) h, @8 U1 _/ p* p6 o" B
            # 将最后一个特征输出替换 序号为【6】的分类器
    ) t: D1 T7 h2 n' t% V1 I  ?        num_frts = model_ft.classifier[6].in_features # 获得FC层输入: k* o7 z$ T( h& p: ~- n
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
    ) o% L0 x4 U9 G6 D: @% i$ h        input_size = 224
    " }- {( K6 r. z: y5 o. w6 w  m2 D4 v* j& ?% t
        elif model_name == "vgg":- W0 p+ Z. S' y; n/ {
            """
    6 Y  L* t: f! N        VGG11_bn. G# M$ r+ S2 \- `# n  G
            """- G' H* c% ]) F
            model_ft = models.vgg16(pretrained = use_pretrained)# H1 E' _' S3 v; G) N
            set_parameter_requires_grad(model_ft, feature_extract)$ G8 `8 w3 S$ ^! g( V
            num_frts = model_ft.classifier[6].in_features7 w& ]* |0 u: w5 U
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes)- d! \# ?# N8 T) b* u# a1 e4 d
            input_size = 224
    ( W( i& N& ?; l% a
    " q5 n* G9 j; x: i7 V9 e. [    elif model_name == "squeezenet":9 }8 E/ Q# D" ~! G0 K
            """7 k! t; D) _& n8 M
            Squeezenet$ t3 }5 u3 G1 Z
            """5 I8 p$ ?) q2 n& l$ h; k6 J( u
            model_ft = models.squeezenet1_0(pretrained = use_pretrained): H; G; M: _" x! w3 }; p
            set_parameter_requires_grad(model_ft, feature_extract)
    ! v5 c) w' s" N+ R" f( I        model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1)), ~9 L$ q$ X2 r2 ^  M
            model_ft.num_classes = num_classes
    ' t. P+ S6 w4 V2 O, B" @        input_size = 2242 q* G0 t: w' `

    ! ?0 R& t" o% t9 v/ i    elif model_name == "densenet":! h8 _! T: f! V' n4 Y! L
            """
    1 R# P3 I9 `6 c- _+ Q9 x        Densenet
    , J) M5 x8 F! a" k" C! {        """
    ' C* V" L5 j! k: k: C! w        model_ft = models.desenet121(pretrained = use_pretrained)5 r5 O9 j/ k( `8 e
            set_parameter_requires_grad(model_ft, feature_extract)+ E) n. o$ T6 V) v4 U. Q7 B
            num_frts = model_ft.classifier.in_features
    0 y# h2 Y4 H  y        model_ft.classifier = nn.Linear(num_frts, num_classes)
    + p  c$ i. F" t9 x! y2 o4 [        input_size = 224
    . y2 Y4 \4 V" P  {2 y! v& B, ^/ ]. K1 [
        elif model_name == "inception":6 x6 b, s( O7 j! _- T
            """
    , w4 @, Z( s" a3 P: S9 {        Inception V3
    6 Y% f2 D1 l% g/ g. {        """) C  y- H; k/ E
            model_ft = models.inception_V(pretrained = use_pretrained)6 H4 y  R1 I( K) G# d3 l1 q
            set_parameter_requires_grad(model_ft, feature_extract)7 i3 u# W: ^6 a, Q( A4 W/ ~% V

    : y0 n2 u% i, \. H% t8 ]3 I1 U        num_frts = model_ft.AuxLogits.fc.in_features
    4 r% }/ p' ^/ V, ^8 K% H        model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)2 Q7 |; J/ p) B% q. B* H& w
    & x1 v$ d2 t) E7 _8 |2 U4 L
            num_frts = model_ft.fc.in_features
    , E* S! R! `; j4 Y0 {6 e# H        model_ft.fc = nn.Linear(num_frts, num_classes)6 `7 b! a7 ?2 h4 c6 m8 T3 |
            input_size = 299
    4 E6 v; |( N& Y) {# J* D  s) i# M2 y& }4 v6 t
        else:9 A: c" v, u& \1 e4 y* ]5 Q
            print("Invalid model name, exiting...")& s4 C& C8 }! e7 g. X
            exit(): @: ^% g: ?& ?
    ! M' A, |( y4 u  T% p
        return model_ft, input_size5 U$ Z2 l6 L  k, T& C2 J; }+ ~9 ~
    + F7 K; @- e6 X4 G0 q- H' y
    1
    9 Q/ B" o( I# O  {: v; ^: u2- y7 J+ V5 C) ?2 u" x7 q7 H
    3$ K, J) h) ?& q7 `  i! c' _1 K
    4! I8 o0 ]) f# P  ^& `7 m
    5
      x6 \( Y0 k" r6 {  T8 B6
    & V& U6 a7 E) s! e78 ]1 N, e1 ^5 N" H- ]% E! D* y7 ]
    8, L/ L4 k/ Q6 |3 N9 F
    9
      V) B. B  o3 j9 ]% ?+ c105 W6 ?9 Y: v5 y$ o
    119 @4 ^$ L1 x$ B
    124 D& g: P- F' A; U4 P* ]
    134 y8 q% o  I9 b0 {" _. k0 a
    14* F0 l5 O. ], K1 {) |
    15# e! m( [. t. B( ^
    16( C% S1 U2 ?! y# O5 T( c/ ]6 y
    17
    * z2 E, z$ N. I6 b" i3 W' W4 U18
    " g, l0 ?3 G2 H$ X9 P' u( o# c19
    ; F6 f* Q; V0 d, ], \' t# P% \- `20
    ' S; [0 m( M5 ~" [  P+ a; ?3 P21; G' d. e4 X8 M3 j+ V- b
    22! S- t' S& {+ m* D5 u" [
    23' j9 C& ~$ g4 K+ |5 ?: w
    24- \" Q$ T& k/ }$ Q8 A
    257 ?" w( y: p; O3 n8 m
    26- N; q# ~6 {* M; G
    27
    - t' P7 l6 L3 }+ Q& t& ]28
    $ _; t% `9 j% `# V+ K/ x& J29
    9 s6 K: N0 v  z0 s4 O( g- T5 |# B* ]30" _9 P# e3 h5 Y  T" O
    31' Q4 @1 X( q5 t0 o
    325 \/ ~! g3 y; s! A* |0 S7 J
    331 W* Y3 Z7 P8 q& h# D  f
    34
    ) I- i6 q! ~" M: f  X35/ U! \- O& c" I  e$ [5 o( r; i8 j
    36- V5 _: ^* U$ k( K6 d
    37
    * D( r; {- {& z  L% g: N  [38
    . o7 ?9 I) o+ A, v. [0 F39
    * v* L: w5 n! h! ?. V% m$ l; c: X  j40( m) z! e1 \9 Q
    411 r  n8 j  Y1 d
    42
    ) z' u9 J. E4 i. s# m& A43
    ' D% a1 r8 ^5 x. q2 ?# S' ^: q448 t4 _5 N- z8 W* G$ n, t: G1 c
    453 D( r! V, F4 @3 ^1 |& U5 j
    46% [" C$ T5 i% Y% Z- f, V' s
    47
    4 M2 L" N/ Q0 L6 G( A2 D  ~48/ i! n& b+ n" e+ b! ]! e+ U
    49
    ' _- Q5 W! U# b6 V. @1 r50
    ) k: ]# m& ~. l# B+ c51
    ; h; ~& M5 H7 I2 {$ h) H7 w52. ^# ]& h$ P1 M( r- i  q- Y
    53
    7 H! I$ d7 y: ^+ G* q54
    $ P7 G+ P0 V/ N0 B1 l554 k% [9 E( n& d' {
    56
    / w. P9 \0 r0 u4 x57
      E; \- K+ M4 o0 d- B( V58
    + L- a% i$ ~& x5 e% q59
    # K& `3 @6 ~/ t8 `# q- _' t60
    ! g9 u6 L" m+ P) ^) p1 |( A) D' i* H61+ T! p0 x0 M: C
    62& H# `  b9 ]! ?
    63
    + |! W& E9 X; g7 W8 Q' \64
    6 c- P% A9 ~8 u0 ?" h) d65
    : n! k' O' ~" g7 s# y% N66& ?: U$ }8 Z$ q2 k3 w( H5 m
    67. I( j3 Z9 I3 ^: ?( H
    68* C4 `- R  U* K* Q
    69
    2 A: q; O9 o% y8 D- W8 c70: g. p) J( ?& H% x0 ^, O# |
    71# `) R9 i4 z! Z4 g
    72
    ' o0 v! q/ T" k" f: `. V732 g6 b8 G) S+ d+ T& q
    745 w& u. n( u& r5 p
    75: K# W7 U* P7 X8 [+ L3 P3 S5 O% U
    76
    . X3 \4 E# W# p: r0 ^77- N* n7 Z( V, U- ]1 [. q* I
    78+ v: ~5 _- h3 L$ U' C8 S' J
    79
    ) K& z9 u8 Q, h1 s  {% p: r) }80
    ! b0 g( J4 S( [3 ^; e  Z. d81
    + p( J) Q: U9 x8 Z82, x/ E" \2 u3 O. ^. y5 |, c
    837 y# u+ d8 B0 U) ~
    7. 设置需要训练的参数0 a' Y& {0 }/ `9 J
    # 设置模型名字、输出分类数1 c6 n/ V2 E' Y1 Z- Q9 d* t6 j! N
    model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)
    : F. j) e+ [8 G, v! O5 P' _# h+ Y" ~+ f, N) o5 z
    # GPU 计算
    ; ]: U5 O& r6 q9 w. Zmodel_ft = model_ft.to(device)5 r$ m# x  a/ f; J& Q3 `
    $ x2 y# K& W( Z, {3 m( }& Z0 z0 r. w
    # 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取7 W, m( u0 ]. E7 G; M' R5 }7 S
    filename = 'checkpoint.pth'. N- P; F( t0 U# t

    & A; f) W9 z( T  v1 D# 是否训练所有层
    # t0 `, I( B' s! Nparams_to_update = model_ft.parameters(). j7 E" T4 S; L7 @/ A( t4 Z
    # 打印出需要训练的层
    5 L) k/ S( X9 z+ n! `0 U2 Nprint("Params to learn:")
    6 n. R+ F$ ]7 O( C8 l& g* S% iif feature_extract:
    - ^8 Y9 Z1 s( z. P/ H    params_to_update = []' y6 k/ y' M' P1 `+ C9 g) b
        for name, param in model_ft.named_parameters():9 l* ^! _, d# }9 R6 H" P: N
            if param.requires_grad == True:
    8 U, w& G6 o; m3 D- ^% [            params_to_update.append(param)
    " ]0 N0 E6 C7 X( O% l0 {1 |            print("\t", name)7 A" l1 N) `# I$ C, f; l/ d
    else:
    / Y0 J# U/ r0 H( l* o' e    for name, param in model_ft.named_parameters():+ w! [4 {, c/ O9 D( p! G  V: p/ `
            if param.requires_grad ==True:& A, m7 Y6 j; k. Y2 t; _
                print("\t", name)7 {1 ^7 m4 \7 s% {) s) G
    0 |' w  l0 `+ t/ p
    1
    & q  p1 L4 B" ]: l2
    4 a: T! }. E8 S# F7 c" g0 ~5 B" [, v3
    $ t* B7 j- o1 n- C& A3 O+ ^4
    0 a3 O9 \2 k0 T7 L; J$ _5" x; N# p$ b, u' S, x4 D* u. E
    6
    9 |& P5 n+ W+ S. |. }7
    % M' Q7 D7 C/ c2 E. O/ U' {" z" B& x8
    & m6 k( r# A/ D- R6 C9
    0 E1 t" W1 E9 j1 S& U8 R10* Z0 x) [0 t) x. ?' t' ?. o6 x
    112 \( h) q( |, C/ U4 p" R
    12
    : U- {$ I) K) d6 N& \13
    7 ?, c: l% e2 U) D" ]; Q14) X- b8 C& q5 o! V5 C
    15$ X8 D6 C4 m3 O' _' w  i+ h% d
    16
    1 O/ h% J' b9 X3 N6 T- k; a  o17
    2 R8 B  G6 ?1 {18
    0 M: a! P. |3 r8 G- ]194 t5 ?  O: ]7 m( j
    202 f, t/ S! V- {' v: l
    21* n/ g# q* m8 Y
    22
    6 X  {& Z: i& }) x& p233 B9 Y, u8 }3 h# v1 E. R
    Params to learn:
    ! Y) }( b% _; |( k- ]8 X         fc.0.weight& t9 o) g: @7 X( [
             fc.0.bias2 A3 G0 m8 ]) n0 R* d' j
    1+ R0 [7 `! g! |
    2( g, z2 x) i: Q3 C0 N
    3
    % Q5 [) t$ n9 ]7. 训练与预测+ V8 U0 z# f5 v; n6 Y# K! l7 L
    7.1 优化器设置
    2 I) j: |3 J3 b# 优化器设置( P; ^$ K. i/ @3 K* r: ?
    optimizer_ft  = optim.Adam(params_to_update, lr = 1e-2)
    ( X" x6 e  k) k* @# 学习率衰减策略
    # m" t! R, }. M- M* oscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)9 ~/ T! x: ]8 ?$ E9 J3 Z
    # 学习率每7个epoch衰减为原来的1/10  p" L& g* x$ S9 n' F* d
    # 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算. g4 }' ]0 S* h( \8 C

    * S) F7 q2 K# s% g5 Qcriterion = nn.NLLLoss()
    6 g' ^7 m  N: `1 T7 P4 ]1 ^, x4 M/ Y  R1" N4 \0 J+ E  B# L7 v
    2
    9 g2 Q$ }' n: R3, J( m4 q. J  ]1 B/ I( q8 t
    42 f# p; Q# h' f1 ?
    5
    2 j  y0 D7 O# l+ A63 |+ {/ m3 J+ |
    7
    , c/ e/ I  }: U6 j& G0 t5 Z1 e8) H) q7 F* j9 d/ i! t8 |0 }
    # 定义训练函数
    * ^1 J: @4 w% b4 M#is_inception:要不要用其他的网络) [/ b: q8 _, _. q# A( z2 S& l: b6 Y
    def train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):
    6 ~0 ~% |9 M1 S    since = time.time()8 _2 c# i8 t, \7 D  D
        #保存最好的准确率) R( w9 @, F7 C* h  @& a
        best_acc = 0/ F" b6 ?6 X" `
        """; M# H* c2 v$ V) z
        checkpoint = torch.load(filename)$ C; w9 b$ ]3 ~! E' W! [
        best_acc = checkpoint['best_acc']$ j6 d* a7 g, N0 T7 h2 z. Q: {5 B
        model.load_state_dict(checkpoint['state_dict'])
    ' H) `' d5 k5 J. f) j1 _    optimizer.load_state_dict(checkpoint['optimizer'])4 a% I; D" H: ]* P; h; I. C$ k
        model.class_to_idx = checkpoint['mapping']
      K* X! q5 E2 [, F  ?5 z    """
    * d9 H( I) ]& H8 ~# ?: ^6 ~% X    #指定用GPU还是CPU" J' L! x) A2 Q) o- C. c2 |  g
        model.to(device)0 v0 I, J9 |8 w- V1 l! a7 ]2 X: D
        #下面是为展示做的
    ) J# K/ c; W, h2 c. M: l    val_acc_history = []
    ! i' N7 a/ c5 b$ F! j4 C. W    train_acc_history = []
    2 _4 U% J! ^8 Y, p4 J    train_losses = []. {8 u& s3 [, s: S- y0 ]: N
        valid_losses = []5 o! V1 j& @& c( _# p0 g5 ^
        LRs = [optimizer.param_groups[0]['lr']]$ d2 y1 ?4 J; I# J) e
        #最好的一次存下来
    & N% R7 b# C" {) E1 u' y7 w7 T6 I    best_model_wts = copy.deepcopy(model.state_dict())+ ]) O* X* z+ Z- k& z% T9 F" j
    8 t# X4 S4 f* Y& y
        for epoch in range(num_epochs):
    ; ~# J9 R( X8 Y7 t% A        print('Epoch {}/{}'.format(epoch, num_epochs - 1))9 Y3 p- ~& J3 |. k* x
            print('-' * 10)  l4 e  z2 k% b5 Z/ I; R% s/ f
    7 I5 h9 N: a8 @
            # 训练和验证. F  k4 K  I9 k; b8 ^
            for phase in ['train', 'valid']:
    & M5 G- z+ g2 V0 D' m0 f) Z            if phase == 'train':
    % [$ R6 Z$ Z9 l( Z5 V; T+ u                model.train()  # 训练
    ( p- ~: `8 W3 i+ h0 Q/ k! k            else:7 {' O& }$ Y) i7 z! q0 e: O
                    model.eval()   # 验证4 x2 w, f4 C$ k. T, D, V9 @
    ! L, e7 y6 t6 F
                running_loss = 0.0/ q$ B  H1 u" H" d4 I; M: Y  ~
                running_corrects = 0
    # G; P7 b3 W5 n$ U. c& y( _1 I5 y% c; P6 z
                # 把数据都取个遍9 q8 q1 m9 N7 K, y6 x) B
                for inputs, labels in dataloaders[phase]:
    / |! P" _/ O% T' Z, }                #下面是将inputs,labels传到GPU$ m! Z7 h7 r. C, A- H- [' H
                    inputs = inputs.to(device)
    3 l% V9 S: N5 D5 ~* r' F9 c" x& \                labels = labels.to(device)# ~, z4 C9 b6 i) F% I  O

    ' n- ]7 W3 b1 U) f- {0 U! d                # 清零4 h, l+ d! n- h# S7 ?
                    optimizer.zero_grad()7 G# D% G0 ]/ `# A" f
                    # 只有训练的时候计算和更新梯度1 u6 f- X7 |% }$ M: @
                    with torch.set_grad_enabled(phase == 'train'):9 T4 ?; n7 O+ D
                        #if这面不需要计算,可忽略- b/ m3 J& ?; A0 N! J1 Q# p- [
                        if is_inception and phase == 'train':: ]2 w: F2 F* }7 Q7 K
                            outputs, aux_outputs = model(inputs)' y. R" l) O. B# _8 M' f; _
                            loss1 = criterion(outputs, labels)
    , A  _% i8 X* H& ]" q( {                        loss2 = criterion(aux_outputs, labels)
    5 t( Z7 [) \# Y: `5 G                        loss = loss1 + 0.4*loss2
    ) L/ C  U0 m1 Q7 ^& h% o                    else:#resnet执行的是这里
    1 `) C( [! I# r6 c& U9 {( s                        outputs = model(inputs)( [3 I* y" z6 U8 c5 d
                            loss = criterion(outputs, labels)/ R0 M$ k6 F' W1 |: K

    5 K) W$ f# W$ l* C2 [' r; L' w$ [& S                        #概率最大的返回preds
    % a. o) C; v" ^8 x) h                    _, preds = torch.max(outputs, 1)( C/ X/ [. `  {3 l# ?$ ?* H9 X

    % x# I  ^1 x5 S6 c# |; K                    # 训练阶段更新权重/ ~7 Y- S- e5 ]7 y$ V
                        if phase == 'train':
    # I$ u8 [: W: `" @* W  _8 q                        loss.backward()
    - |3 G" |* C1 `9 l  o                        optimizer.step()
      b* O' k7 q& W/ M2 G* ?5 _- c) J3 `6 ]% D" h& A* j
                    # 计算损失; ]+ }' i: _; y; H$ \
                    running_loss += loss.item() * inputs.size(0)
    9 t; R) D/ o# e0 u0 W                running_corrects += torch.sum(preds == labels.data)$ Q0 _& W1 z7 `

    ; Y6 l3 n, h* b) ]& T- R            #打印操作
    % r; o+ K! Z) T! K9 L  h( v            epoch_loss = running_loss / len(dataloaders[phase].dataset)
    1 [5 x: I3 W1 ^* A) L4 D% O            epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset); O+ J  \3 S: r5 K; [; G

    ) q! F5 n; b% ^6 {7 s/ E
    * R  x8 a! L9 K, C" N) x            time_elapsed = time.time() - since
    ! r7 A% ], y, O# l# `            print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
    1 O( m' [: f# i: o            print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))' K) O& P* {+ h& o# @& x/ i- @
    ( i! T$ h) a. ]2 Z

    3 ?8 e8 L: T" C) D7 O            # 得到最好那次的模型8 M; p; `- S6 T
                if phase == 'valid' and epoch_acc > best_acc:
      G( Q2 B) L3 w0 C: r4 v                best_acc = epoch_acc
    6 y: \/ C5 \" L6 d; k                #模型保存; v! C" G, d3 Q3 ?1 s! Z
                    best_model_wts = copy.deepcopy(model.state_dict())
    + c$ D& A( j0 @8 n( u" d                state = {
    # K' A/ O8 [. t. j                    #tate_dict变量存放训练过程中需要学习的权重和偏执系数$ a" g! f4 }+ q6 c; y1 {/ y7 ]
                      'state_dict': model.state_dict(),7 V9 I0 c' J  F$ ^! T
                      'best_acc': best_acc,; z* e, f* A) W
                      'optimizer' : optimizer.state_dict(),) w* t  ?( ~% e- A' S
                    }" h6 z! R$ ?1 U
                    torch.save(state, filename)
    # h, H2 ?( X( b1 w7 k" a, s9 M            if phase == 'valid':1 I% u# l2 V+ J: E- R
                    val_acc_history.append(epoch_acc)9 u  g" m  W; m* U8 A# |9 Z
                    valid_losses.append(epoch_loss); y. q" @$ ~' i! t
                    scheduler.step(epoch_loss)
    4 I$ _$ [; ^& L            if phase == 'train':' S" V# O, a0 F+ `' h3 I
                    train_acc_history.append(epoch_acc)
    ) W& m7 @' w7 N( H                train_losses.append(epoch_loss)
    1 g/ N' D, o: C: u# ~% L+ Y8 m( r. l4 H7 \, T% _
            print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))8 B2 M: g- Y3 r" X" w  u
            LRs.append(optimizer.param_groups[0]['lr'])9 j* p: J" a' N. `& s$ D  m6 ]
            print()
    # Y6 _! _5 c, S/ d6 i: W+ c, X2 o2 u, N/ _( @  w$ x. T4 b0 C" [
        time_elapsed = time.time() - since6 O+ r" S( [7 M9 t0 J+ M8 @9 R
        print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60)), j' Y; I2 r8 J' _2 e( u4 W. R$ L
        print('Best val Acc: {:4f}'.format(best_acc))0 V) G$ A: s0 U
    9 h+ E/ S6 P6 V0 U$ C
        # 保存训练完后用最好的一次当做模型最终的结果  y2 _# G" ^5 b( K. d7 J4 L
        model.load_state_dict(best_model_wts)) y9 Y, ^( o4 Y, ]7 ?
        return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs , X- a3 A7 Q: n2 }! _* j' D; |
    7 S% Y* ?1 {+ h& }& F$ v3 B

    . r' G1 g5 T9 l, l1: h* R8 U6 r) f+ R$ |, R
    2+ x, ^% I8 {8 D
    3
      M# J8 e: C) ^; G4% ~. V" N) t9 z6 S7 c
    5
    / q* ^- t6 L) D5 G3 J6: W$ K- z9 s- J
    7
    ; `2 }- d) L1 I( ~' N: W9 w/ r% Y1 _8
    8 e3 h8 G) Q, A% P1 ~, q. T% `" x4 a9
    * F7 d  _3 g$ E+ P- D9 V10& u& h! r% R4 p: y
    11. _1 |9 {5 K2 M9 u: @! V2 I
    12
    - `; c- P, C* ?- {) {0 C% x13
    ) n2 c( B8 P3 G( B4 }3 x5 b14
    $ I5 I/ @! \# w9 d158 ^' i0 V; {. U8 [. `$ X$ ?
    16
    5 I! d' l! |. i" B1 J17
    & A$ C6 X2 i4 v& x, }% b9 }. a/ g18
    - S  [4 O; M1 m( [, h& U19
    ' D* n% w/ _+ R! Y5 H  a20
    8 F5 H# L! f0 \. _0 C* u% U' H21
    ; Z! `; y1 F8 z- b  u8 g  ]  r; y; K22
    7 }) R  r6 K, d2 w3 S8 [23
    1 q- `$ l4 i- R" a24
    2 H+ V/ I: t. E5 C1 h. Z8 S25
      O0 a5 A1 E$ j$ \6 {) C26
    ' ~; w4 A: c7 C( ~% Z& h27
    1 M+ }- T( R; z0 J# q4 R$ _4 O28; P' I) f$ S. K2 f* L" E
    297 D3 d' |, h/ }3 e/ c
    30
    4 o7 y8 V% q$ T+ x8 i9 C8 ?! _31
    4 w3 ]7 F( E5 ^32
    , a7 S$ D) V; \# G, m7 @333 C% }: h$ A7 F5 e( a: k
    34% I' d! K6 W4 V; a7 g3 H
    35
    6 b) O9 y; N$ S36
    7 L, V' j7 v3 R. k- A6 e/ n" [37, ?, l: C4 Y/ f/ D7 y6 r
    38: s8 L+ n: ]: k9 ?) T
    399 ~# S1 G6 \( {0 g' E
    40
    , Z' f; U6 d7 W3 H9 M41
    6 p' D/ K, ~# s! w: [% p( P: i42
    7 J, k; Q9 U- M3 N3 {% d+ j43
    + n: j) n% `' H2 ]44
    / d+ J' z) J* T& _# P45  v; H/ |7 q% S1 C2 a+ V- B
    460 x1 A5 v$ d5 p
    47
    + Z! n( l6 t7 p* _7 S0 T* |48( }4 ^# V! S: c" o8 |7 h
    49+ s( F5 F$ D) z
    50
    / {; [: a8 @- n( n! W4 c2 H8 C51
    6 m, F% J# |3 v  L1 _52- v& \; a  ]; [6 \: c0 s1 P
    53
    9 m! p& M0 p: p& w' G5 l1 I* `% _541 i$ f2 B( j$ ]9 |% D5 Q! j
    55
    8 G$ g: K3 n0 s" ~; f- `' a56
    % ?6 s: W# e3 ]# {- [57
    ' ^+ L7 s8 i2 P58) u& z! ^  t6 Z% O5 q' ^8 Y
    59- M5 W5 B6 w# D: ?) d
    60
    ) |$ y2 {  r- i/ ?61
    : A/ o( u1 h/ N" ]! E& t# G* m62
    2 h+ ?  u9 C+ y! b9 U- C2 {+ B63
    1 K6 k7 m! o6 n' e; I6 d  V2 b64
    " i" I7 |; ?! |650 {; B- ~0 p$ ]7 M  I$ m
    66" T3 I, ?( n) G+ w& A1 J" Z7 T
    67
    ; \3 \( Y4 h/ s68
    1 r) d1 e1 L' k# G' N  }692 N7 n( m5 R% r! ], o
    70# ]- @; n8 m+ s1 U
    71
    . A0 K- Y5 H* Y72
    9 C3 P2 J" D" o9 D# m73
    $ b. `$ J" t- a! E74
    " r2 }/ y# x0 m( y; x; B! o' z; P75. I1 Z: ]8 i" [) s/ Z5 c
    76
    : e: o, s0 `# i, S: Z77# H! ?3 ?' x9 o! h0 l
    789 \( f; E- a+ ]2 @& T: C
    79
    , `: F, F4 z+ i80
    ! L, @1 P4 R* b. N; E81
    5 k: n0 f2 `( u- s$ V$ a823 Z/ |2 n0 ^: Q3 `! B! H  U; Q
    83; F2 r" M% c& R3 u* B/ B. {
    847 k1 l; v, E9 J5 _) S9 e+ j' d" U
    85
    ) f. f# z" l7 D5 \: ~" m+ Q86
    % O5 \  C$ z  `" @: V4 C2 O; ~9 m877 M+ ~; y* d# Z* j" V% j4 _, c
    88
    ! |  U: R; L1 [& M: m' Y+ S& v89
    7 p; e9 P+ A( S7 R' {90
    # Q/ V5 {1 B! X/ w: ^: u1 n8 l91
    + S- D; k6 [2 W; E# d4 W92
    . u5 X( M% d% I4 ]0 F93
    $ k2 |5 V. }. L) N" @# F' c3 Q94
    $ I/ R# P, }5 B0 |2 B. \95; l: ?0 _0 ?$ ]& B& M% L
    96* ^9 K8 @* U% B' G6 G3 ~0 [
    97
    % [$ j) l3 ~+ ~& K# b& U4 i980 ~+ A" [6 o: S$ L8 b7 C" e
    99
    4 y" J) J' O7 b- F4 v6 m  I100
    7 Y# E* w+ T# o. o" B1011 }5 ?, m& {0 Q: e; U+ t- D% P& l
    102
    7 Q5 I; l5 Y, s8 \% Y103: @- @  O0 z& `. M1 K& w
    104( G8 p' J) i$ x( J$ u3 N1 u
    1052 y7 c! H& N- I4 y6 `; k! c: Q4 o
    106
    ; D7 z3 z$ ]1 [* K107
    5 R" s" G0 C8 @' s8 ^5 b1089 q* Y' j. D9 v; X" I
    109& P* y7 j  Q1 \6 D3 {. m
    110
    ! y* N" U! D7 n, t% Y' t111" U+ E% [5 k9 ]1 K9 c3 O' R8 i
    112% l2 O: U' l3 }( g( Y; ?) [( u( ~
    7.2 开始训练模型5 [" F' I) j! d8 F1 t- Z# P8 v2 y' [; O
    我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次/ l! n1 @& h& }- I$ L1 a3 d: j" |

    ; `, {# O  O/ c- u' K( F- e#若太慢,把epoch调低,迭代50次可能好些
    + o7 x3 [6 q8 _9 M% y#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合4 b* ^' Q7 k  e
    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"))7 E+ H  X$ A5 i7 ?; k
    % F& h" a5 L8 z8 x' E7 x. a
    1
    6 b( L. l( [1 j9 g; d* m  H2
    0 x1 R4 M  Z. t* A3
    % A3 k; a  p+ C8 W- d) ~' q40 X  E! I- Q8 Z
    Epoch 0/40 B! ~% i7 n2 H- E
    ----------! ^# I# x) w/ {0 l* E/ M& @, A5 k
    Time elapsed 29m 41s
    3 ?6 z/ a" V( `5 @train Loss: 10.4774 Acc: 0.3147
    4 W2 H" }5 A* m3 x% l$ yTime elapsed 32m 54s* F1 d( e3 T  U
    valid Loss: 8.2902 Acc: 0.4719
    . R- x7 J# I- ~6 Z( e) C2 ~Optimizer learning rate : 0.0010000! B: j( C( |( r# u: A

    & _0 O; L2 h# ^  KEpoch 1/4
    , i4 W* ?7 Q' m$ e/ `----------5 w/ F( j: L- q3 R
    Time elapsed 60m 11s
    % k# D' }3 K5 {+ v3 f7 Btrain Loss: 2.3126 Acc: 0.7053
    ! q, W7 h- y/ W2 YTime elapsed 63m 16s5 j& o# y6 u; J6 M0 H
    valid Loss: 3.2325 Acc: 0.6626
    4 k$ _8 g  d9 U: I$ y4 a! N- O  _Optimizer learning rate : 0.0100000# L" q7 f  y- W  u4 R5 k- z

    1 g/ Z0 t4 o* X+ l+ s& W- l8 QEpoch 2/4
      W# l; x" u6 i/ \8 w----------
    + H6 E2 t" C" W' x: ^' ZTime elapsed 90m 58s
    ! J9 i8 c1 K; s% l/ n& _! S& Ttrain Loss: 9.9720 Acc: 0.4734# T, V$ }4 n: l1 |8 u5 k
    Time elapsed 94m 4s$ a# A& |: z, a' U$ j# Q* |* S
    valid Loss: 14.0426 Acc: 0.4413
    / n# k, M* D5 x7 M+ Y2 T" y9 fOptimizer learning rate : 0.0001000
    4 {2 V+ Q/ h# K4 n8 |  ]) N
    ) F. E3 r0 R  wEpoch 3/4
    * j" \  J+ P, ^' P& {- J* O----------
    6 l& _' Y( m/ `, {6 }! tTime elapsed 132m 49s! p4 w- ]' Z, _2 x
    train Loss: 5.4290 Acc: 0.6548
    % o* g9 Q3 ?5 X" i, Z! U) o, k, ATime elapsed 138m 49s% f( R+ O& Q( Y% c% v- |; s
    valid Loss: 6.4208 Acc: 0.60272 S- D, v% n" @
    Optimizer learning rate : 0.0100000
    : q" v; v0 G" K, f3 b
      `2 N6 s$ w' c6 x8 j% A6 l& w& @Epoch 4/4( t' }9 h5 f6 D9 P
    ----------" U* k' Z8 ^: S/ k
    Time elapsed 195m 56s
    $ N1 v, A# }' X5 _2 K- \9 J1 ktrain Loss: 8.8911 Acc: 0.5519
    # n% h8 _' B+ L% y1 STime elapsed 199m 16s! ^' k: B$ c( `! w! {4 G) r
    valid Loss: 13.2221 Acc: 0.4914
    ( e) V/ z2 p) w: ?Optimizer learning rate : 0.0010000
    1 x5 `7 I$ Y: m% }8 j
    6 [3 D& a3 {! P5 Y' ^$ KTraining complete in 199m 16s; d$ _' W$ Q; l* I
    Best val Acc: 0.662592
    ) d& q, o# M, o7 X
      d! `: o4 P; G4 k- k+ z1( R! p( H; @% U8 }
    2
    / d0 Y$ {& J' R6 ]. `* e' c! s3
    / Z  _  d. e& I& e: V. ]4; J  }5 d  S% ?  B+ g' X
    5
    : @; W; d) j8 {- T6. o! W0 l. Y3 ]' l
    7
    0 k1 i& W( d6 E) @89 A9 R: C% u! {, T! A
    96 C, e" ]" \2 ]# b9 h2 u( V
    10, u/ f+ c4 }% q! |, d; K- G3 [
    111 `+ m" J% j5 V( A
    12
    ; S; N: S* x. V, }13! A+ C8 n7 {  P4 J
    14$ {( q. ]5 b9 L9 U$ n% o4 x% R8 x
    154 {, s+ @' y7 X% I! g! I9 [; h
    16$ U( P/ s' f9 f& J) R0 I! y# b7 B
    17- u: h$ D' _/ F% ^
    18% V4 j; C8 X+ L- A  w. U
    19
    3 q- `" X+ u: Y* o6 R: y$ c201 Y$ n! {4 t+ z3 ^- }) V
    21+ j$ e! k5 i/ Z
    22
    $ R/ b% N- @6 ~& P7 w3 T. q23
    4 ~$ Y( W+ P% }' f* i24
    + j2 S5 p( ^' E, Q% o+ a25
    # y/ ], g0 Q1 b. v26
    * {/ q1 Q# ~& }. Z27$ n0 ^2 I7 N* _2 u
    280 k& E, q- ?" u. K( I  W$ i; B
    29$ n7 @6 G3 f% r1 A- R
    30
    ) s# T8 P. k! o1 o# ]6 T313 F3 Z2 w2 c$ y4 `9 A, k+ L
    32
    ' F5 M9 [0 g  U- x9 [33
    & L: d) a' \' V34) }/ L' H  u# u! V1 o2 k
    35, q% K6 h. X+ s/ O  K& m/ @( u
    364 H+ E- f2 f) t0 }
    372 t% W# v2 l! p6 t( u* f  y/ K2 c
    38, e1 C7 K+ v: [2 e3 Y8 f
    39
    : ?! M6 m" F4 x- K; ^: a" x40
    + J9 q+ q5 Z' j7 H41
    . [3 x0 u! \/ p42
    ! z% y6 }, ]' r7.3 训练所有层) c* T' c3 u+ z  ]
    # 将全部网络解锁进行训练
    & k: o  ~( ]0 }6 \. X- `for param in model_ft.parameters():
    : F1 V  v6 u0 N: V8 x+ l    param.requires_grad = True! H" }& i$ Y0 u- |" V4 |

    + D% f8 B$ q' a6 X0 x0 c* U# 再继续训练所有的参数,学习率调小一点\, a# F$ K7 l2 q# y2 u) B
    optimizer = optim.Adam(params_to_update, lr = 1e-4)% R3 h& R  }3 J9 _( ~/ H( G: i
    scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)
      b1 f% z; A8 O& b; r' x0 U
    1 l( L, R. O3 c: w5 R1 x. S3 L# 损失函数
    9 W4 n% k. h! }9 |; j) s2 xcriterion = nn.NLLLoss()$ V1 j. A* P# H( G" \: Z! t
    1
    - V/ G5 }/ o" \7 \) Y6 K2
    5 y9 S) @! l/ y* c) n, M/ U% t3
    % C% {2 G  x! g4# v; p/ N( z; D, m* y2 ]/ j
    5
      S  T0 W9 L, c- l. s6
      r9 z' E6 g. o% O! s" Y9 d" L7
    6 B5 s! \" y; Q8 x) H5 ]$ T: y8
    0 b( J& p) w/ u3 s& N9
    4 M, P1 u0 l* ]0 T; T10
    . Z$ B6 G  f. V( l7 R# 加载保存的参数
    ! |! b& C. x  c+ Y) {# 并在原有的模型基础上继续训练" n$ d* z2 d! x& T
    # 下面保存的是刚刚训练效果较好的路径
    1 U2 ]" Z: S) Gcheckpoint = torch.load(filename)
    7 T, n9 D5 C9 b( e" V# {# h- s! v# jbest_acc = checkpoint['best_acc']
    ( k/ V6 L  J1 @* z6 n3 `model_ft.load_state_dict(checkpoint['state_dict'])1 k6 ]& M& f$ k1 g7 h+ }$ a
    optimizer.load_state_dict(checkpoint['optimizer'])
    ) @4 w( z% ~( z+ l3 {  a4 U1
    9 u# J+ V3 s" O$ m+ B2/ g$ B+ _! m6 z
    3  f1 I/ o2 K% H# N* u; `
    42 D4 i0 r+ x" X2 Y& V2 ^
    51 P$ S* L3 j  o6 R' w" `+ e
    6
    5 _4 G( S; a0 `; V7
    + T; M( f5 V$ D. u+ v5 _: P开始训练0 D' A& w9 H1 x& M
    注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考2 Y4 J% B0 [, E" c' h1 A: F

    - b! T$ r/ q# Fmodel_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"))
    - H  f7 K3 b+ A1
    3 U6 I. M  m3 w& J" Y$ B* nEpoch 0/1
    " V. q5 T9 k4 c, q----------
    2 d( r. Q! B' Y: n& J2 KTime elapsed 35m 22s1 \  D( M6 J8 m, o5 C: D0 u
    train Loss: 1.7636 Acc: 0.73463 O& Z7 i# `- K' x; ]
    Time elapsed 38m 42s
    * ~8 X2 \8 N& N5 C% }/ n7 yvalid Loss: 3.6377 Acc: 0.6455
    ) [! w. f0 G; }5 v1 `Optimizer learning rate : 0.0010000
    8 q) ]5 j. H( P/ X$ R3 ^2 ~( h; f; `' m8 I, O
    Epoch 1/1. X( k2 G8 X9 n- q9 i
    ----------
    ) z, W3 P5 |: U( P3 b/ J* G1 STime elapsed 82m 59s
    ' U3 N) K, h8 Mtrain Loss: 1.7543 Acc: 0.7340/ y; K" C" Y1 L0 I3 t- |/ J9 U
    Time elapsed 86m 11s
    . }3 h. e4 H; B8 k) f8 i$ Ivalid Loss: 3.8275 Acc: 0.6137
    % k, j4 e# i; y8 [: U0 A: _Optimizer learning rate : 0.0010000
    + S) Z; _  m' H2 N. N" n+ N3 y+ S
    9 S% r9 O; J. C: K+ b! lTraining complete in 86m 11s6 E0 u# z8 g7 K
    Best val Acc: 0.6454776 z8 n9 t8 D0 y( c

    & _; S5 U% S, S4 _! J* t$ K( q2 R1& o2 {1 `+ _& j5 h* Q
    23 O  U! G( t  ], {+ V' y
    3
    1 c8 D7 e: a- \" P0 w/ r: I4
    : q7 a" p/ b. G. D3 C: C2 O" C5
    9 D0 h/ L! H* q1 ?) X; M6" \# v# O/ o6 Z1 q: _6 D
    7
    , }: M: C4 ]9 G+ @8) e7 m0 H/ j/ {* s
    9. F! r" r  ~6 k& b8 ]0 W* |9 P- @$ g4 c
    10
    " Z; H( Q" j. m* y11/ m- f+ A' ^2 [1 U
    12' Q; N- S* }! o* U/ f
    13
    " h/ @. G; ]7 w. a" H" U3 a14
    4 Z$ x' N, {& u3 t15
    . @: s+ ?7 H5 [- h4 W8 y16
    / f* ^1 L0 ?. Y' b& T( v) ^17
    6 p- y. W, V; Z  w( L2 c18
    $ Y& N6 U/ S# j! \5 D  |6 S. S  r( }8. 加载已经训练的模型6 F6 X8 Z% q/ P6 _$ a
    相当于做一次简单的前向传播(逻辑推理),不用更新参数* U: r# ?% H$ Q" o' Z9 n  N: H3 c1 v
    " L" \: h) e6 [# t$ i
    model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)1 R) b  N/ a3 }7 T/ J& s
    ) Q8 S; @5 @& o6 n' W4 @6 G
    # GPU 模式
      D; ~, P% _9 \  z: n9 Z* Omodel_ft = model_ft.to(device) # 扔到GPU中' ?3 x; ^) D9 Y/ F& d- r+ m

    8 n+ ^7 x: z. J. v# 保存文件的名字: \& h) a9 V% `; n; Q2 w6 W% m) h
    filename='checkpoint.pth'" K3 i. I0 t. _

    2 ~% \$ u: [; ]1 A4 W  E( C2 p' P( j# 加载模型
    4 w0 t* h( k5 m- w+ w- dcheckpoint = torch.load(filename)8 v; b& ]- W4 [
    best_acc = checkpoint['best_acc']
    ; l  V0 H" ]  R! f; n# j* Bmodel_ft.load_state_dict(checkpoint['state_dict'])
    " W3 Z4 M, }4 D# |; R1; s% C) j, O. R' p/ V
    2
    - D2 D' k' r; R31 Q  G  r' I1 J: R7 R0 Y
    4
    . M- e' ?: ^# M0 `3 g4 Y5
    : k9 x- I# \( K5 L9 I6, h0 z4 Z1 X: e# Z$ }( v: y9 J
    7
    : s2 T% v3 M; F, ?84 |# I5 X6 W$ A6 P' u! n
    9# x! c4 l( M+ O) e  B# q
    10; o  m$ a. X, p8 n8 f6 j
    11
    ; h! a+ V- z7 G0 ?. N0 Y  S4 B12
    3 |9 g/ c3 o$ F* O<All keys matched successfully>
    ; y  ]# c7 q8 |0 M3 p+ l; r1
    - \7 e& p* e" {2 t, S  adef process_image(image_path):- z% @6 ?. |5 F3 p+ k
        # 读取测试集数据# X- M* q2 v( B8 s
        img = Image.open(image_path)% }) _  h" h; D; W& I% A: B- I% P
        # Resize, thumbnail方法只能进行比例缩小,所以进行判断* u4 ~# C  p; y* D3 h3 I
        # 与Resize不同
    5 Q& b" n9 b( a' z6 ]    # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小
    ; p4 ^5 L- v- x    # 而且对象调用方法会直接改变其大小,返回None
    4 f; O( I: H  ~( |) I1 \2 q7 C+ }3 i    if img.size[0] > img.size[1]:) Y" O1 S/ `0 [( a: A9 k* t
            img.thumbnail((10000, 256))* N% {$ x$ L+ K1 J1 T
        else:
    7 u) _* G5 _( j7 |3 `3 Z        img.thumbnail((256, 10000))
    1 u) Y! l; W, b; |
    2 W' G7 T+ S% r4 H    # crop操作, 将图像再次裁剪为 224 * 224
    + F+ U6 N% n! b) ^    left_margin = (img.width - 224) / 2 # 取中间的部分
    $ n! {, _( i8 W8 g    bottom_margin = (img.height - 224) / 2 2 r9 p0 X6 I1 |( k0 D# x% T2 t
        right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度
    & p5 G5 y( {6 |$ U1 O% X& o* u    top_margin = bottom_margin + 224
    7 W& T/ [& R# o- Q
    . S& E2 f* y2 i    img = img.crop((left_margin, bottom_margin, right_margin, top_margin))% m+ k+ V3 m. _( T; d3 T
    % n% m; M+ B3 j. \0 ~7 {5 u
        # 相同预处理的方法
    , v2 L$ ?4 z% C5 E2 u* S    # 归一化" B. u+ C& r- u/ d9 F4 r& d. J
        img = np.array(img) / 2556 c* Z# N- ~( f# Z! `0 ~
        mean = np.array([0.485, 0.456, 0.406]), G5 X* }1 k; {7 n  `8 h8 |" E
        std = np.array([0.229, 0.224, 0.225])( S% N8 `% J7 U% k4 L
        img = (img - mean) / std
    4 w/ Y$ M, K. B$ B# v+ p( y, h. _7 S
        # 注意颜色通道和位置
    1 }$ q( Z( _0 v! U    img = img.transpose((2, 0, 1))
    4 @5 \2 z, _1 e- i1 J* Y* R$ o8 p4 c6 n$ R- ]
        return img
    5 _$ _/ ]8 r/ `: n, k
    ( p. Q: E( l0 |3 i5 g5 J+ |def imshow(image, ax = None, title = None):8 T( g( s" g/ p9 K
        """展示数据"""7 h( N6 m! J  w- v) K+ N3 q! f
        if ax is None:$ {9 s" D) R$ X% {5 q, r% C
            fig, ax = plt.subplots()
    " ]) ^& s& w5 D& s, C: p" p
    5 }! p9 Z8 O9 [) O2 W% s. q    # 颜色通道进行还原, ^& e) s: h9 d, @8 o
        image = np.array(image).transpose((1, 2, 0))3 F2 a% O! C4 v+ u; B6 l

    : f  y' k0 Q4 U  j8 W+ ^/ b    # 预处理还原% }8 k  ^5 W# _
        mean = np.array([0.485, 0.456, 0.406]). v1 z$ j$ G4 V. f
        std = np.array([0.229, 0.224, 0.225]), P( e. U1 B! y4 L
        image = std * image + mean
    ' J8 \. V- p! I& S6 }    image = np.clip(image, 0, 1)! x" l3 F* o* `5 Y

    2 l& b* k) p# O0 F& x    ax.imshow(image)
    / Q7 G4 D$ A  `; a5 L    ax.set_title(title)" @2 n! I2 W- n, U2 c3 k( _

    & s4 ~* T! i  K: J5 z5 J; d8 {    return ax5 I/ x' T8 C' a& L" A/ s

    ! h2 J% z) v6 H" Ximage_path = r'./flower_data/valid/3/image_06621.jpg'
    ) |8 I0 f' m6 n+ _# Z0 l: Q& Himg = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理
    8 `9 c2 f1 B  n1 }/ V7 _imshow(img)# j+ Q9 f, c! M! {

    . B0 L% D: ^" \  r- V. y1
    3 e4 Z5 V% D2 H7 A2
    , D  u# ]4 J8 @! o9 i8 b) e, {* E3; b# K8 i# q" ?+ N* A+ X
    4
    2 |. E3 V6 K2 D+ k. ~54 l4 |+ E) [# }6 R6 D
    6
    1 W1 k. [1 k2 K/ ^  ?& L/ _73 V- b" q( M; j6 \5 C& u
    8$ i8 U0 E& `3 ^) }# A
    9
    ! \4 k$ @* E. b102 c0 F& B" ?0 `9 |
    11
    ! b5 D. x% a9 L- ^6 p; V! l; x12
    9 s% y& u0 s2 f. P" ~; ?& c13
    : v, C: `3 Y5 h  s! M14
    - s# r! |' |$ J: n* R, |$ {2 m152 J/ _" k* c5 c0 ^3 `5 k
    16
    : }- O/ e' G, I* ]17; w; N8 V0 @  w( W7 v
    18$ F5 A8 D- F1 ?6 J
    19
    3 D7 ~4 y# i& K0 m7 F205 {! N- m% o+ P) o2 g) W
    21
    + ^1 k8 {9 f& A& |22
    8 B/ Q) F8 b1 u# h23& i7 h8 P  m. l5 q* M: H0 Y( ]* E
    24/ w, m# I* o& a: {! }- j# n! m. G6 w
    25+ g' L1 W$ x" V6 L: c- Q
    269 B# t+ Y0 `9 N9 N7 t; R) T7 z
    27- V/ ~2 ^7 M5 B) f
    282 E7 R( U/ `  W% l
    299 D, N2 F; L7 Q- ^
    306 W9 l) J  z6 q/ I6 Z, |  n& N
    31) ^- o+ L& c6 M- e6 s0 ?
    321 @' Z/ m7 E7 t8 U; j5 q
    338 t2 \* d/ |9 `
    34
    5 A* F+ T& O8 I35
    0 G3 R, C* E4 M36
    6 e) a, G; M9 U0 g4 j, B379 h: e* s+ H# [6 f, `/ ^
    383 M$ H; g0 h% P* U8 z
    39
    ! E7 E3 H) G6 B/ N! f( o+ j5 g406 C6 I: D! i! Z1 i5 \4 a3 c+ w
    414 O4 N& o  o- N' Z
    42' p, A$ g5 {4 ]2 y% j! `' ^
    435 A4 f: O, C( C3 }- ?8 B
    44' f( K+ @- ], n% N
    45
    + A& f9 N- p3 h  Y462 U5 Q7 u$ I; s! r
    47
    & W; C5 C2 O6 [* b- _+ O48
    ! U: B1 V9 V+ u  q2 ]- ?49  B* W$ m7 V2 n# |
    50
    2 z& J2 G! A  E# y9 W/ W1 u( f0 o51
    # Z# X* A& a, O* z4 ]52
    3 u% G) X+ y7 L5 N$ |53
    4 l7 v% a' y2 M: |. y- B4 O54
    7 R# M% o5 b: E7 u6 z2 L# Y<AxesSubplot:>/ F" D( r# \* W+ _: z$ _& p! @
    1
    ; f( j! Z: A- q, J9 ?
    % ~# t/ F" y8 ~, ]上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确( k$ o% R6 f. x

    7 @5 a. q: u! w. \3 f7 I1 Jimg.shape, @) M# V! ~9 I. s, r$ M
    1, d. e3 f9 X5 s5 ^. z2 [" c! P$ U3 ^
    (3, 224, 224)
    ) K* _- a% x+ E7 {1: {1 @3 K  P+ `& u9 x1 h7 G
    证明了通道提前了,而且大小没改变
    $ y; O: a6 Q7 C; C2 P
      \  `% g( h. d, Y& ^9 Y* Y9. 推理
    / [" W) D/ C2 Z- I5 B# h% Qimg.shape% T" T$ Y$ }6 [, C- Y
    + O* T/ w) F2 x( ?- F9 U( l
    # 得到一个batch的测试数据" h3 {1 b1 @5 f! w
    dataiter = iter(dataloaders['valid'])' i/ q) |% {) i- {
    images, labels = dataiter.next()4 l* O: e+ g# g+ v9 B; `" r

    + c) J$ l* m5 i9 f! ^model_ft.eval()
      h& y' x! J- d3 w& V' a
    3 ]) x. N9 f% V& A1 `7 g" kif train_on_gpu:) W  {) Y& s# _; F" k* l, @  z
        # 前向传播跑一次会得到output
    8 o9 ]4 n0 |( o7 @    output = model_ft(images.cuda())
    9 Q7 d1 ~" n, S* u' Pelse:
    3 f) L$ `: z, w: F# ?8 B; O    output = model_ft(images)
    + }! _& c0 Y, F* \* f$ L  }0 s7 @% P5 k) I0 O+ S3 @
    # batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值6 ~/ _$ w- i& f8 H. ]) ^
    output.shape
    $ h" u0 K( s, v7 N" G! S$ |6 V+ d3 f( E" K
    14 w8 p: V; M) }( g: E
    2. \' _( \7 |$ Q: v
    3
    8 W5 j9 |: [  X3 d! U4  u% s  q6 j* c  W- p4 w& S8 U/ t
    5/ C$ `; W# [9 W. v( m* w1 X
    6
    + w: B5 N6 k& V4 Q# j7
    $ P1 g# G2 D7 z1 W5 z" Z. W3 \% _( \8- b) ]- E1 d0 K8 U  h0 R2 S9 G
    9$ y6 j: m* _: q# _3 T
    10
    8 E* C6 R; d' i$ p5 u7 m# h110 c6 J/ U' [5 Z
    12& B2 @$ G2 z5 w0 C
    13
    6 c( j0 j6 F! W3 ^9 T14; E3 b2 G$ i# s, P7 p
    15
    9 v8 [9 w6 r2 f! z0 g$ m# k16( m4 ]3 \7 C) `1 A8 B
    torch.Size([8, 102])1 _& S8 ]- T* P% B2 P* u% |  I8 B
    13 Y( C" \- w7 ]* M( j4 T
    9.1 计算得到最大概率
    6 f8 R# w/ D/ s( k) |4 \) __, preds_tensor = torch.max(output, 1)
    6 D* _7 W, r$ g7 Q) h
    8 p) \: r; x8 Q; F, P6 Dpreds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量
    3 [& f! t4 b4 U4 Y; @# A1
      K. c: t" O# l$ I' _2. v4 `% J4 U: h$ Y7 F  G% P# h
    3
    , r; w& d1 j' @5 e9.2 展示预测结果
    " m, U2 m: O1 ?! E* Jfig = plt.figure(figsize = (20, 20))
    7 d% Y! h9 U% V8 `$ j* g. Mcolumns = 4$ j6 V3 a; Y. U+ t  _. Y0 g9 f
    rows = 2
    - [( L& C$ v* v5 H+ c0 U2 K- B$ H3 z$ K; j
    for idx in range(columns * rows):
    % Q8 V7 j2 ^; h. l) i    ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])
    4 ]/ b+ n% f7 b" A+ p. M6 |1 b5 L    plt.imshow(im_convert(images[idx]))" U7 k1 e0 t2 w, Y$ G
        ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), $ D8 j/ G  T: r) s
                    color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))  ^* B1 f; v9 A8 T4 {% ?) T+ W
    plt.show()
    : ~7 X8 T1 Q  ]4 N8 L5 D1 }# 绿色的表示预测是对的,红色表示预测错了, e9 z+ B4 G- B, r+ a
    1
    7 z+ d" `1 p5 Q8 k, a$ h, z2
    * P! c: R7 A/ e1 |35 l- _* B, E* C4 V+ v: H
    4; n3 a( O2 m0 G5 a" o& w$ i4 K
    5. ~, p* e% p! x/ @  m
    60 A, v. x2 D" G, e  |0 R
    7: ]5 i4 t$ n. A2 n
    85 E/ l' O3 ?( x* A
    91 I( N, y; p& h* K3 P9 K1 y- [) V
    10: {1 u2 f9 s. Y" S
    11
    7 p4 j& q1 w9 ?7 C5 Q6 }+ I6 Q  x7 t6 \1 S# T
    7 C# S* M3 M/ v  y+ y, Z

    6 U# Y1 O3 P) d# f) r: T- o————————————————
    / _- _( W/ B- K$ p/ Z版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。% y$ ?( y. G! |% Z
    原文链接:https://blog.csdn.net/LeungSr/article/details/126747940- B; Z6 }! J( ]& _+ Q
    : y1 M. n* t( x/ ]

    " z% z0 P, p0 X# a* V0 x, P( ~
    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 16:06 , Processed in 1.573639 second(s), 51 queries .

    回顶部