QQ登录

只需要一步,快速开始

 注册地址  找回密码
查看: 2772|回复: 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)实战案例
    * z' p' O1 g. f' B- m1 k: Q: G! I. j7 ?
    文章目录; p3 x: O) S- d/ x! F) O. R
    卷积网络实战 对花进行分类! G* s0 I8 i( [/ o
    数据预处理部分
    * }. N" K' t1 V9 w网络模块设置
    ' _4 [8 o0 j9 {( X网络模型的保存与测试- p% Z7 g- l# E) R' }
    数据下载:
    0 Q* t5 ?+ F. c5 ]1 q1. 导入工具包
    0 M/ e) I, n7 V) z2. 数据预处理与操作
    1 Y7 c7 Y0 s! ~  e6 s' C3. 制作好数据源. v4 h( U0 j: U: F  o
    读取标签对应的实际名字
    , K5 b& H# e' Y- s7 Z6 |5 ^4.展示一下数据
    7 f1 n4 Z% R! w5. 加载models提供的模型,并直接用训练好的权重做初始化参数8 }* }; ~7 M4 Q) s  X
    6.初始化模型架构/ }# r4 G7 _* R' x9 d4 U5 Z
    7. 设置需要训练的参数
    : O' n5 N4 k) W+ u% l7. 训练与预测5 W% E7 P, n# K9 a  W+ g
    7.1 优化器设置' X2 b; Y) \" S3 I* k+ ?
    7.2 开始训练模型
    ! C7 w- A) h8 Z6 }7.3 训练所有层7 e4 d8 G2 `. |* E5 A5 v- ^7 ]; F4 i3 d
    开始训练% _8 ?4 B3 U0 I9 C  V
    8. 加载已经训练的模型  n( x: d$ H2 Z2 g, Y# c2 Z
    9. 推理, U9 A( A! @4 |
    9.1 计算得到最大概率8 c) l% U% s' h' L
    9.2 展示预测结果
    ) G, T/ k  U* [5 Z: g写在最后
    6 K4 m8 Y- {8 I卷积网络实战 对花进行分类. h4 h: W/ L; @- E& g5 q
    本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能4 l- B  f$ T& @& A8 S) C
    * u% x2 u( _; F6 @
    在文件夹中有102种花,我们主要要对这些花进行分类任务
    5 I0 T2 W8 \$ J文件夹结构& k3 W" V( l: u% `
    ( B+ C/ k8 Q1 H6 Q/ a1 E; I
    flower_data$ m7 L, v2 j# I9 ?% B$ {
    ) \( b+ d9 \- ~
    train) ^) `4 [) w8 r- [
    % k  e# m' B5 {1 t6 r! U, w
    1(类别)
    4 r. t: f" X6 R- s4 a2* v" u7 X5 D0 [3 J9 ]( i/ h
    xxx.png / xxx.jpg: U: T  I/ [( y/ H5 J
    valid! Y6 C1 ^" J0 U

    / Q3 w, Z9 n3 F' r, w8 P) |  @主要分为以下几个大模块! F' p- m) Q5 I( t
    " M& e7 u/ D/ C/ x2 f. m& D
    数据预处理部分; R9 o/ |( |# G
    数据增强) I- Z: A3 ^# s, q; I& t8 [
    数据预处理% v, ?+ I- S' l: a
    网络模块设置7 c, \( a( m3 E
    加载预训练模型,直接调用torchVision的经典网络架构1 j' K: l8 M  s+ ]6 I" f8 H* c
    因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务
    : c6 b3 D; Y" [9 j网络模型的保存与测试
    2 d& t( m5 `) N2 E# e; [模型保存可以带有选择性+ @( @! c; C; L: t4 f; N
    数据下载:% }+ D( A; w0 R8 p
    https://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset5 K4 ~$ e9 D  Y. W3 D, V2 L
    - _% ~; f$ K+ H! B" N7 V" I
    改一下文件名,然后将它放到同一根目录就可以了
    0 K; r; i0 [( D+ g6 B2 e# H0 z% [0 b. }0 E
    下面是我的数据根目录8 J  `7 ^+ _* ?

    8 x% P! i9 q' ~; S  T2 k/ j( c( M/ z
    1. 导入工具包
    * I; E, g$ I, E; `- n- f( @import os
    $ B: S$ x% M8 a$ t9 `, [( q& _import matplotlib.pyplot as plt2 R# b; l$ T8 @; F
    # 内嵌入绘图简去show的句柄
    " l6 {! f7 M7 j%matplotlib inline
    % w" ?7 T3 S& p% b) b* @2 ^6 Gimport numpy as np
    . D# n0 `! \9 g3 K: v1 {7 gimport torch
    % y; I( v2 f6 j* {9 m& D) V/ ~from torch import nn
    7 ?8 N/ g" U0 C9 W7 z8 i0 i4 G5 @0 O4 I5 h# ^2 }* X  G+ ]) A
    import torch.optim as optim
    , m* t% U# X# A0 f2 X' Vimport torchvision+ g7 b) G. J7 m/ ]+ e
    from torchvision import transforms, models, datasets9 S1 O; j2 w7 `

    0 a" J1 |3 u. ^4 V* f: [  {, Himport imageio% D1 r1 ?2 E  k
    import time
    ( ?* v/ G) G% F! T. X0 yimport warnings
    : B( k/ x; ?) @6 q# \. Zimport random
    % I$ |: Z2 C8 eimport sys
    ! i, |4 }# G+ W# D5 @# nimport copy
    , Z  `% u' K+ \" gimport json6 O3 Q) z. d: t/ y& r& G; O
    from PIL import Image
    ; X/ o% D! ^  ?) z2 d! f; ?9 t
    4 c0 s2 D2 ]8 j, e' g) U
    7 c7 d& S5 @+ ~+ e/ z. J& K  a* j1
    ( k1 P, p+ I( O2
    ) @0 z2 G5 i! x6 `2 w9 S3# B/ |; W+ n/ d$ Z
    4
    ( N2 G# m. U8 ~! X53 R: s3 x8 X$ A# L$ i$ |
    6  n$ H# c/ N, n5 @
    7
    8 I  `/ W4 h8 ~8 [% `5 k( R8
    2 p4 Q, `; x' A9
    8 `' c3 a2 g( s2 B: U10* B2 K- z7 O# X# M' b4 Y. h7 k# f; r& w; }
    11
    0 Q2 }1 n+ ^8 j! ~12
    . K) |7 q2 Y! L6 k% H, I13
    / {' d9 d! U2 k! N+ ]4 t# J, B14
    7 s6 R  r' P, L1 R15
    . \5 E4 X# J2 z0 ^2 i& Q+ e16
    % H2 }6 X  d% H" O( P+ a17
    # g) c; s% K7 y9 a9 p18+ I. B  V. |; |/ h; m8 N
    19
    ( M% s2 ~# s) M1 Z+ U209 }6 H5 X/ w! r8 o- C" M7 W
    21& {4 D* H* U; t! ?
    2. 数据预处理与操作
    " \% P3 i; }& @& v! \#路径设置6 u  }+ k0 {0 W$ S9 Q3 k& R% B. [, l
    data_dir = './flower_data/' # 当前文件夹下的flowerdata目录
    5 r0 p8 m" K  E2 m: Ytrain_dir = data_dir + '/train'
    * \1 L" o0 ]4 y5 kvalid_dir = data_dir + '/valid'
    5 p$ h6 L5 v7 D! @2 ^& B1/ q$ j. q$ U3 H" Q0 r, @4 w
    2
    . L6 E* A" d$ Z, l! _) t" G, r+ R3
    ! t7 D+ M% r7 E" T5 [7 Y47 J  Z; @! `7 C9 K2 c5 i+ q7 V" o
    python目录点杠的组合与区别9 }4 |& \0 s6 Q( h
    注: 里面注明了点杠和斜杠的操作
    ; @8 W7 m* w5 e3 p9 D( ~6 K. N6 e) T/ H
    3. 制作好数据源2 X* l$ A6 P+ c2 K! L9 P9 s
    data_transforms中制定了所有图像预处理的操作2 ^! a; B& U3 O8 Z7 ]5 }* V
    ImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片# P- ?! q( j: L2 h3 r
    data_transforms = {7 R7 q8 [2 X: Z. H
        # 分成两部分,一部分是训练
    9 K. s  h+ G  [7 R' e5 D    'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间  N. O- C# A( f  }2 o7 ]6 j4 w
                                     transforms.CenterCrop(224), # 从中心处开始裁剪; N; @/ N% E' N! t
                                     # 以某个随机的概率决定是否翻转 55开
    " U5 S, r4 \9 i                                 transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转% r1 p* @9 k. O( I* L
                                     transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转4 A# A- R* E) \
                                     # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相/ D% d4 R1 {1 F3 q( ~7 Y6 U
                                     transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),
    $ z- r; w- s, p5 x* h* ^" B6 {  S# y" S                                 transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB' }6 }% Y' d3 G
                                     # 灰度图转换以后也是三个通道,但是只是RGB是一样的$ k! C4 N& Z. J$ f& L* O; o
                                     transforms.ToTensor(),, S. e( s2 m* a6 d. Y
                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差
    ) h+ O- `" u8 [+ H$ `0 B8 o: q                                ]),
    & F4 O  U: v$ k% I4 i3 u    # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
    5 t+ v, \, m& H; t9 u0 I7 ~9 }0 `( r    'valid': transforms.Compose([transforms.Resize(256),
    ) r/ {: d4 h0 U                                 transforms.CenterCrop(224),
    3 j( x/ i- p+ N                                 transforms.ToTensor(),
    " f6 ]" b- h; q7 z( k                                 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同
    1 G0 Y2 a# |% g5 c2 {: r                                ]),/ A" S9 T7 M, H) x7 `1 t
    }0 \' d% G& @& K3 Z  ^- [) r, }, c
    3 A! a! _; Q7 a2 i
    1: q4 m6 ]% [% M+ {( V
    2: L" G6 G4 i1 q$ [
    3
    % [+ u' z6 s4 Z# Z5 X7 g$ L+ n4
    + ]. D4 @0 ]2 b! f, f$ \7 [* L5
      c" S- ?+ V5 J/ |  t) q  Q# s, m# o6
    # d& W  j* y0 g) W: W" c3 Y7
    3 m# H- w9 U9 X' A8
    & q7 e; a( L) d0 p# o9
      ^& T# b2 ~; ]" ]7 ^/ d10" s- T  @: X- _
    11
    " I; d& E% N- H- l3 g12
    1 e% _+ c" M1 h! }139 r$ z& a( e: G  V& p/ z
    14
    3 p. Y  t0 D0 R+ a  W15
    $ D  z" c" z3 d3 g1 l. m  C165 {, g/ z& O3 l4 Y9 |6 |( [
    17
    5 u  E: D+ ]' Q, e3 X. c18
    2 R5 g' i! {' ^, w5 ^19. g. {" X4 ]0 f% h- v1 x/ S3 O3 q) x- C
    20/ F" m3 W4 z& g  d2 G2 `2 y5 X
    21
    / p+ J, A. I+ e7 Q4 G# Xbatch_size = 8
    7 }5 O! E: a" J9 b/ U2 _2 yimage_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}! I* X- s0 B: j5 L- i
    dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}$ u4 a5 F4 [! y
    dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']}   D7 h4 t+ k- I6 ?
    class_names = image_datasets['train'].classes
    + G, i  ~% T8 l9 w1 }
    5 ]  K% h4 h: @$ T  Q* X+ \#查看数据集合
    ! u, F, e' T# k- L* f8 Pimage_datasets
    4 ^+ Q# e# ^8 I3 ?
    ) r) L7 b  F4 I4 }* n; Z5 }1
    ! m4 K, M7 S, N2
    & w; a) J8 N3 j& f5 h/ r36 Z4 k! g% N  u  t
    4
    4 o5 s1 A0 W7 l9 v6 p56 R  N' N: a' c( l2 ]
    6) P, U* H# U$ u4 ]1 L
    7' S, [# S* w' q% S, R1 O1 h
    8
    6 a0 j- u$ C0 p4 t4 J/ E$ Q95 q9 O3 |0 L( `0 E3 [2 B
    {'train': Dataset ImageFolder/ W8 ~: c6 p2 w0 j
         Number of datapoints: 6552
    ) ?5 G8 S! B% o5 F     Root location: ./flower_data/train
    0 Y/ N# R+ t2 T" ^+ z; b     StandardTransform
    / F4 R: g2 u' K. V) q/ z: Z  o. ?' z Transform: Compose(
    * Y, k9 u6 D8 c# S5 Z, q# D; c                RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)' i$ ~9 x: C) S) f5 a: X
                    CenterCrop(size=(224, 224))
    0 ]* R( e1 N& F# A                RandomHorizontalFlip(p=0.5)4 @3 ~8 D( Y$ u1 U" _, Q" Z
                    RandomVerticalFlip(p=0.5)6 ^# B2 W% E- S% p* \& s) \  b
                    ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
    + m' I2 ^5 W" s$ W) Z* I                RandomGrayscale(p=0.025): @8 E  j: S6 A! T1 D3 @1 \$ R
                    ToTensor()
    ' s( x# g7 i* @5 G) `; @3 {/ F1 |                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])+ X4 d  G2 \/ {% T
                ),
    ; q* r8 W) I2 ], k/ j, d 'valid': Dataset ImageFolder
    $ x: F; J( j( l0 w; R     Number of datapoints: 818
    % z+ F- L0 p& r" B; f     Root location: ./flower_data/valid* {7 m: ]- D3 z, |7 A
         StandardTransform
      u2 }2 k$ T9 L7 E4 z. a& W- E8 { Transform: Compose(3 \- ]8 ^, y" {
                    Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)
    / f+ [) e# K7 M2 ]3 [. K3 C! t                CenterCrop(size=(224, 224))6 ]( i7 i3 }6 f" x: t! H
                    ToTensor(). [6 N' p9 z& E9 ~0 @; t
                    Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    2 I. S, b+ r7 r5 a  {- Q$ U: H; F0 v" ~            )}9 I% ^- I; B/ ~; s, {0 |$ l6 L
    * y' Z3 F- `. x8 M& _* \
    1, T! R# Y: O$ _, }, e
    2
    7 R1 k) [# D- z  w; p3
    : D% G/ y6 t1 `( g% v0 Q4
    : M! _  K+ ~0 e, _& i. U5
    ! n  i" v$ c! s* u# D6
    & C8 m1 {' S' k/ B- t8 y) p# c7
    ; l  R1 j  x9 l% g8
    2 T6 i, y: ]% L94 q6 _8 R4 Q. u4 @) }* r
    10
    $ W- L  ^9 |6 Z0 z: X& I11% [5 D$ \! r' {9 p
    12
    ) R. z# e$ U6 f0 F  @( k. U13
    / J9 ?" Q9 u0 Y  P% H14  L) f( H# H( v; V( D+ I
    156 s; I- j; f. G
    16
    ' `+ M( k, m6 x/ ?6 x171 z+ N+ \1 P" p0 B+ ^9 `9 T9 }
    18) _- P8 \/ \7 T% `( Z5 Y: `2 M
    199 h  h; [: x" O9 ]
    20
    3 G1 Q1 _) b8 f9 k/ y21  G5 _$ V4 s2 R# t2 ^5 h$ w
    22
    $ e" c- {0 u9 {- K235 e+ ^$ J7 {' G5 v* [' ^% I; ^, F
    24
    4 Y7 }6 ~, ]) S/ O- t# 验证一下数据是否已经被处理完毕
    8 e2 h  e  W- W, B; J+ r, O- ?) ]dataloaders
    4 Z) W' ?# I6 T10 V& z  H$ @5 D2 E! E6 K
    2
    ' w/ K' Y. p$ S, A# R{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,5 S' v, K8 y- z0 O; M, v- M  d
    'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}. [2 z0 l: a$ Q  N
    1
    8 l" v) e3 g* I6 @9 ~" m, Z2
    6 S7 P: N9 L  J4 j0 W& J! A" w+ Rdataset_sizes( ?- h2 z  O1 A. [
    1) o# Z* T9 P) {' E3 U
    {'train': 6552, 'valid': 818}" o  c' |, c( x0 G
    1
    1 o4 o: z, }; W( D读取标签对应的实际名字" O1 B& [% }( a8 U6 e2 f. Z, Y
    使用同一目录下的json文件,反向映射出花对应的名字1 K0 F" D6 N  I. X

    0 L1 J9 F) ?. ]. k% P4 zwith open('./flower_data/cat_to_name.json', 'r') as f:
      I; M# ?$ E4 [; g# r% X    cat_to_name = json.load(f)- E% h- h9 |* Z! a% S6 J6 O% y! t0 k
    1
    / s# l3 Z- P5 Q$ S! \" A: r6 X2
    , M( L8 v7 I+ {+ Kcat_to_name) R- E  {. @9 @6 e( `+ _+ Z: P
    1* O' u- m$ v, U6 z) X( ]4 i6 S1 g
    {'21': 'fire lily',
    1 B9 _1 M+ r2 q% h2 o '3': 'canterbury bells',
    3 w# i# `; H% Q7 s7 U# @7 S '45': 'bolero deep blue',! p& m& l* |# e2 |# k6 V
    '1': 'pink primrose',0 B0 V+ c: }4 \7 k, c  z
    '34': 'mexican aster',. a0 n3 k1 }0 E; h0 J8 U
    '27': 'prince of wales feathers',
    , B8 u% K; l8 N# N; D% c0 v '7': 'moon orchid',
      b9 G! Q$ N* |1 E) P) O/ a '16': 'globe-flower',
    0 d  s: a+ Y" |, z '25': 'grape hyacinth',
    2 W/ M% q+ b5 M! s9 E; j' G '26': 'corn poppy',
    , R1 U/ R) h0 g) k: R '79': 'toad lily',/ @' @1 X/ w# |4 K4 q& a! H* B3 }
    '39': 'siam tulip',
    " ^9 ~5 m* t' @+ x# B '24': 'red ginger',# ~- V( X0 {4 t8 [+ }8 S' c! }9 G
    '67': 'spring crocus',
    8 a( j$ e) T8 w6 |, T" i '35': 'alpine sea holly',
    : @  H  b4 p+ |( h- k '32': 'garden phlox',+ f' z, b3 G( j, C7 K
    '10': 'globe thistle',6 O. V+ w6 O& L7 A: Q/ _4 A
    '6': 'tiger lily',/ K( W* \/ b  b- t  U
    '93': 'ball moss',
    0 X, q# t/ _$ c# O% g" T; r '33': 'love in the mist',
    8 h" v0 z) X1 _ '9': 'monkshood',
    ! r: `# ^$ X, U) k$ v '102': 'blackberry lily',
    1 K$ G1 s+ Q! e$ k1 b. l$ J- m '14': 'spear thistle',: U6 \3 v* v& L# L& i2 {4 o
    '19': 'balloon flower',
    ) x" h& P% c* m& Q# r '100': 'blanket flower',
    8 a* j1 f: j  L0 C '13': 'king protea',! i' B- O% H$ m: R3 T
    '49': 'oxeye daisy',- x+ ]6 Z* q" B6 \3 }$ z0 d& B; l
    '15': 'yellow iris',
    5 v  h; t( \2 {* t$ D/ i- V4 I '61': 'cautleya spicata',9 i9 s, I4 d3 g: P" s
    '31': 'carnation',
    / @9 x! q+ _3 q, a* P '64': 'silverbush',
    ( v5 s/ ~  D6 p '68': 'bearded iris'," H1 X. ^" m( g# P: M5 z
    '63': 'black-eyed susan',
    8 G8 `9 Z; W2 B( z( @1 U# g6 ^ '69': 'windflower',
    ) V/ Q: F7 G7 ^9 p$ ~1 A '62': 'japanese anemone',! C8 b: M1 n; W9 w
    '20': 'giant white arum lily',/ Z  e6 r3 `( C1 T1 v
    '38': 'great masterwort',
    6 v5 ?; v8 y  ^& P: e# b '4': 'sweet pea',
    7 |/ _+ e  d- G* u) B9 i; t% @2 X '86': 'tree mallow',
    # h& `3 S5 s6 e" \7 G '101': 'trumpet creeper',7 f2 ?% Y7 D6 v, I/ o  U
    '42': 'daffodil',2 |% f- R. z2 s6 M+ b6 O
    '22': 'pincushion flower',8 U" y9 K! V; i
    '2': 'hard-leaved pocket orchid',1 I( C; ~! ]" o/ ?1 @) [
    '54': 'sunflower',
    3 E6 u7 v" C3 T0 Z! M7 l '66': 'osteospermum',' N4 O. v9 a, p/ }) [- l
    '70': 'tree poppy',
    ! }* f( \3 ?% v3 F; n, Z  F '85': 'desert-rose',
    ( }6 n/ `7 o. Z# f8 r: W$ m: @ '99': 'bromelia',: D- q+ `+ D. D% I! B( _$ ?% _9 J
    '87': 'magnolia',
    ( l; @" [' O+ e6 o! d4 S7 c '5': 'english marigold',
      U9 S, k% W* F0 k! |( x1 Y '92': 'bee balm',
    * s; B2 y" p' \- Q) D '28': 'stemless gentian',; m$ M, z# ]+ R6 |4 ]7 H
    '97': 'mallow',
    " E2 {: Z! r# v9 H; o( \  j '57': 'gaura',4 K# Q8 P6 E$ L7 o
    '40': 'lenten rose',
    9 x# g  l6 V6 s9 g, { '47': 'marigold',
    ! ^0 }+ w) U) |; |0 L- d" s- v '59': 'orange dahlia',
    ) u4 F: [9 v9 j% N3 u '48': 'buttercup',7 Y7 p  }7 c: G0 i" H
    '55': 'pelargonium',; R3 E* U, n  V) _  v
    '36': 'ruby-lipped cattleya',
    ( X5 E3 U4 d* ^! H5 L0 J '91': 'hippeastrum',
    " u4 k6 b" O; u1 S9 e) _, z '29': 'artichoke',' Z4 V  `) r2 C5 W
    '71': 'gazania',3 h/ v0 e& I3 w- X) g1 Q" y
    '90': 'canna lily',! A2 ?% p4 H4 F7 B5 q5 m8 _) m& T3 l+ M
    '18': 'peruvian lily',
    3 E2 L* m+ J2 V6 T9 h3 s4 C1 y '98': 'mexican petunia',
    # s4 f1 r" O1 F. U1 F9 ] '8': 'bird of paradise',/ Z- Q( n) q+ P3 `. H5 e9 }* C# H
    '30': 'sweet william',
    ( C! E$ E  d# A5 [2 K" n '17': 'purple coneflower',) I, u5 X. p8 b6 h
    '52': 'wild pansy',+ }/ k- p. V; @/ P, U. U
    '84': 'columbine',
    % M% f7 b0 m9 o '12': "colt's foot",
    9 t8 Q1 P* w4 ^& U7 z+ @7 B* `0 j1 v '11': 'snapdragon',
    6 C* H6 |. Y2 X  V8 } '96': 'camellia',$ M- s7 G) Y: _4 p% }: Q
    '23': 'fritillary',
    " s: ~; Q) P, O. L '50': 'common dandelion',
    . `6 W& {4 g1 K$ R$ n5 K6 Z '44': 'poinsettia',
    ! ~  X, J& [9 H# A  F& ] '53': 'primula',2 T+ t# v3 N2 \" h- m
    '72': 'azalea',/ r- o5 i  L; Z1 u/ d* C
    '65': 'californian poppy',
    5 c  C9 [* k* W9 G* F8 ^' f0 ]7 m '80': 'anthurium',- @( l* V* Y$ e' M
    '76': 'morning glory',/ K# S8 b* g. c: _2 T& ~
    '37': 'cape flower',  Q- K- o, S- Y& r: V8 V# T1 m
    '56': 'bishop of llandaff',
    + j9 ?( G7 ]3 t# c '60': 'pink-yellow dahlia',
      {  l' |3 x  B7 P6 G% d '82': 'clematis',
    1 y' N* `/ g5 W: ~1 V. W '58': 'geranium',
    ( E9 |" d) j) j '75': 'thorn apple',; i1 ]7 M" t- j* v: B" U+ U5 b. ~
    '41': 'barbeton daisy',
    ' w. P4 ^, _) S, L4 g '95': 'bougainvillea',
    : b, o. _7 p+ f2 o0 @ '43': 'sword lily',
    9 d5 `, |! D! A '83': 'hibiscus',
    * m" T, ]& V- l+ v- [ '78': 'lotus lotus',
    9 ?1 U; y; d: u3 ~; R8 H '88': 'cyclamen',/ I/ D, N' U; |1 o" u
    '94': 'foxglove',: {1 w) f. q* ~8 h' W6 e# n# |
    '81': 'frangipani',
    % x# F8 c. I% ?# J '74': 'rose',9 L2 U/ m4 j# u. I  G# G; G
    '89': 'watercress',. E8 V3 f8 X8 ?0 B; ?
    '73': 'water lily',. Z  b0 `9 x, ~3 {9 I1 R+ N
    '46': 'wallflower',
    2 Q3 |& L- I3 F) {' y( o6 W '77': 'passion flower',
    1 [1 P! |( S. j% ~$ v% ] '51': 'petunia'}
    3 `' e# b" q; v6 Q
    # v) E# Q# y2 w11 o( D4 k1 w7 K. {: J  t
    2
    9 R% e7 r4 |! Z; V( ^( M6 {9 l34 ~; m9 Z# A  r9 s0 e
    4
    - a, D: j3 f3 L52 o9 c" `& O4 h, h5 T- g% i
    6
    6 t& p, ]% [9 ~) h9 x7% K# i3 B4 ?* B' s* J* H1 _2 r
    8
    8 s; t1 b  |) G, A* {, V91 O* r. l- w+ q7 J
    104 y! X5 S& {+ ?: Q& c+ n0 i
    11- G; T! A. O! M. W- F
    12
    $ Y; z0 k9 q4 l136 E; U. Q# D$ z1 x" Z% h! t
    14
    * e* Q$ ]7 Q& O8 z' w0 w15% @, J6 F. r- K4 R! a5 S
    16
    ! C3 x6 a2 m, n% N17
    / ]) _7 ?( g& I* m& a9 E2 ^2 f18
    8 m: z2 J, m; r$ u1 k, E* g19* w; M/ n7 H, b9 o1 E: ]) @; u& l7 M
    207 J, p) ~" d1 P7 e. z( L6 A
    21
    ! f( @: ]8 U4 }. ^7 v22
    ! U9 e0 O! F( |) y: ^23
    ! @6 a* O" w; P- }0 J24. ]7 Z4 k+ ?; ]
    25
    + a+ ?3 f3 }7 V- C2 Q% }26+ F' Y( h$ B1 ?) p, T
    27
    & o/ \, O2 e; A1 a2 R% }28
    9 x! T; J: x, v  J29
      `1 @" g) T' t/ u$ z( |30
    4 R# s1 F- q7 J1 N8 \31
    * Q1 b* G! x% @32
    " o. A# f4 U. h6 f* X# B6 m2 P# t33
    ' T" Q2 Q  g/ i) G34
    : b6 e# v8 C: y. W) K35( y( l" h6 F" X+ u
    36
    / L* X4 T5 e, d4 X9 u37
    ) a. x$ ]* N2 w6 q7 t38
    % U1 J/ s* ^4 x! x- s# ?1 p39
    " }% ^5 B% c5 P( d; O' t; v40. \- Z$ S0 T0 o  W8 d! {: b
    41
    6 f5 @! T5 N# h; t# U2 B42; \4 T9 A; G# u: m+ @0 ^5 T
    43- k* |8 r: [3 T# f* D' f
    44
    4 H; u. Y- e  h45* v0 x! o" [* c6 T
    46" H+ u( F( ]& Y) {. X7 ?4 t3 C
    47) x" m5 X. V9 M3 ^) a2 J3 u# K
    48
    . j  Q2 _9 A6 s) u5 D0 F& h498 V$ a  @: @: P1 g9 [
    50& Q# ]1 H: O0 o) N
    51
    + j1 s" V4 A, x5 t0 k  m2 R$ l52" r3 U; t9 g; z- f1 k
    53# b) H- }$ i2 I" T7 l
    54
    $ o, |' H! n4 C1 K' b55, m, ?" F4 Y0 _3 ^, s, D2 [! u
    56
    & K6 E- O+ l& T( m8 Q574 t# S( V7 X3 K. e- V2 ~
    58
    . \0 Z% R2 H9 i$ e59
    * C/ h0 g% W0 V) d60+ G  A5 I7 S7 F2 ^7 M
    61
    # n/ s& r: m. {# S9 t- e2 J& S62
    7 F1 L) D2 K( y3 p$ w$ e0 w# z  C637 N6 X+ T, C9 H# c
    64
    ) d  a9 G7 v$ n' Z: [659 z3 m- K0 u  M: V5 h3 W7 \
    660 \! j# e& J8 m. ?' y% P: d) X; m* B" |
    67, A4 _9 R0 Z% A5 _1 H
    68
    6 R& W: [1 ?" F7 n7 O& Z* H# D69
    / m5 j6 z: {5 o# F; |3 i70  k/ p* V5 Z" z; X
    71
    8 F# Y% T' d* r5 r; O0 X8 b. M72
    . n7 ?: O  f$ `# \9 l8 b5 p73
    + Y8 W1 X7 u& K/ M, V& Z74" j7 }6 E7 J0 E/ r7 z
    75
    8 [/ y6 d+ v3 S! {" B76. R) q* o4 ^' I* [
    77
    # {3 m# E* ?  ~6 a$ p  N7 x# i7 K78
    / Z* n4 f2 S1 G8 q0 |0 \/ A79
    1 i# p* `* s+ [80! d- l; t$ P# ?
    81
    5 ?9 C1 U2 g$ @, t/ k1 [82" X" _% ^  h9 H4 d
    830 L7 l7 q! m9 N  ]% L$ |* m
    848 J- @3 {6 k  b+ p, g) q1 |3 |
    85
    $ I' C  G& O% x6 M, I1 w6 g86" a. o. h' x; @  @, q$ Q
    87# y  w2 j- h9 O& m8 `
    88( q- C0 b6 i9 Z- q: a! }
    89! U8 @. L& N% C& n$ }; c
    90( Y* ?. K7 S! A7 ]( w' |
    91# I+ [8 @1 u3 T! ^% S
    92' m0 |  ^' y( |! G2 i
    93) t4 L3 _4 ?$ @% s/ N, [
    942 ~: Q* S2 g$ M4 O( c
    957 o8 y/ L, b. v( j
    96
    / P- d$ G7 I' y- n3 S# O97
    # C2 M, d) K$ e98
    . C2 G; z" P  V7 b999 a, Q, {! f5 m* _& S5 E
    100
    % p1 \8 Q' I  V, s! Y" O& a101* z# B/ J# x1 N" R% i5 \
    102* S/ _  A  m# Q5 i2 l) ~# R
    4.展示一下数据' g; H4 i1 a/ d5 i
    def im_convert(tensor):
    6 k7 J  p, N( k% Y& d    """数据展示"""3 y. W3 L9 ?/ F% ^' ~" f' `. v  s
        image = tensor.to("cpu").clone().detach()
    , D' r# j' }( x$ q+ b3 ^; V( R. F    image = image.numpy().squeeze()
    : M+ O% k' U& f& ~    # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图3 U& W3 D- u5 ~: r0 {# b7 ]
        # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)
    # Q8 N! }* v* v7 n. X$ b, h! V    image = image.transpose(1, 2, 0)% R! W' @3 u* s1 Q, z6 h9 u
        # 反正则化(反标准化)
    0 m3 ?6 H! |6 ]0 |) W  n4 F. m- i; P( l    image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))
    6 b% i8 U+ ~% L7 z% o
    7 y1 U9 N! t1 _. G1 k    # 将图像中小于0 的都换成0,大于的都变成14 j) x) c3 K# u( g
        image = image.clip(0, 1)) S4 x! _! T, b. q

    3 [+ K! ?9 {6 g  l- a) L    return image: h) c! G" d% |; a. [' v! @7 N- e
    1& M' k) c- f5 n2 J
    2( d$ @/ d' y: V4 K7 i
    3
    % f+ e' ^6 Q) J4
    ) X, F; A; m! }/ y; E58 d# f' x( N3 d) U% R. }6 S
    6; b9 z) H5 E6 l) q7 f, [# _
    7+ V7 S( \! A: L* y
    8: n  D- V! M9 S$ r6 n
    9. l/ W1 r( f* |% O2 S2 Q
    10! V1 J- c  ^, n( |) n
    11& Z& ?" E$ ~! R  c  u! v
    12
    $ J) f& j7 c6 U136 y+ p+ E% N+ h0 S8 o8 S
    14; J" g* a6 n/ H; O4 K
    # 使用上面定义好的类进行画图
    % J3 z" J1 B& l% u* c9 C# nfig = plt.figure(figsize = (20, 12)); v1 l( Y8 o' {! S  V" K2 {
    columns = 4( G$ q: V, l( W( E2 L& Q
    rows = 22 \5 h: M# T, z; T) k8 ?

    * t' N4 y( S# _+ W' h; E* q( F. l# iter迭代器
    , S8 k, S7 w/ k8 `8 W9 }, K) j# 随便找一个Batch数据进行展示6 J3 e8 Z6 F, X/ _: S! B" k$ b1 _
    dataiter = iter(dataloaders['valid'])
    : d0 n9 E5 B. B' Q% d; e( Xinputs, classes = dataiter.next()
    8 Z) e7 d1 c. K6 m7 o% }* [5 L
    4 _, ]2 N- [3 u' Z; a5 w: afor idx in range(columns * rows):; B+ ?4 a8 x& A$ m3 Y5 s. p
        ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])
    : _( o, H2 n0 W6 x    # 利用json文件将其对应花的类型打印在图片中' M" m( ?- p: A  |, W3 u+ }
        ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])
    $ j- O* f8 m6 v- ?0 v    plt.imshow(im_convert(inputs[idx]))% L% W, x1 ~5 l1 `; O) B
    plt.show()
    6 s. ]( }2 v3 t4 }$ g% q1 G$ ^- b. M1 Q$ H' p. o. I' s
    1
    ; r1 e: d+ D; l/ r: m( |2
    - |: e4 B9 ~0 u! E3
    3 B. ~" O+ h  E4 `2 z2 b4 C5 D4% {% K2 E5 J7 ^4 t
    5! \" [  \, ~+ b4 N5 ~
    64 a; F; f, M5 c8 h; Q; N9 S
    7
    . h1 T3 v1 s( ^4 G0 N. k  E8! i# ]1 v& ^7 X
    9
    ( i9 K- _* V* @3 w: [: y10
    % l5 y- E- C' W3 m' M11
    ! H) z2 }5 C' Y  I$ p9 A124 X) @/ A, G0 O+ X1 e
    13
    * S! ?' L& {. ^. N. |2 U1 k0 s14
    / ?- Y( k# g: K; d% R; s0 }) i5 G15
    & _) [2 @. _  q, D16
    ; }, H& @5 R9 Q  t; V+ U) F- y4 O
    1 I3 t9 H+ M4 u- y; C$ n1 {9 u3 d' N9 F8 S/ E5 j* r
    5. 加载models提供的模型,并直接用训练好的权重做初始化参数9 N* o( Y2 `  r1 o
    model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
    ; e  G- l1 A" P) w- ?# 主要的图像识别用resnet来做
    3 f! Q0 ]2 S5 _2 X# 是否用人家训练好的特征& ?! I9 q( U6 a
    feature_extract = True
    " u7 P2 E6 N$ \1
    - e  {! s/ R9 v* W! {2
    0 p& H/ J! f) K! c' U: y4 W3
    , u6 h4 ]/ }; h; y% a/ L, X& w& l0 J4
    1 l' e2 P: Q. F. M" \# 是否用GPU进行训练3 A3 V$ S/ i' N; t& g
    train_on_gpu = torch.cuda.is_available()( N/ V2 A( t+ H9 G+ g) m2 L

    - X: j+ v6 u" Y" Aif not train_on_gpu:  L1 Z9 U, I, z0 f/ ^
        print('CUDA is not available.   Training on CPU ...')
    ( Q0 P! x# {0 g/ }$ Xelse:* K1 p, A7 i, ]6 o& L6 w7 [2 t
        print('CUDA is available! Training on GPU ...')
    4 e6 @' V7 c, e9 _2 B/ h6 ~3 J; o4 {9 c% G
    device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')4 i+ R1 A" R1 M
    1
    . E1 Y' n/ u+ R+ N3 F; ^! }& H$ |2
    - \8 ~' c5 M; H( C2 Y3
    ; B7 u. c, W2 A+ ]4) C9 Y$ v# Y+ c
    51 g4 m8 d" z* V. [. `, @4 K
    6/ J& F$ K) E% g" A: }- I( m
    74 }6 }1 R4 y1 M0 P+ M' Q
    8, m7 y: D( I) a% B0 e* W9 B8 b4 ~
    9
    + O! x- u+ x( p& Y3 t5 W7 UCUDA is not available.   Training on CPU ...8 G+ C, O9 k/ ~( M0 e( X% r
    12 e7 h+ H- M( F# b
    # 将一些层定义为false,使其不自动更新
    1 D8 n$ r2 }# O( V! N) f& odef set_parameter_requires_grad(model, feature_extracting):
    8 s+ O7 i* E' }( h# _    if feature_extracting:
    $ k( Y9 _' A: D        for param in model.parameters():1 n% g+ X7 }8 ^- s2 v9 d3 z  d
                param.requires_grad = False
    * P2 g) g+ s# P* r! i, D) X4 ?1
      K3 W) i* |2 I0 h1 }& z2& D* f  H3 v* c: x* F
    3
    6 ?# ?3 ?* @9 ]+ Z1 \+ ~4
    5 i" d9 H# q( M7 E5& h7 C. Z% f5 L7 p4 B3 F
    # 打印模型架构告知是怎么一步一步去完成的
    ; d/ {$ Z. t, b0 Q8 u# 主要是为我们提取特征的, v" Y+ h5 [1 |  s% Z' D

    0 o$ {# N4 \( Smodel_ft = models.resnet152(), d8 d% y+ B) r) K6 b6 V, \
    model_ft
    6 R" _3 u3 R* Z7 Y  o5 ~' Z1 Q1
    6 P, H; W. I' a4 G9 M' i8 k; d2' m9 i( r- x2 h; y" B) f
    3' M/ E" B" |9 [$ h0 ?0 F0 P
    49 l- J$ M  u; n1 T" O6 W. Y+ T
    59 o" {, S! p& ^/ R: N
    ResNet(
    4 A' \& I( n7 B' S2 o& W/ ^- `  (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)9 s  Z( k. C( C7 M$ v# f
      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)9 F) y% J# V8 K/ K; e3 _2 J) ^$ R$ ~
      (relu): ReLU(inplace=True)
    / p* z! z6 ]/ d/ U5 M  (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
    1 v6 q6 N# \: S& @  (layer1): Sequential(
    ( @. {% d  P& t$ y" G- d( l    (0): Bottleneck(
    ) X; e; F7 j. r; w. ~8 t" P4 s9 f0 Z      (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)! }: Q4 _. A/ P+ G- ~4 b
          (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    , f2 o" X2 T/ Q. a      (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)0 e1 J$ Q6 H  z. s4 u9 J6 w
          (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      V+ x: V8 m: D# p* d1 A. ~/ p      (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)9 I# c9 q5 b1 l. v# x, f
          (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)0 H) z6 ?. x; j" H( X
          (relu): ReLU(inplace=True)  X. }" {! G  S8 T
          (downsample): Sequential(
    + H+ D8 O2 B$ B2 q' u6 c4 m        (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)- d. q* E+ u  m: {' i
            (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    1 r0 o/ F- w1 K. p, P- u      )
    ' A1 M0 M, K# ~( Q: n6 j    )( u( A7 D$ O7 ^6 m
    中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。6 j, M. v' ^! P) X/ x* H/ Q# u
        (2): Bottleneck(/ r# o0 r9 T. ~' |
          (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)( c1 K2 E2 j- E5 M* m% I( q
          (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    ! T* K, a9 T& X* q. G      (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
    ( U6 F1 l- j! v7 s9 A9 ^- o$ A; a      (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    / r2 R4 ~+ h% A6 @! w      (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
    * Y- {- I, L% N) O+ j. g8 Z      (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)( \! T) c: ~% G# h+ Q& G" i" Y
          (relu): ReLU(inplace=True)6 W2 b+ ]4 F2 N
        )
    ' O4 x5 _& ^* I7 X* Q7 e  )7 p" x; \' ]: `
      (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))
    $ K) d& G7 G0 K) {* Z0 u  (fc): Linear(in_features=2048, out_features=1000, bias=True)
    6 g4 u* e5 }" D% ]/ u9 k0 i)
    * k5 G# C  O2 P6 G. _' a# {
    $ o) L7 Z. r, q/ |! A" v/ {1
    ' S( S* b# y5 z2 Y7 [9 z2
    8 t+ W. z& e$ @/ E, E% b1 h- g35 W% @1 J2 [7 q1 a
    4' i- U+ g3 t% a5 L6 \2 H# h5 B
    5
    7 K6 X+ ]6 y5 p1 o" S. j& O69 ]4 b: z) D" W  R
    7
    4 y" C. Y7 k% X  m8
    * f( C# G7 r' p2 n% O" H5 W# N1 y94 I8 b! G) \$ j4 X# X, ?) W
    102 _) Q- A! Y: p3 m$ K
    11
    + V! C2 D! c$ |& w: D12, E" w' i) ~1 Q* t
    13
    : ^9 i2 h  j! y: r: H9 v, `: F; R/ e14: k/ p" o0 K# {& h7 Y& {
    15
    # B) ^4 L9 _) J3 j6 |0 S168 p- R, b7 |9 Q2 D
    17
    ! M, o" n& @3 O( ~18
    . S. y+ h0 D  \) F. B, C( j19
    8 D) Y. D9 J7 h9 o5 N" S; X) l20/ G8 C' A9 v2 V/ l
    21
    ! p$ P4 h1 X- H22
    4 F0 z% q$ I$ r$ z4 A6 o- m23; A8 }  b' r3 u' h2 f( i9 J2 l
    24' J8 u, o, Z) j1 `( Z) i
    25
    # J$ t4 b3 I+ g% ]263 E1 J# M) G$ B1 s! g1 G: `' @
    27
    ' {& B( Y% j, c# @$ W283 @/ ~. u0 ?5 }
    29' ]3 Z3 m  k. D# }: C# B! B: K* [
    305 t* m# _6 c" q0 Q. `6 Z
    31+ w4 v4 _& E2 {# `1 r: C+ @! S1 d
    32  Q5 H" r) D: y1 R
    336 t; O) e; P+ c
    最后是1000分类,2048输入,分为1000个分类/ M' m# V* @6 V6 R% l
    而我们需要将我们的任务进行调整,将1000分类改为102输出
    % z1 ?: \& K9 b5 {2 v: [) m% M& u7 M$ k3 X, o( _
    6.初始化模型架构
    + A& j0 z4 i$ G. v2 u步骤如下:' _' M8 h. ?  _) D) W( x
    1 @$ n9 i8 ^$ m% [* m
    将训练好的模型拿过来,并pre_train = True 得到他人的权重参数7 R+ n: i# |$ L- M- W% G* v
    可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False), r% }* p& [3 o% e/ ~- Z
    无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数) x  x% L. ^! `- r) D
    官方文档链接1 D* j4 r9 K) g0 O
    https://pytorch.org/vision/stable/models.html% ~, E- q8 {( m- Z* A  _6 \" P

    / r3 v) R4 j- e/ O0 R% i  M- _# 将他人的模型加载进来
    & t. E- u" Q- g# g5 O( X# ]1 idef initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):
    * h' D- L# Z7 G6 V* \0 Q5 J    # 选择适合的模型,不同的模型初始化参数不同
    - T# J0 r/ S- k    model_ft = None
    6 a6 U+ E; Q# z1 z; @6 x    input_size = 0
    ) c' h: k- V3 v3 b" a% ?7 H. N2 m) ~$ n# e' O
        if model_name == "resnet":
    3 }! N1 a# B/ v$ @- |        """7 c* ]1 l4 E8 h+ Z* Z
            Resnet152" p( A- c$ q4 F6 I
            """
    : U( l2 a' Z% i& h! D! U! M! i2 c0 _) B( \, o* W
            # 1. 加载与训练网络0 K' c5 ?% x: _
            model_ft = models.resnet152(pretrained = use_pretrained)
    ; h' Z9 G) f- C2 i& T        # 2. 是否将提取特征的模块冻住,只训练FC层
    3 Y7 T# u1 W  G7 B7 j        set_parameter_requires_grad(model_ft, feature_extract). F. B+ Z( j. f% K0 p$ j8 J
            # 3. 获得全连接层输入特征2 B0 w: P6 `- Q  ~* z" p8 W
            num_frts = model_ft.fc.in_features
    3 z# b# }; ~. s2 ~        # 4. 重新加载全连接层,设置输出102% A6 S, _6 w0 f
            model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),8 y" u; J& }8 j2 m2 n
                                       nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1+ A6 c! n7 ]7 \& P/ T5 x
            input_size = 224" a) B; d' {( a' u: C

    $ |! U  I  R5 }3 K8 |    elif model_name == "alexnet":% f! Y4 `6 u$ s3 p3 r+ y4 ?* n
            """
    , t; n9 r$ l, L" z) D* Q( s9 _        Alexnet7 \1 M. v, k9 T+ V+ C
            """
    ! n& F( Q& t- Q' x7 x) c8 H$ L. d        model_ft = models.alexnet(pretrained = use_pretrained)& x" n# D2 b5 k# a$ N; z
            set_parameter_requires_grad(model_ft, feature_extract)
    - ~* {3 V4 h4 A& ?, j
    ' o- O/ s3 N2 }+ r% j4 t        # 将最后一个特征输出替换 序号为【6】的分类器
    9 m4 I% ^$ y5 K* i" L! ~. c1 `  A        num_frts = model_ft.classifier[6].in_features # 获得FC层输入
    + K' l7 c! _9 N2 M4 A4 w5 S! u1 L        model_ft.classifier[6] = nn.Linear(num_frts, num_classes)& Z5 f/ R+ o8 c* ?7 Z8 a" u
            input_size = 224
    1 E* g& [, c0 p, c6 ^0 n/ m8 K
    . S  I* z! E. ^: V: f* u7 S    elif model_name == "vgg":
    ! I) ~# V7 W0 Z* W' z: |        """
    / t6 O# P4 B# K! d1 a! h        VGG11_bn7 Z& V1 |3 s+ m3 ^/ p1 c) i/ w
            """
    9 X- d7 v. e: U9 r% U# Z, R- u        model_ft = models.vgg16(pretrained = use_pretrained)
    # p2 [0 D1 H2 r2 w/ v6 J( R. l        set_parameter_requires_grad(model_ft, feature_extract)
    # d; t! l1 P. X& g9 d2 [# A, r        num_frts = model_ft.classifier[6].in_features4 X( t1 }/ Q; z& {: k3 L; h
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
    - X* a5 ^+ v3 j        input_size = 224
    3 v, d4 `  u9 x+ \3 T7 ~1 w7 `3 n" r' U* E# W6 L6 A& Y1 G! c
        elif model_name == "squeezenet":1 f4 w2 M5 m; l& X% ~1 H5 E" z8 x
            """
    5 i% ~) M0 m/ }        Squeezenet
    ; ?; P, q2 d/ i" N& F7 {  R+ o        """& Y- k. R; F  P: K3 b
            model_ft = models.squeezenet1_0(pretrained = use_pretrained)
    % f8 l  A) `5 n$ y        set_parameter_requires_grad(model_ft, feature_extract)
    ; ?! I. h! l) O) J! a* `" C        model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
    $ J  D& q- z9 \) S2 y  U: `8 {        model_ft.num_classes = num_classes9 M: U: I& t2 Z7 h3 U
            input_size = 224
    " J! w, ]' f) s9 {5 s
    ! U% Y' q; `2 z0 J    elif model_name == "densenet":: R* f8 t- y+ U; P
            """3 X8 p  Y; l" T& p
            Densenet& |6 ]1 K1 T0 B
            """2 a. o# z  z- z2 c' y
            model_ft = models.desenet121(pretrained = use_pretrained)
    , z+ }% d$ V9 }' u7 l% }1 A; N! k        set_parameter_requires_grad(model_ft, feature_extract)
    ' G3 `, W3 \# Q) O& A: A        num_frts = model_ft.classifier.in_features4 A  I+ ^3 i1 N$ S- M/ y6 f
            model_ft.classifier = nn.Linear(num_frts, num_classes): C# N$ h5 r1 }3 d* R
            input_size = 2241 J9 p4 C: G. y) ]: K- X

    5 G. B' e6 I4 b* p1 d" ^/ Y    elif model_name == "inception":! G; {" d* F8 ]% \6 X; w$ q
            """
    9 G; ?! @5 W- y* P        Inception V39 B& H$ b( j2 _3 u
            """- @! q- Z' m. n
            model_ft = models.inception_V(pretrained = use_pretrained)
    ! I6 [* X1 r( ]. V' U/ ~% F! x        set_parameter_requires_grad(model_ft, feature_extract)7 q/ v2 b# L! P/ M9 n

    & S1 d0 E" l2 f! I" h; l% V& E: ?        num_frts = model_ft.AuxLogits.fc.in_features
    ' V" T5 t' ^) Q4 @; p! X: H        model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)
    & L5 V: L) e5 j  K+ t" S" ^( e4 ^- Q- w- `5 ?2 m6 Z. V* G
            num_frts = model_ft.fc.in_features
    + a( ]$ I2 y: c) a% W        model_ft.fc = nn.Linear(num_frts, num_classes)  ]7 [0 B# \+ x- `! d- ^" o7 Z
            input_size = 299( J% f0 l: T: o

    ' A0 n4 M; q1 P& S0 e1 i/ V+ I( o    else:  K& |4 }  U1 D8 g: Q  r9 P( f
            print("Invalid model name, exiting...")/ @" A0 S  ^/ i9 a# v$ z9 A+ U) C
            exit()
    5 k+ d- C9 f, H) F! p) z# c; f/ W6 B- ~& ]0 q" Q
        return model_ft, input_size
    - @1 W8 Y1 c) _9 K+ b
    - [: V3 m. l! D* F- ]- A/ O6 y1( q$ a  m% I/ q
    2- }, A# N& U/ j0 W( t4 h+ f
    3" M# @  P7 U  ]8 y* v9 W8 [( y
    4. u, K, z, h" C$ }$ L
    54 U# m3 V+ z) @: G- V* t# W
    6
    + ^9 l& f: X$ ~77 _5 n0 {2 L' a: n2 _( b
    8
    ) P# a6 t) y0 J/ g  X6 q1 s7 D! X8 U9" U! b5 w# b. _9 T  r! ^& B
    10
    * J: X/ ~, c1 [( a11& ~! k) E; q4 `& b* s$ \; l# Y  h
    125 j) _* m0 e8 R( }' ^) ^. p/ B
    136 P$ n- M; R7 o5 j! A8 ^' f' v
    14
    8 Z; {. a6 O4 z; H15( J6 D0 A: Q; g4 T" B: ]& p# p
    16  Q. b: w8 T4 a  T$ Y' R8 A
    17
    8 m% ^* J4 c9 R: L5 J180 w& b. m7 b8 X* d1 r: i# p
    19) O; x2 u3 ?* X- E% R
    20, |6 r+ L. Q, W
    21
    ; k# _2 }* M& b" O( z228 L* T/ c) o1 t; ]! ?
    23
    & G* x# z- R8 v) A+ O; s+ ^3 H24
    6 n9 j. f2 {" e# c# a25
    1 ^; c4 Z: A/ w+ E26
    2 i9 B% _4 o8 R3 l3 M273 ^3 M; m; K1 I1 @
    281 @- d* [+ O! q! p/ F* }, T
    29& B2 M3 q# t7 P0 k: W  s
    30
    3 P5 m8 x; i5 K/ s- `& N31/ c$ b" P( ^8 @/ W* ^* t
    32' j8 l, Y) V# N
    338 b; p3 x5 n$ W6 [( I! M
    347 e2 k# r8 D& m! E6 i" t
    35
    0 _4 P& G) o4 q+ E36% ^. t+ N8 A+ o1 g. z0 f4 a' ~
    37
    4 f% q# P6 N) P  B38
    ( y5 r; |+ |" i* t. y3 N39% y5 L# k+ v, U) _
    40: s1 Q" ]+ s; P! O$ `  y5 |8 T4 J$ W
    41
    8 p, ^( \5 d- a2 O42
    , e! j5 N& s3 n9 O43
    . }/ H3 R4 C6 w' n0 N44# y; u: V! L0 s$ \# M
    450 T2 K+ Y9 {+ u. F
    468 R" }" E% W, y$ [) O9 ]# x0 X7 v0 q
    47
    . }/ I4 j3 g' O% s# v8 c- T  G48& N0 p. j! P. j& W: Z& ]. p
    49
    ! A6 n) V, k: H1 U- h( ?50
    0 b* @  G5 |1 a1 @% [+ x6 Y51
    " B$ q+ L5 q$ X1 O  i52
    ; R: S% M. T$ O# e53
    8 D9 M, t( e( \+ x54* l5 Z( ^* ]' |/ p4 N
    55
    ( T# D# Q% l5 C5 h  C% v8 ?561 O) ?; D7 v  j5 c3 C3 h* m; N
    57
    ' N1 _" Z0 y1 N$ v1 X58
    8 ]4 Y( a1 C  m+ O' O59# V. m2 C3 q" d
    604 O/ W2 B! F6 p4 c3 t
    61
    - Y% @0 B0 ]+ Q0 K2 }623 @( D7 o- y( V6 j
    63% O0 A: v& _. D
    64
      S, K. s+ I, u65+ a6 q0 I8 g" A7 G# D
    66
    - ^! ~1 B4 ?9 F3 G# i0 [$ i67
    % W. `+ k& L; b, ]. q1 D683 t2 W  q, L( x8 y
    69) x' C3 ^* Z$ X9 ]$ r
    70
    6 Y& z; k  }2 K" s) O( h& y71  n; N9 u3 ~* l4 d! D. ~
    72
    0 @# \& C. {$ k1 ]73
    ! p" {. _" g% l* \0 D8 P74& b7 X& q4 p! `% N" a# z& L9 G
    75$ r7 x; w0 |" j$ [# _7 M4 c
    76
    " b( K0 ~; N) c( s4 C" D+ H77/ u- L1 k2 X* k3 q# s5 E9 t
    786 d  @& Y1 Z+ a7 }: K
    79
    % I# \) ]' L) q5 ~0 K5 X. r" n7 P806 n7 ^" c! s3 p
    81
    8 J; g% p  h% h8 z7 R6 d( V82
    ! }: f7 y; K3 W. y9 }# |; B3 U83
    $ H. F/ ]" L, B* s) v* G- D7. 设置需要训练的参数
    ' e' g5 Z  Q6 f1 T# 设置模型名字、输出分类数
    ( C9 P$ ?: i9 c: `, X" Umodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)
    % X7 y1 S, Y9 `7 K0 H5 @% @  y- W5 i; s. U, Z0 M1 d! ?* i
    # GPU 计算$ U4 i0 m4 p  [% n- y- y9 S
    model_ft = model_ft.to(device)
    ( w* R8 {0 J3 P: b- A. w
    6 ^) R9 m0 A; e* v9 c# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取
    # j3 y$ Q& K7 u( pfilename = 'checkpoint.pth'
    2 x  R2 x, \+ L' H1 t6 Z  L2 D2 ]3 n5 E4 V( x0 [9 H3 [  R( X2 L8 n" ]3 Z! h
    # 是否训练所有层5 Q  N) _# P- ?* A
    params_to_update = model_ft.parameters()
    ) x  ~6 Q7 B2 f- q# 打印出需要训练的层
    5 C" N0 @/ Q8 K# B& @print("Params to learn:")! ~! y9 F  P& E
    if feature_extract:: ^( c, Y) U. X. T7 z
        params_to_update = []
    0 P0 Y& w% c0 N6 L7 {# \/ x    for name, param in model_ft.named_parameters():6 N+ |. H0 x4 H8 m
            if param.requires_grad == True:
    - u& G: S- b- V  Y- D& m4 r/ I            params_to_update.append(param)! b* D% B* l5 `7 e
                print("\t", name)8 D5 c# W' c3 O4 h
    else:" X2 f1 c8 k" p$ \4 j. d/ \2 O
        for name, param in model_ft.named_parameters():, C' j, i: t8 K8 T/ y- [
            if param.requires_grad ==True:
    ' }! ]. |% ?" e! U$ [            print("\t", name)
    / S: [2 {6 B5 U* [4 X% ]  E
    5 b/ C' {4 c5 I3 T7 i) ^6 \7 o9 d4 W1& f) h$ g8 f5 s7 o/ L
    2
    7 w4 o2 G' I0 K; ^3
    ! N$ A7 f0 x) Y" G0 D+ o  @4
    $ V9 f' r1 P" A* v# K0 v) P5
    % H. }# H* u' d- Z. ]) h; I3 v68 f2 j' Y( J4 @; h
    7
    ' B6 b7 C" M" v5 a+ c+ `. E8  {/ V9 `- L% c- L! F; g. G  n1 X7 Q
    9
    1 \$ p" N9 q3 m5 f10
    0 o( f& E. D) z' T1 p) _11# L' G  Q6 L* f' t' W, G: y7 Y+ K
    12
    6 a# T- T( A; ?6 y# S/ V  Y  n13  \% T4 T2 _4 r6 e2 X- `9 v
    14* |1 D  A: `2 f1 d4 H! ^8 C' m" M
    15
    , C; U. v9 S- n: k3 K$ z161 [3 Q. O" d2 y. s
    17
    0 D9 e: y( _% p3 L, p3 f# p* z% r; m18
    2 T' V4 A$ v( q19
    ( I3 p( m6 C5 e8 ?20  q4 W8 A! ?7 f$ b! K/ b
    212 [: e% ~2 e& I6 K7 t
    22
    3 L' L; N2 |2 P; U0 f23" T- o' ^8 L9 K: m  F! U
    Params to learn:" z3 M& w& i) f" M: i+ }( _
             fc.0.weight6 _- j8 Z4 C( X2 M  y' t* ~( U. E
             fc.0.bias6 ^7 [# S- H$ }8 ]
    1
      l$ E5 W5 G# G* x22 v6 h$ C) i: @7 f, @' F2 C. B
    37 D" e) x1 p5 F8 \6 H
    7. 训练与预测
    * A! |5 v, ^. e7.1 优化器设置
    0 E1 \( B: l6 X+ j; i2 @3 ^# 优化器设置
    % H: d4 s  A; Q& O. g$ W0 @! soptimizer_ft  = optim.Adam(params_to_update, lr = 1e-2)7 s% V: Z5 g3 h: k
    # 学习率衰减策略
    $ ]/ e' M* j# g( Z) U" m8 lscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)
    ) N* `: Z% I3 n: q. h0 @. A# 学习率每7个epoch衰减为原来的1/10
    2 B7 L, s4 Q7 r4 r3 x( ?) f+ u7 N" ^# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算+ n% B+ o) G" y" ~5 C5 O- h
    % z0 v! N' `2 e# o
    criterion = nn.NLLLoss()5 ^& J" |% J2 V! t3 H% Z
    1# x( q9 F) n- B# X  M
    2
    7 a$ B! W  J; J# T6 y3
    - V! T& q% Q8 e) a/ }46 K3 W' J, L8 e5 b, a7 B
    5
    # Y) `& p9 k' p) R6
    # m' C( [) G% e  o" g" B7
    6 z  Z( j7 ?- k/ g- z8
    7 T) ?6 s. A( X( ]8 b# 定义训练函数
    2 K9 N1 v) ~) y#is_inception:要不要用其他的网络
    : t  @9 z3 J5 bdef train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):
    , l' i8 h: l& G# M) _- m: u    since = time.time()+ L5 z1 A9 T) k( @- O( e
        #保存最好的准确率: \8 U) S- `9 t( x; m; H6 h$ E2 K
        best_acc = 0
    , J' ~( G3 M: [; k    """
    1 {# Q; W6 `7 L3 [8 f2 B' H    checkpoint = torch.load(filename)- a4 @8 z" h- e
        best_acc = checkpoint['best_acc']
    , Q9 U1 n) n4 ^  F# ^8 n9 D    model.load_state_dict(checkpoint['state_dict'])* n5 v2 b( I; `" N- z+ M' \
        optimizer.load_state_dict(checkpoint['optimizer'])
    ! x. ^4 N$ ^8 p6 W4 b& M    model.class_to_idx = checkpoint['mapping']
    + O4 n2 Z0 F" G    """
    + x1 D- P+ z& Y; x0 g5 F1 D" V/ I6 i    #指定用GPU还是CPU: I/ W0 f8 Q6 k
        model.to(device)# b: V8 t  {/ ]4 r: ?: t( h" A
        #下面是为展示做的, o5 n( c4 I/ p1 v1 ?
        val_acc_history = []
    ) s: f: p+ N7 ?! {* e/ D6 B6 _5 B    train_acc_history = []$ W" ]" r# j6 ]8 y# i8 b0 E
        train_losses = []& M3 ]  S2 x* I$ `1 k6 P; T
        valid_losses = []
    9 {  x4 i1 B( H. L( _, b% `- R, W    LRs = [optimizer.param_groups[0]['lr']]
    ; V. p& f! {: A1 U    #最好的一次存下来4 S  h, J, \7 z( J; a
        best_model_wts = copy.deepcopy(model.state_dict())' P2 E  `7 X; \5 I5 ~
    $ k5 s6 t+ @4 {
        for epoch in range(num_epochs):
    - z1 x4 r, @+ t5 D; {! z3 A        print('Epoch {}/{}'.format(epoch, num_epochs - 1))
    * H& V+ U8 F. f# o) d, g        print('-' * 10)
    , e0 s5 y# ?" p. \: d: @* k
    6 Q. V  g5 d% c/ w        # 训练和验证# C% t+ m5 ~* y: q3 j  r
            for phase in ['train', 'valid']:& B" }# m2 [9 z9 r) ~# H+ s2 P
                if phase == 'train':" N0 B' r1 y# T- U2 m# w) w
                    model.train()  # 训练
    7 ~, j/ f- l( m            else:; O/ L; B+ b+ C  j
                    model.eval()   # 验证
    ! M$ o9 u7 J2 W' Q2 c3 J
    2 z1 d! Z9 X6 Y+ y* X            running_loss = 0.0
    7 w  L! @2 J0 n7 e            running_corrects = 0
    5 w  q( j  [, T8 D# w! Q+ ^+ x6 M7 @3 E) @' x- h
                # 把数据都取个遍
    " T, n+ m0 J" Z- t5 w            for inputs, labels in dataloaders[phase]:
    7 l7 S7 d8 v8 c, k                #下面是将inputs,labels传到GPU
    $ Y/ J9 ?% U* ]/ m+ B# v2 q                inputs = inputs.to(device)4 s$ u- L8 @+ [' h% Q
                    labels = labels.to(device)
    " q2 r: \2 z( {' Z0 w& }+ T: K1 y6 o
                    # 清零
    % p) |; M9 S8 n& J  g; K5 [                optimizer.zero_grad()# j8 w+ D0 R3 ?/ g9 l1 s* |- h
                    # 只有训练的时候计算和更新梯度
    / g( A. v4 o1 ?% i6 \( z                with torch.set_grad_enabled(phase == 'train'):% ]0 B5 g/ A! G
                        #if这面不需要计算,可忽略
    3 n8 Y5 x; l2 D, C9 K8 w: k                    if is_inception and phase == 'train':: ~6 t) |3 H1 y7 L8 F* A) J7 g
                            outputs, aux_outputs = model(inputs)$ _. m( B2 R4 i
                            loss1 = criterion(outputs, labels)
    ( g; \. G/ A9 ~                        loss2 = criterion(aux_outputs, labels). X0 D. A, h5 I  ?
                            loss = loss1 + 0.4*loss2$ ~7 m5 h) P/ b. a$ K% k# r
                        else:#resnet执行的是这里' D: i9 \2 J1 U3 }. c9 B1 ~# u. x
                            outputs = model(inputs). T8 N/ ^8 l- O8 Z3 a2 e6 V
                            loss = criterion(outputs, labels)
    ' u% r: P7 Z  x* |# G% v: x/ |) t$ G( f; G3 w6 e0 y" ?* {' L3 C
                            #概率最大的返回preds" n8 `/ x6 J- t  w
                        _, preds = torch.max(outputs, 1)
    4 V, A# Q( l" V% c, l& ?" S
    " O3 z2 F( @. t+ J                    # 训练阶段更新权重6 M9 ?+ D5 x/ E# @
                        if phase == 'train':
    : U" Q/ I9 ^) P                        loss.backward()' ^, k  h; g% g% U
                            optimizer.step()
    5 u( Q0 U! Y/ e$ B
    + \9 @, C( `0 g: o' S                # 计算损失
    / G: k7 C* v8 n2 Q+ t' z! E& W                running_loss += loss.item() * inputs.size(0)
    # }1 b+ e: {% X4 ?( B  a' G                running_corrects += torch.sum(preds == labels.data)
    9 T! p; Z- U1 I6 {8 A- t! R7 ?5 @/ _1 ], j$ p
                #打印操作
    7 o1 Y" J" T' z- d9 p7 q5 @; l            epoch_loss = running_loss / len(dataloaders[phase].dataset)
    $ p6 `' P" Q: G& L* X$ P3 ]            epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)
    " ]! k$ O" x$ X. t1 R% k0 t7 E9 _6 f

    1 ?3 @% g2 A" m. {            time_elapsed = time.time() - since
    # K0 \, a6 t1 E8 m( d; m            print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
    , s  |( t; t* P3 }$ t            print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))% [( a* ^$ P; W0 o/ X' f

    : r# @8 @3 k$ c6 D/ ~1 X7 o# A6 I' @" V; m: w
                # 得到最好那次的模型) s6 Z/ [& I5 m0 c) N( k
                if phase == 'valid' and epoch_acc > best_acc:
    - h" T: _; p; |  O9 N- r                best_acc = epoch_acc
    8 z8 A" M* G" X3 ]8 k8 m( T                #模型保存
    , ?5 q" b% V: `; d5 X                best_model_wts = copy.deepcopy(model.state_dict())
    9 H8 i5 `- Z0 d                state = {
    . K4 O" v$ [9 m! x2 \                    #tate_dict变量存放训练过程中需要学习的权重和偏执系数
    9 Y2 |4 ]! X8 [5 j                  'state_dict': model.state_dict(),' a  }( \2 B; s4 j7 \
                      'best_acc': best_acc,
    2 T; o9 s1 j* m2 ^3 L* @                  'optimizer' : optimizer.state_dict(),1 o* I6 z, @& I% y! e( [
                    }
    " X$ J! A" n& m- Q1 o                torch.save(state, filename)7 P5 F% v9 L% |1 A% {* D5 t+ D
                if phase == 'valid':
    ' G" k. X9 ~4 ~# L# V* `, F                val_acc_history.append(epoch_acc). K5 M6 s8 n( H/ Z7 v8 B2 I: n
                    valid_losses.append(epoch_loss)$ u- C  ^% r  C. y
                    scheduler.step(epoch_loss)( J- \; x7 ]/ O$ G
                if phase == 'train':' w" j4 i! F7 g4 [! |8 U
                    train_acc_history.append(epoch_acc)( k9 p9 H2 I# u
                    train_losses.append(epoch_loss)
    5 \  t7 R0 M: z# h4 `8 _# w! `- t) N. G" x% J0 q
            print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))
    : H) r* H3 r: R. @  [, Z0 _        LRs.append(optimizer.param_groups[0]['lr'])7 n* V$ {( h6 H2 X
            print()$ K% k" _! L/ Q3 S

    . Y. g( A7 c( U4 L" W! H+ P1 }# I- s+ |    time_elapsed = time.time() - since
    7 i% B. S. g( w% [" f1 Z    print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
    5 s3 X, b/ K7 e5 A) [1 c4 f    print('Best val Acc: {:4f}'.format(best_acc))
    9 s( {2 T3 a8 o: T  H# J! b
      ^: q8 x( c" }8 _* S& S    # 保存训练完后用最好的一次当做模型最终的结果
    ( D  S" J0 n" y1 O# T& ~    model.load_state_dict(best_model_wts)! ?+ S, T% |0 D- r% W9 {
        return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs
    5 d2 d5 b3 n7 v" ^9 L
    : N8 L0 m& r, G
    ! ]# k  j5 c1 o' {! z: F19 H' c3 g7 S- p- @2 l9 X
    2
    : E* S+ d7 v# k. U- }3
    3 k* L( m. n0 w5 s4 K. i4 V4
    + q: @# ?4 r$ k- Y, O5
    " z2 X  D. U: n* ~* @6. |; M- I; c% o, W- d2 \
    7
    9 j( a! |  S8 G& i- s8
    ( G5 _0 E' H& L: D& A+ I9
    6 G3 B* [$ c/ }10, x0 _' A. X' O" r, V6 F8 q6 ]
    11
    0 v  P# C% E2 B: I  U12/ F5 k$ u8 U( h/ W/ e3 b, \
    13
    ) M  s- W) j) \2 T3 K& f  H14* j. `+ N0 a7 B4 r, [& ^
    15
    + G1 D& N& i$ D/ T" n; @16
    : m% N9 l' |$ v1 n171 S5 V9 _4 C) K& e  f
    18
    , R, I% }$ d& y5 n. R19
    0 b& G. n; H6 {! U203 i. r% [% ?$ q/ @% I" O
    218 \& t* n% s1 q
    22
    ! G7 @8 {# r* @  M8 R9 {1 I1 V; A23
    ! K# D0 ]1 J, l* N- C9 o8 |1 v24" C# q. y1 x7 j
    256 W/ T& Y; W% K( M2 G5 b" V
    263 f1 [# L; A. U& |4 P
    27' W+ f" V5 s- ~( ^5 }
    28
    % s. m/ \" r' S& n7 j3 |$ a29
    ' b' i& b0 |! L4 x/ w* q30' H6 G+ q8 |4 ?+ R; A% x
    31) o, M0 c. ]8 Q; G' E6 F: F
    329 p! ?- A5 e' E' g& e8 N
    33
    7 G+ C; X5 C; b  B5 }34) v4 z5 H4 W/ g5 C7 C) p& {, ?8 e
    353 h) l( c3 f% E) E
    36
    $ t; G' G6 j$ G3 x1 k+ _+ [37( ^; j6 p7 F7 Y% f/ \
    38* ?: L4 e4 T8 F4 P  s1 p
    395 ]0 w4 u+ Z' {0 }4 m
    40& M6 }2 W1 e8 @- V* ?
    41
    ( ?" p2 Z) j* P4 H5 _. Q+ S; N42
    # k1 N; K, [& G3 i+ s43
    * I4 Z8 m0 o: j- @! ~( ~448 j4 H8 J+ q" V9 B
    45" r7 s0 d' U0 n3 Q! e5 Z7 \
    46
    5 c# C: |! u7 z9 b" ~+ p47) d" K. K; \* E) J$ t: l, m
    48
    - c0 V0 |% F% P' J5 h49
    5 C6 a5 N+ A: \( R" o) a50
    4 ~0 O, w0 x2 B8 V# E51
    6 G0 i4 u4 [" J* r7 s4 X4 @: A0 \52
    ! V$ q  ?" f% R53
    + e: l! s! v" c- O/ z8 H54& h7 T$ _- P( d9 W2 D
    55/ Y% }: l+ {: \. |! E6 f+ s
    56
    * S! Z% e% k6 ^- v0 A57
    : X+ _2 F6 X) ?. T& ?1 g581 M; S6 E; o* Q- }$ \) s: B
    59$ p2 n0 z, R' {' f2 y
    602 c0 N; V, [. I6 K
    61& S; B3 V" H2 \/ `  }* F
    627 n' f& P* ~( v
    63
    / g& s1 a  }4 o+ D64
    5 e9 `/ |9 T5 V$ K3 g' C65
    2 D$ g* a# U8 P0 Q/ w66
    * i/ @" x! e8 I0 ^' \" K" E67) z5 J8 W0 A0 j+ R8 x- l
    68: K! P: G% j0 z
    69
    ' S" {/ ?7 U2 n- k703 y+ |& D* a: Y+ v8 _/ W
    71
    , s: O2 Y( E( [5 \2 D2 D. X% J' ]72
    8 `* e. j7 ~% O7 z5 I73
    2 ^0 s. W" x1 U  g: }74
    & _; G7 t5 c3 x, J+ z75- m( [# b) I1 o. I: v6 c4 f$ Q
    76
    " O2 I! x2 a# @9 m0 Q77
    4 |0 [0 J  M+ F, A- N787 P: r' p! L7 @; a
    798 ?% r% }, U' U2 c
    80& X# E" w' l3 p
    819 @! p3 l* ~  ]/ X) y/ r% q/ _. R0 w! \
    82& ^8 d. m. B& r/ A  @8 F) @
    83
    4 ], Y( h- @5 e) T$ ?# Y840 N9 }  J: R* T0 K) S, L' ?" c
    85+ _' A# D8 i3 Q, G
    86% B/ Q  K+ |9 }  l9 X7 X
    87
    ' O; j' c, u  M# n* q+ Q! O88! {" U8 p8 w( ?, r
    89
    ) U1 w% q# t1 {* B( H90
    : I+ P9 B* M5 i  @912 Z5 Z2 q; K( C  {2 K5 g2 I! |
    92  A/ ]$ L9 a6 Q8 e7 M* \9 f
    93: h* \/ U( j5 y* D3 X! Z
    94  u1 H9 V! Z7 b( t
    95
    * N, c, a, A7 v0 {4 D2 }( P1 e96
    ( d% y& g0 `) Q( ~0 {) r97& b3 y) r4 S% R/ u) I
    98
    3 L; q- H: X7 f7 K2 g99! B+ p  K) z  c. L- Y
    100
    + x! H4 f% n' J4 m' x101
    7 w; d, M8 L( f" a102" f" Q+ a3 i( a" `5 F, u
    103
    ( Q& _" G2 U) T! o2 `. D/ A104
    * r$ y3 T3 ~' Z- Z; i8 j105. U. ~' S6 o4 Q* I; g
    106
    7 t3 s2 E- ~8 y4 U: I  ~4 |107$ m: j2 c0 f  m0 ?8 ^
    108
    . U; W2 h9 x. L0 M: y! F% P- t1092 h" F6 c! V9 A* b; d
    110
    6 j3 \& W/ a" P* l) E& w$ c) \111) }6 h/ Z% }5 q4 }5 x# V6 r2 K
    1121 |7 Q0 B# k8 d. Z" D8 q5 ~& J
    7.2 开始训练模型
    & o: I& Y/ @( F' R! p" S我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次
    % B/ e  p! F1 m7 b- z( c' j5 C2 }
    8 b$ d# J. A6 b0 V( ^#若太慢,把epoch调低,迭代50次可能好些
    " q# j# R8 J1 K% @8 f0 Z#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合
    8 D9 ]% y" X% `, o* d0 Jmodel_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"))
    6 J' M3 G, P0 V2 v6 @; W1 p8 e/ R5 M7 T# S) I4 a5 U6 H
    10 S! i7 @0 E2 r* R) ]
    2
    ' Y! c5 U1 f  b( w, c7 P/ O3
    7 P& e  a- g0 a6 u( m47 a5 Q/ l8 g7 R* a
    Epoch 0/4
    % w& ?. l: ?5 J/ Q, D& [( {----------
      z/ I, S+ q- I4 ^: aTime elapsed 29m 41s  }+ u# V3 U  U6 u8 K
    train Loss: 10.4774 Acc: 0.3147, i4 U# x: {) G: y  c
    Time elapsed 32m 54s+ Q! W& c! e/ }" C/ M$ S; ]
    valid Loss: 8.2902 Acc: 0.4719
    2 T# b! o' N  }0 H& SOptimizer learning rate : 0.0010000
    8 O8 C; c, o1 f5 }  L! ~8 K8 t" e' y
    Epoch 1/4/ ?" @- x1 e1 G% ?
    ----------% |: L' S" q+ Q2 B- |. N
    Time elapsed 60m 11s
    3 r0 ~. a3 T( O5 r4 Vtrain Loss: 2.3126 Acc: 0.7053+ n9 k9 A0 z, m7 ?7 A; Z# l
    Time elapsed 63m 16s
    - l: S- Z+ B9 f" |9 ~valid Loss: 3.2325 Acc: 0.6626* `2 R# h) s) c! V8 I
    Optimizer learning rate : 0.01000004 _* s  |/ j$ o- I4 Q# T2 q1 a

    4 w+ ]" n, y, ?3 ^6 @  g9 E$ J; NEpoch 2/4; Q; V7 Y* J% {8 H& M* l6 @% K
    ----------
    " q5 O( ~/ }  h5 o+ C$ XTime elapsed 90m 58s
    5 `8 }7 d$ V6 _( Y+ B9 Otrain Loss: 9.9720 Acc: 0.4734: C" Y) \* S6 g
    Time elapsed 94m 4s% F- Y6 x% G! K7 X) n4 @+ l2 h
    valid Loss: 14.0426 Acc: 0.4413
    & h# c2 @  t* K! g- L; COptimizer learning rate : 0.0001000
    2 \& h* @9 n, I$ Y: g6 m2 ^, t
    % B; }* f- U( f" M' y) ]Epoch 3/4
    4 O8 i' x7 ^4 O" ?( [- A----------; ~; Q- N  |2 J
    Time elapsed 132m 49s9 v3 g2 ?. ?: `4 W. g( H  Q
    train Loss: 5.4290 Acc: 0.65489 h/ m7 B/ A8 W1 ~# Q
    Time elapsed 138m 49s
    ! G8 J( H$ y, \! M; @valid Loss: 6.4208 Acc: 0.60272 ]- ]' Z9 s0 ^# X0 Y
    Optimizer learning rate : 0.0100000
    / K( ]7 P. A* [8 |7 u
    : \  f7 f! `' ?% v. C- h* kEpoch 4/4/ \8 i' j  y0 g
    ----------
    ( d4 e# ?) ~, i' x1 zTime elapsed 195m 56s
    - h" m) q$ ?/ v' ]' ltrain Loss: 8.8911 Acc: 0.5519
    - s6 H9 m2 W: ^) I, kTime elapsed 199m 16s
    9 U1 a4 B1 i0 |4 P2 svalid Loss: 13.2221 Acc: 0.4914
    # K) N4 n. H0 d0 j, j* COptimizer learning rate : 0.00100003 V# N4 D/ h( C* W6 @
      V9 c8 m, F+ I! c# ^
    Training complete in 199m 16s
    0 E  I- c4 o* w/ X+ n3 zBest val Acc: 0.662592
    5 X* E2 N1 R) X) b3 |& W  o) r8 `" A  Z7 n+ q
    1% A( a# i* d$ ~( h
    2
    1 A. X' O+ r. w38 L2 h9 R  ]+ o5 D
    4
    ; ^# X; L% d! _0 G6 |5
    7 l0 M- B" {! {* `" U, g3 f6
    . |. J- ]% u# t7 U, Z7
    1 T" x5 }3 A% B( E8
    8 |( x2 a7 _9 J. y- ?6 S9
    2 W7 R: n0 K/ k109 g; ?' j, U; `% N
    11  c* [: A$ O( l9 x
    12! J# |( s3 i7 j- h" ]
    13- R: P( h4 Z6 n1 F
    14
    8 I4 Q- v' O0 s' h. ?% G$ Y15  R, I- c& r% i+ P' _
    16, M* \0 @1 @# D% q* h
    17
    & m1 R; W  T+ h& E9 {1 `18; U' D2 D! C' l+ \4 {
    19+ X# `# e/ _4 ^  L0 O" t
    20
    4 m" B8 J( H7 a4 p) f21& d6 Y, M/ z5 Y" ^
    22& D4 j$ l8 l# b- F
    23  P5 b$ [+ L: l  J0 \& S
    24
    1 v0 f9 ]/ I9 e1 K( U$ c25$ w1 ]+ l- O* X  Z3 n! q0 B
    261 k3 c* Y5 `4 `' }; `. k' q
    27
    / \  `8 i, O) O8 H3 t6 L28+ a6 v# K7 Q' G0 @% b. P3 f9 m/ S
    29% U" u; ?( W$ O# V. A5 R" N& ~4 l
    30: x- M; z; R- e4 a
    31
    . [& ]2 u: m) Z* z) C( K9 }  F32
    0 I) Z  X& ]% ]4 F+ o& v333 X! W# d7 ~( U- ~  Q2 ^( h4 V6 k
    34
    5 N% k6 O3 J9 [  t35' J" Y! s5 ~6 M9 Z& w% z
    36! T# E" T  M/ R$ I- Z& n* W
    378 M" O! v* R/ D
    38
    9 c9 I% b$ B- s# E399 M6 y/ _/ [- F# f9 Y$ e* |
    40
    + H- J( ^, p' [. W8 ~41& X* ~/ C  _2 e/ q% {% A
    42
    4 s, e, Q- q5 q* J* v7.3 训练所有层+ _6 |$ g3 r$ m6 Z  a
    # 将全部网络解锁进行训练
    : {$ b' b8 Z+ B8 Sfor param in model_ft.parameters():
    / \6 x+ J" R! v  \) z$ P, a    param.requires_grad = True* N! e4 t) O: }- L
      b# T$ X! |2 E! j5 h1 f0 j5 B; Y9 I
    # 再继续训练所有的参数,学习率调小一点\& F( r4 C9 v) {% D8 b! w0 a: O
    optimizer = optim.Adam(params_to_update, lr = 1e-4)0 F8 R7 g: Y) x6 F
    scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)% Y) f  E! v/ V4 G* G' P+ t

      h9 X1 V/ G$ ]& S4 S' G' N6 T& N' N; `# 损失函数
    / |3 {* `+ E) l+ Y& C0 B3 U' a0 jcriterion = nn.NLLLoss()
    / z  P% B& a  A! H1
    3 r1 [) R+ q6 B) U; k( c21 R4 ^7 o$ i# L6 W) r4 h7 G
    3
    5 |" V2 R! p8 G% l4
    ( G& X, u8 D1 B8 B4 w+ I+ l55 C+ D" a: K6 \9 J2 [
    6: V4 a1 ?! w4 s$ j$ N
    7) ?) {1 O% E, t
    8
    0 l: {2 _: @3 a  O  }+ e* O! J9. ]0 M/ j* A4 `( n6 c
    108 k+ U0 T6 F$ z+ a
    # 加载保存的参数* X3 x5 J: {* j5 {
    # 并在原有的模型基础上继续训练( E; s' O9 q. ]# s; [' I& y
    # 下面保存的是刚刚训练效果较好的路径
    0 l- ?+ i1 [/ o0 T6 k6 l5 Z" vcheckpoint = torch.load(filename)
    ! W; H  s. q2 j- o6 Dbest_acc = checkpoint['best_acc']
    0 q/ w% z! y# v4 E( m- B4 Zmodel_ft.load_state_dict(checkpoint['state_dict'])& u; l3 m5 U2 M1 ?
    optimizer.load_state_dict(checkpoint['optimizer'])/ n9 I) e8 e- V5 M
    1
    1 d+ c; s" p$ P/ {# Y2
    5 E4 {% K# }( {$ L) e: ?  M+ Q3
    0 n$ a, U5 s5 A6 I4
    : }: R3 C" U' E4 s/ N9 E5
    / s0 t/ w* ^, Y5 u! c6, v" O3 L; w7 Y3 v
    7
    , v! s. {2 g, h" t2 f0 p开始训练
    / A3 j+ ]  p7 x1 _7 ?9 A' F4 a注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考6 V% k6 a. z# {& u9 u9 ?

    / u. l" U+ ?/ A( T. Gmodel_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"))
    3 b0 y7 q0 t9 t( D1
    / Y- i2 B0 |% {. E& Z/ LEpoch 0/1) j  p) n7 Z0 g; K+ k) z2 e( l
    ----------8 h2 y& Y7 S) |; [$ R- \9 {' w( U
    Time elapsed 35m 22s3 `0 e/ u1 N0 `: |& ], ?, \
    train Loss: 1.7636 Acc: 0.7346
    + i# a4 o' ~* @8 t' h  c2 lTime elapsed 38m 42s
    8 X6 G9 z) x+ a4 nvalid Loss: 3.6377 Acc: 0.64552 j! B$ I& |5 u: q0 P
    Optimizer learning rate : 0.0010000
    ; H' x- c7 Y2 j2 j* x! c" g$ H+ A7 j2 K
    Epoch 1/1$ f1 j! c7 W3 q
    ----------
    3 G5 R, R, }0 `Time elapsed 82m 59s$ S" @* I) o; R' s  F* C
    train Loss: 1.7543 Acc: 0.7340) ]" U) c9 p# C( G9 _1 U
    Time elapsed 86m 11s8 ]% R3 p8 B  B$ @
    valid Loss: 3.8275 Acc: 0.6137# i1 f9 P' R1 X* @9 q- W3 P+ [
    Optimizer learning rate : 0.0010000
    ) d# Z, S% M4 ?5 q3 T! _! X2 \, ?7 A( B1 O
    Training complete in 86m 11s
    ; @: R7 K' ~0 ]/ i2 BBest val Acc: 0.645477
    ' D' f( e  b+ M* w$ D
    # ]( I( |# _5 z! L; _1
    , l" i8 K: |8 p/ e7 w! e' A2" I+ V8 f5 h6 B5 t0 P
    3/ t5 E' l- L. J! {
    4
    8 e; ?6 G! z- p: k* |# |54 v: D# l1 i# X' J  o2 H
    69 f' m; U! e4 \, ]) M, m! f  _( f- I
    7
    5 }) `- K0 S/ A( y& c8 [* N% s8
    : ]( s. Q/ J) v6 m" h2 O2 K9. k% }% I' J+ Y. t7 _, J
    10) ~0 a4 P1 |/ G
    11
      k; y: z: g  f+ ], |$ [- J  ^5 P12
    ( x* s8 R: n* O13& ~3 k7 E  z2 z' ?0 t, A
    14
      v7 d0 N) B1 N# O% i15
    * l0 y4 `3 H7 z7 K8 ?4 W3 K16
    5 I/ ~( {0 u  y& J: U, ~17# O6 w, o1 r; ~
    18
    5 i  i" N6 p* V- Y" x! n8. 加载已经训练的模型' i# A) y- Y; Y/ a
    相当于做一次简单的前向传播(逻辑推理),不用更新参数; |, [) m" _* E: L& m& |* c
    ) M2 L! u5 X# V) F* _( p( |2 F
    model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
    , R+ g0 F* B$ c# ^
    # m4 V  q' f7 _) T1 `# GPU 模式
    2 M! n: V! K6 I. |$ o" Z, p- u+ i. h9 imodel_ft = model_ft.to(device) # 扔到GPU中
    - P8 u3 N& v6 e7 Q. t, N. }$ }! w! R) V
    # 保存文件的名字
    , h* {: d2 V8 f3 l8 ofilename='checkpoint.pth'4 k8 t1 F! J+ s/ m1 ^) o
    - b4 F5 V" ^2 ]
    # 加载模型
    6 r: E. u$ i# a2 ncheckpoint = torch.load(filename): k+ W9 f  X# f9 L  d7 B
    best_acc = checkpoint['best_acc']
    ) C. {1 w: H8 A3 ~5 }/ Y+ l& Fmodel_ft.load_state_dict(checkpoint['state_dict'])
    8 C, p1 z: }2 }" u' r1
    9 ^* ]( j; g: |; e2
    1 `5 z8 c, t" J2 y0 J$ g3 N: Y7 m. \3
    + T& q7 O. m5 e; ?6 x41 G3 h) C5 K* E  W
    5+ ]7 \1 i& C" _
    6. W. |& S9 |& j4 u( A4 u* Q. T
    7
    & X: m, I) T3 A8
    + U& w  n# G+ ]1 N* O* r9
    + [% P# ^: X5 `( k2 `# r) P. e10
    6 T; D: p5 k3 H& m11
    6 x: O& d0 B9 `12
    # ?6 D6 O0 T* J1 S<All keys matched successfully>! `8 {0 f0 g1 f
    1: T6 @/ x! @5 h( X+ p5 Q
    def process_image(image_path):
    9 |5 B- s" ^. k    # 读取测试集数据3 {4 z7 V" j  X) a
        img = Image.open(image_path)
    ' j/ J$ \6 d5 |1 y4 v% A& o3 C( Y2 m    # Resize, thumbnail方法只能进行比例缩小,所以进行判断
    4 e+ X  B2 F$ {% P' H    # 与Resize不同3 F1 o7 h- _* H5 l
        # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小6 }' L9 L) |% k) S
        # 而且对象调用方法会直接改变其大小,返回None
    7 h' c& R4 U1 e" m' S$ l    if img.size[0] > img.size[1]:5 W- Q7 }5 N: V( D
            img.thumbnail((10000, 256))$ r& z% \; I0 P1 z$ a  r
        else:
    $ X1 ~' U( v5 O8 Y6 Q        img.thumbnail((256, 10000))
    $ S/ N  L. x  \6 d# L! K* R. E- K+ L! X! Y+ ]
        # crop操作, 将图像再次裁剪为 224 * 224
    9 P' ^! ]0 v. Q" Z    left_margin = (img.width - 224) / 2 # 取中间的部分
    1 ~- {* A" c- O5 i* W2 x    bottom_margin = (img.height - 224) / 2
    # A4 Q7 q8 C% v0 G$ a( z* S/ U& }' y0 _    right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度
    & m# ^; r" Q( P; N, ?5 K    top_margin = bottom_margin + 224
    , t, \2 \+ L2 A7 G+ e3 h- Z
    ' u. M0 S6 w  Q/ Z% t- C# ?& S    img = img.crop((left_margin, bottom_margin, right_margin, top_margin))
    ) d3 d; e7 w) _7 j1 D: y- z' {) Z- s! `' x) U
        # 相同预处理的方法- x# x# V, ?# _) \) G5 w
        # 归一化, @2 o( p' ^7 \( i9 b  E  x
        img = np.array(img) / 255+ _) l4 W  [$ }- s  Y) y% ^2 L' g$ U
        mean = np.array([0.485, 0.456, 0.406])# n" f" F" T3 M' A8 M  W; c( @
        std = np.array([0.229, 0.224, 0.225])- ]% m( b5 j) R* N) i0 C3 d
        img = (img - mean) / std
    8 P' [. [6 n3 f' t7 M3 r1 R6 G
    ' `; x! [' n; S7 M5 U    # 注意颜色通道和位置6 `% w; W' E5 m. u: x6 k1 k
        img = img.transpose((2, 0, 1))
    / k: Y8 s5 A7 ~+ o4 W% [, ]* m# h% L* u1 J+ m/ B4 W: M
        return img5 _) }( |9 Q, g' s, i% D- ?7 g

    + B: A4 w. v) rdef imshow(image, ax = None, title = None):6 o2 B$ S2 Y& a/ w1 b8 _4 E9 F3 t
        """展示数据"""8 ]% `. w! n! u5 C! l$ R
        if ax is None:
    7 B; P2 `" g# v  k( \        fig, ax = plt.subplots()
    " h; \8 l" O( m) q4 x  G# @# B$ d' h$ f
        # 颜色通道进行还原
    " M; I3 q3 M3 C9 v8 @    image = np.array(image).transpose((1, 2, 0))1 z- n; G' `5 M( r4 b( m0 ^1 j% Z+ X
    ) a" X) F1 x1 x! }7 a' p. n
        # 预处理还原
    ; G/ E$ S' C7 L& @5 I5 p9 I    mean = np.array([0.485, 0.456, 0.406])
    $ I; O  [# H& z# h9 k  u    std = np.array([0.229, 0.224, 0.225])9 P! e/ K- D+ E, E0 t, x
        image = std * image + mean1 W% q% Q! f: e9 o
        image = np.clip(image, 0, 1); m* A- A6 V/ v5 J; T) g- l+ H9 V

    ) C* ~, e7 S! C5 d) p2 Q0 v    ax.imshow(image)' X) b8 `0 v- n9 l. B
        ax.set_title(title)
    & p$ ?9 \9 }* f# b  C1 ~; {" n6 i# a3 ~! Z/ |; M/ g, @5 m
        return ax
    ; R5 U* w1 H& _7 z) g3 U, a, o5 U: p! v- W; r- \6 Q+ {
    image_path = r'./flower_data/valid/3/image_06621.jpg'2 p' E0 M) K& y
    img = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理' d% T7 H8 t0 b
    imshow(img)
    5 y, z4 P7 C  l; B3 D# k3 y6 c
    7 a' e; O' \  g( L$ B6 x% s13 q! K% l  r, a8 W
    2
    9 S9 v# V$ c- j3
    % ^! i; u( g7 X7 K4  \6 h/ u$ o) p: {
    5
    " j+ E0 L+ U. }" }- Q# E+ {6
    ! V+ W9 V# {; K% J  z8 r3 t7
    ) W) k& ?+ ~1 D8
    ) s9 c. H+ A& S' F1 J. x  q) i90 k* F9 A$ f+ A" e
    10
    , {% X- y+ G# x. b- r0 q) [11
    0 U* S" @# q, T) U2 T9 V( g7 L125 J. V) |3 W% s: U. L! f% t% S
    13
    ( r  L+ l# `1 v, v14
    % Y; c- t1 F6 e15
    " I% H: q' j) [* D167 \/ ]. o/ b1 T5 h: K
    17
    0 ~8 }1 `$ j: l186 u9 Y6 R- Y3 V; R  l  x) |% F" h
    19
    . s( H- @( j9 W& D20+ ]$ \- J: Q. Y
    21- c7 O4 |8 O, d# }/ j8 Q9 e
    22' |) t1 S  A% N) Y9 F
    23- r( k: t9 y- {4 G3 P! g+ \$ d
    24+ H" D/ T, Y% a/ ]0 x
    25) }" v5 w5 J4 |+ Y
    26
    ! u% I8 z1 s& m  ^! o; L27- i* J1 Y- f- c$ D
    28
    ; K: N7 K7 X+ e' v292 d+ I; X" B: p+ }
    30
    7 i0 o% N; \1 v31
    4 {8 f& Q5 E$ C1 z& e# z32
      f. f" q7 C% g6 b$ T  ^9 G33  B# X! A- _- i+ X
    34" V4 l# s; Y- C
    35( M! o* L, k% ~" m+ f
    36
    ; D0 |" |' d2 a& G7 h5 y37
    ; Q6 Z8 D3 l/ I. I/ H3 p: a388 \+ f) @+ h1 N& d
    39. J  L7 |* B2 ^  [0 Y. N: ]& v" O8 s
    40* v0 _6 d# V8 x9 X4 @7 H
    41
    ; Z) k8 O* S$ G2 g2 n/ y420 D, w0 H  N2 |; x2 q6 R, s' v* _) [
    43) _5 V* z. D  M7 G5 {( c  N
    44- B- k1 r- ~4 r8 r. U6 ?
    45
    8 b/ y( O+ h: W! ]46, l# T3 B& {; i  Y" Y
    47; K3 B+ z: A' N0 B" @! U8 z: _% f
    488 s# l6 V- f& z" u' P2 O# N
    49
    + H7 X! D  L9 B- |& D& D7 I50% m. ?1 a5 o7 L+ U/ v: I
    51: ^! r# p9 F& I& q1 c0 |
    52% K- S* I2 S4 u/ i; R
    53
    0 x  X1 l' D+ w/ G2 c54
    % D. s4 B( [8 Q1 `<AxesSubplot:>( [9 G+ h$ n8 @. J5 r& P+ X; k
    1+ B  a; ~2 e- P8 a
    - O: v4 M5 D" @# p6 o
    上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确/ m, ?' {0 e. F

    . [& t, J  j9 t5 Dimg.shape
    1 `# X" O# ?* o, V; }1 S( H' b1
    ; ~' |' E5 w7 `5 c, t(3, 224, 224)* z$ l' a+ e/ O* x) a/ P
    1
    $ ^4 T8 K: M3 H1 L/ l证明了通道提前了,而且大小没改变
    , K* H7 K: C( ?" r
    : A& Y1 h  y% O9. 推理4 I% P7 ^  r' }6 I6 @7 P
    img.shape
    8 E/ P& i; r4 ?2 n6 z4 n/ c- g
      V4 @( V! D+ t0 x' p; y% G# 得到一个batch的测试数据6 F3 p( L6 J3 j) e/ p! p
    dataiter = iter(dataloaders['valid'])3 W& o, T2 X. [# N7 s! a6 X
    images, labels = dataiter.next()3 W) E) c+ ^# a, e& A+ W+ R. J0 m
    5 v, F3 _) h  v. k# ?
    model_ft.eval()
    2 A( B* Z- j3 E4 {1 v& k9 F. w, O2 \- G  c
    if train_on_gpu:
    5 T: V1 Y$ [5 C6 K8 i( _0 D    # 前向传播跑一次会得到output
    * O3 k0 @4 t! \  m& T( U    output = model_ft(images.cuda())
      ^) q/ U! M9 |' a* {. y& H1 velse:
    5 ]( x6 ^0 O5 F6 x9 }    output = model_ft(images)
    " Y/ C- J  _. c) x1 z8 |2 q5 f7 h; k$ L
    # batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值
    2 p: F5 W1 I- f) D7 `! @6 Goutput.shape! P2 k, z+ V) H! F" p
    ; N% f9 v' O4 @4 U7 T& V
    14 R. m- Y0 n1 _+ f8 S
    2
    2 s4 f( ^# S/ F/ x" w/ j2 F3& E$ o+ Z. j& K# p, O
    4
    ( g, U$ p. l/ d" Q+ d5
    " J/ K' z5 i* B$ W! a0 ~9 r( U6
    3 s4 f  M, T7 M% g3 F% e# G& @7# z2 D6 U' _7 @7 @8 f
    8
    2 j8 [- D! R% q8 W4 _/ C9
    / Q- u# l. m* t0 o, G2 t10
      f- w) Y1 J( L5 i# g7 |5 O11
    0 Q0 n1 D. I$ e12' p7 S. U" f1 U( k
    13" }3 |9 |5 [+ Y' b0 U& d
    14
    7 \; \( a; d- H15
    1 ~5 f2 B2 V# F, W16% {' {* K5 v# f3 w& k
    torch.Size([8, 102])1 [# |/ K, Z* B; l6 _
    1
    - f+ t" O2 O" `9.1 计算得到最大概率
    $ n: s: {! n3 |5 u/ f_, preds_tensor = torch.max(output, 1)4 I0 |. D& g2 m8 [

    ) N$ w( @, z8 h# x$ O' Dpreds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量0 C2 ~2 B  V" e& y' ]* J% ]% K
    18 ], R/ p$ r6 _+ L+ s7 r$ F
    2
      h* ^) Y" X0 ]% s3
    ; ^  m; l* y+ D4 A9.2 展示预测结果
    ' A5 b( Z5 k8 x( o2 F0 {' Y$ wfig = plt.figure(figsize = (20, 20))9 `, r8 y$ l9 D1 [
    columns = 4; G3 V+ c) B  B/ k. Y
    rows = 29 `: k4 g4 {9 L4 b, N5 }# V" K
    9 `7 G1 Z$ f# }2 a2 J' |
    for idx in range(columns * rows):, M6 s7 }+ @, G+ S- E0 ^  K
        ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])
    0 q' j. |. R  q    plt.imshow(im_convert(images[idx]))% }2 `5 Z0 o( R; I. k" A
        ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]),
    + l" v7 u6 D- q0 V, w3 g  M, b                color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))
    ( k3 v- a0 G( p' [, n% I; z1 Vplt.show()
    : W8 T: H' Z( p3 `# G# 绿色的表示预测是对的,红色表示预测错了
    6 S3 r, K% [. Y8 w2 C2 R1
    % A/ B% d# p) k2 r2 B! B& n7 X2' B& m+ s0 r" u- \& g& Y
    31 Q/ q0 J& D* f9 j6 t
    4
    * S* s( c$ ^$ E+ c' K) O; D59 s5 t( ?% j7 V) E+ A7 H! d% k
    69 z5 y' Q3 u# n/ Y6 {
    7
    & l7 c) R1 h, j% u$ [- _8+ _6 C+ Q5 Z( {9 a: |7 X
    9/ \) r* h$ S; ]& q. U4 h/ c% h1 g
    105 x6 T7 p7 K  N/ [* _0 @1 ?8 x
    11
    . g$ f6 G! O$ f/ \9 [' G2 @) _5 j. A  w  Y

    " r' d+ q8 ~7 p$ w. Q% A7 T! _( N* E  b; q* h7 c+ X
    ————————————————4 i$ x, X1 K7 y* R- J. A
    版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。
    - `( u' l$ h2 c: J  D原文链接:https://blog.csdn.net/LeungSr/article/details/126747940
    " F; n: j3 x- U7 w! L6 L; ~8 c: A$ R: a% \. p3 ~! t

    + O6 g0 l% J+ X
    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-7-28 15:34 , Processed in 0.562926 second(s), 50 queries .

    回顶部