QQ登录

只需要一步,快速开始

 注册地址  找回密码
查看: 2773|回复: 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)实战案例% k1 \- y' M( M2 A9 R
    + r1 N3 f  O4 m% K
    文章目录8 M+ U. S$ ?% Y+ T4 d0 y
    卷积网络实战 对花进行分类0 A3 e. Q* U4 w
    数据预处理部分" C- ?7 k/ W( i1 C
    网络模块设置
    5 n& ?* R, M; h3 B' y4 }  G: Q网络模型的保存与测试
    ' K" e' ?+ l2 O数据下载:
    2 n$ {+ x: F* V  w$ E1. 导入工具包$ h+ O8 ?9 E0 t
    2. 数据预处理与操作9 @8 c* D5 i/ |, w- ~7 R) S
    3. 制作好数据源
    * ]% W6 O, a0 L; |读取标签对应的实际名字
    ! m. t- p$ E4 r4.展示一下数据0 _- }  _, e- S% a* G" K
    5. 加载models提供的模型,并直接用训练好的权重做初始化参数
    " O4 E& j+ \: c2 o6.初始化模型架构$ l' a5 F9 w, M: Q4 o, l, Q/ ~0 t! o9 T
    7. 设置需要训练的参数0 q1 {4 q; a9 }0 k0 z3 K5 h
    7. 训练与预测
      ?3 O# M1 n3 H* n2 O6 G- C7.1 优化器设置& [' E* j0 L' i7 k; R1 H' o
    7.2 开始训练模型  e' F$ A* V1 e8 Z' n
    7.3 训练所有层
    4 |  L" f# q- B开始训练
    1 a+ n  F9 c3 O3 R8. 加载已经训练的模型
    * {  I5 a8 L" Q, t6 X9. 推理: P5 e. c4 C: `, ^6 o7 y
    9.1 计算得到最大概率: Y7 D+ l% }6 B
    9.2 展示预测结果
    ' N" ?5 I8 b0 q0 h- a+ \写在最后* x7 [- i/ S2 i3 w/ K7 f) ?3 L8 u# T
    卷积网络实战 对花进行分类5 p( t- _- F- k, N3 G5 E
    本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能
      `7 O& c9 k2 N; P' _
    ' D6 R$ L6 m7 [5 E  f5 ?在文件夹中有102种花,我们主要要对这些花进行分类任务1 J8 K3 B) ^# \- Y3 F+ ?* R
    文件夹结构( m0 B. P% ]1 n5 g& N, J0 l: @' N
    $ r/ G# w: ?* W& W
    flower_data
    ) H! _2 N4 h6 c! A3 C$ b2 L. f8 p- K# k- ?  X& ]+ j$ `6 m% R* H
    train
    * S6 n4 y, `# B0 R2 y) C
    ) F, P# D9 u, c5 q4 Q3 }  b. a1(类别)3 r" [3 `, u3 M- P! h8 q' K
    2
    & m6 r0 C' \% O& mxxx.png / xxx.jpg
    4 K5 j. L' _- `. f  q2 F/ N& uvalid7 @" H/ Z1 v4 |8 `2 S
    # t' q7 I; C9 W$ c$ O6 u0 M
    主要分为以下几个大模块
    ) d0 ~/ J' l. [, I( \; X8 m0 G) }- m6 s5 ]9 @2 u) c
    数据预处理部分. [& M8 y8 s- U$ `6 q+ _: U
    数据增强% K, a. n5 ^. J' G: a: _
    数据预处理
    * d+ b0 k) U# f  s网络模块设置
    4 Y+ T: l4 j4 h9 w: \8 r加载预训练模型,直接调用torchVision的经典网络架构& P% f$ R" P* T3 o
    因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务  a- G1 q8 o* v; g
    网络模型的保存与测试
    9 w9 R. \4 F4 [模型保存可以带有选择性
    5 L8 @# o/ \( w( s6 u+ }数据下载:
    4 w' Z5 v1 I. F& H1 Ahttps://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset, e7 i. [8 [# y7 z
    2 P6 H6 b( W, r
    改一下文件名,然后将它放到同一根目录就可以了/ s. x; M. S, \/ r$ d2 y+ z
    ( h  V" d  Y( R# w$ V
    下面是我的数据根目录
    % j, h! r4 |8 J& Q0 a- h& g; S8 T9 D: X# S" ]# m7 k5 i, w
    5 H9 N9 O2 P5 ]+ B- D1 i2 x2 x" O
    1. 导入工具包
    7 O; G  r( ]* |" r; Limport os
    - Y: i+ y7 I% O9 |+ ximport matplotlib.pyplot as plt! ^! g6 y; F  Y$ B" ~
    # 内嵌入绘图简去show的句柄
    + s; b3 L6 Z# z& f3 z, g9 X: s  q%matplotlib inline
    6 x- w3 N6 O0 k! I" ]% Ximport numpy as np) P2 d1 r' J2 j2 A% g5 B0 I
    import torch
    9 S1 ^8 F1 x9 F2 \# Q) pfrom torch import nn
    2 ~& C' G, h3 x+ C. m- z+ R+ G4 Y/ i4 Z# }8 _  b9 j7 S* E  }
    import torch.optim as optim
    7 v" W' p+ {" ?) e$ U7 |, m& limport torchvision# p1 [: [2 v, B* [* p
    from torchvision import transforms, models, datasets
    " S' A9 e$ a- R; n. t5 t) h
    2 f  G# M6 a) T4 Qimport imageio
    . @; Y2 y1 E. N! Iimport time
    : r7 c8 r. D# O7 e6 U: Dimport warnings
    , t! \! ^/ E8 y- t4 u" bimport random
    ) c7 }$ p9 S) t4 @3 l1 simport sys
    8 ~+ y7 O( N/ h1 {" Q% Kimport copy$ S' s* t7 h# O% o# z8 e! ^
    import json$ ~0 C1 I+ y5 O
    from PIL import Image: [! v- J8 C; z# a! c- K2 g+ w+ Q

    % t6 O9 v4 p! x9 c* H1 u, O
    / z2 M5 }2 ?& `& n2 R1
    # K% e/ [0 Q3 w2
    ' E& s  j8 o: n( |9 Z( R( {* P& f3
    ! ]. U( ^( k% b( E# m  h; H4$ z* v* [; y; \" q
    50 v' a+ W5 f& e# O* Z' ^' ~4 T
    6# |( K6 Z: ]8 \7 H$ O$ ?1 t
    7
    ! k# c1 c5 V( c  H- o8
    9 o& Q: _; Y9 c' D4 _9
    3 b6 {) T2 O# a5 L7 F& q10
    4 v" m7 c4 `3 d9 b4 D11* ~! R" }! o/ R6 _3 n# _8 U. i
    12; a/ B1 H* O4 H# E7 V( m
    13. P6 a( T- ~, K% T9 i* t5 E5 p
    14
    / U8 ?% n! ]* |% R% \15, V( ^4 Q: p# L
    16
    4 P6 M  l4 n9 ^$ l- X6 X, e17& c. N% }$ O5 h2 Q0 ~# K( E9 M. K
    18* A* n; n  q% H( E: `
    19# R: t, S9 Y: N6 K5 m
    20
    " P+ s4 `/ A5 d' d: Y21
    7 ^$ L5 W% q0 s0 s1 H: B" }2. 数据预处理与操作, H- b* C& t5 u6 Y
    #路径设置
    # |7 Z- [8 R3 T" A8 n( Q" Pdata_dir = './flower_data/' # 当前文件夹下的flowerdata目录! x' w; o+ B& J
    train_dir = data_dir + '/train'
    : i# c) _8 W0 t& B! L+ Pvalid_dir = data_dir + '/valid'
    3 N" O: j  e1 {8 `% P0 ^1& I' t8 X# g9 b# ^% r8 E
    2
    , t  a& K! |5 w# q3 W8 t6 K, J3
    5 ~% w5 n" L1 I6 x; q4
    # f: t3 f, d. H( x$ apython目录点杠的组合与区别$ @* U6 ?8 x: y2 u8 `
    注: 里面注明了点杠和斜杠的操作" h5 J: A  f! f% O

    ! \% n  G3 R7 A1 D0 B  M! ?3. 制作好数据源5 N+ r: Z# u5 L6 v1 d
    data_transforms中制定了所有图像预处理的操作6 e7 f! E. z4 L, y+ L
    ImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片( |/ f/ H; P# n' e
    data_transforms = {
    ' S% L1 G! s2 M% m    # 分成两部分,一部分是训练7 |8 R( Z! w# r6 H
        'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间0 J% p4 S8 B# L5 ?
                                     transforms.CenterCrop(224), # 从中心处开始裁剪2 Z, N& e+ D1 q5 t4 }, N
                                     # 以某个随机的概率决定是否翻转 55开
    - e7 v9 `- I: i! q, Y; ]2 q                                 transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转
    + K5 N8 q8 x4 t% L. N) F4 r                                 transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转( U/ ^: R* N  J) W, Z: O' B1 Y& w7 ]
                                     # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相2 s+ ~+ i. r( o- H& Z  v
                                     transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),) v5 o% Z! P+ |! g
                                     transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB
    ; z5 N$ I7 B" \/ u) n1 O. O                                 # 灰度图转换以后也是三个通道,但是只是RGB是一样的. I& R' {" q2 k: K: ?6 n) }# e+ x) G
                                     transforms.ToTensor(),! G$ v! i" H4 |, R. v$ ^- @4 n
                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差4 K$ _, T4 B+ \' ]
                                    ]),
    - p) }& E5 j. _+ w    # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
    ; b! o* g7 X* c7 o+ n5 ?- [    'valid': transforms.Compose([transforms.Resize(256),2 D" T" F& n' h5 `& R9 Z- E  D0 Y6 ^
                                     transforms.CenterCrop(224),5 w$ k- ~$ y3 }( h; E; L9 B0 m
                                     transforms.ToTensor(),4 v4 [, o7 d4 I9 U! ~3 a. c5 E4 l
                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同8 m: h7 B  m% W0 ?( M
                                    ]),( P, \# C7 j& ]% X# F2 p4 ~. r
    }
    4 `3 U# o$ W( G! `. `, O3 D. [6 A9 a4 P+ M3 T" m
    1
    " \7 s( C1 ]& s- U2
    7 K! w; _% ^1 j5 L* }. P/ q( s/ w. J3
    : ~" T' a8 L  \5 W4* R4 N! W# m+ d" d8 n+ E
    5
      O) \5 \& l4 r' {  m" E( d$ _$ r0 z61 d  n4 c6 R1 z5 v: L8 S9 R
    7
    1 Q7 H* h. c( K: K8
      p' k+ U# \8 H9
    ( t/ g3 L: o1 Z" n2 e' p1 p10
    - j1 E; t: {! D) b/ I11* C9 E% a: ?" l( L( C- x
    124 e, G8 G9 R6 i5 w8 M
    13
    ' f- |& D$ T+ u- Z/ }' E14
    & _0 L% z% |, `6 G& R8 e2 v15. t/ T! p* C# D& t# W$ d: r
    16
    4 ^* g7 b& v/ \! b175 u# c0 C3 S0 U/ i8 ]& x7 ~$ \
    18
    # r4 {8 U/ R0 r; P& x$ O19
    / b) x( B6 o4 D/ I, A20
    6 s) T; O2 L& Z21
    ; q5 N; }1 W1 o# }6 {batch_size = 8. H; ]. v( t" @9 l) D
    image_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}9 c( B. z8 Q, \
    dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}( Y, @: l; \" E# ]3 l- f5 ]5 P0 f' t8 N
    dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']} ; ]5 u; G+ R2 a  S
    class_names = image_datasets['train'].classes
    ' _/ b) E# G  U- d
    4 K* G0 r: k" a# L& S4 o. @% W#查看数据集合
    ! c4 g, b6 a) k3 I4 I" aimage_datasets
    + f; |9 ~  b- |( ]& n" v5 j8 v3 f1 ~( S, Q
    1
    % W" |! ]# i" N* C* ?) T29 `7 i. O$ |& {+ m9 W' m
    3
    ; Q3 J4 ?& V. Y- W5 n4 q4
    4 ?* Z& b) [) `  G& f, ?+ H: x5
    5 ^6 O: z$ a- g0 L, j6
    0 r/ {( |) y9 l8 \/ |) ^) l9 H7/ r# ^# N4 ^9 r3 G& X5 |- w
    8
    5 i" ^9 g& ~4 r/ K2 H# B) _2 ?/ z& q9
    ' d, N9 u4 c. W{'train': Dataset ImageFolder* {& c3 a5 z7 l# ]& M6 K* W
         Number of datapoints: 65520 e8 c5 U: z! @
         Root location: ./flower_data/train
    / D$ K! @5 u" n  Z9 |( m+ s7 B+ s& T     StandardTransform
    0 W, s* y3 r5 P/ ?: ~ Transform: Compose(
    ' G) n, D' H1 W/ S                RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)" ]% b' u5 x6 G
                    CenterCrop(size=(224, 224))
    0 ?/ i' b' t; O                RandomHorizontalFlip(p=0.5)% B6 i& D3 I4 D+ I/ p9 q* g
                    RandomVerticalFlip(p=0.5)
    $ Z$ E3 K! |, w5 J5 @                ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
    $ i; T% L5 @( O; u% z8 i7 ]                RandomGrayscale(p=0.025)
    $ @5 ?2 o5 \- m+ [4 Z                ToTensor()
    0 g& M! M) R$ Z3 K6 {                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    3 b5 M5 S6 ^) X' O7 y" Z            ),- {  f% z+ u1 A. n! N6 A* M
    'valid': Dataset ImageFolder
    * c- x  V- [  Z9 O! |     Number of datapoints: 818
    & @. E) u: R" L' W/ f7 p3 w9 L6 h/ U     Root location: ./flower_data/valid7 {- \& }4 b0 A' H; A  y& l
         StandardTransform/ V* m  t& O5 [0 ~: D# k, Z
    Transform: Compose(3 [( G" }/ H2 o- M( e) ~
                    Resize(size=256, interpolation=bilinear, max_size=None, antialias=None), `: \  F4 n* I3 f
                    CenterCrop(size=(224, 224))
    + n2 v& d7 T  t/ P/ k6 @                ToTensor()
    1 k  T) P2 G, T2 z' r, a; m                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    0 V" F" R, ]( _* {/ T            )}
    / ?1 w2 \3 M- W& r+ A1 e& ^. X
    1
    ) S. R0 w) t# l/ A4 A/ k% J25 j' F- F7 R) y
    37 F" R8 a) O( Z9 }
    4
    & s  v: K: }4 O5
    ! }7 g, g0 C7 e, O) s; H69 Q6 Z" W& f- [& a
    7% T. F- o$ o4 t- j. R3 W" F$ Z
    8
    * s. X2 k' V% H1 H/ W# \7 g9
    + m1 M# o3 L  A" Y. e4 D10
    # b5 P$ e' E9 m) S2 u2 y11
    7 V6 D8 k; I' s- }! b6 n126 a% M) _# i, \! l. v$ y
    13
    " `. G6 [% C3 P+ k) W14
    " R8 N2 T/ {0 a  {4 ~( l15& Y' ~2 G% I; Q/ y
    16: o( F, m* B  k5 i, f: j2 R
    17
    3 Z3 U5 ?2 `$ l18
    1 ~. f) J( g, N" r1 S19- _( E6 x  ?% Q. |8 w* o9 I
    20) m6 J# [0 [  v! q1 e2 J# N+ C/ t9 I
    212 l" G) X1 ^$ y
    22) r6 a. B  \) O( }
    238 `! b- \0 e) Y- s) |2 K
    24
    ' G% {1 g( F% k2 K* P+ S& ~/ o# 验证一下数据是否已经被处理完毕3 y3 K, u. l4 q+ f
    dataloaders
    2 h; L$ \* N7 V1- j# T5 q9 b( F0 r4 [$ R
    23 \$ x, w) c% [+ _
    {'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,
    9 s9 T# L, c, W" e* W 'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}5 N9 K$ D5 w5 J% O7 o' q9 y
    1
    ' K! z2 {8 h: j29 O  \' g" G* z$ Z
    dataset_sizes
    " h0 ?0 l* P3 g+ \- Z1
    0 c2 h. U1 j0 r: H5 H5 S{'train': 6552, 'valid': 818}
    " c7 e7 Z! {! j1 e11 J8 j7 j  B1 {+ |' J
    读取标签对应的实际名字# N. q3 f  E2 U8 l  ]
    使用同一目录下的json文件,反向映射出花对应的名字- H/ F3 c& O* F) J' H! s5 Y' e
    4 I* @1 d4 Z( X+ a5 ]
    with open('./flower_data/cat_to_name.json', 'r') as f:
    # W/ Q; B* m) c% M" P    cat_to_name = json.load(f)0 B  T2 }  q8 s
    1
    * u6 l2 ~1 |0 K$ E8 U2
    + V( _/ _7 e- P# r. x9 x6 ^cat_to_name3 {. c) {! z. `8 `: `
    1
    , R0 n# T1 U$ m  H9 a. ^{'21': 'fire lily',# z( h4 u0 |1 Y" ?* P+ N& O0 u2 n
    '3': 'canterbury bells'," u. L& {6 H# K; _0 u4 J2 @: ]
    '45': 'bolero deep blue',/ Y; R0 I1 e' A8 k6 A
    '1': 'pink primrose',
    + S5 G, r& z4 R9 T! @6 _, E# d) f# U9 h '34': 'mexican aster',$ W1 X0 V: v  D9 _/ J3 h1 B  g
    '27': 'prince of wales feathers',
    - x2 W$ u. F# k5 q '7': 'moon orchid',! C' G3 X9 f  `( w9 T
    '16': 'globe-flower',% R' D; r/ u# H# ?* J( ?5 {  w* M
    '25': 'grape hyacinth',: L' I! ?7 \* x3 ]
    '26': 'corn poppy',$ @$ I! O1 @& b5 Z
    '79': 'toad lily',7 x4 I/ U8 K9 g2 h
    '39': 'siam tulip',% A# v! @; J3 r2 f8 k
    '24': 'red ginger',# X# |  v! v# _. d% a4 V4 s9 W
    '67': 'spring crocus',
    / X/ `5 \& n3 R8 F" p '35': 'alpine sea holly',9 K  d0 [- f$ j" j4 ]/ B$ Q
    '32': 'garden phlox',' ]; O- d( X9 j
    '10': 'globe thistle',
    3 }  |# f& ^; s+ v/ J0 c '6': 'tiger lily',
    8 i) w- X! y& m- O '93': 'ball moss',
    # J, [5 t) h, O' ?. a/ u6 _ '33': 'love in the mist',
    6 Q. t/ [0 P4 r8 `6 e/ f% v '9': 'monkshood',
    6 v3 Z$ V0 t, O) u '102': 'blackberry lily',
    5 q4 a# L4 W# a% j- N; @1 I* a% H% K '14': 'spear thistle',2 j, Y' R0 T+ ]" {# {. Z3 X
    '19': 'balloon flower',
    " q7 x, {: `! h( l6 Y) D '100': 'blanket flower',
    , V/ H' g$ F3 j, Z '13': 'king protea',+ r# ~9 S* M6 Q9 v$ \+ k0 s
    '49': 'oxeye daisy',  e. S: {- J8 u  v3 R4 Z& w. J
    '15': 'yellow iris',0 J3 d  E! }# M/ f% @! V2 u
    '61': 'cautleya spicata',
    & O$ B1 K+ ?0 J4 K '31': 'carnation',
    ; H+ h6 u2 r; R' ]. {+ n '64': 'silverbush',
    - \; {; I- `% v- A3 G" Z '68': 'bearded iris',
    % B# n9 K$ T2 V3 H! F+ @ '63': 'black-eyed susan',
    / u7 ]" V2 T- n4 Q/ c '69': 'windflower',
    ! i( b+ M! ?# Y' K% O '62': 'japanese anemone',# B  H( w! _' s# k/ I$ f. i
    '20': 'giant white arum lily',
    . B# w2 l/ _6 [ '38': 'great masterwort',6 Z  ]5 l" n) u# ^% p' i/ E' {% h
    '4': 'sweet pea',: q' n. v; K' g
    '86': 'tree mallow',
    3 o2 {5 q- W1 {, ] '101': 'trumpet creeper',- p) m: j/ D6 H* j5 b, B
    '42': 'daffodil',8 l; J) N& @: S% \5 \1 f0 [3 B
    '22': 'pincushion flower',  v6 J# v5 m* G/ n
    '2': 'hard-leaved pocket orchid',% Y# W( w1 S0 Q" n8 I. N
    '54': 'sunflower',
    1 f5 R' P) \7 N4 p3 \$ J( j '66': 'osteospermum',
    % h3 M$ q; z  {- a# a '70': 'tree poppy',! S# L7 w8 q  @% e8 L' u$ u5 H
    '85': 'desert-rose',
    $ `  K  b. \( R '99': 'bromelia',5 b: E& x( @/ r" z
    '87': 'magnolia',
    $ w! G, X0 K1 y& b5 B '5': 'english marigold',
    " Q; r* y; f! G3 Z7 r '92': 'bee balm',
    # `0 R9 ^: F9 h/ d- P; w '28': 'stemless gentian',  i9 U% x5 w# H# K
    '97': 'mallow',
    " S4 J$ a* V( ]$ r9 y/ O1 M" \ '57': 'gaura',
    " G$ A; m9 t* h7 }: U! H '40': 'lenten rose',
    # t/ F) b* r) u. k$ s) l '47': 'marigold',$ a3 j% C/ n, p" k/ S
    '59': 'orange dahlia',  L3 w  o1 N+ l5 c& m
    '48': 'buttercup',4 a0 r/ w) C& h9 r) q9 {. m
    '55': 'pelargonium',/ B) x3 ~2 a* _
    '36': 'ruby-lipped cattleya',
    5 Z: Y7 q' g, V9 l1 i8 ^: {# U5 i '91': 'hippeastrum',8 _. v) _5 T% E. G' J: K
    '29': 'artichoke',/ u9 \! `3 N4 M4 u* M
    '71': 'gazania',5 F+ d8 e, S6 W7 O: T# ?* r
    '90': 'canna lily',: ~( L; Y9 Z* Y) K) s) j
    '18': 'peruvian lily',
    ' \8 D3 j3 Z6 H/ `+ x: S+ C' S '98': 'mexican petunia',$ R0 N) G4 q6 T- `
    '8': 'bird of paradise',
    ! @0 [( Z4 N$ V3 j '30': 'sweet william',0 h" W/ @" v8 I# [
    '17': 'purple coneflower',9 j- `* D) Y7 q$ q. T9 h+ h
    '52': 'wild pansy',& s: z0 @: c# r) S0 ~' E
    '84': 'columbine',' ]9 D/ {# F7 r, U+ y) |  ~/ I
    '12': "colt's foot",
    4 Y" F& n0 f( M# |. O4 S '11': 'snapdragon',
    5 [$ t( C$ a1 m" y: e5 J) |% A7 | '96': 'camellia',2 A. W+ {; v( J. X/ F" E
    '23': 'fritillary',2 K  V3 o7 i7 d  m" d5 L/ ]2 D
    '50': 'common dandelion',6 ^6 d" L5 [% }; w
    '44': 'poinsettia',/ W  l  k' L# i$ E  @; b
    '53': 'primula',
    2 f+ b1 E! u& I '72': 'azalea',
    ! y( m6 \: ]% b '65': 'californian poppy',. g/ Q3 ^' n: K6 J
    '80': 'anthurium',
    1 I, T) m& U4 d- M! n$ g% M '76': 'morning glory',
      a$ {  O4 x3 R: g4 D- K, y' _0 ? '37': 'cape flower',
    , L) j# x( r9 ]" r+ _1 I' T" Y7 A '56': 'bishop of llandaff',3 X( a. _& G  H9 y6 L0 f
    '60': 'pink-yellow dahlia',) M3 l' K0 C9 ]$ i9 z; g, O
    '82': 'clematis',9 r# E; d3 I# W
    '58': 'geranium',0 ?" g: ^* i1 s" y; l
    '75': 'thorn apple',
    3 Q* l+ B5 n' g8 B, K3 h '41': 'barbeton daisy',
    + b. Q/ W% n7 w/ c0 @ '95': 'bougainvillea',
    ( [5 M( K0 c! ~' b" v' p, B2 s '43': 'sword lily',
    1 J8 ~- ?* e0 ]/ P6 c '83': 'hibiscus',# ?" B8 d- g7 N( z( T; y5 G+ K7 {
    '78': 'lotus lotus',
    ! E2 g$ [. F8 \7 ^# L8 m; q- f '88': 'cyclamen',
    ' g  |! s* b* w6 _1 w' T+ {. v '94': 'foxglove',
    7 `8 p" s! I& h0 G+ G) o* P '81': 'frangipani',
    5 v2 h2 O1 m* b4 I) m+ }9 ^& q" \ '74': 'rose',1 \. A5 A9 E" m! v" e
    '89': 'watercress',4 B8 x% v5 Y# n( [$ z
    '73': 'water lily',
    8 j- R! p$ l3 H8 M/ J  g '46': 'wallflower',
    . @( c- R( _. f- j4 f* D& w6 U '77': 'passion flower',
    1 F; t" I. q$ O$ K7 G* W6 G' r '51': 'petunia'}
    . ^8 a# |6 A' {5 U- A7 J1 q! O# T! k; f* T2 M
    14 H* L: b, N, i5 V9 R( I
    2' \, [& o% F. N, ^8 `1 z- {+ I
    3' |/ ^: Z1 p; L% t4 D7 x' B3 h
    4
    0 o/ z# O8 }% x5 ]' x5
    4 V# R# \2 x/ v2 o6 b2 n5 E6
    ) \9 U. e. {5 k! U1 a+ C' j* ?7
    , B1 F6 y; E3 F* o2 }  q- y+ e8. E. x/ Y' X5 e0 Q
    9
    / i7 x/ ?- U. l3 |1 X10
    & Z, g. |! c% Z# W4 t11* o' u6 d, ^' w
    12
    $ X7 I; L1 T! |. i, [13
    " n& A+ U/ `: T14
    * ^8 D$ u: a9 F) w/ C. X15
    5 L7 s2 g% m( i6 p; G16
    2 u% \( x: _* a; Q173 f* L9 j- c1 e! l. V
    182 ]# F& Z5 _" O6 ]& E! f
    19: z3 _+ h9 j0 j
    20
    8 T4 {- f; n3 `) F- M0 Z, h21
    5 F% z2 B+ R6 ^7 Q22
    * H# k) S, |. ^9 m0 N. B/ @/ M) I23$ R# D. r. b* {. A1 d* T2 d
    24- S8 F$ G8 J! f9 O+ M- N
    25
    & X1 v6 U0 Q( ~' H/ G26; {$ V5 N! s% j) {5 E1 {" K" O  i4 E
    27- Q0 o7 ]/ O4 U4 D4 b5 y9 Y- R  R
    28
    2 e* L; ~0 T0 u+ N) Y# ]8 F/ |0 e29/ j: j% [( t( z9 I, {; m
    30
    # i4 b# A$ B3 e3 R2 u* Y& }$ `319 ?" R" `! n! R
    32
    , [/ Q. k# r/ q33
      R' X9 w. w8 f' }6 ?+ h34- u% L7 C$ J3 P7 H3 J
    35
    ' ~6 f- ~0 ^: Q) j; b360 @( p0 L& b  h9 n/ {& c+ f
    37% }7 Q$ F# E! m! U" y
    38/ `! x0 m" K7 d5 s) k! e" [: P! [
    39
    # I6 f; W' }+ s/ O( j; m40
    % E  i  p+ v1 v4 I6 p414 o: _' Y+ x* g, W: N7 {
    42. S& `: ]4 a' S  |0 O" `
    433 A; [: b, [/ b4 T8 F* J
    44. Z- [  s/ i3 l- r4 r5 n: {
    45) W6 G: ?) d* \9 \
    46% l- }  D& u: I! ^$ V; R
    47
    # v$ t% j  R0 m; q: g0 h! i& ]48. R1 u# O: h  H) @
    49
    9 a: m3 @. w$ Z4 N6 |7 P/ K50
    , L. b& o, S0 a7 L' h. J1 A518 ~% V2 O+ z! t$ M3 u+ z
    52
    ! E" Z6 _6 F  b0 B- T4 r53. i9 C4 E3 |- ~+ \, }
    54
    # G0 g# W, O: F9 E+ T* G55
    3 e. ^9 Q- `% y) d56; M, n/ u; G/ ]* r/ y1 C+ S3 o
    570 B3 q) k* ^% N2 ~2 d* R# x8 s
    58
    / A  k' l* t; H2 q590 T0 Q6 c; b& Z# E, G, J& m! B
    601 |4 `1 z( O& B% Z5 s  F0 z3 l: Q
    61( w2 J" }2 s% g2 I
    628 H0 T' d, b9 X9 b" ?. I6 Q( ~: M
    63
    6 ^1 P: O5 C4 y: _* ?2 q& J5 \, G644 c9 B# n& {* r5 @# g! O
    651 \' R8 r$ L. T0 j3 v; Y6 B' x
    66
    & X; ~% t/ d. g67
    6 @: `" O) F2 S; B: B685 u/ }* q' l. |9 E4 D6 U9 \
    69/ u% ?' R2 f5 `0 D2 a1 `3 J1 n
    70. t) V4 @$ v  y  Z( |5 N
    712 U. V9 T+ h0 h5 g  _, f% \; _3 C& L
    72# G+ M: v" |: ?! b$ \
    73* C' |* [0 J: M, h) C
    74# A5 V- Q- q7 N- r! s
    754 |* P: U0 f- E( [1 i2 Y
    76
    0 D, U# w' u/ R/ Z  s/ h% \4 W% X# L779 e# @9 `1 y3 F! v7 f' G
    78/ T" w4 k% \& U( I, o! _
    793 s0 \5 F* v6 y) D3 Q
    80' r7 l' o$ m# z/ G. Q' e8 G. I/ M+ U- x
    81) e/ s. q/ Q7 U* `' C
    82
    ! l$ \; f, R8 H6 Z0 \' c+ ?835 e! o( a, E3 s! l3 m" O4 Y9 H% {( q
    84* ?/ B) s3 T  W
    85
    3 k: W4 C  p$ B5 G8 @. a86. s( }& G. o; P( W" K
    87
    2 Z$ r" m8 N- A88$ L8 A/ T( a: Z% r# H# X0 b9 E
    89+ Q! I! a% k8 ]. Y* \* t" D
    90) d+ z( j+ J) {8 [
    91$ K& Z# G4 m4 G/ R5 k/ }
    92/ L  }* f' X, c
    93
    0 I: _/ u* H1 v9 l5 E; K94
    0 w3 H+ G. e8 v: A95
    ; b3 v* S6 r2 B1 \96
    0 k/ @9 P" V- Q5 y' j* Z97* T0 K' m& b' g$ [; j) \, B
    98
    9 {8 w: o+ R) c9 K1 Y99% h9 N0 i2 ]+ ~4 Z5 C
    100: q+ Z4 U% e4 l+ d
    101
      z) s* f/ g1 q9 r, _102
    - W) R4 d8 T" F$ \& Z5 l4.展示一下数据. S& r  B, }; v( J: o
    def im_convert(tensor):
    + T* m% Q7 N7 Y1 w- W! |5 y, w    """数据展示"""- f4 ^# r8 Z8 C
        image = tensor.to("cpu").clone().detach()
    $ W! ~8 a8 u3 ^    image = image.numpy().squeeze()
    / l. j) h. r: C$ o+ n( R1 N. n    # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图, w; J  V# ^3 P% u8 z$ T
        # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)
    : K) _0 y0 C; ?+ E/ J, e    image = image.transpose(1, 2, 0)0 l4 a+ O, ?2 Y/ v
        # 反正则化(反标准化)9 E  P" d; A$ m6 x
        image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))) v" ]9 [2 g6 X4 j# f% O9 A$ }$ _

    . K, w2 `8 K/ {. q3 }1 l& b    # 将图像中小于0 的都换成0,大于的都变成1- ?2 y* Q4 O* j" W! A
        image = image.clip(0, 1)
    / L. ~8 U$ _% H1 ~  V. s4 v& |: n$ @$ q: Z! ^
        return image
    4 Z9 a8 P# T1 ?9 }1
    ( s3 b; J% n5 E7 n) o0 f2
    5 W2 d5 z8 X2 U( y1 X3
    % j# Q% c2 Y8 L9 s/ L, o. e40 |1 E- W" A9 P( ]' ]+ f$ d
    5
    ; e" Z+ B8 p; D  l7 a* ^6
    4 ^! @) M& D; Q9 f2 ?. y7
    6 [+ h2 g" @# ]. a/ r# t* K8
    : H" r: n7 _% T$ X9* D( w; y- j& ]4 f3 ?" m
    107 D, V- u  Y. ]( [9 Q$ U
    11
    ! k: z4 ^( _2 ^/ {- Y: t12
    ' h( [5 l: I" P' w" o2 M9 ]13
    " E$ O& {% V3 F- ~/ O! t3 O  C3 B14
    7 h9 z" t  r& b6 T5 t9 G7 L% c/ a# 使用上面定义好的类进行画图
    . f9 G% i3 |7 H9 ]fig = plt.figure(figsize = (20, 12))1 q* ?( A0 E  W1 m3 t6 `3 [5 w) q* ~
    columns = 4, \$ k$ h; G2 o" `5 s5 v
    rows = 2
    8 q2 R. r: e4 G8 \, E7 d
    . W% t9 p; K; M  O* M: G3 Q# iter迭代器
    ! z9 c0 Z% c$ Z3 B) A$ f& \3 Z7 j4 C# 随便找一个Batch数据进行展示
      B: B7 Z; \( i! ?7 q: s2 g- F6 Sdataiter = iter(dataloaders['valid'])& h+ k$ C0 X, O% ]" K" d
    inputs, classes = dataiter.next()- q# z  P% j' @7 @
    # N( F! ]; V4 Z
    for idx in range(columns * rows):
    7 W  r2 n, W6 _1 N# s. c    ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])
    6 x0 I* Q8 ]" P5 H9 {" G6 `    # 利用json文件将其对应花的类型打印在图片中
    % X2 F/ S& b# |. e* _( a+ N$ T' @    ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])
    ! h1 f* g% @# T    plt.imshow(im_convert(inputs[idx]))
    0 W7 G; ]! \1 [  t3 Dplt.show()# Z4 F; w6 v, Q
    5 I/ E- }* ]8 H% d7 |7 c
    1, x6 g) G4 U7 G" \; r- m! X
    2
    0 V- m5 D. k, Z5 _6 j% K  v# @( l3. T/ G) m, J% B2 {$ v7 v  I
    42 Q7 }& t- P4 }$ b! I1 W
    5! Y" i: \1 \: K
    6
    9 h) H  R: Y) m! W# d& `7
    # n4 T4 T. ]$ h6 F; f# Y- }9 l8
    5 c& ~( K& i  V9
    # x/ H2 ~& P0 N; O. u0 g# Y10
    2 L4 m; }" N* }  N3 S7 i3 b/ l11
    - y& \8 N' |1 S+ a% F$ ]( V12* l1 j7 J/ Q# |  O7 `" R+ A+ a
    131 Y, h$ L- d9 I8 Q+ p1 ?- ~
    14
    ( Z- o0 [+ s, B" {15
    * C) V5 J# S/ k: _16
    ! q# F. d  B. p; N$ W, |) p5 P% \/ z$ M# m
    5 f' `0 n1 m3 f7 N2 r% l
    5. 加载models提供的模型,并直接用训练好的权重做初始化参数
    : U6 O: p5 h5 I5 ]model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
    - N2 V, F" C# w9 g7 [& z6 J. G# 主要的图像识别用resnet来做3 x6 h& ~. d5 A9 `
    # 是否用人家训练好的特征, ?% t' H! b" R
    feature_extract = True1 N9 {) b0 Y6 c& K2 z/ G
    1
    6 s+ G7 r" Q: Y& m' r$ n5 M  w2
    - ^2 f7 ^" E& n/ `3+ v9 J, ~# J/ U" e$ ^( \
    48 D, E) V" n0 m: X. [
    # 是否用GPU进行训练
    . h% u9 b: z- x3 u2 c  ytrain_on_gpu = torch.cuda.is_available()
      a/ g2 W/ {+ {) q1 x* J9 r) A4 N3 w
    if not train_on_gpu:" o& C0 w* n3 y% T# a" F
        print('CUDA is not available.   Training on CPU ...')
    & d* ?7 F8 T& L+ q% D1 A- {0 _$ b" yelse:' V# z. L3 s1 L) ~* H% [) _
        print('CUDA is available! Training on GPU ...'), ^5 X6 e) W& g! S- s5 c
    1 U7 z: B" s' g. l# H# W
    device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')& l# j3 R/ Z! M. n9 n" M
    16 a) k* L5 j  A) F3 O
    2
    0 X. U9 o+ H/ J1 K! X# ]33 f3 W5 @. B  l2 {7 P7 b
    4
    9 Z# l/ c; w1 ?. n. X5 N; P. r: K& p56 B( U- W  G/ T7 h" K8 W' C1 }
    6
    0 [8 E: t; ~. f7
    1 p, o/ l; ?! @% I' L" W8' F; O) f* Y  k# M
    9
    3 f; ]; @( G2 V4 @. e8 A: xCUDA is not available.   Training on CPU ...$ x- f& G/ n, W) S& l% G% o
    18 Z3 r* l7 N3 d9 N6 L( g# y6 [2 |
    # 将一些层定义为false,使其不自动更新
    1 z! i- i; J/ a: Y4 W: ^def set_parameter_requires_grad(model, feature_extracting):
    6 b2 V/ {# M- ]( g  x) q1 i* c: F    if feature_extracting:5 l7 D  S' `6 j! `: e
            for param in model.parameters():
    9 u9 _1 E9 K$ d# U) q& ~            param.requires_grad = False
      R% k! N- t- V  |5 U5 X1
    , Q. z! a# b" ^5 }2 f& x' q0 X- d24 D$ d# M6 M5 l4 m1 T
    3% Z; I/ X  K  w% }# e" @
    4" f6 A0 D0 U! L+ B% _; f, F
    5. d9 f/ `/ S% m
    # 打印模型架构告知是怎么一步一步去完成的
    + P9 @  c; w- c. g# 主要是为我们提取特征的, A! w/ \/ O: K0 O

    7 s: t/ N4 S6 h8 h2 E4 }  cmodel_ft = models.resnet152()
    4 k9 G) N& O3 \& d% ^+ Jmodel_ft  d' G( ?: v" l) F
    1
    : L% h  d$ c, ~) ?. A2' H. y% d. b) {  T$ {
    31 P3 Q/ U" ~) G8 ?3 K
    42 i5 A2 Q, J" ?9 B1 `
    5* ?: U3 Q5 X# H& P
    ResNet(
    ( o- z& r3 D& _' J1 A3 B  (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False); P5 R$ F4 ^# ]9 \% n
      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    ) _) M7 k6 y: m# O( e: J  (relu): ReLU(inplace=True)
    & Y3 D; I  `. q% E3 K6 Q  (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)6 V: v' O- U+ ~/ Z5 ^$ }4 f
      (layer1): Sequential(
    ! A4 V8 \. o7 _8 X9 ]    (0): Bottleneck(6 N3 c8 ^% |2 \+ `+ R0 |
          (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)
    1 F: J2 x: V4 z6 }/ f      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True), u8 }* {+ C( {1 Q# Y) r
          (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
    0 ~2 o9 ^" o( H2 c3 t7 w      (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)0 t) M$ ?, ~+ \; V
          (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)6 E) R! X2 @' H- [1 t
          (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    " ~. G  B; `8 ?# X' p      (relu): ReLU(inplace=True)
    / h( ^9 y7 q/ i. v      (downsample): Sequential(, P# p! F. ~' v; W
            (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
    9 E! ?7 h) m( t4 Z) m3 H$ l' `6 _        (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)7 O$ g! @1 v2 q% {/ Z  h
          )+ I0 ^" H* N2 R* l
        )4 S1 \& I/ `9 o/ k! ?4 z8 c4 |
    中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。" ]! c' c, w2 A: e( n* A/ C
        (2): Bottleneck($ d% I( X" y+ m1 m
          (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)
    - n6 G) @4 w) i  W8 C8 ^; h      (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)4 w4 Y1 R- J& e: n
          (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
    ) }$ H  n/ Z4 v3 v1 \1 e2 Y& B      (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    5 ]7 B+ W* ^* J, e8 W3 T7 @4 i      (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)( ]( P, N2 @# Z5 J+ g9 g+ A
          (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    6 \8 n* X' ~0 X1 P) Z- M4 a      (relu): ReLU(inplace=True)& t& l- n5 Q1 p4 W: U
        )
    ! {% l% Q  g7 e) s" ?) |& x  )) m7 e4 Y0 S. r( _% e9 {& C
      (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))
    0 z6 U" v2 {' z- x6 e7 I  (fc): Linear(in_features=2048, out_features=1000, bias=True); u- }1 g4 E+ R4 L- M2 {
    )
    * `# u! T- i6 ^6 m0 }% }) M1 K$ l2 H
    $ N; L( K) A3 |$ T15 w# ~3 `+ O: }
    2% i, C4 D4 t- w! @+ G" |
    30 u) Y7 w; |, C2 i
    4) T3 v0 C$ i' U; ?
    5
    6 S6 |" g: O# P+ Z6+ {( k: J+ k3 F- I0 x
    7
    + J: R7 f% V& u, T! q8
    - e; U; `6 R& k4 q/ o- c- E) j95 N8 `% }. p5 H
    10, ~. B$ u6 _9 H+ K% Z
    11! O# i7 P. n) d- u0 m' Q% }- \5 N
    12: b+ V: `5 }9 \' {' W& E3 V9 N
    133 X- U0 p  z9 T6 W* P  u: b5 |* x
    14
    8 {8 U8 K8 D( G1 H6 j  q$ k15
    4 P0 J! [- d# n16
    7 K1 o  E( V3 V175 G/ [8 C9 p3 U- i! S/ Q2 f, z
    18  G/ E5 p8 u) `
    19. ^( e/ s2 v: U9 e% V
    20$ l9 S% w/ r2 `/ t7 P% O8 d6 l
    21! h4 {9 X. ?* g& J3 S  V6 a# l7 @3 }7 h& W
    22
      q5 b+ Q6 r. p* [0 K236 V; N8 i6 y. _7 \  ]+ W6 z/ \% C3 e
    24
    % _5 w3 A. a! `! I" J9 e: M25( R& U% ]' O# ]' m" ], d& B
    26) I' E; q: x/ |9 @- n
    27
    & e! i; P1 ^7 P28
    1 e4 S* `2 @( Z29
    & r1 w: B7 R7 |, k& v) ?% d* b305 d' [# \* b9 G0 h5 P* }
    31: w5 }0 i' e5 m, m3 x2 q) b9 ]4 x* U
    32
    4 n4 W+ {% x# Y7 y! l* N2 a9 Q33
    & ?: c6 S& r, L1 z最后是1000分类,2048输入,分为1000个分类, W( ^; L8 g$ Q, m
    而我们需要将我们的任务进行调整,将1000分类改为102输出! ~  {; x+ H6 W8 f' f. D
    + J4 q& I; S# a# L# C
    6.初始化模型架构
      A) N! G- W* k) t步骤如下:
    1 f( G  s2 X/ e$ z; ~/ w; ~3 D
    5 X, ]+ h. T5 t- G+ t将训练好的模型拿过来,并pre_train = True 得到他人的权重参数. Q) {' S7 i; ^( i) G5 F
    可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)* G/ W/ e- y' s; X* z- k' @' [
    无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数7 {3 b; `0 }6 C9 R! a
    官方文档链接8 r4 k0 t, U: B5 }, B
    https://pytorch.org/vision/stable/models.html
    6 H: {0 r; i" a4 @% @* G
    * t' E0 E1 R, v( g# E; o# 将他人的模型加载进来
    + _2 {1 |. h, _+ F' s. I/ L- [def initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):) J, N/ J( i: a1 Z1 Z: T( a
        # 选择适合的模型,不同的模型初始化参数不同* f$ F& s2 ~7 ~# B+ @
        model_ft = None
    8 w# j0 S2 A/ a1 H* F- K! k. v    input_size = 0
    : D+ ]4 V3 i4 i8 l" o
    $ m- N! o/ y" ~, V' \    if model_name == "resnet":5 q" H: b% _: R6 c! c9 L3 }% G
            """% X, `7 n" p9 H" _: n: m4 M7 c
            Resnet1522 U+ B- Y4 ]- k- ?5 G
            """. o$ d+ C& n+ I1 Y8 X
    1 n, ?9 n  l: Z
            # 1. 加载与训练网络( `: E5 ]. T' o3 {/ J, r
            model_ft = models.resnet152(pretrained = use_pretrained)( O8 j* a- H6 j' V" O0 G$ L* d
            # 2. 是否将提取特征的模块冻住,只训练FC层
    7 t! k8 H! V' R        set_parameter_requires_grad(model_ft, feature_extract)
    6 h( q6 Y9 E" V        # 3. 获得全连接层输入特征" T7 Z6 u$ E3 b) c6 _9 U; o
            num_frts = model_ft.fc.in_features
    + [" I, _2 f& a3 e1 t: X        # 4. 重新加载全连接层,设置输出102+ P9 v# B' r# p/ t" }( H
            model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),
    - V4 }& A& z8 P/ y% A                                   nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1" s* r6 X0 Q8 P* ]4 J& u7 h# k2 f
            input_size = 224
    % W. x* T# }8 T. s, n3 U" H& s' D. I. Z$ k1 i3 u0 W# ?" F7 c" T
        elif model_name == "alexnet":6 s; D, |( K7 `$ S6 n4 v
            """
    , a4 i8 O: P" q8 _2 P0 r        Alexnet: S' v7 p5 }3 O) b7 b
            """
    2 _' _3 L( a# r9 [& b8 I6 F        model_ft = models.alexnet(pretrained = use_pretrained)5 n' q: t2 @6 O6 M. A- V6 ~, m: c
            set_parameter_requires_grad(model_ft, feature_extract)
    ; w9 H- G' _" d8 r7 O. E
    ) ~( x0 Y8 {) h) B! ]/ B: }        # 将最后一个特征输出替换 序号为【6】的分类器
    * V' j3 m# `8 b* A: z        num_frts = model_ft.classifier[6].in_features # 获得FC层输入) Q) T* o, d* ?4 @( y3 o+ T
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
    4 o( k" q! ?9 l, A; n        input_size = 224
    / Y1 v$ a6 J0 T1 n+ |. ?% a6 T" m, N& S6 L0 H6 Z- P* @
        elif model_name == "vgg":, @# X- _7 C( {. @$ q, M  E
            """% i+ ], m4 k7 m
            VGG11_bn
    3 ^9 N5 L* f3 h5 i  Q        """
    $ u8 b! O; I- D& r        model_ft = models.vgg16(pretrained = use_pretrained); k/ b% a8 j5 i9 G% `
            set_parameter_requires_grad(model_ft, feature_extract)
    ' O2 A! f9 U. E2 y        num_frts = model_ft.classifier[6].in_features$ b( h/ X; D' {+ @* k# e. p8 D
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
    6 c9 W7 V; U6 D2 i0 n2 ?        input_size = 224
    7 `+ G6 A7 L3 g+ Q: P, K; C# _3 I# t8 b" b3 P4 E7 N% S" u
        elif model_name == "squeezenet":7 _8 h8 ]  g  n- }) J
            """
    ! d! |2 `: C& g  k% ?3 o        Squeezenet: D0 M3 r( s  |- T) c
            """
    & ]) x0 P: ^4 c) [        model_ft = models.squeezenet1_0(pretrained = use_pretrained)
    : R' a9 |5 t. x+ z6 T3 Y. |1 i  e        set_parameter_requires_grad(model_ft, feature_extract)* |  `, y! J: D; j" t! A
            model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
    : ]- E! O; X' d* b$ b2 p        model_ft.num_classes = num_classes5 \+ b6 f) Q% b3 G
            input_size = 224
    % a0 W, u, A2 A* D9 _. e' p& Q
    9 f) H/ K% G+ b6 i: O/ S    elif model_name == "densenet":
    ' _+ Z% n5 c7 C: V( y+ H        """, o/ O- ~* ^# w5 B' B/ K% A/ W5 k
            Densenet' L2 `  U- R/ o7 Y5 N! x& B: M
            """" T. f9 z+ ]  `6 A
            model_ft = models.desenet121(pretrained = use_pretrained)  f2 s. n2 e9 W& r% e5 Z) _0 g9 P
            set_parameter_requires_grad(model_ft, feature_extract)
    2 S  T8 R* o5 X5 Y        num_frts = model_ft.classifier.in_features
    : W( s2 y4 V8 o- g2 ]        model_ft.classifier = nn.Linear(num_frts, num_classes)) q- S4 p9 S$ W( S' c  a" q
            input_size = 224
    9 o8 y$ n  p6 O7 c' |- b6 |9 ?, w/ i2 y9 I3 c& Z
        elif model_name == "inception":  ?; S* }0 s- g$ z/ R
            """+ x" M+ h; X( P, S9 `
            Inception V3# X" n% |) `: [. _0 [  g$ e* h
            """
    0 e$ E4 p$ y' |+ ~- [3 C        model_ft = models.inception_V(pretrained = use_pretrained)
    ! T# w/ p* F, i9 u5 K        set_parameter_requires_grad(model_ft, feature_extract)% [) B0 X  G2 Z& n# N

    ' W& h. |, Q* V0 M7 y0 [        num_frts = model_ft.AuxLogits.fc.in_features' J4 l) Y' F( E& P
            model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)4 I+ g) }0 ]. e2 ~1 d, H

    7 Y: k7 J7 d5 I' {% ~0 C4 [        num_frts = model_ft.fc.in_features1 S( v1 U/ v; e! ^. N7 Q$ C6 ~" j
            model_ft.fc = nn.Linear(num_frts, num_classes)
    0 n4 x+ V% G* b9 R        input_size = 299
    , W1 i  ?6 r; W7 |- v8 Z) A9 D1 M7 f0 [6 |
        else:
    1 y+ N: }, N* ^% r        print("Invalid model name, exiting...")
    $ M* K) g, G% k7 ?        exit()  {  O2 L6 {$ p! s! E
    ( j& e, n, H3 {8 Y
        return model_ft, input_size3 h6 W+ X0 f5 f4 U0 Y
    - _# w+ t. y+ l* W6 w& C) \
    1
    " ^. N9 [* I6 ]6 B. C2
    9 b+ M/ F/ h: C8 H" c3
    ' q' a( L+ U" a4+ h+ z; b: C1 e( {
    5! }  i2 }2 J& G& C% l. B. E/ x9 ~
    6
    * W0 x: h+ e  k, A7
    ) r! K/ `/ M& F4 q1 T8% d4 X) _. i! Q1 ?" i7 j- F
    99 T6 d/ W6 A3 w! J1 ^
    10
    0 d8 e" ?4 p* M. m5 O  d11! z1 Q8 c, @* S7 a3 |- W: ?# [/ h
    12
    % l) X# b- R' t132 F+ J# D* D  `- o0 j8 m: J- n4 ?
    145 r" _. |/ X" O. k6 M! U; Z, {* z
    15/ x1 P9 |- D2 X% L; ]
    163 s$ y  s! w0 I
    177 J/ y2 G3 Q$ g& C# n; z" A9 J9 O# [
    18
    2 |5 E  N( N0 L( Q/ u, c$ u8 L19
    & [: M. g' w+ x7 m20
    9 y! }; t' c  f3 g& g+ C" u217 Q8 _3 p! \9 x3 E! k7 ?$ O
    22
    3 X) o$ a* @( p$ s( c* V23
    + T3 E2 h% F) s5 _241 V3 w) d, t1 n3 {) n. D
    25! B& |1 @! H/ L6 E! u/ f
    26. J1 I% ^/ ]4 ~+ S. \
    27) M* d6 n. _1 }, ^7 F, A) @
    28
    ( O, T- E4 z- H29" p. m/ L- M9 r" \3 ^9 d  D# T
    305 ]" i. L7 M- k# X5 l( k
    316 \, r3 n# T8 f9 h( a
    32
    - A/ f% \& T/ j/ \$ H33
    # R7 b) E4 e5 N" q7 F# [; N34
    $ W, N% U1 g; G, d$ a) o357 B. v  \. W* @
    36
    8 S1 L) a: [5 t& U37
    & C! N8 ~5 L: p* }* z) R2 }% L0 h384 G9 h- @6 Y) |: Z9 ~
    39
      M7 X# o$ H! P* m+ \40
    2 V  E& u  I3 U41
    $ _# Z. t; w) f; A# A1 x42# g5 W5 Q) r/ D
    43
    8 J7 T' \9 ~5 e0 ]8 A- f444 ]- R0 [2 i- R$ y! q! @
    459 A( `( f6 x. L! [+ E$ a
    46) R3 V! |! H8 ?6 F& g
    47
      [" M1 r' t+ }. ~6 `% c1 G48
    1 ]$ B% M1 h$ O% ?+ \3 q49
    1 A5 p' `) F; c4 i0 j50+ k, ?) \3 B7 E1 ^& u$ p4 n9 a4 Z+ t
    51
    ( e& W" I( j8 U$ M5 ~0 E9 F6 S* H52/ C* P* S  u* S9 {' O9 K, I
    53
    - k$ E2 H6 C- a" G* I' e0 Y540 s; B, q8 w+ [. g' y
    554 I8 k/ k7 F3 Y. R
    56$ ]( v2 f+ w, `0 h, Z: k2 ]( C
    57: t5 b% N" z0 x. q9 r: A
    58
    ) o. {6 p4 n$ l( D6 `& C" {. g599 M- E; J7 i' M( A
    60
    6 h: T8 G0 t; s7 M1 n9 i61
    ( f- _; F  _5 c3 `62$ n5 @0 z# E$ @& x" b
    63: m; K4 R+ T8 V  I7 U- \* D' K! b
    64
    # b+ z( M/ ]8 O' o4 k1 ^65
    0 Z# B* p! ?+ I  J66
    2 D9 h: S  x  R% y67
    5 q" Q3 d/ C+ J  M68
    0 b) Y9 u0 ^1 J$ Z! ?1 I: }. |, m69$ ]- m' z! C5 C2 i2 r
    70/ d# S4 d# ?1 `$ o, v1 I
    71( ~/ ]; i, D1 j* N7 V" |
    72
    0 [3 Y, q: M/ J& x$ q# g. Z  i73
    " q! f$ f9 z# J' S74' @, l7 f6 l1 L1 C+ A
    75
    $ X  \: B4 I$ O' \9 E* F9 f$ J/ C1 u76  l7 f) y% @, }% b- w7 `0 J
    776 `3 ]2 J! p( E% O$ @7 k2 Y" a
    78: c% }* X3 f# }8 \' I5 g, c
    79
    / H* ^4 m: w1 Q: q% F' [( }3 Y- l80
    + H4 {- Q" N' q  m4 l81* L4 L/ N) ?6 i  Q0 h$ K% \5 Q1 D
    822 f4 w9 S- Z' e( M
    83
    " p* S/ [+ [2 p; G7 A7. 设置需要训练的参数" }# _' U& B$ a4 C, H  @5 R
    # 设置模型名字、输出分类数
    # k' x8 [* D0 wmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)
    & }2 d) a; X4 y; y" E$ Y: H5 S! b3 u# V5 A5 Z+ F% K
    # GPU 计算5 V& k7 N& `, A; B& v" N9 o' K( J8 K
    model_ft = model_ft.to(device)( {5 s- \9 I) |2 r) S6 }4 W
    + j' S. b) @: @' v
    # 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取4 A) h: I$ Q2 i0 O6 _& Q
    filename = 'checkpoint.pth'% B7 L+ w( \2 ~( I1 _6 a+ R5 U

    / S  u# s' y: Z# z# 是否训练所有层% F" B: D# @. Q3 ?4 p- Y
    params_to_update = model_ft.parameters()7 z: A, |  u& I: P1 v! e3 K0 l$ @
    # 打印出需要训练的层
    ' t* ]: c- }* w( d, L/ o" Tprint("Params to learn:"). A- D! J% w* x) ^
    if feature_extract:
    & O, w" R9 K3 D- ~4 K" F% z    params_to_update = []& O% O6 W, [: t3 g0 n
        for name, param in model_ft.named_parameters():0 a$ w/ p' y: T
            if param.requires_grad == True:6 [- W$ W! I0 V9 B" t# g
                params_to_update.append(param)0 R. x2 g+ V  _% N. F% E0 c
                print("\t", name)2 F7 e4 v( h  B! m) x
    else:
    - \' k3 k# T: h* g5 ]    for name, param in model_ft.named_parameters():
    8 ~+ r9 Y; K* R+ d        if param.requires_grad ==True:, H" m" N  M, [7 V2 k$ U* K
                print("\t", name)% M# i- J# {! I1 o
    4 [6 z! J9 L+ Q
    1
    / @3 Q* f3 T2 P- ^# @2& }$ [1 ]# a2 w
    3$ G& c8 x5 T* U- b" y6 F1 ^+ `. D
    4% y, x2 `% P, L7 R/ n2 R. M
    5
    ( I, S. r/ s  s( M2 ~2 q9 b0 ~6
      f3 ^$ [) V2 ?, ^6 j1 ^8 E7
    * E) _( \; V+ ?8  K$ F3 |' x# k! J- f
    9
    . x$ b- ~( T- g+ Z10
    ( f( ~! a$ S- L5 s' D11
    ! u# a7 X8 r& \# R$ m- L12
    * q  a5 t" N& G: }3 Y- N13
    / u; ^( E8 R$ j  f: ~140 f+ a, c1 |7 R0 t0 L
    15
    7 Y0 M* C: v- V) K! g16
    & M; q3 |% ?$ F1 R: A174 ]% a8 w0 g+ b  _5 C; Y) U! f7 q
    18) X  ^5 P  m6 ?0 V0 b/ J6 N1 Q( N
    19
    . F4 ~, h# _7 _) @) J3 f203 X; K7 X9 Z  V2 P% C# j5 D
    21! K7 C) `! `" x3 ]1 t" A! {* V" ~- D
    22
    ) V) K0 J9 J/ ]" c$ _23
    + t8 T2 W8 M- V3 T$ u  A' NParams to learn:
    # `4 u, G$ P) g. M3 y7 `& m' c  R         fc.0.weight  `% d! L5 W( V4 r
             fc.0.bias& k5 I5 _& p% ^$ T+ f& m, R  x1 ~$ X
    1
    7 M' C! P  s1 R2% R$ f; [2 z3 U4 A- w: f3 |% W
    3
    2 p8 N# Z, _2 e8 D7. 训练与预测# Q) O* k# o3 d2 S2 o
    7.1 优化器设置
    $ N6 ^* Y4 N: u; N! \, J# 优化器设置
    / `! d3 i# \3 o( Aoptimizer_ft  = optim.Adam(params_to_update, lr = 1e-2)
    , A( K5 [, ^2 j, H! w# o# 学习率衰减策略# n5 Z6 ^# M$ }: f+ E5 q
    scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)
    ! t4 b4 Z, p/ e# 学习率每7个epoch衰减为原来的1/105 x: Q8 P0 I; ^3 ~% q
    # 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算, Y2 [# U1 B) g. D% u' o8 |1 W% k7 D
    2 u# _# L9 `$ B  v0 w* H% U
    criterion = nn.NLLLoss()
    ; @& t# q1 g+ x! M/ _/ X1
    ; t1 f& X2 C. ]- P2 E; G, L2/ u0 w7 q- l! @  X& L; H* C' _
    3
    ) h+ [: V# f% V& d# T1 H  z4( [% b/ R$ q# X: R) z9 \0 f5 s0 {
    5: O. J; h- N4 }9 o. f6 k+ ~0 v/ O
    6, N! B( `" B8 j
    7
    4 H6 |6 h: }- w1 T7 Q8
    - `2 m; R8 B( o1 K9 y5 F$ t+ F$ d# 定义训练函数, @& Z, U9 x* c4 ~6 J, h# u" z
    #is_inception:要不要用其他的网络
    1 o6 o5 O4 W* S. X5 xdef train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):8 r& z0 F8 t( m7 N
        since = time.time()1 y7 C7 j& `7 f1 O
        #保存最好的准确率
    5 q6 m* R6 H. }9 [    best_acc = 0
    4 _* G" X% w  h. d    """
    7 o. B' c2 W2 [  _/ a% f- a2 Z    checkpoint = torch.load(filename)- M' B4 a2 R+ w
        best_acc = checkpoint['best_acc']7 t2 A5 v7 c: N
        model.load_state_dict(checkpoint['state_dict']). t& k  @* S3 m$ a! D" K
        optimizer.load_state_dict(checkpoint['optimizer'])$ L% }; C, _3 a/ O0 X- p0 s( b
        model.class_to_idx = checkpoint['mapping']
    6 B# G& s; t( [    """
    . }! B4 J; ~) t5 e/ _1 a, f    #指定用GPU还是CPU& r# g) @9 y" R" O) q
        model.to(device). v  q( Z/ J! c; O9 O" g3 i& n
        #下面是为展示做的
    ; y# F* n- L8 s6 v( N" S5 o: C    val_acc_history = []
    $ y2 F, g. A& t; {    train_acc_history = []( j! f1 f8 ~8 Z" Y
        train_losses = []/ V6 O& F# j4 U* i% R4 y
        valid_losses = []
    4 q* c, |) s: t) |4 e$ J  j; D$ I8 R6 v    LRs = [optimizer.param_groups[0]['lr']]9 k/ r' r% O7 C  p. s
        #最好的一次存下来+ i& a0 t: I! G4 ^
        best_model_wts = copy.deepcopy(model.state_dict())% s. ^( J- Z+ h/ I# [7 K7 W- p

    ; u4 b' z) D  F# g; A6 B; [    for epoch in range(num_epochs):
    1 j" S3 N+ Z/ r, K- t        print('Epoch {}/{}'.format(epoch, num_epochs - 1))1 f( m) l0 p, l! t9 p" q: {
            print('-' * 10)
    3 c& N6 y1 C4 _) D! ^
    6 f; A4 i% J  i# e$ O8 n        # 训练和验证
    2 p) \9 P' x/ H- [        for phase in ['train', 'valid']:
    ; {+ J* U+ a6 i  y$ E+ v1 P            if phase == 'train':- J& W8 q& [+ [" Q% {% A
                    model.train()  # 训练0 J, C& ~: C! b. Q  d; A
                else:% J: M( t7 R/ {+ \0 K4 T0 b
                    model.eval()   # 验证
    $ W4 T, Z8 d: M8 p; ^, l6 F
      k( a8 c5 p; w            running_loss = 0.0: k8 |; y0 H5 G% x  J3 r+ [: Q0 U
                running_corrects = 0
    * i/ ]3 E, ^+ u0 n) \0 `, k
    % h  G) l: u3 h  k/ i/ ]. d            # 把数据都取个遍6 b; h, d. L/ f9 _  V* `
                for inputs, labels in dataloaders[phase]:
    5 w: J! e, u& |- m3 R8 Y0 F6 [                #下面是将inputs,labels传到GPU
    0 N; Y8 I/ V& G+ q6 h( _1 H, c+ Z                inputs = inputs.to(device)
    7 i7 M4 b4 }) M% S8 E9 W! s                labels = labels.to(device)0 {: V/ F1 v- s" |, u# A4 Z
    3 C. X8 v  l! o0 M, M3 a
                    # 清零+ g! }2 W  l- d8 W5 G6 z
                    optimizer.zero_grad()2 a' M$ R1 a3 P" f; C
                    # 只有训练的时候计算和更新梯度
    ' U0 d* G% K- D                with torch.set_grad_enabled(phase == 'train'):. t4 o2 W; s& J* Y# Z7 Q
                        #if这面不需要计算,可忽略+ a9 Y; o& v+ E! f
                        if is_inception and phase == 'train':
    3 u8 ]8 \2 @6 A; a$ q7 [                        outputs, aux_outputs = model(inputs)/ P& ^; K, ]; P7 l7 h# Q! U# K
                            loss1 = criterion(outputs, labels)& t; M; S: {! R7 S% q( }
                            loss2 = criterion(aux_outputs, labels)- d4 [- e) f: W( F  ?
                            loss = loss1 + 0.4*loss2  W. z, K6 z. m6 A- a  w; v3 O" b( C* y
                        else:#resnet执行的是这里1 L3 X: [$ s# f0 Q( U. u
                            outputs = model(inputs)" F9 b% _- S  @, e- p! F
                            loss = criterion(outputs, labels)
    % y2 d! Z5 C' f6 i" x+ O6 _1 u, ?: r+ [5 a8 J" W8 e
                            #概率最大的返回preds5 x' o$ G, r2 U' \& b
                        _, preds = torch.max(outputs, 1), T& U$ V5 u/ x9 L+ I% G2 F

    9 h. d, O& v! G- i; X+ A                    # 训练阶段更新权重
    6 Q7 g$ i! ]4 W2 j+ ~                    if phase == 'train':0 ?1 `# P, B: a' g# ^
                            loss.backward()
    ! ~: C+ X! t* D$ g; T                        optimizer.step()
    / g' s( w/ Y  ^' B- I
    . L3 k+ R0 I- j                # 计算损失# }$ n. \9 D- V  X
                    running_loss += loss.item() * inputs.size(0)
    & I+ _8 R/ w1 e6 R' ^# }& {9 T3 Y- c                running_corrects += torch.sum(preds == labels.data)
    9 T4 X) w9 k" [' x2 R* ^
    5 m7 W$ I& ?$ n9 |% q            #打印操作- B) V" ~, j& e) s
                epoch_loss = running_loss / len(dataloaders[phase].dataset)( O7 r8 |; a4 o9 t7 a
                epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)$ H' D& T" e- ~( c) A  s# J5 x
    & G, g. ~, D' @7 P( K% j

    . @2 A3 V/ A9 ~: V            time_elapsed = time.time() - since! h0 a5 l  h0 E, K  C, x; q) h  ^
                print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
    , m% i' |) g5 a. V. C; I1 b            print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))
    : _  f% R2 [" H# F: r7 C/ t" \+ N; T- P6 m3 O
    + D  J4 b2 h6 d* s9 F
                # 得到最好那次的模型
    ) j+ D/ ~, {/ H: w            if phase == 'valid' and epoch_acc > best_acc:5 m5 H. ^8 \6 V+ l& L  A
                    best_acc = epoch_acc* c9 V1 \: @6 \) R
                    #模型保存! v8 V3 u7 h, w" F. u
                    best_model_wts = copy.deepcopy(model.state_dict())2 }+ a7 L8 |* w( `$ X" K
                    state = {
    3 p& K- y9 M! |                    #tate_dict变量存放训练过程中需要学习的权重和偏执系数8 k. i& o( ]+ G5 H) _$ i3 ^" U
                      'state_dict': model.state_dict(),* }' R* _3 p1 d: G3 s
                      'best_acc': best_acc,
    2 U: y: l/ r5 R, g                  'optimizer' : optimizer.state_dict(),' R) B, W' w# A, _
                    }- O3 G9 O# w, u# b3 I
                    torch.save(state, filename)
    ( I8 `2 P. K- F! m0 c5 e- M4 h% S            if phase == 'valid':
    $ S. n* n) o& y0 i- Z% {2 D                val_acc_history.append(epoch_acc)
    # X$ `6 P% o: u" u                valid_losses.append(epoch_loss)
    - [* p  C9 d2 m" ]. F2 G                scheduler.step(epoch_loss)4 H% k% S( a* K1 v2 c; B5 U+ h
                if phase == 'train':* R7 m  Z. a% P) z
                    train_acc_history.append(epoch_acc)
    5 ?$ x1 b- t  Y  D) s2 m9 i' f6 d                train_losses.append(epoch_loss)
    6 _/ M$ ~6 |$ ~. y# w6 y
    + Q3 G" m* J+ a! o        print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))
    * @  l; e2 O/ D7 d" Y        LRs.append(optimizer.param_groups[0]['lr'])
    3 O- ~: M; J5 i( e+ l        print()$ n( ^+ F8 q" O
    9 V. p0 k8 M( w7 a# s
        time_elapsed = time.time() - since
    5 ~4 E# T2 i9 f  M4 A; t& ~. c    print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))3 X( `7 _  a  E+ A* d1 T) e+ A0 _' A! ~
        print('Best val Acc: {:4f}'.format(best_acc))
    ! Y+ |; J5 U  ?- ?2 I. G: T
    ' K' J4 x% f5 u' E$ f' l8 M    # 保存训练完后用最好的一次当做模型最终的结果, o5 x/ p$ C$ N" |# U
        model.load_state_dict(best_model_wts)
    * i1 ]% h' C; e$ ~2 _4 X  O8 F    return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs
    ( e! J# T# f1 K" d# J( I8 j7 n/ r' t
    8 A+ ]' q( V; A  d$ J
    1' S: p/ R( j, C$ r* I/ K' H5 d5 U: W
    25 c9 Z+ G* |1 p' Q
    3; U$ t1 e) \: A7 T( s- C1 W
    4
    % t4 e/ w" Q8 b" ], Y4 y6 R7 ^: C7 _5
      m/ O. g0 {. n: L& Y% {0 _) U: J% h6
    3 F8 v# G% D. ?. V4 y6 K- Y" D7+ Z5 x4 M3 \# @3 }' `
    86 f- B$ i8 G; q& o
    9# p) K4 e, c. U  ~$ _
    10
    2 z' l' [7 V: @) Q11$ P3 d! D4 V4 N- c5 j4 M+ Q- O
    12
    1 n, V1 a  \1 @; `13
    4 @& k; W+ ?0 t0 q) L- }- t14: K. g: I, j- X2 s5 D( d5 }
    15
    $ L/ p& W9 g6 k% f* y& R9 I' w. w, B16' U3 l  ^' a, q3 J' K# @
    17! H; Z. o  V' Z& ?& d2 f$ P  E
    18% e  M. f$ B5 t+ v  B5 z/ S3 p
    193 H* ]5 i( q- v% R/ \
    209 J  [% \% X9 I8 k0 I  L% t
    21
    9 k! ?: r4 f- |; \& X+ j3 l22
    ; Z# {; g, V' g" w/ c; Q23' |; I' O9 U  i# g$ P' e# ~
    246 ~8 d9 b6 o! {* O! ?% T
    255 S0 q; [1 T: P
    269 [2 F0 G1 |0 Y. E8 Q. n+ T
    27
    * v5 W1 Z+ u* g5 g0 J28
    3 z4 u+ h) V( M29
    8 U5 a* ]8 V  ~% c( F307 [( j+ X' A. e- y' Y6 ^
    31: O- O$ n& B3 N# b' ]  D% {+ Z
    32' \' j* s9 [+ P8 Z" t% ]9 ^! X
    33  r. n8 f( h$ R; M9 r- }
    34' a2 `4 S; }& _
    35/ @  L2 ~- K) m/ e: t( M
    36
    6 \  \$ Z! L% W, R37
    1 f( I& b/ d6 g9 o8 f+ H* w  B' ^38- k5 Y6 @. }' F
    39) i( p* H5 U0 h0 ^% v! F
    40; X4 `$ G3 V: w& E, s8 N3 h" Z
    413 _2 f6 I) k# s6 H" W& v2 U$ g
    42
    ( H* }9 a, L6 c43
    : `. ?# Z; T; s& n" H44
    $ ], C5 B( I# m0 K7 I45+ s, T  G; y, l/ A0 M$ T4 }
    467 Q& [7 K$ {6 N+ r5 a
    47+ N9 J2 H, k4 A- I- m
    488 m' d4 @6 F! [9 k) x
    49# q* h3 k5 `* s. Z
    50
      r. O4 e: u, a0 c) G3 b" p51& a% u" B2 u3 \2 K) ?  p$ G
    52$ {/ }  h8 p0 [, t
    53$ S' N* g5 `# }
    54! `$ b' @2 T! O) M
    55# p( {8 f/ ]3 o' n
    56+ N% S- g' O7 L
    57. |) }3 K( S* x) {0 G, z1 y+ p
    58
    ' p7 S: s" O9 y1 S' ]" N, }59
    % l; t) E4 p# |( b606 }8 v' d/ X0 o: d, u+ \6 |' Q0 o
    614 B8 l3 O7 E$ a6 q1 S
    62& |! Q9 \0 E( B9 E5 a( C' [4 I7 M
    63% U) H6 Y4 q6 H+ t+ r4 t4 S' d; q. H
    646 t: f# \8 U# y* A
    65% v* f0 @" ~' [& p
    669 B& \5 V7 ?: G' Z. A
    67
    # @' s/ s6 ?* w4 [( m0 e: Y4 X* h68% p( _0 P( @2 _% `! X
    69- o1 o$ g* ]" R9 u0 }8 `! O' H7 `% d, ~
    70
    3 h, m8 J: F" x5 U3 T0 x* K71& C7 l4 T/ w% K# V) B" T" ?+ {
    72
    + d3 W" E9 ]! j4 S' y. S73
    7 {2 V$ E" ?; M' V* ~74( l3 g  P  W  V( x6 M8 l; {% o( u
    75
    . ?# o& I: ^" t/ A76
    . d4 J! ~  |: g7 G  q77
    " j) E4 q' N( D: V% M78
    . \  _) Z! w" m5 e) V! w1 P9 p" \8 }79
    & a1 n; V9 {3 r6 P80
    ' N; t5 b* q; U. D2 }81+ K5 A* `4 u/ d" s: T) y0 x
    82
    + H4 b1 M' z8 _; k6 `/ Y5 `, ]! d83
    6 N. o+ c* b: D0 Y6 O2 B: [2 E4 k842 r" D. q" X7 q" f9 X6 G2 i7 b* x
    85
    * T2 K3 ~3 E$ m  y+ d86
      u2 ^2 Q! D. F% P; \( X87
    $ `7 d4 A- Q0 }& }$ K883 t- e. L! W5 a$ ^
    89
    / Y' M" I% i  O  \90
    " ~+ E7 O% a2 P4 i919 D/ o) I3 T. X( J( A1 \
    92: u% T: _8 K$ ?$ [, m" O
    93
    ) G# \8 i8 }5 w" ^2 |( ]( C) r8 d94( Q( S; n" O. p
    95
      C$ b) ]& \/ K9 h96
    ; w# T1 @; }: [972 \1 V+ p! ~- d/ X
    98; R1 H8 f& x2 ]6 ]9 m7 Y1 n8 [. k7 x
    99
    + g" S$ ?9 n8 V- u) h# x0 \100
    * T0 W$ n# J# Q8 h1013 j; T8 D8 n! |& l
    102
    $ s. c/ Z8 r1 b: A' w- y103
    ; q4 U4 D0 v, F; F9 p% s) y* [' W/ v104& h; W2 r! |$ C. p- f& d
    105  T9 @( S+ \+ N; [) B
    1063 G  \. I. o5 Y9 a4 X
    107
    4 `* _- E1 a' ], T108- b+ h2 Z+ ?+ G& Q% o% y
    1090 j- X/ R. q: s1 Y! M+ F& a
    1108 d& f6 d2 S9 K1 y
    111
    " i/ W; N+ v' I: C. D) f# n112
    7 y8 s( c2 m: l7 k/ f- a7.2 开始训练模型
    , s  v4 K( l" M" z% K# v5 J2 y我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次$ g$ o' d. ^" R0 Q

    9 C$ N+ ]# M& Y9 I#若太慢,把epoch调低,迭代50次可能好些4 b7 j" |% i6 f- K
    #训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合  z' |, _9 S3 H# F1 d* d' l
    model_ft, val_acc_history, train_acc_history, valid_losses, train_losses, LRs  = train_model(model_ft, dataloaders, criterion, optimizer_ft, num_epochs=5, is_inception=(model_name=="inception"))
    2 t1 j( ~$ J7 W; z$ _- v* j8 f7 n  }( v* y. K
    1
    ' P% b# H! f  ^$ k6 q: c2
    ) S$ L2 J6 Y# _3 t2 f39 E5 Y' W+ E* Z8 D
    45 i9 i3 D# F# e' ], Q" F; z
    Epoch 0/4
    $ e  X' }. t; {5 g" H7 j, J" o$ e----------
    $ p) A" D! e) E6 d, f4 pTime elapsed 29m 41s; k2 K+ P2 I4 p* V% Z; p3 x
    train Loss: 10.4774 Acc: 0.3147$ d& i& V/ K/ y5 G
    Time elapsed 32m 54s
    , i0 m8 w* Y- I8 [9 Svalid Loss: 8.2902 Acc: 0.4719- a# v$ k+ f; d: k7 b
    Optimizer learning rate : 0.00100001 B% d( ?- X' ^* _

    / o9 t0 h6 D: G1 B  L; nEpoch 1/4* L+ v4 I" m$ O/ }3 {3 C
    ----------) ]& l9 R- k3 K$ @7 K/ Q$ \5 ~
    Time elapsed 60m 11s
    6 c4 x$ A5 @- k  Ftrain Loss: 2.3126 Acc: 0.7053
    ! Z& @9 G* ]0 k8 DTime elapsed 63m 16s8 q1 S# q& q' c- W( m4 j7 [9 E# ?4 v
    valid Loss: 3.2325 Acc: 0.66260 y6 E4 r; ~# L4 s: u- s+ |) j
    Optimizer learning rate : 0.0100000
    7 @2 i( T  ?# W- I  A: s. D
    5 J8 e. a8 x$ O& e9 d5 B, cEpoch 2/4
      I* y1 s- c3 z2 X4 X) `5 z$ U----------* z; Y1 W8 {5 k; B2 }
    Time elapsed 90m 58s- P, [; a* y( A5 J3 W
    train Loss: 9.9720 Acc: 0.4734
    ) c3 j0 h+ v1 f7 xTime elapsed 94m 4s
    6 d/ G7 q* J$ @% @valid Loss: 14.0426 Acc: 0.44133 I% X# b& N# m% p
    Optimizer learning rate : 0.0001000
    " V9 _! A4 |$ d. S
    $ p; n7 D$ C' ]Epoch 3/4+ k# x6 J) Q& b
    ----------
    4 O9 @0 e' F& Z' t# yTime elapsed 132m 49s
    . H: l& T0 S0 c7 }train Loss: 5.4290 Acc: 0.6548* A# P, M7 q# O  K$ P
    Time elapsed 138m 49s, H$ E0 V' \; B& a
    valid Loss: 6.4208 Acc: 0.6027
      O4 e+ M9 i' e' XOptimizer learning rate : 0.0100000
    & A) W! Y; V5 b
    % }3 ]1 T7 @1 lEpoch 4/4
    $ W8 ~4 G+ A6 W# p! \$ \& ]----------- k6 H8 {6 R) M9 t9 Q* ?$ z0 W
    Time elapsed 195m 56s
    3 u% R4 ~' I. t2 ktrain Loss: 8.8911 Acc: 0.55198 c) r. y) I  ?8 c' t
    Time elapsed 199m 16s! c# m  J1 k* K- s
    valid Loss: 13.2221 Acc: 0.4914
    ( d" r0 V( L! a  W8 G1 YOptimizer learning rate : 0.0010000. k& W/ u+ s$ i+ o- t' y
    / H6 N) u  O0 X: M
    Training complete in 199m 16s7 r) @% {  b2 V7 E9 V
    Best val Acc: 0.662592
    , O1 l0 z( L; B" n" q3 z( X0 |1 S5 K
    1- U6 L, u  E" e. i6 p; }7 z
    2
    7 G5 W) S" N& t% A- B; X3
    : K. T' ?8 i6 p) O4" ~5 A$ H. ]! m, a6 O
    5
    ; i" x, \2 I$ {/ G- Z6
    ( {9 ]3 r" K2 [* l+ }3 P: x7
    : V( U2 y* ?2 T) d) B5 A8
    2 c( u2 J+ s! ~) [( V  j* ^" N  o! `9 a9( d1 ^* o) N4 |9 |5 j5 \" X4 R) r
    10/ X% {8 L1 g9 ^+ _$ O4 I4 P: ]! m0 k* S
    116 q& j% K6 j. A; L
    12( ^+ k( V7 K4 u' |$ I1 u0 b
    13! I. L0 R2 x8 g& E+ w# J
    14
    % x- }; p4 w: T% x! Y; t15
    ! P& r0 K% L9 P0 I16* F; @: V- Y" V/ |9 U
    17
    " ?6 E$ C  f+ s  O* I4 W) G18
    , q6 \2 `. H: W, H3 J6 w19
    % M2 X. M$ s) Z6 N% k0 V7 X20
    9 e3 c/ |) t7 g. C2 k6 ?' f21
    ; G2 f9 B1 }! y6 p" N22, `7 z6 u9 p* b4 }/ \# p+ N& x
    23
    9 i0 f" ~; g, a241 P6 M. V# T# ]& E3 P  v
    251 }$ B, f' A0 |# r, V% P
    26
    4 B! Y+ p2 I9 j27+ h4 y4 u2 C- t
    28; @' i& P; T$ I' A) v2 S5 M6 Y% K0 ?$ ~
    291 Z  ?# X" H( K  H/ v8 f
    30
    ; [5 T7 @+ b0 z  D& H* C314 V+ J) ?" d* ^$ J7 s
    32& U* V5 t7 A: A+ Y
    33
    0 c% {6 c7 Z8 a, ^+ n& t34. w: n, b1 z  N5 k
    35# d7 F" N/ S) a7 h
    36
    $ o6 }. g" c' P37/ d7 Z" |/ G# W/ D5 S0 v6 O5 }
    38
    : J, X% w$ O5 K3 H3 p39
    7 w4 l0 ^; O" [, e40& m5 ]6 ]) c) \! o- k" W3 E. }
    417 f+ A) |1 U# m7 G4 v0 [
    42
    5 w- c% D; D7 F* F7.3 训练所有层
    4 h! p" i9 B4 s0 F) E  @# 将全部网络解锁进行训练( Z+ j6 \) [) I, M# H( u' E5 \; p1 |
    for param in model_ft.parameters():
    3 l# ]; F  ^1 j$ F7 q    param.requires_grad = True
      C: y) d( n! k5 l1 g9 D# a0 F. W6 B' h9 c0 m! l
    # 再继续训练所有的参数,学习率调小一点\
    2 E. O. h3 _( u7 qoptimizer = optim.Adam(params_to_update, lr = 1e-4)+ k3 C; ?. z+ ^2 D
    scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)3 f; t. D5 Q# G& C7 z; ]
    # A9 E! ?. u* g
    # 损失函数% i* y2 y& p  ?) X
    criterion = nn.NLLLoss(). }/ Z& V9 k3 x# B3 I. K
    1$ A4 r0 F+ F  E8 ^. s
    2
    / K' z: ?& C" v3
    5 H; u' Y2 V- }5 [& e7 H4
    3 H0 j7 ]. O" `5% A6 G2 v! k! e: n+ d. A; F, t
    6
    7 Z/ q- l' A/ p0 H' C+ p8 q; ]$ P! D7
    1 D7 a7 _8 w' Q' }8, D2 X9 U% L% i) O7 w
    9
    $ i" Y8 u  t: g) w" ]5 i10
    $ Y, ]+ j' ^6 G3 m1 m  W3 n. A4 \& _# 加载保存的参数/ |; T; D9 k! n" f# \
    # 并在原有的模型基础上继续训练
    - T: N0 @: C3 Y, ^3 {2 a. L# 下面保存的是刚刚训练效果较好的路径
    * j; [3 v8 T* S' p( f* tcheckpoint = torch.load(filename)
    ( H' X% q3 l  P2 y0 s3 z5 ebest_acc = checkpoint['best_acc']$ z; ?. ?" M% R8 W7 q$ Q
    model_ft.load_state_dict(checkpoint['state_dict'])
    8 \3 K; |  w" foptimizer.load_state_dict(checkpoint['optimizer'])+ R1 n6 m0 _$ j6 Y
    14 T+ {0 s0 }' P
    28 r2 m" a5 @, T! q+ l* v
    30 c" U: ~2 d& F$ N/ J' Q( X
    4
    % w& m2 H( t4 B. N+ J: X56 W+ z) T# |' ~/ O0 |: _% [( H" _
    63 z4 S$ H; x5 w+ m/ x
    7; I  H( v) Q7 ?3 S1 _
    开始训练
    , C& O- W, u$ S) `0 m4 a9 t- T! [! [: q' Y注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考5 u+ Q: N: f# {: j7 W
    8 b; \7 T# D" Z7 ^1 z
    model_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"))  o# u' y+ W- V7 o4 |
    1
    ' }8 P# J' n( K  M6 u7 jEpoch 0/1
    0 i8 [& M4 w& ?( D----------$ r8 K/ @# E6 p% ?4 `, }
    Time elapsed 35m 22s
    4 b! C! m. a$ u2 S3 Mtrain Loss: 1.7636 Acc: 0.7346  a) B+ `0 m+ w3 _- I8 C7 V+ J
    Time elapsed 38m 42s
    1 O/ g7 y* g' H+ {9 rvalid Loss: 3.6377 Acc: 0.6455. N5 e! r- p- S( {
    Optimizer learning rate : 0.0010000
    ) }! a( d' u$ z0 {  v7 ]# T) Z) P' g
    Epoch 1/1
    4 V, r" P' W- o# }1 }3 ^& z0 [----------
    - E  A! G! G# h- D- b9 kTime elapsed 82m 59s
    : w* ~( [$ M2 q% P: O; q+ ^* b  ytrain Loss: 1.7543 Acc: 0.7340. V* ?* n! j: N: D5 U
    Time elapsed 86m 11s
    : R. s% Q1 }) e0 [0 }7 ^+ Q- q0 Xvalid Loss: 3.8275 Acc: 0.6137# M+ n- v$ ~4 }1 b, ~( x' K* E
    Optimizer learning rate : 0.0010000
    ) R& b* A0 d. Y2 q  R0 ?- ?9 l4 G7 ~5 [% J  r  E
    Training complete in 86m 11s9 A* P3 a. @& {5 w$ G  A* |5 ?
    Best val Acc: 0.645477
    - n3 |. e7 `9 C" M: p! B
    % z5 ~& l( D* t& T% Y16 o0 u) |6 q4 y* D
    2
    ) \; H0 e% u8 l$ W+ p' \3# L* j. C, J! D( I! I' c, ~3 W6 f
    4
    2 k, y5 R( C+ Y7 j5
    9 o# t3 i$ ]  ^, b6 R/ a. E4 |$ v# d6
    . ]' A( T( ^# [* e7
    # k* R# M! S3 ?) Y1 R0 v5 g1 p8
    1 M/ _( e- _2 ~* z1 K# B9% d! a* r# z6 z/ K  i; T
    10$ g6 F( c. o+ W
    11
    0 L" b4 ]. u- @* r12
    ! L3 V2 r3 [0 A$ {9 J- }13; V# M1 b# E. {
    140 q+ x  i1 T. V1 g( ^6 f# n8 J
    153 k9 C" R  h% z: ~3 j
    16& ~1 N) h( n- e8 {3 `% T7 Y, |3 B, F1 ~
    17
    " r* @+ I8 b1 x18
    0 v7 \- c9 U* A) u; b9 Z" f8. 加载已经训练的模型: p; a2 t% q* F: Y! G+ q/ P
    相当于做一次简单的前向传播(逻辑推理),不用更新参数8 o" q, r; K( F) S

    % j/ t! N) `- ~. Xmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
    0 i" ]$ Q5 Z+ _$ [2 p' `; {3 o6 s! o: R, W5 r8 F7 n
    # GPU 模式- B( e  H! C, o: D/ l1 _1 x& M- W
    model_ft = model_ft.to(device) # 扔到GPU中
    7 K: {1 g7 f' T* ]( F
    $ }& W6 b6 @: ^& B3 D  B! a# 保存文件的名字
    1 k- s2 G9 T# \0 G5 C" nfilename='checkpoint.pth'
    # v! i6 {5 k5 _' j5 T% t& ^
    . Q( }2 z' w9 x6 e# 加载模型/ q: W& d+ o3 v
    checkpoint = torch.load(filename)
    2 ~( V; a' a# O9 kbest_acc = checkpoint['best_acc']
    3 l" h& g/ z& X' F/ o( A+ e  Imodel_ft.load_state_dict(checkpoint['state_dict'])7 F3 ]4 [+ U  `2 \2 f# \
    1" F9 d& f  N% j  X0 n; _
    22 @6 K8 E. b  T/ p9 w, j
    3$ V7 e/ H5 g; p3 `7 V) Z$ r
    4: F9 o$ T; V  r) g
    5
    % y, U4 b6 R, _% H! v* s1 E' H6. F% N2 ]4 c' I, f5 j
    7' d7 G7 g4 R0 q+ I* e, Z8 \7 n$ ?" g
    86 c' b5 _; w0 O6 }
    9& ~7 Y1 r3 P( e: W& S* D
    102 h2 M$ F4 c( m/ l5 _1 d
    11
    % J' x  Y5 I( q0 }, n3 c2 N5 Z12
    3 F- f- V8 u6 Z9 L- a+ s$ ~<All keys matched successfully>( C1 z! ]1 x3 h9 k% r/ N6 h" y% q, y
    18 j2 h' C. a" f3 S! K$ J
    def process_image(image_path):) h; |- }. n4 L' [4 S) U# a$ ^
        # 读取测试集数据6 P" R' ]5 Y$ [9 @, N! r9 n
        img = Image.open(image_path): u- h) d1 Z" H
        # Resize, thumbnail方法只能进行比例缩小,所以进行判断
    $ e+ l9 K/ r; Y    # 与Resize不同3 E1 E* W; {# o" z# d
        # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小5 o6 A' f) I& ]! e: |! k
        # 而且对象调用方法会直接改变其大小,返回None
    & T' z4 _4 I# \8 ~+ o# [    if img.size[0] > img.size[1]:, c3 U5 c/ A' w7 ~2 J6 a8 S
            img.thumbnail((10000, 256))) a0 e. O2 h  j0 g2 a  T; S$ b
        else:% y( X! I( k, }7 d. q% }! L4 e7 L' `9 e
            img.thumbnail((256, 10000))# t3 X" B& T; U! A8 I* i
    * c- ~9 l5 _4 I" ?
        # crop操作, 将图像再次裁剪为 224 * 2247 H7 {& }/ E, G: {( P, |2 T
        left_margin = (img.width - 224) / 2 # 取中间的部分
    . X9 G5 v5 Y$ W, b7 A5 l    bottom_margin = (img.height - 224) / 2
    0 b4 |1 J! H( X- W0 R    right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度
    + r' w) a! ~' ~, a! f    top_margin = bottom_margin + 224
    / N: F5 [5 R9 f$ C( e: K6 u  W4 z" R. J+ U/ o  L
        img = img.crop((left_margin, bottom_margin, right_margin, top_margin))
    6 y0 ?( V1 K! {$ y) b+ x9 F3 ?8 l% F* o, s. |  y; ?* x$ o
        # 相同预处理的方法& Z( d  X, e, |9 b0 O3 K
        # 归一化, ?$ g! ^+ B, |5 {1 T
        img = np.array(img) / 255
    % F+ w( M5 F$ Y7 L$ W    mean = np.array([0.485, 0.456, 0.406])! a2 p; G) H% B( a3 o# m
        std = np.array([0.229, 0.224, 0.225])! Q$ L1 L+ W. c$ `
        img = (img - mean) / std" s9 G8 k" `, x! a5 T" {3 Z. O

    % w+ i6 n$ A& r+ ?3 r" H( p. \    # 注意颜色通道和位置9 Y0 T+ [8 _$ u3 Y
        img = img.transpose((2, 0, 1))
    8 e( W5 F; z* r6 |
    & v, Q* s; s) k+ g    return img
    # y" e* i/ ~# Y4 Z5 x# N' A4 `! d: R/ E/ }; Z" O
    def imshow(image, ax = None, title = None):
    : I; D( e; p- @% e    """展示数据"""
    0 |/ S) i4 h4 H# K1 |9 q; q4 m    if ax is None:( c" A# U$ `1 _% c0 x. _% n
            fig, ax = plt.subplots()! l) Q' [6 M7 T3 i
    : }2 \& S5 l" c* x, T
        # 颜色通道进行还原
    5 L+ P, c% ^8 d; V% _    image = np.array(image).transpose((1, 2, 0))1 \% Z/ u9 J8 v3 t! h& t
    + m6 F$ h, n. H9 k+ c$ _) \$ A
        # 预处理还原/ M  `) {: X# ^* l) \
        mean = np.array([0.485, 0.456, 0.406])) @9 W: u0 S' B$ P' f/ j+ z
        std = np.array([0.229, 0.224, 0.225])
    ( i( S. k9 ]$ l* }, G9 P    image = std * image + mean; u' D! Y+ b" y% [
        image = np.clip(image, 0, 1)" ~9 U/ S! B; z4 H
    ) u' a, b$ D7 ?
        ax.imshow(image). q) l" b) S0 @9 d" h
        ax.set_title(title)8 c. i" o; r' n! ~1 S4 |' H
    0 b( U* q% `9 j1 P" s
        return ax
    1 [- l; s/ G' a
    + G9 u, e3 |/ E% S6 g' @( t8 ximage_path = r'./flower_data/valid/3/image_06621.jpg'
    ; _+ _. X* p( s. u" q& G! S: L' \img = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理
    8 q0 m8 |/ F5 Y! ?; G8 e( Vimshow(img)2 M( H1 z, J5 |$ ~, P% w2 u, m

    ! T6 G4 K0 r9 Y! k( ~1% g; p/ h& T3 p' b) E( J% H5 `
    2
    " G( H, P" d' i3" M! P0 Q. I! a' X
    47 P9 t* A' V, p$ S+ o6 r
    58 f$ W, [- M2 D- E
    6
    " G( r' R+ a) O" n/ ?5 _; b77 D- P  j/ h( Z- x- V, V
    80 B, `  I0 |% j  x, ~1 ~+ Q
    9& i  W/ C: \: N
    10+ h! j+ c; Y! M6 E# d7 W# G
    11; p, g3 J/ q* j9 T2 G6 E& ?2 y
    12
    ( d# O$ C2 `& }! h% j* P8 g132 Y& }$ z  y0 w
    14
    ; a) k3 |4 p( @, }' ]% B/ x9 n157 K* s! T9 J2 A4 G, ]4 A5 y
    16
      z( d+ ?- m; r6 h- |7 X( H! x17
    ( Y) {0 h0 V- C, _1 H/ [8 G7 _18
    0 R( m+ Y/ L+ h19- C, t/ v" `) u. _; u' _
    20
    9 [6 N' d" ?6 D2 K  {& I21
    : M- l# y# [$ g  E4 I22' z8 \: Y1 F, h2 D8 W+ c- V" j
    234 X. J5 W$ g; ]& l/ Z0 g6 _
    240 K. e! Q$ Z2 t# J+ E
    25* _- M- c, |. p8 u. S
    26
    7 {4 W# W- _4 l  ?27
    # }- {2 o7 p! z$ z28% F& Y2 s8 B; s8 W3 a$ X
    29
    * y* s' k% Z7 H* J8 }- d% [+ U30
    3 x' M: f' K$ y' x2 j319 d* @; j; _; H' F* ~8 m
    32
    , X& c( Z8 o' l33
    * r# V( [, }( k, @( G34
    9 S; U0 X$ C2 Y35
    1 A" `/ O" n# c/ G" T7 l366 s' K" \3 ~* n3 ^4 Z& j
    37; d  h7 t* u# k+ h& N+ V! D
    38
    4 y2 p. M% r5 ^  a39
    . c9 {* m- X1 t- S7 ]40
    1 P) E0 A; x# \* i1 g6 B# b" z$ H41. [7 h) ]: G' K
    42
    5 A: B, q4 ^' r2 }$ a2 x; a43
    ' c. k9 T9 c6 M2 E, [44
    # z/ Z) g$ _7 I45" n9 L9 e8 Q( h5 E- w* F
    46) |9 f, Q, @: N" S* O& b; F( L
    47
    & K1 @% ?* v/ D+ X3 e. p0 B/ ]8 Y488 J; {$ v% l$ w# L8 \
    493 }& [) x0 o+ f$ @, ]4 K4 i
    50
    " n0 Y1 |( ?" l, ?, Z8 S3 A51
    " {3 v; Z4 \# Y4 R, w52
    3 c; k5 g$ f" Y- n! U# G3 D535 k- s& C( ~9 {3 X
    540 g( Q) B' d2 `, f" s" j# j
    <AxesSubplot:>
    2 P9 j/ S* W' x1. h& N/ B" ^+ C5 P9 ~
    " b5 @4 q( u0 F( S' H
    上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确6 \  D. D( W/ K) u: w
    - w/ C; W5 ~2 Y9 v+ s2 H$ t
    img.shape" W. E- P! I/ F4 B
    1
    / j- \6 K" f* c(3, 224, 224)* q  d7 V# c. p% k4 T8 X1 F. G
    1
    ; F: K- \* `) V$ O5 ^! A( F证明了通道提前了,而且大小没改变
    5 b7 Q0 t1 B4 q+ s, b8 n5 A/ T& @8 V; `- ~
    9. 推理4 c# c% O: i; o/ g4 r
    img.shape
    $ A, p4 R+ w0 m" g& ^/ Y( w
    & K' n0 @& ~% D) o4 N/ k1 }0 c& K/ T# 得到一个batch的测试数据
    + g% p4 l4 b) d+ m7 _dataiter = iter(dataloaders['valid'])# {! a* T9 Z# Q' c+ ^" B# ]
    images, labels = dataiter.next()
    6 J* U/ a5 }2 q
    , l% h# }) N) i# l# ~, V1 t8 R3 Cmodel_ft.eval()$ x8 }* I0 X. \) X$ m/ h, X9 r: x

    2 A3 \& V* U; U: T  jif train_on_gpu:
    7 T* J) S( p0 Q2 H3 S$ h+ ^    # 前向传播跑一次会得到output1 C2 h. N; e  J2 g( `# {; p
        output = model_ft(images.cuda())) z) M2 v4 F  X8 |- \9 ]) r
    else:* B8 L1 Z: T) G& ~6 u
        output = model_ft(images)
    2 U6 Q6 `" i: M; z/ ?& C& R5 c  W2 w+ E) m0 J9 L) I; W  t6 Z# J
    # batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值
    ! l: F& i- u* s5 p. C5 H3 @output.shape
    * i8 t- n& @5 _7 a" V; A  Y: n8 X1 p
    1
    . `9 K. D) \. K  Z6 E  a- r3 [9 R2+ z# I. `/ q, L! D9 K
    3
    3 j. H4 ~. G: h$ i2 A, u/ c4
      W- g: B3 j8 {2 M5
    ! L5 y" d7 T0 I4 r1 h6
    3 Q5 e& p2 j+ ]7
    : [# ~" _2 M9 I, r. @* J8, H. B! f. @2 H! _1 X6 R" S
    9$ `$ q6 T/ S* b- b
    109 C2 n! \) }0 ]0 I
    11
    4 o9 y3 E* ^6 O. o& O12: B, x. n& T  w7 U& F
    13
    9 ]9 c. F: x( @- p$ J0 o14
    / e3 b- g5 T4 u15  [% m5 n7 w& }* c" i9 `8 R
    161 B9 e3 M5 j6 m4 q
    torch.Size([8, 102])
    & h& Z# U" Q8 Q+ F1
    * N0 b( E/ M( K4 K0 _9.1 计算得到最大概率! ]% v9 L4 r/ w( }; f5 A
    _, preds_tensor = torch.max(output, 1)* _6 P- o) Q4 y, C  D! u# E) U. F

    8 K) Y: t; u& Wpreds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量
    / l$ T5 d! ^- J% b" m- v8 P% {5 `1 J; }1
    . H; S6 o9 v4 n- H/ A3 P2
    % R0 d% x9 Z; l2 z3! ~, k5 z" F0 G+ U$ x& p1 U
    9.2 展示预测结果! u3 I8 k- |9 P7 m5 V
    fig = plt.figure(figsize = (20, 20))) r3 |% T) {" s- [* X  x
    columns = 4
    7 C. U+ w% }, [. _' prows = 2
    9 v0 }6 _7 ?& A; b) E- h( x) m
    ) T6 f+ F4 E# Qfor idx in range(columns * rows):+ ?! c% R- a3 H+ l3 g/ C
        ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])) q2 F! k* M1 E' R
        plt.imshow(im_convert(images[idx]))
    " f0 s: V3 |9 j1 E0 U    ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), ) u5 Y+ K. m% B0 p
                    color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))) V! q+ P" c& P9 p" ]% W! ?" ]
    plt.show()
    " \5 g& @9 b# o* k9 v# 绿色的表示预测是对的,红色表示预测错了
    - ~$ r4 C0 f) u* r8 W1
    - v7 g$ p# `8 D* t' x  I2+ y" M6 n4 I8 C; E+ p, l
    3( z# T3 S' q0 P% M
    4( t* x1 I- a. {. W8 ?5 e4 b2 y9 ~- s
    5. v! G! [: ^4 y! `/ d
    6; w* C  H; |* ^6 g" f9 ]
    7
    , v+ Y' m* p* W5 s8
    / B+ f" U2 W! N( F- z' G9
    + r( ~( X" K; K) k104 A5 ?9 ]8 @2 w# `
    11# j9 S: |+ N2 U/ ^7 [2 m* x
    / K% K* ~6 I. n& \! D

    ; q" P3 ?' O& `9 T0 B5 z1 N* p5 [  V1 Q- @# G
    ————————————————
    . R3 l( w0 w$ b9 f' M5 V版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。  R5 M7 M% n% s2 v" e& I
    原文链接:https://blog.csdn.net/LeungSr/article/details/126747940$ A. X2 ^( u. ~. |3 o. b& X
    # P2 s( J) a$ s

    4 y: |9 L& |3 z4 R6 p6 k
    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 18:09 , Processed in 0.424322 second(s), 51 queries .

    回顶部