QQ登录

只需要一步,快速开始

 注册地址  找回密码
查看: 2836|回复: 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)实战案例. p- l* f& a" g% `
    3 `% Y$ |. y* H2 _; ~# D; Q
    文章目录
    2 B; J* T  G2 k! C9 n. x4 ^/ z卷积网络实战 对花进行分类
    " X5 M( l/ q; {1 d* G( i% `数据预处理部分
    * P5 f0 f: G0 P/ @7 K5 `5 @网络模块设置& \" l+ _! B5 r2 t2 L! K0 i
    网络模型的保存与测试
    ' Z# |1 D0 O3 I- a" ]& C. }数据下载:
    & Z0 C/ A) @6 ]  _1. 导入工具包
    ! T0 W; ?- Q4 b3 k2. 数据预处理与操作/ ~8 T! m: j2 x# s4 ]1 Q
    3. 制作好数据源0 v) i+ L) \$ @5 Q  {
    读取标签对应的实际名字
    4 `, D& s# @2 c# _8 a# n: X5 C4.展示一下数据
    0 q2 t5 m5 u5 Z$ A+ }% q5 F5. 加载models提供的模型,并直接用训练好的权重做初始化参数6 p; K" Y; w' b/ J, k( P0 k
    6.初始化模型架构
    * T6 l# o; L; l6 O7. 设置需要训练的参数
    7 ?% ?* y- W6 {# T" t& r# U* J, X7. 训练与预测
    $ b! D" P' Y0 m& r  @7.1 优化器设置
    1 J) L! b; s3 \- n- U, C+ @7.2 开始训练模型
    4 N3 l4 _* w% c2 R" J  I7.3 训练所有层& g. }$ `% f: E) s& L
    开始训练" C) t& Q! j' g' d4 B
    8. 加载已经训练的模型
    8 j: G0 F( t5 {! T9. 推理
    : z% P0 I0 \( V( l2 C# `# U8 X9.1 计算得到最大概率4 g7 L" S) g$ L
    9.2 展示预测结果
      H8 E/ ~; C5 T% j. {& f: A1 _写在最后
    ; ^+ m6 ]7 F' C4 P; F2 V卷积网络实战 对花进行分类! t5 I, S1 a- w' }7 ~9 n
    本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能
    0 y) M. N8 V, S6 r4 d* ^: j; ]% h. U3 S& T$ l+ q
    在文件夹中有102种花,我们主要要对这些花进行分类任务
    0 f/ U. p. T2 W$ w$ s* a/ T5 E6 h文件夹结构
    ; F" z+ @9 `5 e: [
    ' c) O4 {9 o1 q3 Z* F7 ]& F5 b, Iflower_data; t, [, J5 z) M0 [9 H7 s* m
    ; p* y" o* b; W
    train' Y! K. C& K# X/ i/ q( s* e
    % y3 x4 G) |$ U. |2 ?! w
    1(类别)
    / f4 E1 i; W0 D6 ~0 t$ E9 w4 T25 z, ^7 M3 z" I6 z" e- w
    xxx.png / xxx.jpg
    9 a9 J) M9 K' B; a$ jvalid5 x! q- t) j7 B" v
    * I" A2 G% G7 I3 O
    主要分为以下几个大模块
    4 k% G7 m" b: T( v/ z, F+ S6 ^
    ' ]+ s; }3 L" e5 e+ J数据预处理部分
    8 V( h0 m9 {9 x$ [, P5 r5 l数据增强
    5 T* f# }+ ~0 _+ }2 v" S数据预处理
    5 S2 t1 t( V' k1 D8 _% K2 d6 w7 E网络模块设置3 {8 Q9 s  h. k. C( x/ b
    加载预训练模型,直接调用torchVision的经典网络架构
    * R. q! q( p0 B+ r4 o: S因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务
    $ n) @* H: B- [& ~1 b+ t网络模型的保存与测试
    8 c5 ]3 m' Y6 t模型保存可以带有选择性: Z- F: |& ?) `9 T, y
    数据下载:( i! S- q, \! \$ X$ e+ y/ e4 c6 h
    https://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset8 L& \9 G+ b$ l1 w" W2 C
    8 a: x6 b" h: G, s3 e$ z# R( U
    改一下文件名,然后将它放到同一根目录就可以了
      b* A! q0 W$ y# W/ x$ V/ g  y
    ) Y* P% t1 V& x下面是我的数据根目录
    % ^5 W$ T) x" r: z  Z0 T5 U; |" [( a, G/ r' A# e( Y; y
    + S/ ~9 O( d$ b1 l
    1. 导入工具包7 M; [4 T4 C9 @- n, f& l
    import os
    $ ~- x" N* T7 `% p" ^1 j3 _' Iimport matplotlib.pyplot as plt
      A. a, y& p, q; `% s# 内嵌入绘图简去show的句柄
    ( b! n4 y& q; }+ L%matplotlib inline
    & `; e2 y6 ?4 ?( U* S7 x- u9 Q- mimport numpy as np9 p' q6 `3 m0 h0 H7 J# w6 g
    import torch
    ) y! {7 w7 G, |. efrom torch import nn% k! N$ {7 y* b* i0 X

    : i3 x6 |2 y4 S0 R3 n5 r4 Fimport torch.optim as optim
    " @$ W  o& p0 u! @! o  m& Aimport torchvision8 F$ r4 D0 k% h# g+ A4 H# h
    from torchvision import transforms, models, datasets6 n  F' d$ K& l" B0 ?
    9 _. n! z5 d6 \, E1 x% [% E+ `8 Y
    import imageio
    0 F* [+ g6 j; K; K6 Yimport time
    3 h: g$ p  R) U# J8 a: \; _* Uimport warnings
    $ }( O( e0 u9 p( G6 ^0 R" H- N4 Wimport random# R7 y$ \# h; A1 l
    import sys
    , S+ s- d5 Y) Z0 x  cimport copy! V) W6 j6 d' s; @  F" b
    import json
    & o1 @+ i# B  A& V( ^from PIL import Image- C% g! k5 h& L; M

    % |& J; I1 `5 s: O8 ~- ]. Y, X/ {7 Y. d* r& Y
    10 {, _4 h2 i5 N. i) S$ ~; C
    2( u+ z/ I' O* J0 u
    3
    . L" d& A7 d6 z4
    0 m, ^- p; w; G; y5 z. N- N5+ |/ p  o5 H) c5 z( X
    6" Y5 u& ]7 b2 N1 g7 D6 O
    7; c! f6 s6 k( C/ W. R7 b) n
    8
    ! ?5 P+ I% @- |4 f2 \1 s) Q9! @- R3 Q$ `7 B0 Q4 y' @
    10
    % {8 z# l2 G% g2 W# ?' {+ b; z, m3 P11
    / ~- a* J3 x! `9 ]6 D4 E% `' d2 G! m12* X  v4 D- U" f0 A  v# A
    13
    9 F& W4 c& q" G7 H+ ^: i14
    5 M* ?$ Q8 O1 |0 i6 v. ~, o15$ @/ S% r. Y7 D! v& I
    16- H  `. K/ W& P3 g8 Z/ k
    17
    # z* \& L' I9 E2 v* P+ x18/ G$ `) U$ u9 \& L6 p* y: m
    194 N  n, ~3 Q0 c- m  j9 [
    20- w3 T* a0 H4 I" g: R0 M
    21& d$ q8 j; o$ h1 \4 u& _
    2. 数据预处理与操作
    . V% N/ u' r- K. k. K4 N0 `#路径设置5 Z+ `$ [3 P- Q9 ~
    data_dir = './flower_data/' # 当前文件夹下的flowerdata目录
    ) O2 ]0 L$ Z% Atrain_dir = data_dir + '/train'
    & ?9 U/ U# M# ]9 p/ J/ U3 Hvalid_dir = data_dir + '/valid'3 I6 `7 _( }/ V% [0 {2 s: o
    1. Z* a5 M* I* D3 a4 @8 H0 _6 b
    2
    - m8 r9 l& V$ f! B+ y32 D9 j9 [5 Q: ^; ~
    4
    4 c% ?' G5 l4 ]8 I5 lpython目录点杠的组合与区别
    2 I9 a3 P6 _  W5 V注: 里面注明了点杠和斜杠的操作
      M4 L7 a: l; U# @: U$ Q5 N( w: R5 B* Q( ~% P
    3. 制作好数据源0 f. D+ A! t7 T1 E' @( @* [
    data_transforms中制定了所有图像预处理的操作& b2 C( K: @9 L  Q: V+ T
    ImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片) b* _: r# S0 _+ G
    data_transforms = {  b2 e0 g; q& s* g( P4 w; w
        # 分成两部分,一部分是训练
    ) L, d, ]6 Z. O  D, q    'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间
    # j: `# J# x1 G$ f                                 transforms.CenterCrop(224), # 从中心处开始裁剪
    / ^# E; e) {/ i+ E4 |                                 # 以某个随机的概率决定是否翻转 55开
    2 Q) j9 _: |0 i  C                                 transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转6 |" C1 H6 _( F9 a) W3 y; c
                                     transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转* x0 s- k; m/ \. T) E2 s& ^
                                     # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相1 c7 S! `! H: ~0 V
                                     transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),
    3 B- P8 Q+ r3 I                                 transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB
    : R+ s, l) g) N* D2 C( d                                 # 灰度图转换以后也是三个通道,但是只是RGB是一样的
    $ D! ]% O0 m; R7 s' `( q                                 transforms.ToTensor(),
    " j' x2 p8 k) A2 ?" F$ B0 X! P* A6 H                                 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差0 q" u1 o( c8 v8 W
                                    ]),$ {3 ?) s! Y: Z7 W% T0 w  p% L
        # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化& P8 l  \' t  c
        'valid': transforms.Compose([transforms.Resize(256),
      Z9 d1 }% {6 V5 i                                 transforms.CenterCrop(224),
    - }  [: g9 f+ L$ e! L( K: o                                 transforms.ToTensor(),
    ( @& ~0 u3 D  w4 T) N                                 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同; X5 b& o3 o+ q
                                    ]),3 R. i5 a$ J, C6 `
    }
    2 q9 a& b6 |0 ]/ y: N2 N
    7 a- B8 g: u: ]0 |3 B16 o& }9 |# h7 J6 N: @4 B
    22 H; M, f4 D. g* H3 Q
    3
    1 s% {2 F# a+ m6 f+ ]4 z$ i3 ~4& X( l# g* R/ [% `2 K
    5
    + V% E2 t- z  _' d# o9 k5 Q5 V0 q6
      B  e1 ~8 L$ C$ F) K9 Q7! @5 _; W0 _$ k/ G( N* t
    8
    : n9 E+ c" q$ V* ?9( t3 G( b3 z, N. s+ S& r0 ~0 b4 ^
    108 H/ m! F, f9 V; W( m
    11
    % r  d, L7 [3 q* \& p8 J12
    0 ~7 R9 h; }8 [3 C% x" s135 |% g4 ^0 u3 w  M# c2 ^# F( J
    14
    ( f( _0 b6 `2 m! N& B15
    & i* {6 O  j9 Y" I; C- B162 ?4 n% k' {$ y: s4 p
    17- w9 F- c; J* G+ m8 R+ P& I9 M
    187 J' |9 @/ {1 ?( v  z" q! J" \
    19
    & p4 [. `) \! g+ M9 ^20, q4 X4 Q: R: q  u! [/ z: ~$ f
    21
    1 u  G6 R. u. g* @; nbatch_size = 89 u% R) L2 E$ ~! o/ W; Z! T! V
    image_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}
    # k: ^9 T& t5 v5 ?dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}) N  B2 F: A1 o$ q& q- V8 K5 w
    dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']}
    , i, j' c/ D! }% E" `# J& O6 L5 rclass_names = image_datasets['train'].classes5 ~$ J/ A% L7 X% c4 T

    & l5 [  r3 Z$ E6 m#查看数据集合  F( p( ^5 x7 m$ v( A4 J
    image_datasets
    1 x, Q: g9 d: h) ~6 X; y- C$ i( m' ~& A4 J3 M) y9 e/ g7 h7 r
    1
    3 \" N) Z1 C- P% M- T21 a2 ~/ r0 ^* l7 D  f& j
    31 v( g5 F6 H; j5 o, o# d, {
    4* S% K% N$ V7 i( a7 t
    5
    : J# X: ]' S4 V4 e0 [5 a6
    ) n3 \3 b7 D0 s7* ?1 ~/ ~- o7 M$ P- q) z
    87 S( e& \; L& b3 X, p0 c3 |1 R# O
    96 I7 U* n( b- c0 E
    {'train': Dataset ImageFolder
    4 c5 _5 k2 |# b' Y) O* P0 S5 L     Number of datapoints: 6552
    + f9 U4 @: Z) ~4 U1 D$ T     Root location: ./flower_data/train
    0 ~; @. \% A- _+ x  P, U     StandardTransform
      M  v# C4 y" t! O Transform: Compose(
    7 x9 b4 T) U: ~7 S1 @                RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)1 F, H3 h* [; Q$ w
                    CenterCrop(size=(224, 224))2 X6 O6 C& |1 q1 R4 _8 T
                    RandomHorizontalFlip(p=0.5)$ U  B) [- a: C( y) n; }" w
                    RandomVerticalFlip(p=0.5)& w7 R; W+ i: a4 N; O  s
                    ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
    ' f  n9 }" K* V- |                RandomGrayscale(p=0.025)9 E# D2 n& h- A7 R2 f1 l8 I0 {5 ]
                    ToTensor()- h5 M/ v1 t  m# O( f
                    Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])4 [$ d! |. r9 O# h6 r
                ),2 z$ p* y% B, J* t  R: t- a0 L
    'valid': Dataset ImageFolder1 ?* ^/ `  o6 k4 D& D* P
         Number of datapoints: 818
    & _) k- u( z) z6 I     Root location: ./flower_data/valid
    8 n+ Y7 k2 a4 h; m     StandardTransform) {$ f9 j! c$ `6 \8 a  {
    Transform: Compose(
    / L0 }) s8 D. w4 h4 M                Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)
    4 a- N8 t1 N8 X2 N2 a% R                CenterCrop(size=(224, 224))( p& N) K3 m+ {3 Z2 O) h
                    ToTensor()3 V/ Y, u0 ]0 G5 o7 ?' s" T& L0 X
                    Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    / k1 g; F$ u, R( u( X            )}
    * {( ^& o3 V$ T8 Z9 A$ `% T% N/ D) h. i, e! z# _, @. ?" ~
    17 c: s7 q1 M. u0 F' j/ j" m
    24 i& p) I) ~; w
    3
    * Z$ v% X  w4 t7 M# S49 ~: g. W4 E7 d
    58 z8 n) ^, c$ o' X  {. U6 U
    6
      V3 J! t2 w5 @% E8 F) ^7
    % G+ _  L3 y& v. J) |8
    6 R; B& H; `9 W0 y9
    3 L1 {8 s7 q3 o0 u# V+ k4 g/ ]9 h0 @10# J9 _* H" Q, M: ^! Y
    11
    3 p9 Y, q5 {8 A$ a9 l8 Y12
    ; c& ^# @- U  C$ f13; E, R8 `( e& Q" V( j7 Z  _% @
    14
    ! ^8 A  x$ b  M) G& N, D9 p15' G& Q" T$ P% [/ v: m6 l# M5 ~4 J4 p
    16# Z; ^2 g: P3 B2 r9 K1 i( {
    17, g0 Z6 ~! v9 l8 ~. G$ p
    18
    ) r4 X: w5 v' x" B% Y3 U+ b; S19
    0 p& \: i* n& k! X: p  _, I20
    2 _. K8 e  x% ]  s; l( N, b21
    . D7 t; T1 M8 ~  l$ `22  w* T% z6 e) u$ n7 E# ?
    23
    % V' g2 y$ S9 F9 H! C24# ^4 t% z2 F3 w) a% s/ U+ T4 P
    # 验证一下数据是否已经被处理完毕
    ! G9 Y' {  F, O, I+ N4 m- Ndataloaders/ }9 F# y3 g2 n" D) C
    1
    % q, A% _" W* S6 Y3 q% P, l1 f2
      Y6 y0 M6 k# n5 Z. N8 m+ b{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,
    9 d/ P' }# e0 I$ K7 k! p) A- l4 _ 'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}  M8 y$ L4 b! J5 w" d  Z
    1
    + [$ s& h9 M, i4 G2/ c, j& c  O! t% Z
    dataset_sizes
    & K9 O/ {; }/ N11 h: V8 e" W  p( }3 L
    {'train': 6552, 'valid': 818}- U0 j- @, x4 v, |: P+ W- e& ^, O
    10 ^. _: M! h; o
    读取标签对应的实际名字! P, `* X7 O3 [8 }* f( ]" C
    使用同一目录下的json文件,反向映射出花对应的名字, Q" p4 A! ?! a1 w# f* f& @; J9 W
    6 {% d6 m& P: }$ H6 ]3 ~4 U2 K
    with open('./flower_data/cat_to_name.json', 'r') as f:4 U, ~9 g. |3 ^& ?. p! M) G; t
        cat_to_name = json.load(f)
    - R* g$ }' h9 n+ ^/ r, L" [1, l! m+ s. g$ r- k
    25 a& C( W7 I; [7 }
    cat_to_name
    % n) Q: S2 Q5 b; h% v* @1' ]1 D# p% c  w& O" N: k
    {'21': 'fire lily',4 J4 T% S7 O, Q# A
    '3': 'canterbury bells',
    1 X& ^$ F) T* ?+ |6 g2 ~  Y$ E '45': 'bolero deep blue',
    5 e# V* G  V" D$ g' V7 J* s0 a6 m  I '1': 'pink primrose',
    & V, S9 M4 d0 U( H. i '34': 'mexican aster',
    , {$ G$ \# U7 R2 c/ u4 ] '27': 'prince of wales feathers',
    3 X% t9 R- l/ x, ]- n) `6 @ '7': 'moon orchid',. @1 d% V1 }% c" R$ w
    '16': 'globe-flower',* n6 G: @- ?: @" n# H  A3 I
    '25': 'grape hyacinth',
    / r1 l2 Z+ r( G5 B1 \5 b; E# O: a '26': 'corn poppy',7 N! I# J/ a$ K0 I
    '79': 'toad lily',
    4 h- H& p% O+ J% C- b/ { '39': 'siam tulip',: O1 o! u0 ~- c* {
    '24': 'red ginger',
    # p' P) D/ g4 ~7 S '67': 'spring crocus',
    6 _8 c# [$ p, ?3 C) r4 K5 w '35': 'alpine sea holly',
      c* t) M+ ?# H '32': 'garden phlox',6 o# T; E* H' c( j$ a
    '10': 'globe thistle',
    " n; q( T& O- H% B '6': 'tiger lily',  K: a; P* E, o+ Q' E
    '93': 'ball moss',4 P$ {% ~7 c5 E9 y
    '33': 'love in the mist',
    3 x/ T; i2 W* i2 ?  e '9': 'monkshood',5 a/ Q3 @9 f8 N6 [
    '102': 'blackberry lily',
    9 E/ D. K' K& r, K '14': 'spear thistle',
    - k, W5 Q& Y" ^% O9 _4 b. d '19': 'balloon flower',
    ) Y$ Y3 F4 [" g1 x '100': 'blanket flower',) i* K8 e$ C+ `
    '13': 'king protea',
    4 n; P# n6 D8 O: M6 M% N '49': 'oxeye daisy',
    / i- g8 F: O3 i4 o '15': 'yellow iris',0 ]- L7 n' f* n. B+ l0 p8 s
    '61': 'cautleya spicata',5 Z3 o/ F; }# n7 p  g
    '31': 'carnation',
    / |# Z8 i' m2 \- V% Q6 b '64': 'silverbush',* `7 M% x( N, _' W6 ?
    '68': 'bearded iris',* K, l/ m* V: K+ s; j9 J' n5 s+ V
    '63': 'black-eyed susan',
    ' a% s. j$ |0 k: P '69': 'windflower',
    # [6 `- ^, M+ N& v9 l '62': 'japanese anemone',# X( b" `. e$ E+ V/ l, k$ B( l
    '20': 'giant white arum lily',0 S# V; p' j* E: T
    '38': 'great masterwort',, i( K' P! y" j( ?& }. i+ t
    '4': 'sweet pea',: n( X. k$ Y  [1 `2 `4 i/ m
    '86': 'tree mallow',
    9 L$ ?( k3 F, s" @$ g+ ] '101': 'trumpet creeper',  p$ C  N% L4 b. P9 `
    '42': 'daffodil',
    , j: J) W) w! Q* @ '22': 'pincushion flower',
    4 \: j7 P# B. r% y6 } '2': 'hard-leaved pocket orchid',
    ( Q$ N6 H6 b. [0 c7 g; g7 u '54': 'sunflower',
    8 V: F2 E! ~! N' g  `1 p) `+ t: S '66': 'osteospermum',! F% E. ^6 a( @
    '70': 'tree poppy',$ b- ^, U: R  ~+ W
    '85': 'desert-rose',
    ! q& D& r8 q( o0 A; K  U! k5 [ '99': 'bromelia',. Q* n7 O3 m. |. m! g6 E7 J" a
    '87': 'magnolia',
    . n6 q8 |2 j# U/ c2 T '5': 'english marigold',- v. v3 Y- i9 e/ t$ |) P
    '92': 'bee balm',
    2 W( N4 F: W7 A '28': 'stemless gentian',
    3 z( i9 F7 O! R) l& E  t '97': 'mallow',
    - I. j' r$ z( b  k '57': 'gaura',
    9 D8 c9 @: \4 |7 d '40': 'lenten rose',
    " q9 k& j0 P; v. w '47': 'marigold',
    / ~5 f9 o1 V8 {: j '59': 'orange dahlia',+ T/ }. s; x' _
    '48': 'buttercup',- u) V- `+ v. V* I
    '55': 'pelargonium',4 F5 s" E( n/ N) ?
    '36': 'ruby-lipped cattleya',
    , f# t+ Z& F$ G  F- x7 A '91': 'hippeastrum',: k6 k- ?0 }9 f# x. }5 b9 s
    '29': 'artichoke',
    * _. S4 K! d4 T5 p7 J. Y '71': 'gazania',
    + z: j. b" L7 r* b0 \8 P '90': 'canna lily',7 p; _7 R. L" `/ w9 y$ i
    '18': 'peruvian lily',
    7 @) E( ?# ^5 \$ M6 z '98': 'mexican petunia',
    $ [  B3 Y" B+ n. \ '8': 'bird of paradise',! q$ k$ b) m7 h# a! Q& d
    '30': 'sweet william',
    & [9 F1 V0 a  J) n; k6 _ '17': 'purple coneflower',
    3 F1 Z! h2 U; F- @7 { '52': 'wild pansy',/ Z/ ]. a. l2 B% b0 v. V
    '84': 'columbine',
    * A, ~% J" U* { '12': "colt's foot",
    - o$ \2 f' F2 ]- i2 N '11': 'snapdragon',
    , f3 w# X0 r1 p '96': 'camellia',5 r4 a" L0 u3 e# ~
    '23': 'fritillary'," w; ]9 _+ T7 p( ]! ~7 s$ J
    '50': 'common dandelion',2 m/ Z# o4 F% {
    '44': 'poinsettia'," B! Z! O2 L! V/ V8 q# ?) _* @
    '53': 'primula',
    $ a5 `) \5 b) Y! L& ~( c/ p '72': 'azalea',
    & |5 U; d( A* u4 [( m: | '65': 'californian poppy',0 R8 ]0 _) z: [# s
    '80': 'anthurium',% C# n* ?4 j, G! j3 ^5 F
    '76': 'morning glory',) l# `+ \" c' K$ M) }3 U# P3 J6 s: L
    '37': 'cape flower',& Z5 z& C$ X/ Q& i
    '56': 'bishop of llandaff',
    # `& }  N+ o7 Q+ i) ?8 r: C9 i9 J '60': 'pink-yellow dahlia',; ^. F- D8 {. f3 l/ x4 l
    '82': 'clematis',6 s5 J. ]4 d; B4 t& I
    '58': 'geranium',
    7 A7 v- f6 v, p" ]6 c '75': 'thorn apple',
    # a4 |& f4 {0 j9 K$ ~4 J '41': 'barbeton daisy',! F. o% p3 H% H5 W0 Q
    '95': 'bougainvillea',! \* x0 c6 a6 P: ~1 n/ O5 X7 c" f
    '43': 'sword lily',
    ( a) S6 G/ `+ K5 T; { '83': 'hibiscus'," N# H4 y; V$ e; l( n0 [
    '78': 'lotus lotus',  _: O6 K4 S  {8 ^$ n
    '88': 'cyclamen',
    & |6 g7 j& o. p8 {+ } '94': 'foxglove',
    9 w6 X) j4 K5 ` '81': 'frangipani',  W' p( e. p: j' J- L3 e& W
    '74': 'rose',
    4 K5 ?4 ^( t9 P9 C- U '89': 'watercress',( s9 r/ I( C# |
    '73': 'water lily',# [. g& b# S9 D0 U7 T
    '46': 'wallflower',
    : d5 y/ [* M0 g. u+ p4 t- n! { '77': 'passion flower',* ]) @" n" q, V$ w$ j4 _# d
    '51': 'petunia'}3 N- b# `+ T* S5 C+ t8 i$ u0 V
    ! F# T, n# ?9 P6 f
    1; F4 p' t& C+ s3 h0 Y4 b
    24 N) u5 q( U% m4 w$ S7 S
    3
    5 k. a9 B  v7 E, m6 f4, q/ v, D* L' T9 b1 m. z
    5
    + r% `& |5 k+ k9 ^6: M3 v) d# w8 I/ K' m- t$ _9 K
    7
    ' r( J, d& E" n6 u3 X# ^! o) o83 Q0 t) a! y2 S8 P  e
    9- g. F2 k; {+ M. }% q2 |
    10
    7 r" o9 A& b7 C" v, n7 [11* ?6 X8 `  @/ l4 Y
    12, B+ |* x. T1 ~4 H3 ?+ r# g
    13
    ' h9 K9 _! B0 \4 i- U7 c148 z1 G+ E8 K7 n8 r4 g" M  K/ X
    155 M9 y" v- Q7 D$ u
    167 |* `& Q( S: w: p* O0 ^) q! Z
    172 I" q% [4 Q# ~) b2 |* b. |! {
    187 _  i  _+ U% w9 J7 F) x
    195 M$ v" u8 ]/ x. s
    20
    , r* i* f* M7 \& b* Z) o/ u* S21
    2 g2 `; X8 |5 a0 T4 `8 y3 {. T: o22
    ! n1 p8 o1 `. Z8 ~: w$ x233 ?3 R, o' O4 M8 a
    24
    / U# l% m8 x/ t' c7 l255 _' y$ u& `1 c, S9 g" ?3 Q! S
    267 u5 a; i( y+ B! C5 s
    27
    ' ]4 S# l% r# C8 q289 S; n( S0 d- D1 \6 y$ B+ d+ D
    29- M" A; T. E: R+ I+ Z- }
    30
    $ X4 u  G6 q+ |& y310 p/ Y' l4 p  _; T7 Y9 p1 S9 l
    327 J0 P* ]  [3 [- p  w6 h9 {
    33
    ' T5 p# T+ R0 L/ @0 Y7 B, }343 l5 N0 L$ y8 o
    356 Q# C9 y9 @! b! }% ]; V7 g0 w
    36, ?6 }4 \% Q9 b' s* @9 y. @$ x9 _
    378 K$ U& u9 [" g5 k" B
    38% ?" e' c5 g! @/ ]( m- q! q: M
    39
    , Y! Y% T. Z0 X) N$ q40$ i- y1 R% T; w$ f2 s% @) C6 }
    41
    7 {6 `* M5 w# i3 ?42
    ) R4 d( y* Y5 |5 U) `43+ r' V( ~0 p9 b
    44$ Q/ j) E/ i8 z+ W$ e1 T3 ~
    45
      L& k2 c2 d6 F: r& T. t/ _463 _) |3 P( t# A5 J$ r1 U4 C2 V2 p
    47
    ' @% [1 K$ G, K' U48  L( a% O8 W& L
    49+ o% |" q8 h) b
    50
    : O) ?  x8 M5 O7 Y51
    + {1 ^9 S  h/ `9 G7 J52$ a0 k# U! Y/ p+ n+ ^1 o) A
    536 A9 I; \! ]/ B
    54/ H7 }/ V0 O2 w! V* E! Z- ~% Q5 r
    557 j/ j0 o& D: Q; O
    56
    2 L/ h, ^% v8 J, n9 O  b57
    2 d+ V/ G) L) I  t58) f' T$ |! U* _! k1 e1 J
    59
    3 X0 g$ L# ]" G9 J4 H& S0 H0 T4 r60
    * l9 |% x! G# q: j61! u1 u! S+ s& O
    62
    2 H) D. u1 ]. n9 l$ J1 j639 |, `; o1 g2 H0 U0 \
    64
    ) M9 ]! R* q0 b65
    % K6 O: R) \9 Q5 A0 C; {$ k66
      g0 M) n4 I3 N. A+ N67
    - X6 h# G" |( P68
    " K- q  L2 o) D8 ~7 i- z; _  G9 P69
    % ~' `8 A* I9 F: E70* v7 U9 \" Y/ S. l3 m
    71
    5 {7 F. i( Q% Q  z$ G0 W2 o728 c9 U: C6 I+ z4 E  B  T
    73
    % h+ O$ O- v0 ?# O: g74
    ' \  C5 Q+ @6 b3 ~& z" `753 v4 a! d$ J& H' {, K+ ^( L. c
    76; ^; q8 P3 |* w
    77! B: q0 g$ e- O- l, p- W
    78
    2 e8 w3 @$ J7 e5 ^79
    , ~" Q" K0 o1 f# I! p$ @' ?80
    8 Z1 t0 D- [& p5 [81
    / X3 H8 V8 U% i- k' Q0 U& B/ O82: O; u0 G8 e( o: z/ \+ v" T; n* u2 x) U
    83+ Q( L* l& ^9 [0 l' d; v2 [- h: S
    844 T  R. e: Q6 B& E; D
    852 h9 ^; n; F0 L! F
    861 n  l) y0 O0 o2 B: |9 m* q* `
    87
    8 ?# _& {9 U6 g, k* }% }% o, ?88
    6 X' D" ^/ l; x) _( d' N+ ]89
    * o" [. b1 B0 g% X0 P90
    ) ~8 |, j& y6 x- k9 _4 H91
    + V- `$ h6 P( L; ^" j+ j92( E) O6 P9 u# L; o
    93
    # K) y! s, A) ]94# d, ?2 t" ^2 B. v8 K9 C+ s( C" I
    95/ c7 n/ m  I- h) \0 E- i
    96! _  }5 t! Q: V8 g" {/ W: W
    97. p1 ]/ v, Y" t
    985 W0 w0 V+ s* c% Q
    99
    , d* G9 x0 D: S1003 g9 }  C0 y5 [4 y" `1 v' ]7 I
    101& X$ _2 V9 W. |0 v% _
    102; K$ l9 ]6 h( `0 A
    4.展示一下数据8 B" j1 c+ `# p7 K- [9 n/ P+ z2 I
    def im_convert(tensor):
    9 D/ E4 x3 T- O. E    """数据展示"""8 ^3 e% o9 C( r* b/ O2 k
        image = tensor.to("cpu").clone().detach()( l0 q2 r- A* i& B
        image = image.numpy().squeeze()
    & @# K; J9 Z' c+ r  M- S    # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图7 [' A0 Z0 g4 D3 y) z0 C
        # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)
    $ U: U1 T& L4 V. w, O" O% M8 x    image = image.transpose(1, 2, 0)
    ! T- G" i; m" ?- f, Y    # 反正则化(反标准化)4 o; k+ ^; ^7 }- T! ], O9 l; j
        image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))
    / Z5 X5 z1 d5 |  ]) i9 c5 X5 X% U: n' D0 m* u4 |4 ]- A( W4 A
        # 将图像中小于0 的都换成0,大于的都变成1
    , R  l  s$ v* h& U" [    image = image.clip(0, 1)
    / x  f# g  F' C: \5 L. e, l$ j  ]- c  c) Z9 r* M
        return image5 E9 T: G* V2 S" ~, Z
    1  m; g# @* z( J8 q( E" f
    2
    * Y, N+ J8 e5 ~3! M$ c0 y. M  k0 f* Z! L
    4+ W& B, E3 v6 J5 s1 _9 J. U
    5
    - ^6 @8 G$ a& l; W! ?6
    ' u6 d( L( X- ^. Y/ [7/ W5 O1 H7 ^  L# }% Y1 v4 s# U
    86 n( }  i  X0 E3 @8 H7 ]
    93 b6 V# Z+ P& E) U; W
    103 \4 n5 m' B- c8 q
    11+ L- Y6 @* j: n2 A. T- C
    12# f; I. d% D/ ?0 E( M& `
    13
    " X+ r8 q. M; r- w14% n& s4 o/ m/ _) F' R9 @
    # 使用上面定义好的类进行画图
    3 K/ W9 m: U) V$ h% i& [fig = plt.figure(figsize = (20, 12))$ y5 R# }$ [0 c
    columns = 4
    : K2 n! i! Y: `) M& V* B# grows = 2$ L5 A+ c; g* u

    * ~8 u. `8 ]4 p# iter迭代器
    ( `2 a: b+ K/ E1 R: n! Q# 随便找一个Batch数据进行展示0 Z0 ^; E( `% e
    dataiter = iter(dataloaders['valid'])
    5 N9 P' }* s  f  J' `inputs, classes = dataiter.next()
    ! D; U. t. Y+ ]. x1 p& H* D" d; B
    ( A% b0 j) ~' n& q) n5 I4 {for idx in range(columns * rows):
    ) Z2 ]9 X8 J$ U! e    ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])" p) T5 X. f9 M4 F1 V6 j! C
        # 利用json文件将其对应花的类型打印在图片中
    ' |" L2 T/ ?( t% k( x& s5 G    ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])
    : I! _) D  U: G    plt.imshow(im_convert(inputs[idx]))  t9 p. Y# J7 |6 B6 g
    plt.show()' M0 e+ B4 P4 S! D/ T
    ) g3 J4 Z  {: j% A& [; T8 O
    18 \3 ]# W) m$ j: H  e  ~6 V6 i
    2
    ( ]/ M" F: {4 V. ?/ ^# ?3( u6 D/ l6 y% p  y+ ?0 _. `$ `
    4; k# x7 G8 G7 g* b) @
    5
    7 A7 y& `" W9 s8 [1 `; g2 A6
    + a, f" \0 c  ~. y& k0 a+ a7
    # q2 Y* l! n+ |2 u+ P& S# W( F8
    / R& [0 c% @6 A: |. O4 [9$ e$ b. p! F- J7 U% H& D5 U: O
    10
    * N% m' y  V. N9 k* o- L6 }' u5 [11
    1 F: b, t# O# c! b" y% D12
    : _6 b, `0 w% Z: T+ J  M13
    , q* H. O. ~" _14( s. D4 K4 l2 ]2 z: W0 X+ x$ o3 Q
    15( [: x  b3 l3 D5 I
    16! x* k; ~# u: v0 z' @" R5 q
    ( K1 i6 H. V" G! B

    6 K0 z9 [2 L0 M+ o( Q# f8 j, B5. 加载models提供的模型,并直接用训练好的权重做初始化参数
    9 [, e" A' F) l# i* u9 \! Qmodel_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']; |% c& J3 G5 {6 V$ H( y4 u+ o/ m
    # 主要的图像识别用resnet来做! _+ `8 H% o- Q: ~1 p
    # 是否用人家训练好的特征7 U2 T; Y- [3 v9 ?
    feature_extract = True+ j( K# {0 X; H' I* m1 v0 Y, I, m
    1
    8 I  Q3 c/ E1 ~" _6 D6 K2
    , l. m6 y; A; R$ Q9 t, Q' b1 m4 s3' [1 ^( |6 @4 R6 \% R7 m- G) I
    4* E+ g" _) g, f& }5 Z+ H
    # 是否用GPU进行训练
    5 b. H! I( }: T9 d# p, htrain_on_gpu = torch.cuda.is_available()
    ' h) d( p% A' l* Z. [0 F! e" q0 L/ i' ^& p  C# p6 d- F
    if not train_on_gpu:
    7 X( U' v6 E: u: Q3 p    print('CUDA is not available.   Training on CPU ...')
    ' q6 r7 z% z7 ^else:
    9 A2 }" a: F- B6 j. o    print('CUDA is available! Training on GPU ...')" N0 X2 Y4 L. W6 c! S6 h
    , h8 E+ ~! E9 [5 l9 P; {
    device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')8 E( I7 x+ n, n3 `5 k  P# J- q
    1
    7 Z. ]9 k+ _. T/ F3 v( W0 O2
    % K6 b# j6 l6 J) U- ~) i9 d% t3: W2 w! S2 o) P
    4, A9 b7 z9 \% r  S9 X1 Y
    5
    , z! u, C0 R8 m3 X! S2 f& H% |- h6
    . P4 i" [5 n5 ?8 E* T2 D; V79 G( i; J7 T" S3 c5 m
    8
    " M" N1 u' l0 y5 b2 e' j& R, C9
    # z# W0 N; B; e" PCUDA is not available.   Training on CPU ...
    * [1 Z' L1 o0 e% m2 ~4 \1: i+ }8 u  Z( m6 H9 \
    # 将一些层定义为false,使其不自动更新
    0 C2 W0 I* y" ~  _+ P  C8 `4 T( ndef set_parameter_requires_grad(model, feature_extracting):
    * N7 h; L9 u6 z1 Y" C! [    if feature_extracting:
    & @$ R! L+ ]9 ]; G4 N, }& X        for param in model.parameters():
    ! `; c. ^, V! V) w9 ]& m            param.requires_grad = False
    , k) V: T# t& R2 u3 Z# U* I. q1
    % Y* n) T( Q3 ~2
    % t+ |# P; {- B* w1 K3
    6 S' t+ E4 d% u; A( `4
    6 t# {# |0 j+ H. a2 |  y& }9 Q5/ M3 t6 i9 q( n: E. O$ N- _% e! h3 c
    # 打印模型架构告知是怎么一步一步去完成的+ S/ d# w$ j' @2 W( R, ?, |- h
    # 主要是为我们提取特征的7 F! X  `9 f/ x, ^; @! K

    7 a/ z( S. A$ R8 ^& ]model_ft = models.resnet152()
    6 ]8 G  Y* X  H4 s$ G( K! Wmodel_ft
    # d1 M3 @! q* X* O1 X/ \! b# J9 o1+ q9 j8 n! V2 V$ n5 E* f
    2
    7 w# E4 N" o1 l- J3
    ' `' s, w& N8 |0 z; v5 I! z4
    ! O- K4 S* U1 F- @! \+ r1 C' T* l5
    1 i$ `4 T) z6 |7 W/ cResNet(
    ) u  D# o6 N; j- b# b. p8 e; K3 }  (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
    ; H% _, c9 I, h  (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    - }* t  M# v8 P7 m- _& Z  (relu): ReLU(inplace=True), P4 ^  G: N: |" u6 \* F9 s
      (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
    % P; e3 b9 [: [  X4 a$ ?+ V( [  C  (layer1): Sequential(: r- ^* d/ c3 |8 F. I  N
        (0): Bottleneck(
    ( B5 \) C* J1 h" `      (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)3 j9 |1 p/ M- o" ?& o3 @6 O' S  e/ f
          (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)- q# j/ B4 V9 {: R
          (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
    # N3 S& ^. P% k# y2 `2 ]  {      (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)" \- n5 w  c9 K# o
          (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)- \/ }* O, o8 Y% B- U7 g
          (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True). |  L9 I- G+ x9 ?- s9 `( x
          (relu): ReLU(inplace=True)- u/ h) Y. e/ X0 @
          (downsample): Sequential(
    % V, F1 {; S6 ~1 A        (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
    2 t( x2 |5 V  F( w+ Q        (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)1 U% c+ x4 f9 x$ ?; ]. I3 X
          ); P9 I7 T$ C7 F* l" n9 e
        )
    ) j! M3 ]2 l3 L- w6 X% O中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。; k! b1 [! A; {* W9 }7 d
        (2): Bottleneck() ]) H0 y# I! N0 b; g
          (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)- v$ o$ F6 b  ?/ Z  b
          (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)  B1 _, ^$ N( e6 A5 X/ U+ e
          (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
    * n" q( A3 z! V8 l# f, d: q      (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    ! u2 ?% l" k( Z) j% q0 ?- f  P      (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)+ O; [; U, v9 I! e
          (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)& ]. B* D) f# l5 d+ O" u, M& r! _; s$ T
          (relu): ReLU(inplace=True)
    % R  p+ M2 y: ?/ i  F0 C, x    )5 y8 y) E6 ]1 a' l% z
      )
    7 Y- W  e, |& D& p3 B5 |  (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))# s: H5 K$ D9 T2 f* |0 p$ Z- v
      (fc): Linear(in_features=2048, out_features=1000, bias=True)
    9 Q) h) i+ |9 s4 @6 o)
    ; \  W+ s/ ?% K6 m; m
    ) Z3 f1 }4 z/ f* r, D16 j) w+ U: y7 t- f
    2" Q% r3 ]  b" I& d3 R. w6 O: t% v3 l9 R
    37 |) Z. T' W! l% {$ X( J* i
    42 F2 K1 ?. _) h  _  P1 k; L
    5
    # T$ `- B) o/ [) ^% B- Y6- j9 V& [) I" C5 ^/ X+ x& s* b
    70 D" J- T9 v' ]; g' u  L
    8
    - z" b3 q/ B7 d3 z9- W% A& d: `, J  ]
    10
    " d" c5 ]+ w) E4 B7 B11
    . a* u. m' y$ ~2 I9 G12% @; X/ |" Z9 J3 v% {, s
    13. s: b% T% S! Q  S
    14
    2 y9 p0 |* s, |8 K% I3 Z15. e5 l. {! X" ?- ?9 ~3 B5 q9 C
    16/ T- }, ?$ E5 M# f8 l, \- A4 t( X' s
    17
    : S- A, I% z1 R( {0 Y+ C3 z1 @18/ [6 C: m3 n/ j9 J- B
    19
    9 a8 O1 I- m+ f  ]20
    & E$ l' k" ]8 W- z! C- N21; c5 r, G  A( _. Z. l1 h1 o- l
    22
    2 w/ B2 o* m' i+ U& a. K/ _23
    ( o, m3 U+ q8 c% r. r# Y243 d7 C* o' z1 _! |$ O- m/ _
    25  F" I7 e' h" E2 l" b+ E
    261 \- _/ Q0 R: p, `6 v
    27$ l0 A# {+ u2 @9 _! G' E
    28
    4 N$ W7 _/ U0 J6 c1 [$ U. M29
    ( \7 N0 s0 e8 W+ o' |2 E, d30- ^% u/ k! y' ]0 _. b! j4 `
    31. O7 l+ E% N0 ?7 Z( r1 N8 z; E
    32
      [6 ?) L1 y6 \335 V1 L( `3 j  P/ w
    最后是1000分类,2048输入,分为1000个分类7 x  L$ ?1 T+ g3 G: W6 [- F  q. _
    而我们需要将我们的任务进行调整,将1000分类改为102输出
    ' q% H2 [9 H2 {' N) i; k6 k
    : P; ]8 j( W- y% U, K% B% f6.初始化模型架构
    # ?) S0 R) t7 V0 z步骤如下:- P8 S# `0 Z8 m; o
    + E4 Q; ^2 q/ t& W
    将训练好的模型拿过来,并pre_train = True 得到他人的权重参数
    ( u* [' g5 E6 f+ j9 J, L( r1 G可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)
    ; D  k2 O* n$ z无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数; I  }# ^6 w4 B/ _6 G" l
    官方文档链接6 k% R! o$ ?( m. G6 E/ Y  t
    https://pytorch.org/vision/stable/models.html7 \3 W$ @% b4 w1 b6 E
    & B- @6 ]! H3 @; ?$ j* v8 Q  m7 F+ `
    # 将他人的模型加载进来
    3 b; a! U1 x- v* S6 Kdef initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):! M% a, U6 w9 Z7 Z
        # 选择适合的模型,不同的模型初始化参数不同
    ) N$ @' D* p" `, ^: O    model_ft = None& f- }* {/ \9 q. ~1 ~' X
        input_size = 05 y2 w* u( ~' a( A. t

    " h7 \4 A9 e( r+ E    if model_name == "resnet":; e7 ?7 i$ a& n! F1 v
            """
    , j' ]& k0 d/ ^* P( Q  U. O        Resnet152
    ' D/ |2 b$ L& e4 {& d        """& Y- f7 Z  c3 r  j& Q$ p) O$ G- p6 }
    1 \( ^3 [& W* h  ~+ B( J0 r
            # 1. 加载与训练网络
    1 T6 ?* x# V2 {        model_ft = models.resnet152(pretrained = use_pretrained)
    ) @" m% F& k  x: f4 z$ J* _        # 2. 是否将提取特征的模块冻住,只训练FC层
    , |! ~5 F- e5 H6 ?3 l        set_parameter_requires_grad(model_ft, feature_extract)) d9 c6 N! `; P( f4 @  M
            # 3. 获得全连接层输入特征  q+ x+ ~. u! c
            num_frts = model_ft.fc.in_features" g7 N7 F$ A3 U0 t
            # 4. 重新加载全连接层,设置输出102
    , B5 J8 H$ [0 B. ?8 s        model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102)," v2 F: C) P' x% G) ?
                                       nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1
    # V$ J. u+ `* k" O% c        input_size = 224- g8 P% @0 ?* G. W6 i( h
    ! E5 V: D5 x! }3 f- z
        elif model_name == "alexnet":
    # W# Q2 _% U3 B& _5 {        """
    7 X. V& Q( ]& n* d- \        Alexnet
    3 Z2 }; i& {, [) s        """
    9 l6 ]9 C& ]5 ]+ P$ d9 M        model_ft = models.alexnet(pretrained = use_pretrained)
    8 A& f: @+ w8 u8 b        set_parameter_requires_grad(model_ft, feature_extract)2 w/ K- u" Q5 o+ \3 B
    9 |4 a1 F' ]* L3 ~& v
            # 将最后一个特征输出替换 序号为【6】的分类器, G% G$ T+ _' A. [3 g
            num_frts = model_ft.classifier[6].in_features # 获得FC层输入8 T: A, s: m; D' `
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
    / V0 L& n) z/ |7 ~. L  C; Z        input_size = 2240 J- b% V) n: {8 A
    1 r" y- k# k. }8 L" C" J3 }( T9 K
        elif model_name == "vgg":2 Q. T8 m6 E$ e4 i: V9 _! B$ T3 V
            """
    ) R8 i# R7 t6 }        VGG11_bn
    5 V1 S# C& R3 o' c% Q        """
    3 i* x5 K9 n7 c. ^+ @8 Y, e" \        model_ft = models.vgg16(pretrained = use_pretrained)
    3 S6 ^& p1 N. _9 J        set_parameter_requires_grad(model_ft, feature_extract)
    . H, [( |6 h% ?1 [* ?1 i2 _; u        num_frts = model_ft.classifier[6].in_features( Y6 M4 Y5 Q+ x9 U, J
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes)5 t$ w6 [7 h: D$ W# D% f
            input_size = 224
    0 _/ ^! G! q6 e9 z# n
    6 A4 y7 y+ ~4 u8 {  e    elif model_name == "squeezenet":6 \+ y& m' i$ h5 [. M
            """
    + g3 Q" m, J7 g% ]$ C3 b        Squeezenet
    ( x, t2 v. O9 ~9 x( ]+ a        """
    2 R0 t+ t2 b" S1 k  {        model_ft = models.squeezenet1_0(pretrained = use_pretrained)( m* f8 D" M; l( ?$ H* p
            set_parameter_requires_grad(model_ft, feature_extract)
    - b; @* N3 ]' V" ]- \2 _3 Z5 w        model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
    5 E0 G2 F. D9 V) e5 Y1 V9 X! R; `& G        model_ft.num_classes = num_classes
    # D0 l/ V; [7 J+ H1 L' G' k        input_size = 224
    0 b" f2 y7 }7 E  a4 h; I: {! ~6 v# v6 }7 n- @+ Q; G5 u" `
        elif model_name == "densenet":! o( }; o! S* j. K' q8 q
            """" V, H+ C3 v9 _) M- \
            Densenet
    1 |5 i9 |; L# f        """# p; d4 B* i4 \8 ?
            model_ft = models.desenet121(pretrained = use_pretrained)2 Q  K7 p& N9 L0 q" q2 N
            set_parameter_requires_grad(model_ft, feature_extract)
      |. {1 I$ k9 R" N        num_frts = model_ft.classifier.in_features0 }/ i5 ^) d8 _/ O4 J$ ?$ T
            model_ft.classifier = nn.Linear(num_frts, num_classes)
    7 m  {' l. c+ E9 `# G% C        input_size = 224
    5 M2 c+ R7 ^! y' r
    5 m6 i/ J: X& e! V. Z    elif model_name == "inception":, h5 E$ O& L& n
            """1 @* @; T5 M0 }( _! g, F
            Inception V3
    2 V7 W4 R4 @6 l* r) l- M        """. g/ @1 T" L( j+ i+ u+ @# U
            model_ft = models.inception_V(pretrained = use_pretrained)  [* R/ K3 @' x) h* u; P
            set_parameter_requires_grad(model_ft, feature_extract)
    9 x; Y, U6 f' [8 ]' [
    4 \. A" e7 w; V3 A        num_frts = model_ft.AuxLogits.fc.in_features
    . \) G8 k* w0 T. s: s4 o# l6 p, J9 |        model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)
    & F( P7 Z; O1 ]0 b- h8 _4 x' R5 f) N* l' P! ]- ^
            num_frts = model_ft.fc.in_features& v5 L' ]- ]" k
            model_ft.fc = nn.Linear(num_frts, num_classes)
    6 f% F7 L6 M0 j# s* ], O# @. S        input_size = 299, _1 g6 N# c- h, G

    : ^9 F; m! c$ U' z  M' m    else:
    ( V* m" K2 H! x  ~3 r7 d        print("Invalid model name, exiting...")
    9 N/ Y7 y8 x1 f7 g3 f        exit()
    ( K  l" J8 W5 h( A' C+ O, c+ v7 G* ?# G7 {  u
        return model_ft, input_size$ e: |8 ^3 n9 \4 z8 \+ y
    : Y: x  V& U9 Z& X  h( O4 M
    16 j6 r- }( j3 Y& ]8 h0 C
    2" g2 _9 N# z& v% D- M' y
    3. U$ e- D* U: A! s% v4 F: @
    43 \/ b9 a& O( e! C, o8 \& i1 b  j8 {
    5
    / f' x7 U) c* W6: R9 C+ h0 ]5 g, Q, S+ _8 J
    7/ D$ Z) \2 R' F$ S- ^# \
    8
    0 @. V6 ~6 p7 I. W( W0 v$ E9
      _' }) f  V* U* _; r7 ]4 v10
    0 |8 ?( l6 W  B6 A+ Q1 I# N11
    7 D' i: ^! ^+ Q' E4 l" ^12
    & M+ r1 }1 m5 p3 s" M: `" ^139 }6 z" L/ H, P# M. z2 n+ R6 s
    14% `3 J+ k0 _3 \# u
    15: F, q  f5 ^& s. [2 i# Q) s
    168 N& M" ~; F) ?- R
    17: K3 K5 }" y# h1 x2 U9 N/ \* z
    183 ]" L9 R2 \# Z
    19
    6 k4 f; n- I; c& d20
    # n6 \, ]/ e6 T& o2 x21
    ; `/ E: s4 ]' K$ I4 B226 G  c  ?9 b, r& ?  B, W3 x
    23* P- Q- Y! N# Q
    24
    ; O) ?' R; L. B9 T9 o5 J! P25
    % e9 z. k2 k6 a0 l26! U' a3 `( K( G& m
    27
    ' _, X& n2 z4 G" G9 x! d- q28
    + h8 E3 I. D7 a29
    " w+ B% h0 \" K0 X# M30  l+ |# B: }$ a. d- E2 E7 a
    31
    8 D: t  n8 d5 s8 V0 ]328 F+ c) {9 C$ W5 K( B" V- ^9 t
    33
    / r1 O- z& F6 `* E0 o% n5 G34' G& ^$ Q* g) X) P/ k8 J1 v5 l
    35$ P7 U* c% l8 R: X7 q: e1 z
    36
    ( D0 L* v) L% O* x3 H371 [! d/ [4 k5 s. E
    38
      h2 B, L" N+ o" L2 j39$ s' t1 J6 c$ l
    40
    . Y* A4 ]% H8 i' Y417 l  W+ O# [! ?% V% V
    42
    1 M9 g; }- ?! Q/ G6 c43$ l9 M* b) A8 g% y/ @( S! a' o
    44
    0 l  D" |* @. q; m1 j45
      {  g7 D3 J6 Q+ T468 ]; k$ Y: p; c5 K+ e- D9 O; k
    47, p1 g( B5 ?! z, K
    48
    % J/ k9 i! F: I5 h$ r5 M( Z0 A49
    4 [& `( ?% x+ O5 X; A. f) y6 z" y50$ m7 U$ u' k5 Z7 f1 O
    51
    ) Z" m2 r; g( ?% M' u. {; Q7 H52
    3 @9 C* ?  ]8 L53& e. N+ ?$ k( S
    54, s" }" |9 ^6 H& o* B/ s' M/ T
    55
    * B* e; W0 m0 s: Q0 f56+ B# u: J' g+ P+ S0 ]
    57
    * m4 d. y8 p$ }2 X5 x  T58
      E% l1 S. e3 q. w. X! V) x/ R+ r59
    % q- {+ n1 d% ]' \' }. e60
    0 d" E2 b) Z0 c6 c7 _& _: \) c61
    0 Q2 ?6 F! P( p' j% ?62
    ) X2 s% q7 u0 ^+ b+ J- S63
    ' n; u6 ~9 r+ i6 l" a5 R9 V64
    6 o* V  e! X" ?  n! z65$ T5 q# T% i" e
    66* Q, n5 ^) h: l) f
    67
    , T7 G4 L: n3 k0 m: s0 T0 l68& A6 R$ i; L. p+ [5 o8 I
    69
    5 v# n, C& S" y. J2 i& P# j70- D9 e) P# w. l# q8 t
    717 R  y+ f6 F8 h
    72
    : c# c. O/ T  _' F733 t# X* f; M, [" M5 F9 f
    74
    4 I; @" S- S& [; ~6 |0 o& k* C' O750 C' s) j) K1 u
    76, d$ ]. f$ X8 Z3 ?8 P
    77
    6 b* W! G" a" b! s783 O, e! J' G3 g# N( {4 L" y3 C
    79
    0 }8 @3 B0 F5 H' e( P( K800 z9 ]) w- o; L
    81
    6 |: r3 B+ X+ {7 Q! S: p4 N+ ~82
    % R  U& ~" A: W" V" z) X83- b. u! E+ w: R. z9 G
    7. 设置需要训练的参数
    7 U# t3 i2 w. x# 设置模型名字、输出分类数
    6 d; H: `& a: tmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)9 A7 V3 [% W& |& t% w
    # j% y" Q8 G5 e+ }: e1 E
    # GPU 计算
    # g5 b4 C- W& Imodel_ft = model_ft.to(device)
    # }+ Z: ]1 ?' B* \& j' D6 R" a( p; L' x# g; J
    # 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取, @' x: o# X/ p0 Q* O
    filename = 'checkpoint.pth'2 g) O3 X# q( Q; U( J/ F' X

    $ M( Y* h( _1 X# @0 R# 是否训练所有层
    . z; ]3 w) G4 @1 V( f, d, d" wparams_to_update = model_ft.parameters()# A" A9 S+ r& J# a) _
    # 打印出需要训练的层
    ' X3 Z8 I7 l! d! |7 ]4 Jprint("Params to learn:")4 _4 I& l) _% p! g2 X% x
    if feature_extract:
    / p" ]% C8 S1 E) O2 C    params_to_update = []
    / ?4 `( S$ U$ N. d( m. G    for name, param in model_ft.named_parameters():! L6 j4 B3 R- I$ m: E
            if param.requires_grad == True:1 B1 |3 c& E# O( ?7 E: h
                params_to_update.append(param)+ m% D% k3 A* ~1 h
                print("\t", name)
    ) l! _0 l4 b2 b& y+ Selse:
    8 ]9 d* D% Y- z& U    for name, param in model_ft.named_parameters():  `; ?! R# W. O
            if param.requires_grad ==True:$ s8 u2 B9 I* M9 _
                print("\t", name)
    5 ?- P! j* H* R( s1 H* \0 M3 Z
    1
    ; z9 `# |- F& D- E- z8 j0 u2; _& [/ u" K! f) H
    3. S: e, b- ^' s9 @6 K. T
    4$ q' j5 ]1 R+ s
    5+ A4 b0 P7 m2 A% S3 q, D, F8 l$ i
    6
    " D* b/ V, T" K; [7! w0 \8 k. [  [
    8
    3 ^  i; c: w9 R6 Z! f9) C6 t2 _$ s4 _- G. S9 Y. }- O
    10
      s* Z, ]# e$ _11. X, C) V  o' v
    12
    / v6 S6 l  [' P! `, f13* _2 U6 L0 W. |) m1 E. L
    14$ s! [1 B! b+ \* R: `
    15+ B9 |! g! T% w: M( w) m7 u
    16
    $ T3 ~* o3 {: E0 S17
    $ Z3 k' b( I' R7 ~& }18) l- [' i2 o  {+ b# Q
    19
      B# X* m/ _, E) A! v  }5 g20! ^' [, m0 k1 A
    21
    . j2 N4 ^! F% K* e& ]( w* f22
    1 Q/ A! \' x- X# y. `23+ b3 `+ b3 y# L# I
    Params to learn:+ _1 p9 A8 n0 n( \# n
             fc.0.weight+ S; d# G7 c* A
             fc.0.bias$ ~7 g2 e, J4 R% d( `
    1
    ( T( `; M) a! @* b! A. b2
    ! q, c6 [- c$ \3
    ) R3 f6 }8 w5 b; o% m/ p4 q7. 训练与预测
    & u" P, i7 g( [9 H: N5 K, t% t6 J7.1 优化器设置
    - P* t9 U9 V2 _3 P6 o4 y* E$ B/ M# 优化器设置
    # V0 M# \: `% J: W; u1 i, v! doptimizer_ft  = optim.Adam(params_to_update, lr = 1e-2)
    8 u' U) u, v  B/ V* j# 学习率衰减策略! @9 s2 n8 l( F! Q" l
    scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)
    , s* t& ^; b9 Q9 \9 i1 C3 {# 学习率每7个epoch衰减为原来的1/10+ a2 e" \& H& @  o1 o  i' L
    # 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算
    ; W6 Z9 e5 k& k6 r, `% o
    % d- t' z* a; x0 p: ?; Vcriterion = nn.NLLLoss()
    - `; n" @9 X; ]# t7 P. J1- O" @; l4 d8 C7 y: p
    2
    ' f4 W9 u0 Z2 W* M& A6 V5 V* F0 Z3
    * q+ C" L  e' {7 p$ [4
    . f" L" K; z, Q* X; U. D2 ~1 l5
    % h, ^: }% X. x# C  S# R6
    ' y* z- c: v8 g1 g4 U7 v1 t: z3 L7' X9 m1 \/ v1 H
    8! b$ c* O3 z: Z+ u- t0 K5 \9 [
    # 定义训练函数
    4 a( o* I* N0 B9 u8 U9 E#is_inception:要不要用其他的网络
    4 E' s9 `0 ^0 S: Q! h1 pdef train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):7 d: X, n* c$ n
        since = time.time()
    # J9 z6 D& K8 ?; ^: q1 ^! P    #保存最好的准确率
    " t. e/ w. S- o) r% q# q    best_acc = 04 R( x! W3 i4 i( C% q7 U
        """+ W& W; c, M: U/ n
        checkpoint = torch.load(filename)1 s1 M4 ~! i/ v" I  z. W
        best_acc = checkpoint['best_acc']
    ( v, F  O) I6 x' U+ K) @    model.load_state_dict(checkpoint['state_dict'])$ q+ ~8 f  x+ J2 ?" ~: \2 C
        optimizer.load_state_dict(checkpoint['optimizer'])$ i1 P+ g3 L& A' X" x2 m
        model.class_to_idx = checkpoint['mapping']
    4 U& i7 c1 j* |! [6 A    """$ w7 |1 Q& S2 B& D0 w% E
        #指定用GPU还是CPU
    + W$ N/ _2 p% y" G1 i, h    model.to(device)% g3 r0 f/ f2 t. t- [( _  R
        #下面是为展示做的0 h& J8 s3 }) B& O" `7 ?
        val_acc_history = []' m* w- G: t' Q
        train_acc_history = []8 V  ]. [1 u5 |* n: C- ?% |, t: n
        train_losses = []1 e' v$ @# a0 |) G+ a
        valid_losses = []" G4 S; Z: M/ @
        LRs = [optimizer.param_groups[0]['lr']]
    8 L9 z" G6 h, b# f    #最好的一次存下来
    ( T  m9 J" `: }/ o    best_model_wts = copy.deepcopy(model.state_dict())
    & a) Z, m! _: K3 c$ E. a" D
    9 y2 C; W. [4 e) ^) S    for epoch in range(num_epochs):
    ! A, ]6 Y6 |# W- X: p7 J        print('Epoch {}/{}'.format(epoch, num_epochs - 1))0 [5 y4 }; j0 B! u: Q/ z
            print('-' * 10)* x4 H* o) V9 ^0 H* n4 L. B$ G9 J, Z6 T

    8 z! P$ u& x% {8 ]1 ~        # 训练和验证
    + [* l! ~; g. H0 n2 m( C& h        for phase in ['train', 'valid']:4 \! A! k8 j7 R1 f! H* o5 D' @/ T
                if phase == 'train':1 q  g* I# \( N% n/ T
                    model.train()  # 训练
    $ u2 S7 x  w$ Q: D5 T/ }0 g/ B4 C            else:
    0 C! C/ r. q. R- G                model.eval()   # 验证
    - g' z1 E$ K/ N+ l2 Z5 Q5 I  @$ Y7 o& ^0 `
                running_loss = 0.0
    6 C  K, b/ f7 x            running_corrects = 00 O( h2 Z5 |  M- j! w5 O, u# N

    . P9 F% t. Y0 p. ~            # 把数据都取个遍6 {# }7 U  E+ A8 c" A
                for inputs, labels in dataloaders[phase]:# ~  b4 i% p& Y% [9 T
                    #下面是将inputs,labels传到GPU0 e+ N' c: ~- _( [3 r! j& S
                    inputs = inputs.to(device)
    # e/ g3 v3 F" J                labels = labels.to(device)$ _6 b) h0 f1 C) Y! X

    $ |# K: n7 Y; a0 R) F0 n! @# g                # 清零' M% j5 Q2 p& }- a2 M
                    optimizer.zero_grad()! T+ J# s9 B. E6 {2 B  m: ~
                    # 只有训练的时候计算和更新梯度; _7 U0 {) i6 V! i. d9 X( S
                    with torch.set_grad_enabled(phase == 'train'):& C5 |. V2 p! L7 R. q  }* S+ W& U
                        #if这面不需要计算,可忽略
    4 c3 b; j% Q0 n% K9 Q                    if is_inception and phase == 'train':* C3 D6 O0 P) [3 L' j
                            outputs, aux_outputs = model(inputs)5 N- J* ~0 h8 a2 ~' q1 G
                            loss1 = criterion(outputs, labels)+ [3 ~2 G5 @* ?/ D% M( K( L) _# g
                            loss2 = criterion(aux_outputs, labels)% F) x* L. L" ^  J# A# u: s
                            loss = loss1 + 0.4*loss24 O& F- @$ R0 C0 \" r
                        else:#resnet执行的是这里8 M6 G4 w2 n3 H& f
                            outputs = model(inputs): _( I3 h% _$ Q. e9 K2 @# w, L
                            loss = criterion(outputs, labels)
    8 B0 U: w1 p- A* y; }* i: w) W- v" u1 U7 s& |
                            #概率最大的返回preds
    5 ^. a1 ~. x2 U7 [1 b9 p  {# w( s                    _, preds = torch.max(outputs, 1)
    2 @% Z/ e9 S5 R  A' P1 g) ?: f
    7 l, q. e* @! E                    # 训练阶段更新权重7 T- L  S2 [+ N* T, i' ^- J- T
                        if phase == 'train':* Z9 }- [% g8 T; i$ N- ~0 m
                            loss.backward()
    ! g6 I% q1 o6 N* ~                        optimizer.step()
    + Y+ S4 {/ b' l( m3 K, @$ u* I0 D& g) |# Q
                    # 计算损失
    - d/ ^+ S& O5 W( I- r2 @& ?                running_loss += loss.item() * inputs.size(0)  }4 o5 ?/ {* o8 \: i
                    running_corrects += torch.sum(preds == labels.data); C* ^0 s" g6 L+ E( W& N

    6 {2 Q% b5 E' n/ I+ Y            #打印操作( j6 N- W. ]2 ?- Y0 Q! x- H
                epoch_loss = running_loss / len(dataloaders[phase].dataset)
    6 }! Z% \% M3 j' c0 H            epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)
      `( b5 D; y7 y- [! X8 |* S9 k, A3 p! \- w

    , I4 C7 a# |8 s* ?3 N: [6 B            time_elapsed = time.time() - since
    " Q8 W% \2 W; p1 X/ P/ z* t/ Z- c            print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
    2 X4 O4 J/ ]0 h            print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))
    6 [7 z9 u( V" E" j$ K  V. s7 n1 m8 F* K9 k# i% E: c2 B8 E

    0 b5 U0 M4 f8 _8 c" m            # 得到最好那次的模型
    , J$ ^' v8 e& t, `/ U            if phase == 'valid' and epoch_acc > best_acc:
    # A$ f4 X" a% u/ {+ L                best_acc = epoch_acc
    - W' {9 P: ]  @+ x) s5 F                #模型保存
    , |! z9 Y8 q( D                best_model_wts = copy.deepcopy(model.state_dict())$ Z* _3 G% W+ e. P& h( C& [
                    state = {
    " Z3 d) z: @* O                    #tate_dict变量存放训练过程中需要学习的权重和偏执系数
    9 |) a8 O" R& M" v- P( n                  'state_dict': model.state_dict(),7 Y9 b  C+ U5 f
                      'best_acc': best_acc,% U! ~5 K1 ^. u
                      'optimizer' : optimizer.state_dict(),. D% b5 C$ ^8 X- F6 @- B+ P
                    }; T; a3 n6 b  U6 l0 U
                    torch.save(state, filename)/ T) I- z. s9 G$ B! O# a2 l
                if phase == 'valid':
    2 R1 ~8 \' C6 q- C% D' j                val_acc_history.append(epoch_acc)
    ; T. A7 }; U7 U! ^. r                valid_losses.append(epoch_loss)
    ! R( H4 t$ P4 r1 {                scheduler.step(epoch_loss)
    8 t- ?6 o- g& a6 b            if phase == 'train':
    1 ~+ c/ y% d, E) j7 ^/ W/ [! ^                train_acc_history.append(epoch_acc)$ _& D) t- d% `% \# ?; e
                    train_losses.append(epoch_loss)
    ( M- j- k) C+ O6 R7 v0 D5 q/ m" A0 g7 O# B( a# E& P
            print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))
    ( G6 l, o, N5 A: p$ f5 @5 q* P+ ^6 W        LRs.append(optimizer.param_groups[0]['lr']); D  a/ S: t* B8 m6 @/ ]1 I# j
            print()5 Y8 Y' ^8 x8 P( X  \4 S/ f; a6 r

    ! J. v0 @; J  P! v7 x    time_elapsed = time.time() - since! l/ i, ?( U. X& E0 w; A
        print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))3 j! o! g  N9 x$ |- J; A
        print('Best val Acc: {:4f}'.format(best_acc))8 c6 [" ~/ {1 r% {  a8 f
    1 N: t: L: x" o7 A
        # 保存训练完后用最好的一次当做模型最终的结果
    9 t! y: }/ [+ s: M/ z8 @  K* B    model.load_state_dict(best_model_wts)$ x' U( g7 B, k+ \9 w
        return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs
    ( ^1 s) N& j4 d5 p- s, |  Y! p" e& }; r6 |- `6 Z$ H4 ?/ {. f

    8 U8 _: u* v% M6 v: ?10 y  d7 S( D  Q, z
    2
    2 \* c* X# Q+ b8 o3 N/ v3$ b: {$ e* z/ n4 Z5 z3 j
    4
    6 c/ b8 }1 Q7 B* m4 ^5
    : u% y- t" A- r+ v/ S- C% A6( q# d$ N$ D+ M4 [/ Y; }- B+ j0 m; T
    7; y; K. l0 L, m+ ~& P
    8: V7 g" Z% p. R1 |( s$ Q
    9
      V+ R0 @* u. z  i7 D4 Y. _10
      }# n% b& B; p$ K0 W  F5 }11, h* ~) X$ a9 e% j
    12
    4 P% z0 u, ~+ y0 f: z3 i0 N136 l4 m2 n: z" g6 h4 J& j) h
    14
    3 {3 C2 ]# g  W( R15
    9 c) ~+ N# j2 j* V16
      ^4 M$ L0 f9 N2 x  Q/ ~17
    2 D9 e& v7 s" \1 N8 P7 w: G18+ h/ f4 A+ z/ f0 B; t2 J
    19# |9 k" v! n% D6 j( p
    20
    8 P" [& D5 r/ e' N21! H" Z& L+ j+ K0 {( _5 x/ Y
    22
    2 j: \3 B1 j5 a& W1 w1 L23% m8 k" ]& }; P2 `
    24
    : f" n8 Y# V' t* o/ o, T# p256 X7 }# q  [9 d" N+ L$ }
    26( j: |+ o$ v1 T6 L" C# T) A
    276 }8 R+ x7 ~! y1 T/ S6 U0 E' _* z$ ?
    28
    $ X6 ]& M8 g9 U8 @- k29
    & `& o- c; y# b3 R; N30
    6 Q8 I/ h* W1 A31( i9 r% f$ s) s8 H
    328 d1 d. w3 }5 Y$ K# ~4 A3 P* q
    33" {' D* ^4 p* Y
    34+ r" J- E2 g% E4 E$ U1 E
    35% L9 H5 p* S& B2 E; ~
    36
    : b: D( ]; C0 V6 n! i7 P379 r" T2 d# r8 k( D
    38& T" Z' G1 i0 F! r
    39
    2 T7 d' \: k8 ?9 R: f0 h  d. V403 ?: T3 {+ p8 `. ~$ A$ X- @" e
    41: x- s- s1 T! F( R4 S0 r
    42( ]3 X2 t: w! j, W+ f) c& Q
    43
    6 L# b2 j: c8 I+ Q$ q% Y1 F44
    $ I+ z. r& Y# g/ ^/ K  V45% X7 G; b! ^( c: o
    46) W2 ]2 V: z6 y4 S6 K
    47( l6 m. W* }3 W5 t& ?
    48
    ) P6 d% E* @+ ?- Y  V8 F49
    + m/ i1 @9 T1 R; R- a6 w50
    7 a- U+ R4 u2 y- L' @51/ L; d9 b* f. D2 Z
    52
    ( I3 r$ p% _9 g4 b537 p0 s+ \( r# D( K2 r* j
    546 F6 ~2 y3 t9 v9 O5 h5 ]1 a
    553 x. ^/ ^4 i% c, a, \
    56
    & U0 Q9 M. j- [6 ~3 _57
    . G' V' @6 V* m58: }6 U9 _; Y( N0 S+ ?- {
    59; Y" v9 T* t2 t" b/ `0 ?6 E
    60( c7 _% h, e& ]) T2 {" Q
    617 z$ J7 Y8 a* U8 n' m' D  |
    62
    # X# D6 ~' U; f" M634 G/ H3 j( `7 t/ K& m" K" Z
    64; u1 Q: |. i+ j' `& y4 {2 X
    65
    2 V1 ~6 O  k2 ~2 V* ]664 v2 i* L, l2 k. A9 P; A
    67; v2 q) v# N8 H
    68
    8 a6 J/ n; O0 q7 i: Y69# q1 p- q9 `/ O, Y0 z, P3 O
    70
    7 D' B$ e! _. r2 S71
    & w, c" S) ~% A$ n5 }726 v& s- J9 ~7 v
    73. O( S+ G& x9 u# f9 ]4 M3 M
    74/ x' V/ y6 b* D6 \; o! Y
    75
    # ~2 i4 y# T( [76
    : O: W( d6 v. r3 C6 j1 C6 m77. [. f/ s+ h5 K) C: K
    784 B! H- l& r6 Q; d2 o0 s: ?5 j
    79
    / [% r, [8 V* _8 K3 |: `: H* K3 g80; k$ W) y" F0 Y5 |! y! o( A
    81% x0 L6 Z2 e4 w9 {7 G, j& }+ H
    82/ {' A6 \: M4 w; @0 y
    83
    ) \+ A- ^, m* M  q9 ?84
    ( r" ^. t. ~9 }$ N9 N' w85
    ( k" z4 k$ G* b; P" u86
    3 n& `  p) x  G) s875 @- E: `2 X. ^; J1 K
    88! ^7 e1 ?& i0 b/ v; x1 _: {$ l
    89
    & o- L* D: O3 e' ]& E90/ Y+ }7 E# Y7 h3 Y/ |0 ^
    91* D! H, d$ l7 E! K
    92+ ^. F3 |! Y! Z
    93$ W. E- {$ @% O$ u
    94, N" [: X7 }5 ~$ m' v) U- o. y
    95
    ' v. G& M/ m# ?0 h0 ]4 J* M5 W96
    - n) {" h" e, Y+ j! Y# F& D5 K970 w2 p6 P) p, ]9 X$ R
    98
    8 ]0 w6 @( T* c6 u997 B: ^. J$ y" o- A
    100+ X: w7 T) w' B
    1015 X, K! u% p# [
    102
    " g2 ~9 D7 x( ^5 R103
    1 b- X) E) o" j$ U) S9 e* z1044 b* \, F8 N8 a/ {
    105, R2 u7 y- O& H0 n+ O" y% \
    106" l" V7 N9 P0 K) i
    107
    + n4 `( [' D/ ]% t) I108
    0 V9 C$ z- P/ B! Z" M9 _8 z109
    % x3 ?, X% B9 Q110
      J& q+ a2 J6 A. f' m& H* h111
    : K  C7 l! H0 t) C; W% ~112
    , d/ ]2 |7 ~1 W/ o6 ?4 r7.2 开始训练模型2 e; K6 E  R3 c  `( b7 x  z, \
    我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次, j2 c. h0 O6 i' L4 ]( T+ t: c$ Z

    ) ^: O3 J3 D/ ~' {( K" S5 }% S#若太慢,把epoch调低,迭代50次可能好些& B; S% f. l* x, d: v  }
    #训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合
    - O  q0 v+ Z" N. z0 n  Qmodel_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"))
    , N8 w/ _. b! O8 R+ _
    " ?& k7 W/ g/ h( T) B1& A  e9 [0 x2 Q  G  H9 @6 Q
    2
    ' X/ P0 i9 ]2 B5 T3) g& U2 t* a4 V1 U6 j0 ]) y) ^" ^2 w
    4
    & k, e. }0 ]3 f8 P# u: s8 I) `Epoch 0/4
    ) b2 u/ B4 ^4 p( t/ j+ [3 [. |9 K----------. T8 k" t$ {$ c( U$ K) m
    Time elapsed 29m 41s& {, m1 t8 t- D  Z4 ]. c  d
    train Loss: 10.4774 Acc: 0.3147
    ! n* q5 }# u* N, Q) ~Time elapsed 32m 54s, c% {/ q1 T2 e$ y$ C
    valid Loss: 8.2902 Acc: 0.4719
    4 {+ C/ ]* Q( q8 r$ o4 T2 U  s+ gOptimizer learning rate : 0.0010000! b; O' o1 O; X
    , S. V9 W1 {$ v& }* l
    Epoch 1/40 Y6 }: `& P- \0 W: Y7 j
    ----------+ [; g% M3 l# c7 |  Q4 ~% V
    Time elapsed 60m 11s
    " y1 I' O, L7 E3 Y) B2 h( o8 N: Jtrain Loss: 2.3126 Acc: 0.7053
    & p6 U7 K. t. t) X4 a5 |) JTime elapsed 63m 16s. a! _5 O+ w9 l8 }& I6 [
    valid Loss: 3.2325 Acc: 0.6626; N; d& H& E  L( j" X
    Optimizer learning rate : 0.0100000
    ' v1 O$ d2 A! S( k  J5 Z. v) v/ f0 v* X: A$ n* d% q: }- g
    Epoch 2/40 N, y& j, S0 y, t! K3 v8 l5 i9 O
    ----------0 ~8 c0 x: b/ }' ]7 l
    Time elapsed 90m 58s- T# V$ F* G' j
    train Loss: 9.9720 Acc: 0.4734/ w  l; h/ R/ l% c6 R  T  r
    Time elapsed 94m 4s9 }+ z( S2 n4 j" y
    valid Loss: 14.0426 Acc: 0.4413/ E3 T# j* a' h3 G) @$ L  l
    Optimizer learning rate : 0.0001000' {$ t9 s( s. C: m2 n" {1 i

    ' i( Q9 b4 x8 {5 C: kEpoch 3/4% f3 j/ W& ~: p1 _: z; @
    ----------
    * _& E, \7 [3 k$ M1 `Time elapsed 132m 49s6 ]3 Y0 l8 i7 X3 y. V. T. g7 z
    train Loss: 5.4290 Acc: 0.6548
    / g6 y' c! S9 I# jTime elapsed 138m 49s* o$ W, P  v) V- E* a" n7 t; G! v
    valid Loss: 6.4208 Acc: 0.6027
    " r$ Y8 x" q. R3 ]& tOptimizer learning rate : 0.0100000
    % _. Q& `& K: C. y- ^4 v6 w8 k# ?' J8 j: s! a* R2 {6 \( [2 p
    Epoch 4/4
    " [1 V! @8 V) P, _$ }7 r----------, ~7 L+ V0 }9 @4 x
    Time elapsed 195m 56s% }3 }" a, L4 J. t6 K$ ?' e
    train Loss: 8.8911 Acc: 0.5519. m3 ]8 ]9 m; A2 O5 A
    Time elapsed 199m 16s3 ?5 q+ n% x3 F$ u( y: E+ K
    valid Loss: 13.2221 Acc: 0.4914
    $ l" I6 v( {1 C9 I& K: `4 POptimizer learning rate : 0.0010000# k9 a4 r. h  r

    , I: n; U2 i$ M5 Q! B5 V; S: MTraining complete in 199m 16s
    # {1 l8 R* P" j# E" ]- wBest val Acc: 0.6625920 ~. z3 N: a2 ~8 ?. {
    / A3 y, u" U$ ]# b# i4 @
    1+ y4 y; j+ x! N- k# e& l
    2
    # G4 X) C" d7 Y3 m3 B3) x' W2 `. K6 x/ }7 A
    4- p$ F* ^% l$ Q) M( m/ ?7 t# e2 j
    54 a8 a1 y% C4 {" m$ b2 O
    6
    " G. I! I$ u6 ?$ g1 q5 x7( Z# k& n0 H& H# [
    8
    ! N/ l3 Y& W  y0 Y4 [* D9
    " C( f$ {% y& `2 w10( ?- E+ Q# E* V
    11; c) l2 m: m5 X! L+ E1 i) U& N; X
    121 E1 I  e" B/ V
    13
    ! M, r0 u: L, ^# [' |% j$ W14, v, p+ i8 P) E4 k3 N- i
    15
    ' m4 I  q$ [: N  k( h168 |, |4 s7 e* N1 F- G) F7 d4 z  X
    17+ \( Q( C- E7 M! ]! W5 e7 J
    18
    ' g6 b- Q( i; U) [) O19" Y# r/ ?7 B7 ]1 |3 I% J
    20  d& y  k0 @- s4 W- u$ Q+ P9 H: n
    21
    9 V" Z: {& K8 E/ ~7 `# V229 h3 F' X6 R+ D- B' E
    23
    ' L3 P3 ^  t0 C- a24
    - l. a9 ^, k+ v4 h0 L& j4 L  V4 Z: d25
    0 o6 r5 ~+ `  k" F26
    ( [3 A6 P  _8 k  F. _  S2 _27
    $ _6 l* r2 u  D4 G28" |$ _$ h7 m! Z6 z% j: d
    29' m; R2 U- F7 U$ r& n6 ]
    30* j8 A- S3 {% P6 d3 M/ G0 c) @
    31, Y8 D$ r- A; ?- ^! V
    32
    6 {- n% m, G: t7 t1 ^8 S$ f337 K" Y  @" t/ `$ r
    34
      w: U/ F! i3 u& T( g3 U8 T35
    # F$ M( m; d& E* f$ ]36* U  \8 l) w& w: D
    37
    # }+ X+ t* L; J/ e38
    ! w/ K/ K: v" w' C" e' ^39/ I9 J( [& d, n8 b% f6 l
    40; M! N$ X% x2 m3 d8 n6 b
    41
    $ f8 i  \7 Q: F6 w' ]- R) I* h+ X2 Q42
    , |0 t4 [1 |1 z6 ]: K7.3 训练所有层! D! n' _; S; L; G, e8 g
    # 将全部网络解锁进行训练
    : b1 V) s$ w& G$ {' q# Bfor param in model_ft.parameters():2 ?6 B4 |& r: j0 d, E7 S- K& y, U
        param.requires_grad = True
    1 g. x! ]) G' _% o8 |# z* D; z9 w+ S) R. O- Q. T+ `, u5 }
    # 再继续训练所有的参数,学习率调小一点\& G& C$ m2 `! t' u
    optimizer = optim.Adam(params_to_update, lr = 1e-4)5 X; e' k1 H. ^6 k
    scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)
    9 D7 r- l% ^* G# W; ^, f) ~# `: ?4 n9 @7 g1 t: w6 s: t/ g7 D' }
    # 损失函数# t- z% ]; R/ k7 j6 k$ b- t8 S
    criterion = nn.NLLLoss()
    5 m9 R5 u5 |) Q8 R6 \8 u3 P1; D& v3 V! _& ~- y( u
    2, r& y7 U. O) x/ e; M: \
    39 \: n) E  k* ?7 c4 r. a7 ^
    47 O$ u. U4 q; U, V# i' N/ J
    5& }- G, r3 J3 t2 }. ?  V
    6
    ! d$ l4 k7 W; {4 Q7 D) @" z2 ^+ v7
    4 H8 [0 S) j2 f  L& U8 R* C8, w, U* W9 W! l/ ]
    95 \0 m9 N8 d* H7 T
    108 o( W2 ?+ y0 S& O& N" h
    # 加载保存的参数
    " e9 B# q6 U' D1 o# 并在原有的模型基础上继续训练$ Q, C* s+ ~; `1 h
    # 下面保存的是刚刚训练效果较好的路径
    ' O% y3 |% \! O; j% A7 [: scheckpoint = torch.load(filename)* i/ o+ G/ b4 y' d
    best_acc = checkpoint['best_acc']
    % `) I; v$ S: g+ omodel_ft.load_state_dict(checkpoint['state_dict'])3 |4 Y7 z$ {. D: h+ i
    optimizer.load_state_dict(checkpoint['optimizer'])1 [6 }, R' Z# @. C! t
    1
    # Y2 R; |; b% |* y4 z2
    ! e) i  `. A- x% g+ j" `4 [3( s! E3 m- w6 \4 Q  ]& M
    4
    0 i0 ~( x7 ?8 x5! F) ~4 m/ F) F  n& t+ n* C- T+ l
    66 P4 V# n; G5 P1 L
    72 c1 H8 s3 q) B. `6 D
    开始训练. \% C: K6 x0 Q
    注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考
    / o+ O4 O' ^3 V. e+ \& ?) N
    * l7 x1 H" R* 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"))
    - r* ?: O, p. e* U0 I  i: g8 D18 E; l* b4 ?4 y& B+ R
    Epoch 0/1- a0 [3 w2 _% F! W
    ----------# G% \7 T( r2 c* A8 v
    Time elapsed 35m 22s, ?# ?. C5 z$ Q6 g: E
    train Loss: 1.7636 Acc: 0.7346
    6 Y$ @) t4 s  S; y+ yTime elapsed 38m 42s7 U+ e. U$ D  q& L% Z
    valid Loss: 3.6377 Acc: 0.6455# b0 x! ^8 @+ ]: f' V8 L. I
    Optimizer learning rate : 0.0010000- ]" c- k7 f" E4 \

    3 k& n! c( e" ~/ M1 AEpoch 1/1) J  E" R( i% V0 \( y
    ----------7 @2 z* l" W, z' C6 V# |* _4 J* X, g
    Time elapsed 82m 59s
    ; }1 b& V# {$ E2 @train Loss: 1.7543 Acc: 0.7340
    # K8 t- |4 I8 ?/ X$ m' NTime elapsed 86m 11s
    % \1 @* L  s$ e; d5 \6 }( Bvalid Loss: 3.8275 Acc: 0.6137
    2 A$ [2 t+ m8 A3 n8 W! I1 r: rOptimizer learning rate : 0.0010000
      D) m1 O# y5 d/ U
    5 y6 D. {6 t3 v6 {- iTraining complete in 86m 11s3 K" g! D# ^  t/ I% v
    Best val Acc: 0.645477
    6 l% ]& t6 ]7 P+ c
    # n  f8 }; j  q2 [& u; f0 T3 U1
    # C5 U5 J, h  K% Q  ^2
    * U! I9 d% o  x% \' [) B3; F" V& v1 O9 ^2 T: l- Y
    4# c, G/ }2 u+ a. K# V" a( ~" X
    5" p5 Q; b" S  `2 ^3 V
    6
    % R$ F) U: L( F. V9 Q; q7
    5 l5 T, i7 q. e  E8
    ! @# w# @5 ?& B/ m5 @/ Z/ v9( E% B" _* n; G# Y$ Z$ g) L  _0 O
    107 U8 u& ~+ f. r( r1 b8 m3 o  d! |
    114 u* w1 ], Q/ J: w9 [0 R' N0 i/ H
    127 l, c+ z  A' u3 d& n" k* Y% a% F( o
    13# Z' K6 Z+ u* u9 c
    14" e2 N7 \/ ^, G0 M! v0 ~
    15
    7 c" i5 R/ `5 m  j- P# V, \16
      w1 D$ t: W- P0 O  T175 Y2 }0 A* p2 J4 J+ u
    181 `6 d+ A7 |& C* k; Y8 Z
    8. 加载已经训练的模型6 R1 Z6 }1 y% s" u- ]
    相当于做一次简单的前向传播(逻辑推理),不用更新参数
    ! m% I% d: W2 D; o; C  i) Q7 S) t% x& o9 D6 D6 Y7 L
    model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)2 c; z+ ]- F6 M6 E8 W$ J
    7 X+ H, t7 t/ _. Z; T
    # GPU 模式' C, }4 T0 @0 e# B
    model_ft = model_ft.to(device) # 扔到GPU中8 f3 u. b/ k' Y; @1 A5 ~2 G  v! |

    % Y# `9 p& \0 W# 保存文件的名字3 {' I7 u- g* Z
    filename='checkpoint.pth'
    4 W& y* v! S' p4 b/ l& u. `0 ], p$ ]. W  L9 s
    # 加载模型
    ' M- B" a: \6 ?5 C7 [1 vcheckpoint = torch.load(filename)# n* s: h$ b2 z  c5 A; V
    best_acc = checkpoint['best_acc']" f- V; Z9 g) h7 G+ A( K
    model_ft.load_state_dict(checkpoint['state_dict'])
    " J! G$ B% M) L; w2 {1 j1% s% X/ ]; w1 U4 z$ C' m" K
    2
    " B+ R5 ]8 K" p- R3' P; e& f8 @- a/ Y; B0 i+ I: I
    40 [, X- q/ T0 O1 z- @* R. V  d$ [
    5
    : K6 ]0 c7 Z0 R' y8 Q6
    1 i  R3 S( P0 f% O+ ?7
    % P- B/ O7 Z. F/ Y2 E4 r8
    ) D3 M) V' F' @1 j2 @0 \2 @9
    * e  T. V) \$ H$ \106 o9 P: Y/ a  c: \! h6 k' H
    117 r) o. P" I* T
    12
    ; o; n6 m! }+ I4 {. [<All keys matched successfully>9 @, |6 o4 ]0 ]0 \% Q; d
    1
    9 K; f) {( x1 Adef process_image(image_path):6 C4 ]3 c0 k* R, e6 C
        # 读取测试集数据; \! s! e: t9 ^/ @! G4 W
        img = Image.open(image_path)
    ( i* ^% W1 o/ R5 z: B: s4 o    # Resize, thumbnail方法只能进行比例缩小,所以进行判断
    9 d! d; P; M  P1 K' F1 [! z, v1 t5 q    # 与Resize不同
    & H! X" j% c& w/ k0 T: o5 U    # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小
    . w3 F. J* X6 E2 [  g    # 而且对象调用方法会直接改变其大小,返回None
    0 Y0 W( t* |) v' Y3 X    if img.size[0] > img.size[1]:8 i3 B0 V' x) R: y7 e2 h9 t5 r! v% g
            img.thumbnail((10000, 256))
    # N, e) @- f+ N  ?    else:/ U7 c3 s. z# s$ i: i
            img.thumbnail((256, 10000))6 `5 ~+ N0 }5 x- O! F/ \% a

      F6 v6 t+ `0 b* H, h9 `( W    # crop操作, 将图像再次裁剪为 224 * 224
    ( o* k; N1 D% m0 y7 A# A& ?    left_margin = (img.width - 224) / 2 # 取中间的部分
    9 R& Z/ |: e# ^" o    bottom_margin = (img.height - 224) / 2
    5 K, ~# I8 i' M    right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度
    - k6 S& Y; |; f0 f( A% H/ {. w% {$ ]6 L    top_margin = bottom_margin + 224
    " }4 P6 z$ ?) j* e# Y& a4 y# d. T- O8 b
        img = img.crop((left_margin, bottom_margin, right_margin, top_margin))
    ' t6 g% F  W- L# i* C
    9 ^0 j( D& m% V; V8 u    # 相同预处理的方法* v2 Z+ s4 N* e% T( h- W
        # 归一化
    6 m5 N! R* }* h2 t: z2 W    img = np.array(img) / 255
    5 z# T5 b, G6 o7 }    mean = np.array([0.485, 0.456, 0.406])0 M: @( f' ?" T7 S$ X) j
        std = np.array([0.229, 0.224, 0.225])
    4 i* Q' S4 C# m6 q7 @    img = (img - mean) / std
    ; c% Z. ]6 m, C2 Y% e; d! x
    ( B- l+ m( @+ e1 |) E) Y$ M3 ^- B    # 注意颜色通道和位置
    $ @4 a3 K: E" I    img = img.transpose((2, 0, 1))
    6 z$ ?/ F$ v' R: d# k% s8 C/ {
    1 R6 x6 X9 ]: N) I" n! K    return img9 F" \. R' F% \% p2 }' U; n

    2 J' |/ k; n. o) O% Q' qdef imshow(image, ax = None, title = None):
    ) ^* W, T8 P, j0 q    """展示数据"""
    + f/ r, ~3 F  D  t) I0 h+ N    if ax is None:  F- e7 A9 |5 s# G# b" u, z) [% E
            fig, ax = plt.subplots()
    1 w9 i/ u4 E/ R3 H) C1 V" z6 W) I+ z8 t
        # 颜色通道进行还原
    ( ~4 K1 d0 |1 S# f0 \    image = np.array(image).transpose((1, 2, 0))
    $ [6 q1 k4 A  o% i. d6 P* J" A7 X, y0 Z* G! M2 Q
        # 预处理还原" {: d% b& u  j) e3 G: j
        mean = np.array([0.485, 0.456, 0.406])
    " Y7 o2 U4 c; \9 b- O: r    std = np.array([0.229, 0.224, 0.225])& r& w5 R* C( L  l9 T3 A3 r
        image = std * image + mean
    8 U, }+ C$ @) ?/ q5 S1 g    image = np.clip(image, 0, 1)
    ! X) e+ \  d6 l" x1 k" x& M5 P5 o* o4 g* f9 D, Z/ g
        ax.imshow(image)7 L1 U9 }+ {1 L) ~0 I, B: E
        ax.set_title(title)
    $ u( @$ x  B0 {, T/ q
    4 W  \) b1 Q* |6 C4 s    return ax
    4 m8 T8 T0 D2 {' D: J% u3 X6 g6 l) ~6 K& G  b# y7 h
    image_path = r'./flower_data/valid/3/image_06621.jpg'
    & G0 p! n5 |6 r3 o: B- L. uimg = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理
    . [% B; n3 ^. p/ C! kimshow(img)# |8 f, b' a  {4 t
    ! j2 F4 A2 _, \: q" q4 \( S
    1
    0 a% G' V% h% B- K# E; C2
      ~1 v1 D$ J& x38 E$ G& N# a2 T6 A3 a" q, c/ g; |
    4% ?% F  \3 `- ]7 e! A  V
    5
    , Y! V# s2 ^9 x: n$ r' u6# \# D0 |6 u4 y( w9 ]* E3 i/ C
    7+ ~' F' o* ~- ^- u0 X
    8( e0 ~5 c% N; T! {" S2 a/ I" ^+ M% A
    9
    * K) `$ L9 ]7 _1 b10
    9 Z' x0 l8 E, T: B& @. T7 D7 ]2 s11
    , a* {+ ~$ w6 A12
    $ `8 w5 f6 Z) g13' F8 @! u+ ]7 k
    14# r1 N5 j" A; j5 B# e1 t# U8 B
    15
    # |! M# K( j, M* y' F16* ^, X9 z1 a- m
    17
    9 A, i& W9 ~0 g; y: q& f1 k4 z8 Z18
    / H3 T% {7 H7 N8 C* Z1 {19
    * o; ~2 h1 ~$ ^8 B201 z0 {7 m+ C; @) ?5 m% F; M/ ?
    214 M+ O. ^% E9 U) n
    228 x, }: G1 k* Y; W; B) Q
    23
    2 l- l, t) c/ s+ G: p, p24
    * q9 ^3 _' J. x, G5 ~/ B4 a8 t$ {25
    ; D$ K$ {) M* a- U( T! {: Z8 c+ j26
    8 W7 j# l& N! }4 {27
    6 s( F2 y( o% q& @1 K28
    3 r: r2 A, n& R- A: H298 c# v8 q2 k# Y. f% B. [7 G% j& F9 C
    30
    2 b$ Q0 X- ?  o# ]; O* O$ A31
    3 K' x, f+ K! W; O; |- K7 V. b32
    ( a& V) ?1 o! E) W33
    4 h4 c- G' y9 _9 `34
    . W7 _1 b+ n0 B. r) n35# U) q% t  v( C- a! L
    36
    % o  O1 |- v9 A0 j* U4 y3 A- N' t37- s7 ^2 F$ v  x  |& \) P9 ]- d: T! g
    38
    & r+ ]8 V: q7 R0 c39, z3 N' E& ?( J, C
    40  Z" H/ V% y, X' f8 O! q6 G
    41
    8 g$ ]5 f3 q$ T5 s# s  j, R* o  Q8 P42
    : I/ A1 C* ~/ |432 A) U* ?9 G4 P* N9 b5 X
    44
    " V7 _5 J3 b8 @7 U0 E45+ t# u7 `- g5 u3 @7 U6 c
    46
    9 N1 |4 N, ]2 v47
    # e/ A" s8 I# q2 S: t. ~48
    " a7 D, z. }* d: h/ ^49
    ) }$ h' v+ T& x9 j6 B4 v+ i50
      f' r6 f5 B4 F1 A" P+ @51
    0 a/ `5 s2 X* R7 c, r5 X520 z1 `1 g4 l9 @% |% S# x# |4 ^7 G; L
    53
    / C; w3 D& v2 m6 ]% _54! m$ i& x+ A& e4 s4 o7 `; f7 `0 n
    <AxesSubplot:>
    4 X7 c" U6 L4 Q( u1
    / X! o( E7 ^1 m  b# p' g
    $ [# C/ y$ O$ e+ o( X. M, x! f上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确$ E3 t" G' ~7 M2 d2 |

    ! \5 h/ L  G; `' s5 p" r$ X6 {3 m% [img.shape" A5 v2 `, r0 F+ S
    1
    " A( }" W: b5 E% |8 m) j4 q3 F(3, 224, 224)
    1 I* Q$ i8 B- J" I1
    . s, g% W& `8 J3 I! Y' \" |证明了通道提前了,而且大小没改变
    * r* l7 Y& Z6 ^4 Y, [4 {- u2 f) {
    9. 推理4 ^" F# |  p, o( Z- q& b4 q
    img.shape
    ! ]3 U7 i' i" i6 e( `& o
    4 K2 }4 P) m1 f3 D% G# 得到一个batch的测试数据
    6 I% I( E! f; f. Tdataiter = iter(dataloaders['valid'])
    3 n" Q+ Q4 Q7 T! }images, labels = dataiter.next()
    * R: G- K1 I8 d" B, w
    1 Q5 z; b/ m9 y6 r) L% v( tmodel_ft.eval()" o( @" Z7 O* t: }* B
    7 N" Y! Y. j/ k& l
    if train_on_gpu:% [* w2 t% m5 t7 o5 M1 [
        # 前向传播跑一次会得到output
    / i3 l1 ^7 c# N6 o' ?& @$ k' y/ ~$ c3 D% Y    output = model_ft(images.cuda())# |2 l' S- L, B9 n/ r  k- j9 x
    else:
    / |8 b; x! N+ s9 L" |! r    output = model_ft(images)' E9 L0 o! K7 Q, N5 E: `% j

    9 Y* l" c5 F- z# batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值
    0 s% }7 G: R  @3 @2 @output.shape! \' _+ |" v# F; i5 o

    + v" H! t* {8 c! Z  z5 h1
    6 u% L* {. ^6 F& u* X$ h- }  c  Z2& l5 a3 j' e1 _5 Y
    3
    3 i% `' N2 G2 R9 L* b. J% S9 J4
    1 z# F; y' [7 I( {5 C) R8 ~. @5
    6 M6 G  W% B6 v# W1 p2 }# ~# Y60 h  F' \2 s* X
    71 d8 [+ v" A' w2 B* v% H) N
    87 @# |" @: _0 y
    9* x2 `/ ^  C$ Q# x9 O  n# D
    10# z2 G9 z1 j; F! p3 f; G
    118 _# U: h& |$ ~$ c
    12" |% ~; B5 m- }, P) N+ R  F6 k
    13
    1 ^9 B$ H! I* t3 g- ], z14
    , g8 P/ S' Y6 c2 l" s9 [- a3 D15; r) o( x# F2 W  C5 A$ b; F
    16
    4 h( O9 D( g' q1 I- Utorch.Size([8, 102])$ B% F4 Y8 r6 K" ]- {, K
    1
    2 Q% }" E: w" r+ ^% b9.1 计算得到最大概率
    ) w6 [2 P* c/ A8 V_, preds_tensor = torch.max(output, 1)
    & M3 \  C) {; i( f" G8 H6 l' J0 a  F  @# @- ^' D2 B
    preds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量/ s8 z, z, J& o, |+ X* q- V% c
    1
    # v. J# {$ i: o! |  g, \2
    & a# g' R, h+ b  B3
    6 T% y5 m( J/ G9.2 展示预测结果' v! a- |# y; t! U/ g1 i, V
    fig = plt.figure(figsize = (20, 20)): J3 ?6 d: e8 I9 Q4 t
    columns = 4
    7 r# P( G3 y1 Orows = 2
    2 c, d6 L0 _& |8 {; q0 _$ w# L: |8 |% h( ?& S* Y* A5 z1 @; e
    for idx in range(columns * rows):% Z9 w0 n; v4 n& r
        ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])
    6 n* b, R4 B) i0 |3 G1 b    plt.imshow(im_convert(images[idx]))" T" _+ D* y7 @3 _- r
        ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]),
    # H; z6 u3 Z: ^( }; }                color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))& Y. F# M9 r( k! m% y
    plt.show()6 m' m, S) l& H8 d% O
    # 绿色的表示预测是对的,红色表示预测错了$ n+ q! T& w4 t8 e7 y
    1; b" a( {% B# {( c: V- Z* o
    24 _7 a! s& c$ Y* g( |) Q
    3; d- o5 z; ^: r0 N) e
    4
    7 `7 {+ V/ W" l2 u5 X' C$ C5
    0 R- e  A' N6 g- g1 f: x* ~- o60 T1 o+ z4 u6 a/ p
    76 {) N9 J) h4 P! }9 K4 a9 M' \
    8
    ( _4 x5 ]" v9 Y- l# F; z: E8 ?9/ W$ [+ y* w% Y1 \; o- i& ]
    10+ \- X* m; S2 b
    11& X& s+ S2 t4 D6 j7 B
    ! Y7 x' \! B/ `, \. X% L% w$ |
    & F* h6 ^: r6 A$ U& k6 V3 ^

    : F& \6 z2 l( d/ [————————————————
    9 B. o/ ]5 R. \5 |, t版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。
    - L+ p- Y9 o. ~( K# y原文链接:https://blog.csdn.net/LeungSr/article/details/126747940
    5 w  n# [6 _; g2 _, G; y# B8 c
    9 d( b1 g' e/ T& ~  A3 m8 Q' x% _0 v/ k. g# A" E9 h+ H
    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-24 17:03 , Processed in 0.479738 second(s), 51 queries .

    回顶部