QQ登录

只需要一步,快速开始

 注册地址  找回密码
查看: 2777|回复: 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)实战案例3 V7 [6 {$ d# g% b/ }# F% \
      M  b8 a! Z* _$ I. s6 k5 b& Q0 O1 P) F
    文章目录
    / W" ~) r& @& v9 S* u- s8 f* Z卷积网络实战 对花进行分类
    1 }& T/ H1 e: @0 G, L' h8 z数据预处理部分
    9 q) ]3 x7 ], u. `& }7 ?  U网络模块设置
    ) c6 O5 k/ t/ Z. p3 n, e1 q1 L% U网络模型的保存与测试
    # r8 G4 X+ E- Q/ E) o% W, M; }% A6 X2 I/ z数据下载:
    6 D% K, M8 E4 u$ [) ^1. 导入工具包
    2 x7 S* {' [; K+ L; v& n2. 数据预处理与操作  g2 \3 l! R6 Q: v
    3. 制作好数据源% I- g4 z) O& d# w: ~8 L: g
    读取标签对应的实际名字9 }1 S# T- k4 E* v! H* {
    4.展示一下数据& j6 |3 V0 c) p2 B7 g1 g. g
    5. 加载models提供的模型,并直接用训练好的权重做初始化参数
    : T$ S4 a1 p6 H, z& ~7 O6.初始化模型架构* C0 U0 J3 K* S; I/ `
    7. 设置需要训练的参数! C5 x# m2 b+ G1 U+ m2 n
    7. 训练与预测& f6 M, g4 p: {! o5 ^7 `% ~8 [
    7.1 优化器设置
    3 H$ b1 i7 F6 t7.2 开始训练模型
    # x! x% y; F! U7.3 训练所有层4 a! W9 R- L" _1 T% ?
    开始训练* g: O! x1 b5 ?, ?  ~
    8. 加载已经训练的模型
    ' `! i* x: b4 R9. 推理
    9 o; q% {( }. {' s9.1 计算得到最大概率# }& r1 c! I& v
    9.2 展示预测结果: J. v9 W" R5 I6 \7 N  F% E6 k
    写在最后! A/ r2 f$ K% J! g& K/ V
    卷积网络实战 对花进行分类/ ~- d7 W- }7 O2 ^0 B; v: C
    本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能8 E# M( v) s6 N. j: W

    ) ^$ p0 K$ L: G$ I4 E  E5 s在文件夹中有102种花,我们主要要对这些花进行分类任务
    + L- B. L- @9 @' A) t& L. }文件夹结构1 L: q, H# t8 L+ X7 O
    4 F9 G* X# }1 t0 Y+ `+ c2 B
    flower_data
    4 n5 E+ j) Y" t* [: {5 w
    # H: @( W" e8 U# g% ~( itrain
    7 H) \: W6 y; k! X/ q" `! g* ]0 V8 I+ b/ Q
    1(类别)) B! _* Z' ?  L. ~) _. R3 Q
    2
    : m9 H" q% K: p9 Oxxx.png / xxx.jpg- Z4 l$ ~# V6 O% T: P2 m8 v9 h
    valid
    # P' T% H2 n  q3 u# v, z
      f0 k* v' @+ p# e$ K主要分为以下几个大模块8 q% \' r1 {1 P% V

    6 z. V. \; d* ^+ }; Q! e数据预处理部分
    % s, l2 @: x" Q数据增强
    * v1 r; ?* ^- [$ m1 b: m数据预处理" o& E4 a& r9 B9 Q& A4 f8 v
    网络模块设置
    # c- Y: ?' W2 ^2 h3 ]加载预训练模型,直接调用torchVision的经典网络架构' f. x& I2 K  P$ X
    因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务* T' Z# r& ?8 M
    网络模型的保存与测试
    , v, p; a" U9 C7 F模型保存可以带有选择性
    $ p3 f) i& f6 e0 b7 e数据下载:" N$ S/ |; v' y/ w
    https://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset
    ; h- l1 F1 P/ O# i/ j% b" e
    " z$ f- b) i# H5 j8 w7 Q改一下文件名,然后将它放到同一根目录就可以了3 J8 s6 I; {$ U8 k& ^
    # p0 v' z7 W  D! ~; [! W
    下面是我的数据根目录/ k4 U' k" @! w/ ~. h+ V

    7 h& ]1 }4 A% M2 S, f7 \% f
    & p2 o: w' c) r# c1. 导入工具包
    ; x; I& M' W  n" p$ ~; `import os
    7 i' W! \( a* r. T0 S# Q- |3 Q# iimport matplotlib.pyplot as plt) o! ^# _; q  E$ ^# g
    # 内嵌入绘图简去show的句柄- w( Z9 n# N$ E: v$ V' J$ Q
    %matplotlib inline
    & q0 N; h% g7 K9 L7 n# \/ Iimport numpy as np' i5 D$ \7 u: u0 H" T* b; P& m
    import torch
    ! l0 ^- i. K- D/ j' K$ ifrom torch import nn' E8 W5 E/ L0 O
    7 g9 B% S1 l( p# Q4 o( \
    import torch.optim as optim. h- f( O; |1 c$ p" o
    import torchvision
    $ A( }; o" J" Lfrom torchvision import transforms, models, datasets
    $ ~* a$ P) U" F9 `
    0 q+ x  X( I  c7 F7 @) j0 W: qimport imageio
    ) Q. G4 P" _* X; ]# Jimport time# S* m2 _& I! a# P
    import warnings
    : @/ N2 x4 E  \6 w- ^. z- Zimport random
    , R) a! Z6 c! _import sys
    5 k6 @3 ?& Z8 u/ F, q, Pimport copy; u" P3 O& H, A* u: c! c
    import json3 S6 p: p0 w/ c$ }- y
    from PIL import Image; Q5 |0 @/ L7 |  d
    , z' r* l; P6 C

    # j: ]$ v1 e5 N3 c% f1
    , A7 n" W! M* w25 G7 Y- i, g) {# x' z, r
    3* E4 R8 ?" Q+ o$ |+ a1 }7 c
    4
    5 l% _7 p: O0 D6 a. l1 b6 `5# n9 u$ W  s" F' J; {
    6
    * b, g2 [6 T( M/ X7$ }3 C9 z/ }1 y; \3 O0 P
    8
    ! B) s) f5 \( G7 U' X) Z* u; T9 \96 F  O* A# g+ q- j
    10
    - q- c7 E2 P, b11
    ) o9 v& I6 m9 }. P12
    ' i- W# e2 ]8 @' W1 q1 z139 H0 ^2 }  G  y- K2 M9 b6 j# `% z
    14
    + ?+ T, ?* ~+ `. t7 q% B9 l15
    ' ?. q9 [( U8 G. Z4 \16% J% t! h2 `/ }6 ?5 B' E/ R8 Z
    17( u6 X/ y' L! L& |& g: i( ?
    18
    1 v) K/ r: D0 r) y5 o: K9 W7 d19
    % M5 R+ a9 B0 d) M/ ^( l7 i- V+ D20
    ) m/ J* [* N- i$ J% \) w& O; Z1 @21% E# j/ F5 J" Z, \
    2. 数据预处理与操作0 g' M9 B0 {* _  x0 Y0 e
    #路径设置
    4 j; [- y) R' K6 ?* U" wdata_dir = './flower_data/' # 当前文件夹下的flowerdata目录# W% f6 }% F4 T- P" q
    train_dir = data_dir + '/train'  U4 `! }, p; ]2 h) \5 s2 A
    valid_dir = data_dir + '/valid'" x- X% {) b4 [: x4 X0 T
    1( C8 M% Q( d" _) s$ R( B
    20 C, S/ B+ [' k* }
    3
    & Y" ]! z- \4 P. r/ Y% k4
    ) [7 @. |/ t$ D( }1 Wpython目录点杠的组合与区别' r0 P! z7 a$ o  G% i; H& ]
    注: 里面注明了点杠和斜杠的操作. z0 b0 _. t/ G* e0 K

    5 Y( R( k% b) l/ r$ [. @3. 制作好数据源
    3 x/ G$ |1 k6 Z: a) `3 T) Gdata_transforms中制定了所有图像预处理的操作
    + A/ p0 E9 I) j2 D  vImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片
    + h3 |% C9 b" ldata_transforms = {  X$ c+ d$ p7 A/ \
        # 分成两部分,一部分是训练) \, ?4 Z3 y4 ^# b& Y* G
        'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间; m& f$ E" J6 P% r& |' B
                                     transforms.CenterCrop(224), # 从中心处开始裁剪
    7 H% r$ B7 z. b: O# V+ b% l1 L                                 # 以某个随机的概率决定是否翻转 55开
    ) k- ?6 [+ n" m( X( C                                 transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转
    " t6 s8 A' x* Q% j; F( O                                 transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转
    4 t3 [$ P+ m2 [+ M- b                                 # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相
    ; M. N' j6 b7 |" Y                                 transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),
    ; ]# r& N/ `& A9 G                                 transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB2 M' \1 S& c: v8 m4 f
                                     # 灰度图转换以后也是三个通道,但是只是RGB是一样的3 _# }# _- _  O7 f5 u* j5 X
                                     transforms.ToTensor(),
    : i4 ]% u' \/ ~' D                                 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差' I1 h3 a1 n( Z6 Y, w8 I1 q3 L
                                    ]),
    5 R3 h, H! a4 K( L- ^: m; U/ f    # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化" w! N  E/ m* c7 N" d* b
        'valid': transforms.Compose([transforms.Resize(256),
    5 y  x, e7 B  p& d( g1 y6 U                                 transforms.CenterCrop(224),
    : s6 y% r9 i. h' O. P                                 transforms.ToTensor(),# i2 H* B5 [/ ?1 r" p
                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同
    , j, ~+ L  ~9 B" \5 X5 p                                ]),: Q6 T. I' j/ \8 r
    }
    ( S8 K# E2 t. x" j+ u' z, a. M7 U! Q) }9 [* P. v
    1
    $ v7 q6 x2 i$ x% e# C* I- Z2 F2
    $ i6 Z8 C; z/ y6 _% q, X$ P3
    2 g+ a0 x: ~& ^4 h5 p0 t1 A  |: c  I4
    - ~4 C8 E3 C' f3 r6 M5
      I, Y, v/ k# h) B6 p6
    - z# s& l( J( T+ \3 L5 U+ l7' i2 r/ R  t& Y8 y
    8
    4 }/ s2 k$ G! s0 q9
    " ]) `' Q9 s! u# I, ^# _& V' w10
    / I! \- \' A9 v- T! ?114 P9 O' r4 Q% Q8 _
    123 A- K/ Y  D+ z5 |
    13
    8 X1 c$ e0 i5 n( t( ], k+ ~14) v+ K2 Z/ D5 A  w' U2 N
    15
    5 P8 [. {4 f* R3 C& g. C16
    7 X- \3 j- z% q3 x6 ^2 }# r1 q# ]- v174 N6 \0 V! w5 A  d8 t" h3 @* V
    18
    3 t( f  n- H  G1 A6 [7 w/ r19
    8 a' I4 D/ Z1 i; L20
    0 |$ P' @9 ?, D, Z( w21
    ; X4 I  ]) m" r* }2 a; m3 q# Bbatch_size = 8
    - }. F' L. H7 O3 f: Qimage_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}2 M$ g6 p4 I' }( B
    dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}5 X3 }! F( K$ V$ M
    dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']}
    ( {/ `2 N; e% a$ h, E  Sclass_names = image_datasets['train'].classes* B( }( W  Z3 G3 ?, x

    & \& }- e; Z; ~( V6 ?9 u#查看数据集合% q( R' B% X+ ?! {- e
    image_datasets
    ( J# ?# N& P$ t% u7 V  J
    + E" `" T; a+ E1. B# M& u( v4 J* @$ B! q
    2" \) n5 p+ o5 Z4 m
    3: z1 K+ y5 ~% Q6 @' u: _$ ?
    4: p- V" I3 C% ^1 I1 l
    5
    % ?" H- H& O* [/ F5 z+ ~% P6; e8 m5 `* u% y6 ^0 u
    7; h8 _$ _! }( x
    8
    : ]  a1 p- V7 P1 m- M9
    3 r9 R5 W% m  F8 X{'train': Dataset ImageFolder7 e; h5 E2 B1 E$ l
         Number of datapoints: 6552$ a9 p8 O, K' \! w
         Root location: ./flower_data/train
    5 {8 c: K. C; F  Y4 m: C  s     StandardTransform9 V5 Z( z$ B( h1 H7 [( s
    Transform: Compose(
    4 @% v7 W% G* r0 x" @( f                RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0): o5 l4 J9 E: X( f: O: T% i
                    CenterCrop(size=(224, 224))" j7 x+ H, o+ L; a8 G5 M$ \
                    RandomHorizontalFlip(p=0.5)
    : K( Q; J4 p. c                RandomVerticalFlip(p=0.5)( K; c0 Q. \# ?' v  l
                    ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
    ' \8 o' z  \1 G6 h" X. A                RandomGrayscale(p=0.025)
    + \; t7 _2 L3 m                ToTensor()3 A& i' l3 _! w/ M" ]" S
                    Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
      W' a* o" ^0 F- r" d            ),2 Q8 E2 T+ @6 S! n2 m. V
    'valid': Dataset ImageFolder& d. F; e- `- p- K
         Number of datapoints: 8182 M* @, ^0 |, h6 C
         Root location: ./flower_data/valid- _$ R% ~& X6 D8 M8 b
         StandardTransform
    * j$ v8 O* x5 L5 }. W Transform: Compose(
    ; ^* N0 x- b. a2 H5 |; f, y6 R                Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)( X: b* j2 F5 u* ]
                    CenterCrop(size=(224, 224))
    0 J# M- a0 k$ A+ h, r2 \0 X( C                ToTensor()
    . ^4 f$ y5 U! y# d0 r: a2 M( u                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    % p$ G/ D6 P$ S+ u% m. S            )}2 x# E4 y7 G0 ~- y+ D

    . [; }9 T: y5 ~8 G+ @+ [! T1 [1/ {7 x. j% @* A, \; Q5 c4 h  B
    2# S  z9 `: m: [6 z; i$ P& u/ i- T
    3
    # W) Q/ i/ \4 U& d5 ^! N4
      v5 w$ K# ^7 M  J3 p0 h, o' x5. f" u( y( h% i: D, N: Y1 E" |( t
    6
    ( E; O. \) C6 B7# k- ~6 ^' N9 P  n! c; l! p. S; [
    8; x; `$ T" _1 x" F
    97 v$ I6 G/ e( U- E4 o) O' v
    10) p% l4 o0 n1 Q6 u, B
    11( J8 S7 }$ B6 h6 p$ L
    12
    ' `* G. X4 m; d4 ?: X# M! z& ^13
      T8 k8 A2 ?; B14
    8 }! l7 x/ |6 i2 y5 [15
    * V* X+ p& C5 X$ T/ H16
    ! L* B- u+ X+ c, ?; K% M- c2 a& O17
    2 ?, h& A2 L6 E9 K6 P" y18, J& R7 R1 g% k4 \
    19
    & A$ G. M* T, ^20
    ! k1 ~/ }' O2 ^; N2 s4 @% v, t21$ B: T( [3 F% m8 Y( H# z
    22
    * p# F. P0 R8 C5 `6 }238 e5 m' \. l( |3 Z2 [. j) Y( U
    24
    9 n, B/ \3 }4 B5 D- |6 c# 验证一下数据是否已经被处理完毕
    & h$ z; f: ?& g; X- S& Q& x% I0 vdataloaders- e, I6 s) T/ S8 J5 }0 P/ I
    1! j. i1 ~: F# S  }$ \' Y
    26 m7 ]" ^, w/ D4 P, {  p
    {'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,
    8 u0 w) T0 e% j6 A+ x! L 'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}0 q) L, {2 \% T4 {
    1/ z/ Z  X, c' c4 F; }: {. A
    2& T% b1 G! F. Q9 G5 ]
    dataset_sizes
    4 M) q' t$ q& s) Z1
    0 W4 _) P+ T2 B& S. x{'train': 6552, 'valid': 818}% z4 S5 N. \+ b# ~2 |+ J# |
    1) ~, G- \( J3 D; b
    读取标签对应的实际名字8 ^# {! j2 B( K: h! K
    使用同一目录下的json文件,反向映射出花对应的名字
    ; z8 B( e3 B- v# M  p% ^
    & p7 t+ \0 t5 B. U3 [$ cwith open('./flower_data/cat_to_name.json', 'r') as f:
    . j6 L% b) f; M+ @# C3 ]5 V2 K    cat_to_name = json.load(f)
    4 y8 ^% d, N: C+ n6 S10 M. _! O' e) {. \2 Y
    25 f# U$ R* B' P4 }! d- ~8 @
    cat_to_name1 _$ e: Z( e2 H3 z
    1
    7 T8 Y0 ?6 e( t+ Z3 k3 y! g{'21': 'fire lily',
    , C3 {9 h6 ~9 q, b6 U* h& t, v '3': 'canterbury bells',
    4 n- b7 C! u2 J) }2 f+ } '45': 'bolero deep blue',7 d9 d9 d8 @, @! Y$ q. }
    '1': 'pink primrose',
    ; |  M' G5 ~% H$ Z: F '34': 'mexican aster',
    ; l7 M( a3 s6 K9 h1 y '27': 'prince of wales feathers',
    6 ?, \( l7 c7 {' N9 ^1 y& _ '7': 'moon orchid',& v$ F0 W6 K# `. o9 V! m
    '16': 'globe-flower',
    " U) s/ q0 Q/ w5 z+ T '25': 'grape hyacinth',
    ! p7 l9 k/ Z! N0 i/ @ '26': 'corn poppy',* K0 k! Z$ W) R
    '79': 'toad lily',
    & C& W$ `! [- n1 c '39': 'siam tulip',4 l  I9 w! r1 @" M
    '24': 'red ginger',
    5 `! d. p$ s2 b6 b '67': 'spring crocus',7 u% P: ?8 B8 A
    '35': 'alpine sea holly',
    7 I" c+ z" V$ G. f! _ '32': 'garden phlox',
    7 P( q  m5 x3 A( c6 |5 Y/ ?& i '10': 'globe thistle',
      z& m7 F- E/ r! ]/ n '6': 'tiger lily',6 @8 u. D6 p0 q, z% e; U
    '93': 'ball moss',
    $ k! ^- C* Z% c  h '33': 'love in the mist',
    2 B" ^5 s" d) t '9': 'monkshood',
    - k  y' e9 j2 q' s; l$ T. M5 D1 t '102': 'blackberry lily',
    * y2 F# V' n/ u; T* ~( s, r6 s' r- F '14': 'spear thistle',
    8 Y2 k1 f! R0 m) K/ x8 c! m" y '19': 'balloon flower',+ |( B# Z  ^$ V1 d. S8 ^
    '100': 'blanket flower',
    . y% x9 G" z# }( `5 ^: d '13': 'king protea',
    1 t2 D- V; H7 H4 b6 p7 _: W- u; Z '49': 'oxeye daisy',  u9 d3 U, l/ \0 A, q
    '15': 'yellow iris'," X: L/ X9 Q* q- P$ ~2 m
    '61': 'cautleya spicata',
    . c4 r. X( b$ J% i( D$ W+ m  F& j' Z '31': 'carnation',7 p3 _! c& J2 b0 E; X
    '64': 'silverbush',
    1 r. H9 [7 D4 d8 U; H8 I/ z" |$ S& | '68': 'bearded iris',/ Y3 y+ v6 k8 m6 O) i+ d% S) i
    '63': 'black-eyed susan',3 K# w/ R+ e* b7 D* V
    '69': 'windflower',8 M0 \3 G7 e  `! b
    '62': 'japanese anemone',  j: f+ u7 J1 F/ |
    '20': 'giant white arum lily',
    2 J! R5 B6 }$ g( y* X0 x '38': 'great masterwort',
    2 m' D7 \9 \& d- v) U" [ '4': 'sweet pea',3 P8 l/ U$ N2 z! q
    '86': 'tree mallow',  k2 N. x/ B  |. F" n. C) r
    '101': 'trumpet creeper',
    0 n" m. ?, X- I7 V3 j" X3 P$ w0 V '42': 'daffodil',
    , \8 w! o3 ~+ Q: x: L- c9 O) l '22': 'pincushion flower',
    1 w+ n- d2 r4 O3 A; V '2': 'hard-leaved pocket orchid',
    ) F% l+ c+ f4 Z$ w3 H0 K  T& k# I '54': 'sunflower',8 \. F) m( W' `3 W
    '66': 'osteospermum',
    2 C. H+ n! N, u  r- j '70': 'tree poppy',1 w5 e* Z& d' n9 ~) `
    '85': 'desert-rose',% d+ L, p6 M# N' Y: n8 }4 b
    '99': 'bromelia',- |3 R% f3 d7 Z0 R9 Y/ ]
    '87': 'magnolia',
    * P; g2 b' U& Y# N) H) q: E '5': 'english marigold',, |7 @4 [! p! `* S9 x9 m3 Y
    '92': 'bee balm',9 g, N* j+ }) d7 G( l' _
    '28': 'stemless gentian',
    1 X' b3 _3 ^" e) ^3 C. d '97': 'mallow',/ G5 K$ E: C, J4 \+ o
    '57': 'gaura',+ W1 s. ]0 m1 k% e, S- B
    '40': 'lenten rose',
    ! p0 e. A3 v9 H( A1 z '47': 'marigold',) u( s  I; [- U2 W; l; h
    '59': 'orange dahlia',$ x  ^- ]  i1 Q7 D
    '48': 'buttercup',+ d; V6 J7 f) H! F; F# B+ O/ n
    '55': 'pelargonium',
    5 T  I1 Y; Z0 U& G4 @ '36': 'ruby-lipped cattleya',
    2 O) ?5 f( i# t  C/ z5 M '91': 'hippeastrum',: |; ^! y7 Q2 \- D
    '29': 'artichoke',; g  ^; ~7 V1 }8 }/ T: V
    '71': 'gazania',
    # t0 @  `! ]& P0 ]. p4 H '90': 'canna lily',
    - F3 ]( _) G  ~  R8 N0 Q1 Y) Y% \ '18': 'peruvian lily',' S1 k" j$ A& l3 F8 V6 \+ S
    '98': 'mexican petunia',. P' c/ V% s$ k. L. b/ n- N
    '8': 'bird of paradise',
    1 E8 }) A8 }8 G6 Y( T) ^9 z- i '30': 'sweet william',) @0 h  X# }* c
    '17': 'purple coneflower',
    * `# K6 C% k! M4 [# \4 f# X% c4 p '52': 'wild pansy'," M8 t3 Y+ e8 F( `
    '84': 'columbine',! Y5 t! A3 R$ `1 i$ v
    '12': "colt's foot",
    ! B; y$ l/ X" u0 j* x '11': 'snapdragon',1 J) f! q/ L5 L4 s0 h  i
    '96': 'camellia',
    ! W+ x+ `# T/ ? '23': 'fritillary',; |; Q4 v6 K' d7 |- z
    '50': 'common dandelion',
    % ~% t- P$ m1 q( p% S '44': 'poinsettia',
    0 R6 W* C# ^, q1 |2 Z* B, r' m# d( v '53': 'primula',3 L3 P& ?8 e. r
    '72': 'azalea',
    # U# i' W8 F; ^' B+ C$ r '65': 'californian poppy',
    3 S/ d$ {5 a6 b '80': 'anthurium',; N" w+ k) }# I$ _, F+ r! A
    '76': 'morning glory',* f( d  B5 K( g
    '37': 'cape flower',
    , L% w% e" ^, |1 p# { '56': 'bishop of llandaff',
    " s- X: L) Z# n5 A '60': 'pink-yellow dahlia',
    2 r3 t) Y( l# S# {! A) o '82': 'clematis',3 V: ^. B! {  |2 W) k0 j" ?" W
    '58': 'geranium',+ \; s; X" p' x2 E! H
    '75': 'thorn apple',
    % I3 l" e2 y8 P* o '41': 'barbeton daisy',% P4 m* P6 v1 `
    '95': 'bougainvillea',/ K' j% I0 B4 V
    '43': 'sword lily',
    ' p$ |2 |7 n6 q7 v# y( B7 y '83': 'hibiscus',
    5 b. \! i/ [- \  I8 G2 t '78': 'lotus lotus',
    " |% |: Z) N( o/ m: E5 ^! } '88': 'cyclamen',
    / c2 s) t/ K) Y+ B3 g4 j' s '94': 'foxglove',! g8 o% f" l( n7 N
    '81': 'frangipani',  M( W( z  a; O# P" y& |
    '74': 'rose',, g8 N" @/ I. S  }" i1 H5 {
    '89': 'watercress',: ^  g, p4 E; W4 d8 q, G% `' h
    '73': 'water lily',
    9 m. f  F- H. J5 i, o- B' r '46': 'wallflower',4 R- J. F2 x- c4 D) N$ b$ y& E' C
    '77': 'passion flower',
    : V. B9 F& P8 D5 q/ x '51': 'petunia'}, F' ?5 p& E9 D1 O
    ' z- K5 l) s( m0 C
    1% m# M. R* E' u/ x! Y; l; q6 K
    21 `8 c1 @0 h* {
    3
    1 Z, ?" K! W/ x" J8 x4
    & i) w- S! z3 R5( V. V/ j: t; |, Y' W" k
    6
    % Z! M0 k3 _( L/ m6 O1 \7$ I0 W4 g2 R( a6 m; J3 P
    8
    & p% t. u8 l0 b8 `; b8 p6 }9
    3 u1 ^8 t' ?3 Z! P2 c2 K100 _1 B9 ~5 e9 n  p, m
    11  \0 m" i) u8 I& z- }
    12+ c; D/ J8 I, p8 k: J4 z" ^  p
    13: R4 v! W0 x; \  L7 D
    141 y/ D+ ]$ Z$ J) j( l
    15/ F5 A" h3 C) v3 i3 _: \
    166 I/ r$ r  G- P/ l; y% k
    176 w4 K5 w$ a4 s; T! {3 f
    18: Z# a( O% t# p; f6 X& i, [0 m
    19
    ' Q/ a* C+ a0 u, M- M, w' J% r20
    & p6 t& a' L$ c% {, ?' H8 R4 Z214 f5 P# T# M5 v
    22
    1 J1 s" F* t) f) B, J* V. G23+ U3 Q  B1 ?$ U) z% q
    24
    $ Y. C* e. \+ o' X25
    : R- L. X1 {3 f- s' v26& d6 m# ]' K' B' I
    276 _' T/ y! P" ?9 \; `$ i7 \( X
    28
    , o- [& k* t& [3 F" y29& `# W3 `; S2 X
    30
    2 f2 O- F% b6 A3 J) I% x  |, H311 n$ X# [; o" l/ G! b7 m
    320 |5 N2 ?+ E* h5 T, |' I
    33, U) Y8 ?: t7 L. [& Q. a
    34
    ' q- F) {  `5 p! D' H# u; {* k' [* `$ N35
    5 y7 g7 l" o* ^  Q4 B36
    ) A( v$ U7 D  f5 h6 p37+ o! b- n; Q' _5 Y4 U; O
    38
    - A  Q9 a4 ^% O, ]  A39
    : @& z+ J  ?/ s5 |. p7 y5 y40
    ! j, L) J0 A# ?0 i7 s7 t41
    3 U8 f6 y0 U- L& u8 b42
    5 I( E( P/ Z; R+ a+ e) x! j/ W435 M4 d0 h+ u2 ]" O" A5 v7 v
    44# T6 y. X2 p5 p* ^) }0 X( l
    45
    2 e  ~3 l0 i" n. k# w5 O! S/ \46- W! p, j# m* P2 d
    47# x; G. X0 a; S! p5 P# `3 B
    48
    * u; z5 ^4 ]( C, K8 P" k) Z% T  a495 F) e# F$ B4 N  S0 C6 Q% @1 K
    50
    . [5 Z- Y1 _8 D511 F% X% Z( d7 E2 t' G) `
    520 i, T6 {2 K8 j- t- u
    53
    3 h$ @, O; y, x  Q: k; N546 J% g/ S+ W1 i9 T3 E1 c0 [1 Q
    55
    * `' m1 `; s3 m% [% q6 @56; n& n& O, X2 V' o0 ~1 z/ b) P
    57
    7 ^$ V: i+ J3 p% V  o4 T58
    5 a: o7 Z2 O6 w" ~! s1 O% p59
    ; s$ v9 d  z+ e60$ i/ e2 V2 u. f0 H1 m
    615 _' q! \7 z0 I" F* G& u; M: ]
    62
    - x$ R. ?* D! a" p. @63) `! v! Z( W/ Q# a6 r) R7 z+ F% D
    64
    4 L/ y6 G' j/ x5 R- b65' A6 k1 a( l- ]; M( n% _9 ?
    66
    - x, \( a7 T9 d4 d8 {67
    1 Q, _* I3 i- t& [) S7 F# ~68) ]; R3 v  l( N# p
    69
    9 U5 _* M. W9 k  M1 [) Y709 V% V' q% W: H- K* M- C2 |. f% H
    71
    & V: Y; F# o9 l) P6 _4 E& _# _72
    6 W2 p' I+ X) u# T* L* n732 {" m/ k7 Q. u7 J. ?& g+ p# }& ]: h4 ^
    74
    . }: Z* T3 {+ E+ N6 Q* a( E75
    # _; \* G+ B8 T# E9 v: Z76" B/ F' e( `* F# N5 h9 R$ o
    77
    * p/ n6 E8 z* f' V78
    0 K( s' f' ^/ t79. D; }" O! K6 c- I% H
    80
    # D  U9 s1 c- T) J! @( X! `; D81( O/ O, M7 A, X2 P7 K; f
    82/ r- e, e; u) r: k
    837 [- ?; w/ r9 K6 u( V8 }) K0 Q3 D( {; W
    84: q+ \1 @& k2 H- g) J1 n
    85: H8 b* |( x- l# d
    86! m6 |% k4 `! {5 J, w9 J+ ]- ~) @
    87
    - p2 b' M8 H3 w6 s  M* f88- ^' d' K" |0 Q# b0 Z0 k! q
    89
    8 b7 y* t# s, s8 m0 }8 Q4 F' u& d90
    $ E% E+ M( L! i" H, p  ^4 N91( F9 g2 a4 C6 D  ]3 Z0 x
    92$ K+ r5 K6 ?" t
    933 x3 B' q" V* }/ q
    94' S% F+ B) j  h. L) V7 x
    95
    1 ^: c- ~1 ], j- m! s96& e$ n: `6 b4 c4 t
    97
    & Z/ {2 A6 `. r( ~9 a98
    5 B4 a% r) q3 [5 |% A) h99
    + J1 O4 M' ]0 p100
    / W2 X9 e7 b2 g0 ~8 J8 |, s101# R6 X! X8 @* O7 n+ }, j& F9 |# Q
    102# M+ ^: e) k  J% T/ }# g0 V
    4.展示一下数据0 k) |* N2 V- X: F: \  X
    def im_convert(tensor):: ~& u% S- w$ T0 f2 q
        """数据展示""", x, M4 l1 v8 v
        image = tensor.to("cpu").clone().detach(): r5 U  G. Y5 K
        image = image.numpy().squeeze()) d2 r. Q; W$ ^
        # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图- T& {* Z  f! n
        # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)
    ) e( t/ Q; S; C3 `) T    image = image.transpose(1, 2, 0)3 a& l  ]; a. k7 g- K( J
        # 反正则化(反标准化)5 U! }% o! n# [6 ?% {
        image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))
    ( L+ Y! b  A) H  y& u- H
    . D9 x$ `. s! h, a( [/ [8 f, T    # 将图像中小于0 的都换成0,大于的都变成1
    8 |0 K9 z) z) T/ Y    image = image.clip(0, 1)5 T2 m. a5 u7 [3 O; b1 W& X
    8 g  G+ F: c# j( O' k
        return image
    4 A  \; \* N# R5 g1
    8 H7 K* X+ J4 z* Z( y2
    . g- V* n' G" x( o$ J. f3
    6 }1 X0 z" a6 Y9 ]0 y) E) B* c3 C4
    1 h1 N' i: K1 K1 B  }1 r# H1 F: t5
    4 }, W0 e9 i( @$ Z4 p6
    & ]/ T% N7 G8 L2 |$ K) E/ z1 M! y) z74 [* O- Q4 Q7 ^- N# {. r
    8
    ! g/ {6 G, S1 H2 J: p9) P% G' N4 t+ A  Y# a3 p( b
    10
    , Z) P# L$ i( f2 q- P11
    % B3 Z% y; x0 D4 o: R2 ^12
    # z' u8 f- x/ F13
    ) V, N7 d3 [% m# ^8 l) m# `14. [3 u. x$ ?+ o% c8 C2 \/ p( s6 F
    # 使用上面定义好的类进行画图3 `/ o; R. N$ E, |
    fig = plt.figure(figsize = (20, 12))
    4 \7 J6 g$ D2 {2 E/ Xcolumns = 4
    ) m6 J% Y4 t$ ^  J8 prows = 2
    2 E0 I9 n" e! f
    1 S+ l7 [$ Z2 ]. |# ^# iter迭代器. ^4 z2 D" m  C! R, U
    # 随便找一个Batch数据进行展示7 _% j: Z$ f, S* n" ~2 L, n
    dataiter = iter(dataloaders['valid'])2 g6 Q3 h9 B! v9 t+ e
    inputs, classes = dataiter.next()5 X- o9 t, J; ?7 C

    ; E) K( X! M. B( W; O, T" Rfor idx in range(columns * rows):
    ' f8 r  J% e2 K3 Q7 o" i    ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])! o$ }: V  |3 G# v0 ?
        # 利用json文件将其对应花的类型打印在图片中
    6 u. w3 y0 K/ \# j7 J3 _    ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])
    , t7 t. t/ K" S$ N; i% U    plt.imshow(im_convert(inputs[idx]))
    # P+ a6 i* t6 n5 [( v& I  y; f0 z; Lplt.show()
    - ?$ F- \$ O/ X6 \% A( m
    & ?* V- f/ m5 B/ R, }( B1
    9 J3 _& \4 k' `6 h$ m& i2" f% R  E" @0 N; U
    3" i; \+ k- }8 l5 q6 ]) l! J
    4
    / E- v/ r. z( _, H" Y* K; @, ?5
    $ a( \. z" }& G3 N, w$ {6
    & s& c! w  z5 s, ]  V( [9 i7
    + u1 D& r3 k" W* R' o0 z8
    - C& ?8 {" Y- B, A. x2 c9$ G+ W* _7 o8 V  S0 I4 z2 r% F
    103 ~1 V" ?) D* F) Y. U
    11
    7 B) ^- U1 }3 `. @+ v9 H& k12
    ; l6 {' {9 }( v7 j# c- D! z13
    - T1 M8 U7 J+ A' `14. {# a. O+ ~$ M% }, n% {! k1 ~* s$ E$ ~
    15
    ( C+ r2 Q/ b7 Z' M2 b; |16# p! ]9 E2 y7 w7 V

    " `  x$ w2 H1 K' C: I
    $ N3 A8 A2 s7 E0 V+ s* N6 |5. 加载models提供的模型,并直接用训练好的权重做初始化参数
    ) C- m, ?2 ?6 D5 T7 Y+ t2 Imodel_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
    2 g; |( l6 m0 y/ q9 H# 主要的图像识别用resnet来做( h  `7 E, r& d- [! e4 w4 q
    # 是否用人家训练好的特征
    . }% N; h" D7 efeature_extract = True$ K) G/ p# H' q/ v0 t# b% Q- v
    1
    : s! {# q+ g! y9 o26 D, R5 r, ?% A: K
    3
    & B- c; j" P+ r* p, W' z4
    . z% r5 E/ x2 l7 M# 是否用GPU进行训练
    0 E0 A2 |8 i0 btrain_on_gpu = torch.cuda.is_available()- _% C8 V1 G: ?. r7 r$ f

    / j9 P) {0 n8 ]* h" Fif not train_on_gpu:
    + ]% ~4 C: G/ I7 m    print('CUDA is not available.   Training on CPU ...')
    + r7 ~2 g8 T5 E* q' felse:3 c8 T+ f& i6 `
        print('CUDA is available! Training on GPU ...')( {9 A( g8 g& b$ R& n- u7 `, f$ d  ^# s

    . a  F7 Q% a" S* v2 ]device = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu'); C. e9 o- v# w( ~5 s
    1
    " ?3 F- z0 w2 ^& E& p20 Z, R7 }; C1 y6 {3 S2 E
    3
    : U" ]8 k7 v$ u) L7 g, M3 s  F7 P4& I8 Y# v7 r9 k2 C& ^* k  K
    5, p+ d: M8 r9 m& M& w9 d4 }0 Z
    6
      p8 v8 e; _; U9 {* ~/ J7
    - n4 \6 `2 O% J; B: k) V87 T+ S* V& _6 V  l8 r
    9
    $ \4 l" z  b4 e$ k4 c. I2 f/ H1 ~CUDA is not available.   Training on CPU ...& c% m0 D0 ?* v9 ~
    1, i) e! D: j4 E, r+ e. j: `
    # 将一些层定义为false,使其不自动更新. S3 T5 \1 C0 @
    def set_parameter_requires_grad(model, feature_extracting):5 T+ Q8 N; d# L3 q
        if feature_extracting:
    8 v# @: V6 o  n6 b  I- s1 R        for param in model.parameters():2 Y8 a) M3 P) v) D& Z
                param.requires_grad = False6 `/ X7 W4 K% O/ k8 w  t3 ?) ^
    1
    9 `2 V0 t0 S9 z" f0 Q" G2
    4 z# s# p3 V5 c$ ]0 Y3 f9 M4 i3
      l( L/ W- S" F; \46 `: ^# ?# U; d
    5" @# O* E! p8 v5 Y
    # 打印模型架构告知是怎么一步一步去完成的
    6 r3 E6 c2 ]+ _( Q+ ^# 主要是为我们提取特征的" y, G6 R3 P3 O4 {) p  k

    % \2 t% |, S4 r( W+ u/ ^* xmodel_ft = models.resnet152()
      p/ ]1 N1 S- u! Mmodel_ft6 X  D- G3 C2 ~1 l
    1
    5 _/ w' f+ T8 R3 p) c7 v/ {2; D- @* h" @: x1 G0 A
    3' k* l' W) Z; n% S  _! G1 ?" y9 j
    4
    , D! z+ y$ [: {6 I4 p" F( O5
    8 g3 D1 ?: L2 ^& z4 ?+ F5 _* X" }ResNet(& i* `$ ~4 }  Z! {$ N
      (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
    5 d3 L# U+ b( y, z' ^0 @! z3 R6 l  (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    0 I7 i0 [1 {0 l6 g9 _6 Q. y# i  (relu): ReLU(inplace=True); b8 z' K. S/ S8 u
      (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
    & Q/ {* r% ^+ \0 i  v  (layer1): Sequential(
    2 W6 ~$ O  g, `8 n' b    (0): Bottleneck(
    ' A9 d7 k4 a! f" l! H      (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)
      ]* k0 _3 q" w$ R      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    , l+ V" W1 d, C& v0 D# ?      (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)" S' p6 q; s' q! u& c: Q
          (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)+ t  p5 M/ \  ^/ n: `
          (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)  K; [# I# u! t. d& }3 [
          (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True); P# l/ X  @, F- V0 o
          (relu): ReLU(inplace=True)
    7 ]5 M# O. j9 @      (downsample): Sequential(
    ! w3 o. [! p# e. L6 ]6 @7 A# r        (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
    " u* ~# h/ G/ U0 ^/ E        (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    ; ~5 J/ J& q6 S! m      )
    ; ~" C; H; u4 x- }2 v( e    )$ E: X- Y6 Z1 _0 H7 V
    中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。
    & C) c" ?1 ~+ P( |% t    (2): Bottleneck(4 T' {1 k7 r9 y) e0 M2 L3 F9 s
          (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)# U1 N) \) I. k
          (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    5 K' l/ q3 c$ `# Y      (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
    5 d+ A/ V! k# v8 J% O, r2 ^      (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      C1 k) ~1 |' ?; D      (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
    - }* O1 n! H+ l' ]# ^1 Z$ _2 E+ D      (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    0 k8 T$ _1 F' u9 L. y      (relu): ReLU(inplace=True)' Q# H: i/ p0 T  {
        )1 I2 a, Z: f3 ?: p6 F
      ): b9 `+ T; ~9 u: a8 |
      (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))
      U2 m' m" m! s  O3 ^7 \7 |/ d% \  (fc): Linear(in_features=2048, out_features=1000, bias=True)
    % V% ~9 ^# i3 C$ h" d4 j  B)
      ~, L% O4 O) \0 U3 @( b& G. G' L8 f! Q+ m
    1
    ! x# ~, Z8 ~: z: M6 Q3 @& i8 l" a" Z% y: `27 A9 N+ j! g' Y) [/ d' i
    3
    2 i* s# y% I( k3 e# ^% Z3 K4
    6 a  T9 F" s. _6 h5 \5
    ! Q9 G0 e* B+ L/ A2 s2 ~  G6 \7 U  q6+ Y# x4 {/ E1 @1 S" Y2 t) T$ G
    72 g( X2 h+ z. {) v9 _
    8& J/ l  y8 K! S- K( i+ {6 I; p' q. @  J
    9
    ' v- K" M6 v/ v; m# R104 y  O; |9 ~. _. ]4 x, N$ f+ }  z
    11( k/ [5 w* h! b9 c" v4 [
    12! T. \9 X& Z4 ]! e2 _8 N* V
    13
    * Y9 @) V" Y4 t. @% s: l: e14) j4 \, R* |* }) V4 ?
    15
    4 ?& p2 n/ j3 c! t( W' s16# C) Y& _  L+ M" k2 ]7 @2 S3 L
    17
    ; V% h3 K7 c: }" \% C6 g# e. B189 f: I* Z5 T4 h( Q/ H& P4 e, J
    19
    $ D# _' G3 t/ ~, c9 v: c2 @206 k3 |/ v+ B, f$ U0 l
    210 q( F: w* \. U. K
    22! a# X: I# r: ~# D
    23
    1 W7 o4 ?, X( z3 H: F: x24
    4 m9 E$ d! N8 `/ V4 q25
    4 C. _. ?% R2 Y$ k" o26* i5 }  x# Y0 Q/ R
    27
    " z) j2 y% _! M( Z28" ~, t& T) j  Y- Q# {
    29
    - u* o$ S7 j  `' g5 \30. j% L% X6 `  w# I/ u
    31
    3 W: R& N( c* w7 r% y/ |32: F/ c2 }$ H: P' [
    33" f* F: A6 G( i2 |" s
    最后是1000分类,2048输入,分为1000个分类
    7 A5 z. m/ {6 c' w9 E4 i而我们需要将我们的任务进行调整,将1000分类改为102输出: p+ ~6 ~5 ]1 f" l# R- ]2 [

    6 A% \, l# r$ m% m6.初始化模型架构7 F$ t  Z7 r8 ]& ^" Q
    步骤如下:
    6 S! O) n) }% j. u; }. g( F* a
    3 l& O) L% A& a+ N# c( {7 D将训练好的模型拿过来,并pre_train = True 得到他人的权重参数  C% i% \% T( i9 w2 ]
    可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)$ a9 }3 a9 s/ j
    无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数
    0 g( G# o# q; z' L+ {$ T/ ?官方文档链接1 R  ]# Q# X3 U1 j% e0 w. M% G0 ]+ j
    https://pytorch.org/vision/stable/models.html
    # E- X% R3 P/ X' H+ a8 w7 Q; I" R( K8 m. a: ~5 A
    # 将他人的模型加载进来
    ) K4 X- ~1 r; l; ^2 a! ndef initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):
    - s: q; [2 {% [7 x( S% I& m    # 选择适合的模型,不同的模型初始化参数不同
    0 q. y8 T+ a3 Z6 z1 h, Z7 K    model_ft = None
      T, f+ z, k! |3 `) b) F: F    input_size = 0
    . }* Z5 u  `0 }6 j
    5 \# y: K' x. u. ?2 w    if model_name == "resnet":
      F; y* K  Z  |% a! f% i9 \        """
    # W1 H* o9 ]0 F* f# q- ]        Resnet152. g; c. s- ?0 J! D# l0 P, n
            """
    % m$ f: ^. Z! j% Q% }' @5 f' v5 T; u; ~
            # 1. 加载与训练网络& G+ B$ s# L" q8 a8 D
            model_ft = models.resnet152(pretrained = use_pretrained)
    # y/ B. s# A' R6 J9 n% A        # 2. 是否将提取特征的模块冻住,只训练FC层' g7 O  o: K3 y7 f
            set_parameter_requires_grad(model_ft, feature_extract)
    # R2 ]9 V: C9 \1 m/ G        # 3. 获得全连接层输入特征
    ; \; O1 z( I" H; U( }! \  z        num_frts = model_ft.fc.in_features  V. X9 _  }/ j. g
            # 4. 重新加载全连接层,设置输出102% G' M8 x7 n# S8 _8 L  s; W
            model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),$ m% a0 C' P( o  p+ u2 f2 ]
                                       nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1
    ' I. e- b( G! \8 [5 f        input_size = 224
    4 ^6 L0 M9 O, v- m. x6 @+ i- D+ m9 W- `, C  _1 }. m5 w6 ]
        elif model_name == "alexnet":: W5 c2 j, q  c
            """4 J3 s( c9 h& a8 t% A) i
            Alexnet
    1 ]7 D+ d& [: Z7 s- N( c! _        """) p& j  [3 F; A+ {$ m
            model_ft = models.alexnet(pretrained = use_pretrained)% E  R; c4 g' \. W  {
            set_parameter_requires_grad(model_ft, feature_extract), ?5 {0 f; q# z! k4 |1 D& q

    + S, m2 G6 F4 N5 g1 y        # 将最后一个特征输出替换 序号为【6】的分类器
    + B2 {  V6 q/ q2 s4 K3 z0 E        num_frts = model_ft.classifier[6].in_features # 获得FC层输入
    % @; {; _) }3 t, n7 p% K" i        model_ft.classifier[6] = nn.Linear(num_frts, num_classes)! o7 W& z2 _& a" \+ z
            input_size = 2245 P  L- T/ H5 h- ?% v

      C. Q7 h, U: X: I    elif model_name == "vgg":
    ' Q% c) P- o; \: J' C* }        """
    ! ^6 s# W- R/ P+ h$ d: z        VGG11_bn0 E" |( n" S9 [( U+ Y& E
            """# q3 \, I) \7 r) S2 _
            model_ft = models.vgg16(pretrained = use_pretrained)0 ?7 H+ g& U' r$ o1 j1 ]
            set_parameter_requires_grad(model_ft, feature_extract)/ z5 U/ _2 Q9 }: U8 C- c
            num_frts = model_ft.classifier[6].in_features1 n( L: d& ^) a$ L# \
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
    $ J7 t: [* c, z& p$ a        input_size = 224
    # h0 B! _$ W( L/ z0 L; a* q5 m& h! T7 E
        elif model_name == "squeezenet":0 s) i0 r& A' V5 h6 h% l
            """
    5 E% Q8 B4 O" G' _        Squeezenet- N1 Y3 O( P% G* j: S
            """
    # Z- |+ z4 a6 k* s7 O/ M        model_ft = models.squeezenet1_0(pretrained = use_pretrained)
    & }7 B# o$ c# N( q- w" o: t        set_parameter_requires_grad(model_ft, feature_extract)
    ! w; D$ c- m8 f4 I        model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1)): n) V, X' L0 O4 @
            model_ft.num_classes = num_classes3 \: T( g+ B$ f6 f. m
            input_size = 224/ V( p/ r8 Z& u# P& ]3 o
    ! ]8 F' n1 P( @
        elif model_name == "densenet":8 E! o2 ?- R' B1 g8 O& d
            """
    / }8 U4 O9 K; y* x# O: I1 M        Densenet
    . F/ {3 u. @; R% A        """; Z: ]" R$ G& U( i
            model_ft = models.desenet121(pretrained = use_pretrained)
    + |2 f% U0 B& }$ E' A8 o        set_parameter_requires_grad(model_ft, feature_extract)
    0 j7 Z3 C& N3 a1 U8 k7 N0 w        num_frts = model_ft.classifier.in_features
    6 [# F8 o4 l! K/ H1 `        model_ft.classifier = nn.Linear(num_frts, num_classes)7 H7 R, c6 U4 m# y" I8 e( Z
            input_size = 224" p5 \2 V: @1 H! X* l( C: L

    8 O9 t. E1 A* K' ~1 z% b2 K    elif model_name == "inception":! i7 y2 C1 R" `  t- d
            """
    9 U  Y% U/ B" H$ Y5 G" o% `        Inception V39 ?% a2 m2 _8 J  k! Y3 t
            """& b) G7 e+ I* A, o
            model_ft = models.inception_V(pretrained = use_pretrained), v0 _$ w* I- M  l
            set_parameter_requires_grad(model_ft, feature_extract)/ y7 r* f. w+ U. W, O+ [
    4 t7 q& o) p  l, \- c$ x$ Z, t
            num_frts = model_ft.AuxLogits.fc.in_features
    ! Z, C0 v  t1 D& Z0 J6 e3 [" r- M0 P        model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)7 _) X  a  y( A5 {
    " C+ x- L/ G, x
            num_frts = model_ft.fc.in_features+ d4 ~  P# F  A- n" p+ n
            model_ft.fc = nn.Linear(num_frts, num_classes), h0 u$ i& W$ y
            input_size = 299
    ; v% j' R- \" U0 j. }" z
    ; {( |) B# W6 \/ k6 ]    else:3 L$ H4 w: m" O' A9 s. T) m
            print("Invalid model name, exiting...")
      s6 z: K( Z. k2 n. ~: [1 M) h        exit()
    7 G/ `4 Z3 z0 U/ I  a1 @6 K! t; ]
        return model_ft, input_size1 s, h% k9 o% l! t8 L+ J' ?
    + k7 c# c% r9 T) @, x$ c& s% L
    1
    ! D& Z7 H3 K  F7 P2. }* a" @! L! @' w) Q, O, t6 _) Q+ Q
    3
    # }" X9 z4 I  P4 ?1 S9 `; ~/ X8 k4
    , f7 {5 C1 F. y+ R5
    7 w# r- v) Y: f9 b4 I4 K- {( H6  `) {" H4 ^* {$ }
    7
    & ?0 U3 s. v  h9 K( ~  [8 B8, @6 |1 h2 i+ c
    9
    + Y+ L, a* y) L0 G0 [$ }! _* r106 i5 v1 M8 L! l  ]& C. H  {) L
    11
    + N! P4 g/ K  L& O1 O12
    9 e: [7 b5 x* Q% z134 c6 c+ j9 v! l
    14
    1 K4 I1 t' c, ~* U+ g% f  \/ G155 |5 v! S8 x9 A
    16
    ) z/ v, f* F# I1 h$ B) T( o# A# Q17
    ) u3 K; L4 o- P* D  X% Q3 _5 G18' o" j( X0 e  L0 s
    192 M- U- {1 S- Z9 J8 P
    209 J0 N8 _( q! M; h7 c+ n) _) u# O
    21
    ) X# D5 K7 H9 |# L$ F- s22# l; b- B- @+ Z+ `" j4 i
    23! d8 Z" ?$ O0 y8 N; T4 U- n
    24
    6 u$ H1 h7 D: G" A4 C; E25
    # z0 Y- G8 ^. D, }  \  Z# ?269 {8 Y( S/ B/ M4 p4 p& p
    27+ k0 V% l$ F! t/ B: Y; a) y8 G
    286 n' o+ d9 v5 }: n: Y5 K
    29% }  ?4 G  S0 I9 v, i/ G& S2 F) A
    309 S% P# s  K1 \% G$ s
    319 H- g4 |! b$ x2 H& d/ D
    32
    1 ]2 C8 ?5 ]6 E' w& S3 \, d7 O33
    % R( t" [2 q3 X; ~9 n34
    ! Q( @$ {' T' u' B/ w35
    ) I: a+ @$ {- q$ T7 `/ E1 g( e" P% A36
    ; u7 r7 Z& ?0 ?+ E0 G. w) a7 \  O37
    + V7 t9 v% u# ?3 A, _; m38: U( S& }' r4 ?; A! R
    39
    + r* i" U% O* d9 Z40
    3 V: a7 e0 p7 d3 V, s' u! m41; b: S# |2 Q. v0 a
    42" A+ d+ \, }: g' r
    43
    8 U/ ^7 v" U: o  n1 `6 {  ~44. w7 j+ g8 Z6 ^# E8 S$ E
    45
    * O" P0 p3 f; s1 [3 H46# H) ~5 ^. N, ]7 d3 N. }: ^( L
    47
    3 P! s, O0 j2 K* ]$ V, o+ D) U. q48
    ' U. a1 w5 Y. z" j$ L9 n9 J: t1 J494 {; q) T; u- L. F
    50
    , z* c( }  J4 s51
    & n- M- Z: |( c) x. E, g/ @$ ]0 D52
    " F- A" L: ?+ d  o2 Y; h/ y3 H% b53+ e" A2 X0 }, X: E8 L0 h" ~2 [9 l
    54
    ( E' w9 J& @% g% ~4 d$ e( |1 J5 B55# h! [: h6 k# m+ `( B
    56
    * P: H# j* O* i8 _57
    : F% v& m5 B: d* C# `# q8 t58
    / N- v& u/ Q6 k. d% _; q& h+ t3 M( {592 B$ y; i  @' P2 m
    60
    + f) j# x9 Z! n9 R4 s61
    1 `  }/ O, n" N* a3 j) a( E# h62: U6 v# }7 \( ?0 V! k* J
    63% T3 R( Z1 O! H
    64; I# ]& j  W$ [
    657 n0 M( U- b1 N7 X  n* ], N) V
    66
    4 D7 o' O+ g1 A1 g0 i67
    % C1 b9 S& {9 r& g68. ]1 ^2 \* h6 B/ Y
    69
    - U, G! f3 y: P" z# f70( K: }/ r& p1 }2 W
    71( ~! p! I+ Y! y7 M% K3 O) H
    72; P/ @! `3 O. a  }
    73: u9 C) u: {/ r8 e
    74- z" p: w8 c# ~& w$ L
    75
    8 r, e# ?9 ]0 @0 l) e76
    & s+ ]2 a$ ~/ J4 E8 j6 b77
    9 _, N: z, r' Y4 f788 v& |, B" ^/ V: |/ `
    79- q9 @" @  z( f' a
    80' G4 s# U, n2 o) [& H/ n
    81
    # d: u, j' v6 [, ^. b  i82& u3 m, H/ W; q2 W& b0 F6 \/ t
    83: |# C  L$ f  D
    7. 设置需要训练的参数3 j, T- d, ~4 V9 y0 ~. d- `9 E' j
    # 设置模型名字、输出分类数
    + y" u& B6 y$ l& ]' l) C/ |model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)
    & i- ^" C. F+ ~$ z! [( I0 H# u2 ]+ r- ^- m1 c  }
    # GPU 计算
    % x* `* m0 w: R  p6 K6 T8 ?7 `model_ft = model_ft.to(device)0 v( m9 R1 g& }5 t5 v5 k

    2 ]4 b, K* Z" u$ {, e/ P# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取8 `9 a+ i' i6 {% k) E. ?" `" ~
    filename = 'checkpoint.pth'5 b# J8 K' r$ G# [: d

    $ A7 n$ u0 Q# O; a( [# 是否训练所有层7 N4 i9 z, l# U8 {7 e
    params_to_update = model_ft.parameters()! a& `, u* B3 G
    # 打印出需要训练的层' @3 b" Q, Q8 X" k5 G
    print("Params to learn:")
    0 F+ F7 R+ o+ g' P/ K- aif feature_extract:7 S$ o* A! Y2 a( d2 F3 p1 j# p
        params_to_update = []/ Q/ N' Q7 \' k* I5 g( I
        for name, param in model_ft.named_parameters():( c1 U7 c; Z2 ]! z0 K2 |
            if param.requires_grad == True:
    : n0 e3 _7 {/ d; g            params_to_update.append(param)2 l' B$ j4 z5 w0 e
                print("\t", name)5 Y' W% x: B7 i5 n; I
    else:" I0 j/ L  U, P: w
        for name, param in model_ft.named_parameters():
    / L) _: \" _9 n- v3 L& d. X        if param.requires_grad ==True:  @2 S6 O2 D) v
                print("\t", name)
    ( C6 [$ G" f' E# E. I2 p/ O5 J# J& v3 B
    1  x$ `1 R1 v2 j& A6 }
    2
    0 l$ ^) v* S8 L7 }7 m7 m3
    * ^/ X7 k( v: e" ]3 ]& o4+ ~) M8 x% d; G
    5
    1 m+ M% p. r4 O- Q/ V$ A7 q6& ^5 n  X. e7 `2 ?. X
    7# [- a) F1 \; ]4 O$ q( B8 O! |( g  k
    8
    $ l) U+ q  c$ m- x$ ?8 ?9% B* q' n7 ~+ Q7 b
    10
    : R5 P( t; m, h8 Z5 l' f3 e11
    & b8 q$ U1 b4 r! R12
    ) M. L/ N% b  E; W13
    + G1 [0 X' o; k& }. [5 q14
    & z9 ~5 \. j- W1 h15
    3 K" m" c* N# d0 k16
    ! [, x* j' K4 o% D' O' H17
    0 m4 O' N* Y' N# f18& j) ]4 ]2 i& s- r% q: N) H9 k: S3 h
    19
    ( I4 l, Z8 Z) t  Z, a, d( l$ r6 O20" b% w( \9 S+ ?5 U( o; x
    21
    & Z: L- [7 G# Z22
    7 w. ~! ~# v  r& I: v4 @0 n* Z+ }6 F23
    ; g( M7 [8 r$ t7 DParams to learn:
    4 Q: D/ G) j8 m- R         fc.0.weight' X$ `5 C4 u+ b5 Y
             fc.0.bias2 t5 X/ y, h, T1 u, W0 F; ?5 m* [4 M5 H3 R
    16 p+ Z1 V7 z5 E( k" x4 p1 l: r
    23 X7 L* {( i' y. f+ f' V5 D
    3
    2 ^! n8 P; r$ d7 j. x7. 训练与预测
    % A6 m1 u+ g) t7.1 优化器设置
    ! H$ U) @, ^' `3 j1 F4 J# 优化器设置) f, Y  N& L5 ?8 ?' M8 l' o
    optimizer_ft  = optim.Adam(params_to_update, lr = 1e-2), z2 R2 b1 Y/ j" Y0 T3 I7 k
    # 学习率衰减策略
    ( l& R  t2 a* B2 @* ascheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)* ?7 f0 S* L( ]- g( u
    # 学习率每7个epoch衰减为原来的1/10
    9 b) x$ ?! x/ g! ^: {0 Y1 e4 O# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算, d: I) }9 |9 R* L: ?
    ) x7 u( q/ n* k
    criterion = nn.NLLLoss()' t7 W3 S4 l9 C' K+ t/ u
    1
    + `8 |7 t7 N* M, I& f" Q2; Q9 e% H; `7 Y4 P$ f
    34 [; b( z' ]3 t$ f" m2 S! }
    4% w' y& u4 ]0 S8 e
    5
    : l; b5 r- Y- k6 i3 S$ e6+ Z- H1 r) q3 F4 t- ]; m
    7  ~, B# H2 Q. J6 p
    8+ P' f7 E1 D; b0 l2 l6 n
    # 定义训练函数
    * {2 _" o& D# \4 z) Z1 S#is_inception:要不要用其他的网络+ o$ f! Q, ?) {9 @. F% Y; q
    def train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):* |$ Z4 l* i( H9 @
        since = time.time()
    - Y4 g! w( C8 F  R7 }    #保存最好的准确率6 Z$ b! X8 ^0 T) M( [! J
        best_acc = 0+ x- g$ p5 l, U9 U0 w2 ?, j
        """
    ( l4 Q: G: h9 w1 n, u9 i    checkpoint = torch.load(filename)( i0 r0 U# s1 x7 |: H# h) f
        best_acc = checkpoint['best_acc']% ^3 Q* t- M% l0 I5 Z7 T
        model.load_state_dict(checkpoint['state_dict'])
    7 \* y1 f! p! C" ~8 F    optimizer.load_state_dict(checkpoint['optimizer'])
    : |) D9 l" `/ a  Z6 k+ q' ?5 v    model.class_to_idx = checkpoint['mapping']
    / j6 z7 y8 o  t) n9 y    """! E. ?6 F2 P' ?) M- G
        #指定用GPU还是CPU
    ( E  W, v' H( I/ |5 k6 T    model.to(device)
    , \* \- l; Q7 F4 J" o( L" Z* M7 W& c    #下面是为展示做的: a, a) U0 R3 j0 h* N
        val_acc_history = []0 S/ e% n- q: ~" m' {
        train_acc_history = []
    + j& D. j# ]% \7 j) b( j/ X, U% z( |    train_losses = []
    7 O% G$ J- }0 H) E    valid_losses = []  ]6 Y, A1 H; o# P, Y
        LRs = [optimizer.param_groups[0]['lr']]8 d5 V7 V$ o7 w" H- e! ^0 {
        #最好的一次存下来5 X6 Q/ K. V% c3 Q* h
        best_model_wts = copy.deepcopy(model.state_dict())
    2 j& u2 ^, @2 G' G7 i& Z7 Q: o8 Z" u( |7 O
        for epoch in range(num_epochs):
    # K! _( H+ \$ o) b. N        print('Epoch {}/{}'.format(epoch, num_epochs - 1))
    * N8 T2 ^1 M7 N4 }3 h" M2 s        print('-' * 10), k8 N( s. j/ [3 l

    / }8 y5 Z# k9 Y8 T6 c        # 训练和验证$ X7 u4 B2 R9 V9 u4 U1 a
            for phase in ['train', 'valid']:
    . Q3 i" U( y5 {3 O1 @# w2 V            if phase == 'train':
    " B, ^+ Z( v& w+ V7 L0 V                model.train()  # 训练) @% u6 j8 X6 T' O" L
                else:
    7 K1 o* t0 J' h6 U+ e3 \                model.eval()   # 验证- h7 R6 u2 G7 I1 M. }$ }
    ( [, p0 L" G4 N  U: P5 s/ J
                running_loss = 0.09 G  ]: T: |9 c3 a
                running_corrects = 0% P9 H0 S3 e; e9 t8 D

    6 M) N) o7 j5 u% b6 A; I$ N            # 把数据都取个遍+ s% `6 e; e! \. M' q( Z8 B
                for inputs, labels in dataloaders[phase]:
    - J: J( r2 p8 B  H( c! r. o% H                #下面是将inputs,labels传到GPU
    : Y, z; q  h! W2 i, z& p: c1 z                inputs = inputs.to(device)3 V6 l6 U, t* p, I2 P
                    labels = labels.to(device)
    6 t* D! {, B% l
    6 n/ Q+ w/ n6 D. k( X! F9 X                # 清零5 X) R( h2 R0 R* q8 L& m* Q9 P
                    optimizer.zero_grad()
    ( R. i% i7 \  t' T: Q4 r. k! n7 F                # 只有训练的时候计算和更新梯度
    7 I+ @& d2 h5 `) k6 M0 ^/ C                with torch.set_grad_enabled(phase == 'train'):8 H: r- p- i# d2 I; h& M% O$ a
                        #if这面不需要计算,可忽略+ m6 ~) S+ [- B0 h1 }( _2 ~$ N$ z
                        if is_inception and phase == 'train':9 R+ ?* K4 w$ {3 T" l
                            outputs, aux_outputs = model(inputs)# a* }& S" ]. G, {
                            loss1 = criterion(outputs, labels)
    : j  X/ k1 \2 o                        loss2 = criterion(aux_outputs, labels)! y( ]# |) ?; y$ w; U0 {
                            loss = loss1 + 0.4*loss2
    * v0 ]8 }6 c. V7 ]& @) e. p. r                    else:#resnet执行的是这里
    " v2 s+ }! L$ H4 C5 m! s, g                        outputs = model(inputs)( T, T( c5 X3 W# x3 J
                            loss = criterion(outputs, labels)
    + u. d. f$ S* @2 U2 j7 Y/ U1 G7 w9 H# V% r
                            #概率最大的返回preds
    4 Y5 A0 s% X  G1 }- n$ U                    _, preds = torch.max(outputs, 1)
    ) B1 V( p/ y! c' L* E. x* J. w
    * u3 D8 a* L: R5 ]                    # 训练阶段更新权重
    ( K: ], T# S1 l+ d0 T" N8 Q                    if phase == 'train':
    % e7 o7 `3 L8 `( k) P# \& p( ^                        loss.backward()
    / _. ~$ ^( E! B, \2 j. h# [6 q                        optimizer.step()
    " q8 Y% G- ^& R4 `5 k  n; {
    , n( d6 [% ]6 _: `' ^                # 计算损失/ x( k1 r4 F( c; ]3 Y: i
                    running_loss += loss.item() * inputs.size(0): T) k/ c: `$ v& \) L
                    running_corrects += torch.sum(preds == labels.data)) E. T  e: |2 q( Y+ s, t

    / m: Z) d( o2 e( Q2 Q( s            #打印操作
      t1 l" i* {. p3 G. h- T            epoch_loss = running_loss / len(dataloaders[phase].dataset)
    ! p5 r; {+ z! p, Z' K            epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)2 E4 x+ a1 H& z# Q7 `9 K

    7 t8 n0 c  b8 w& n, M2 ^: {( ~2 w* H
                time_elapsed = time.time() - since
    1 Y# j/ Y" [  w& ^: i; G            print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
    + C- r! w& {; {) d6 I1 b+ Z( Z            print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))& `: L2 z2 f. K. E  F1 v, l/ ]

    8 t& a) C  T3 W7 h: W
    + R6 F& M; j  Y+ f0 N0 ^            # 得到最好那次的模型
    ; g8 ]7 p- a, ^3 O! k            if phase == 'valid' and epoch_acc > best_acc:! r. P: u9 s4 X' X1 g
                    best_acc = epoch_acc
    6 M1 U- w/ p4 V, Z" c                #模型保存
    " W8 {! ~! n& b; c9 @                best_model_wts = copy.deepcopy(model.state_dict())
    8 `* \0 \8 X% E9 |                state = {: f1 ^5 r2 L: c  p5 u) q! T! q
                        #tate_dict变量存放训练过程中需要学习的权重和偏执系数
    : y* H6 e# f% x5 B3 s5 ~                  'state_dict': model.state_dict(),; ?7 g2 i! m$ P$ z( F* g' h
                      'best_acc': best_acc,
    3 B" P  L9 y( C" E# C  L  H+ s& o                  'optimizer' : optimizer.state_dict(),# }5 {3 c/ G9 j5 d2 H' j
                    }- k4 J+ T3 V0 {* [; D0 Q
                    torch.save(state, filename)- c2 _: l1 o5 J' D' T* T6 N
                if phase == 'valid':" _; K/ e( a# H) Q! i
                    val_acc_history.append(epoch_acc)# k6 D4 p% l* _7 L" c' [2 Q/ v' {
                    valid_losses.append(epoch_loss), W: P, M& f: I' Y
                    scheduler.step(epoch_loss)6 r. q: A. S8 }- X0 c/ u  c
                if phase == 'train':0 _( x  @. D* u. n4 P' N
                    train_acc_history.append(epoch_acc)
    2 ^4 D6 I1 e4 i9 [/ ~9 ?8 ~2 v3 u2 N                train_losses.append(epoch_loss)" Y7 D* v' q; Q: l2 |2 T
    3 d! v! o" }: w3 o
            print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))
    3 o: I! X8 R- ~* w# N        LRs.append(optimizer.param_groups[0]['lr'])
    1 ]; G" a, R! a        print()
    . T. R$ S  w- [8 u  g1 e  |" q4 ~, X# O' u! I: `4 {
        time_elapsed = time.time() - since/ P; z$ e" r# G: i/ Q& x
        print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
    6 _% f) P2 Y0 l, B" E0 _5 P0 c% V    print('Best val Acc: {:4f}'.format(best_acc))2 @% \, |  ~% t: z) i

    * P: G4 o  {# L/ p9 C" S1 K    # 保存训练完后用最好的一次当做模型最终的结果
    . q6 s0 s. a' t2 ?3 A    model.load_state_dict(best_model_wts)4 c, t9 d9 T* d
        return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs 9 C; G" b4 t+ j2 G3 T$ k9 V
    : A5 [4 G0 s# g: u& R+ E

    ) P- A6 _* x+ q1
    * E0 e5 q+ V5 t; p9 {2
    6 g  r# q- N" F0 c4 F) z% `. x4 Y3+ m  r/ [: T( y  `2 [4 u! ?
    41 f0 u8 ?& S) R+ j% G2 S
    5& A; ]( r! d0 [' u) `4 z8 V- d
    6
    * N# v. V+ `! h8 \; B" z8 ?8 [7
    ' j5 `2 i0 I/ l87 i$ X+ q( h8 @# O
    9
    / {3 L) ~- \# W+ @8 o10) U9 e6 X5 b/ E5 L- P
    112 o7 N; r& e' s; _+ ^
    12, A# h, d2 f0 y" ?( A7 [
    13& c3 O$ M% K! r( x0 s6 E
    14
    5 \* n; {5 c0 N5 P3 Q5 v' I3 a15
    ; Q% w, T2 U- H4 u: \, Q$ O  s16
      z/ T4 C9 j2 c( A+ }, p9 L17
    * a+ y8 q" L- z3 U18, Y( _* a1 ^( G0 M+ M
    19
    & T8 R* V+ n' l0 L205 I. B5 E0 W! |& e+ h. _
    21; _$ R5 G7 Z3 G& Z0 r5 n
    228 {2 x  ~) j& n$ a4 ~0 @' y9 x( ^
    23) e: W; ]" b! k1 B9 r& s
    24( x& l  u9 u, c7 b+ O+ U. j
    25, L' z! s# l" x
    26
    6 r' n7 p1 K6 p+ A7 ~27
    6 J! ]. E3 O* m3 x" r& D* C) ^) y28- i- K  k, o0 S! X
    29
    / n& ?. U% k1 y0 v! G8 b2 h30  ]' D7 g: q2 M) u; ?) o
    31. @5 j" r" X* Y& {9 `3 {( _
    32
    + o* U+ b( E+ s/ t4 z33
    + U. B: A. o0 z; m: k( g34
    - R/ G; G# \) @8 G( u35
    4 ~2 f6 S! z( J3 G. U7 l36+ }' p: E. i2 g/ M6 I
    37
    ; _& K- d* u" O* A4 T- g" M38
    , e* n  e8 Z$ W/ ^& l& y/ }% A- t, I39* |1 |1 X( i( L+ |  `2 W8 {& g6 T
    40
    / f9 H3 d1 S6 D* @* a41
    7 d. z2 r9 f& d' c6 y; c$ L, u: ?, C2 z42
    7 p( E; n6 V7 H& @, w43% Z- k, B6 y+ d+ i. H$ j
    44/ v! x# g5 A- K3 ~$ O$ Y1 ^- r  s6 B3 o
    45( ^5 e- @; A1 k( Q! J6 [
    46
    9 t4 D+ U1 x0 Y& g2 p47
    0 i) D  ?/ e  U7 w! O4 J$ ]) G2 F48
    ! B$ l% e& P# u: L7 {$ E49' N1 R8 N6 V$ _% M
    509 W) s3 b, Y6 `5 Q; ]8 n. D5 L
    510 B, B( v! N7 x0 r- q$ L  _8 v
    52
      P( [# c$ [1 u0 x534 z. L5 D* z! p' u
    54; ?* L" J7 m* T1 a
    55
    " T  U$ X. s6 [2 L# p56
    4 [- @8 D6 I5 N0 n* z574 {5 X. ?1 M; Y8 e% M. R
    58
    % n  {% i+ S' U0 y59" }7 O6 [, j- K2 N2 o' J
    605 M2 G* T7 R2 m' r$ W0 X% D
    61
      p$ @: {6 f% [62# L4 I9 V% x3 h. B+ c# Q. o' ]  e
    639 W0 j$ n  }: ^. k) q
    64
    ' o5 l3 G% o) V" Q' P# X0 p65
    % Z2 z4 M* q) v3 L% {66
    # u3 `& @3 c' H' ?. x- W- E/ m# c) P' F67
    : M( Q* ~6 c% o68
    $ g7 N& ]7 |8 v9 k69
    0 z) _- D( e2 J$ j$ f702 _* X7 U* d2 k7 g1 B" C; J$ Q/ i
    71
    7 w* f+ _4 Z  n. ]4 j4 ?8 V$ I& R72
    3 `  a# a% O! \3 z; n: G73$ l& A7 z1 c$ U+ E  W. y9 J
    74
    ( U) N: B- n$ r0 k* Q* Q4 V& m75
    2 [5 V- L) m8 ~3 y+ ~2 [6 w0 y767 ]7 e; F0 J( e# |! b( r
    77
    " ?$ V' Q3 Z% P9 T' O( u78. |6 Y" M" @% k0 V
    79! e( q8 O/ e) S' b
    80( x+ F$ U4 h6 t! O0 x* g0 S' S
    81: I1 i* V3 r3 Q, g' F# D3 Q
    829 a9 M) _3 T0 ]' i' v' X; j
    83; G$ I( h3 l7 p. c; s
    84! J* l% R& G. }0 m4 }2 A0 f; _
    85
    ' i' t, @& t! R+ n86
    $ t! o" r4 c8 e4 r0 u' V* d870 r- D0 Y. R5 @' L' `4 u) X
    88
    + I+ u' \; M3 X* l6 c" |89
    + ?; i8 l$ y! U2 W4 y' [8 N& R900 a# B, n8 U. F
    91* w% J' }( m" {# `& q
    92( q/ T! t! Q6 B, ~
    93# T$ |! J! j: S! Y
    94
    8 j6 c0 N3 o) Q) @, ?" I# ]953 I# @6 ]0 `) f, s/ i- H* {
    96- Y5 }/ O  Z/ F; b! _- t- Q
    97( Y0 t  P9 I0 A7 ~" I" {
    98- Y  T; p9 ~0 x( ^) L% L5 r  g; M
    99
    7 G4 x' X6 c0 r( }100. A& V; S3 T7 R" m2 m, u
    101
    5 N  c) p9 y& }4 X102
    ; M- R1 Y0 |' Q103+ Z$ O! d9 F8 T& M* L. w
    104% s* |3 E$ V  m$ t6 ~. z
    105
    4 i; r. E7 B1 `7 g& g106& b$ \4 E& a- Q, K
    107) @4 k* {( E+ ^- v+ T2 P7 q) g
    108
    ' j: i) W+ e# k/ D109
    - h. V5 J5 t7 L; Y6 H110
    3 V5 j+ g( Y# e111: ?8 y0 I' _% C. ?$ @
    112* L+ ~; ?3 }) U! W2 R! p
    7.2 开始训练模型
    1 T$ ^4 I2 P% r& e/ ^; c我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次
    8 u& n5 u! o4 H" ~' _* p6 b$ u+ W" k! Q4 Q" J
    #若太慢,把epoch调低,迭代50次可能好些
    : g% Z# `9 p8 v* k. }6 h8 Z#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合
    + [8 Q3 s2 Y" V( {; [3 V: 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"))
    ' B/ v1 T' O- V0 E- l4 q/ U' r7 N9 B, g  x5 c3 S& S
    1) K5 ^- Y' _; g$ O6 @  D; y  f
    28 h  x% a  ^. s
    35 k/ `5 E5 `  h7 V+ v8 J
    4: R; ^  X; ]4 _9 W+ L6 V7 h
    Epoch 0/4$ s( P3 c* E' Y& N- a
    ----------
    - K. j5 U* G' z4 aTime elapsed 29m 41s
    1 O2 K, T, {% N5 h0 y9 Gtrain Loss: 10.4774 Acc: 0.3147
    ; B& q% V9 D8 L' {Time elapsed 32m 54s" \1 x3 |$ Q# ^* o  S, @
    valid Loss: 8.2902 Acc: 0.4719& q, C/ g, ]9 q% e$ \" m. T& W
    Optimizer learning rate : 0.0010000
    9 t9 |* o2 Q# `/ z2 l: I- B( t1 ?- t4 G6 D, g
    Epoch 1/4  Z! ~( e5 L$ u
    ----------
    * u9 y4 |! M- \' J9 I& TTime elapsed 60m 11s: i8 I, p& b$ K9 C; X9 L
    train Loss: 2.3126 Acc: 0.7053, O3 G2 V, n# G3 J* F4 W/ C$ I! p
    Time elapsed 63m 16s) O- f+ O& t- I/ \+ V9 y3 Z% j
    valid Loss: 3.2325 Acc: 0.6626. e( }  D5 I" D
    Optimizer learning rate : 0.0100000" u6 |- P  U  A/ N; \$ x
    # _5 d( w5 |2 |; O5 p" m
    Epoch 2/4
    / d4 S1 `( S! |----------
    # Y; {! S6 o6 v0 [Time elapsed 90m 58s
    ' b! F3 X3 h9 Y6 l% k0 Z7 s# Ktrain Loss: 9.9720 Acc: 0.4734
    0 t1 \- [* S/ I& ]; T3 Y; @Time elapsed 94m 4s
    ( Y8 U9 w% P: s' V! fvalid Loss: 14.0426 Acc: 0.4413
    2 w% z4 \: O8 w. g; v7 wOptimizer learning rate : 0.00010005 }2 P1 T! w3 n7 u  J2 r2 s  B( Q

    : a: h" Z" T7 K1 V! [) `( h3 bEpoch 3/44 p' g' T/ s  a& y3 v
    ----------
    & {% }- X. W3 N: o$ X9 FTime elapsed 132m 49s
    " _* u3 o3 k+ z( k0 t% D0 jtrain Loss: 5.4290 Acc: 0.6548
    : k! U4 |/ s5 V% z) mTime elapsed 138m 49s! H' ~/ y7 h1 D  q* P4 b( C* b
    valid Loss: 6.4208 Acc: 0.6027
    * a- i; `$ f" K  H0 P+ Q9 t% KOptimizer learning rate : 0.0100000
    4 Z% V* P3 c( q5 ~' O+ p1 f2 z( x1 S  K3 ~
    Epoch 4/4. k+ W4 _- r% [* h5 v3 ^9 V
    ----------/ _; @$ Q1 `6 w- y' O9 U7 l& a8 c
    Time elapsed 195m 56s% ~  X7 u3 ~5 ?, @0 O+ ^- A5 g7 _
    train Loss: 8.8911 Acc: 0.5519. U6 u% G! ~0 Q4 P/ S
    Time elapsed 199m 16s- W  {9 Y% e. r- ~1 f3 t
    valid Loss: 13.2221 Acc: 0.49143 |6 |9 `) J" u! V
    Optimizer learning rate : 0.0010000+ S! W6 b' G  g8 m# B) q" O

    8 ]/ {4 L; d4 @# J# {% f* s4 HTraining complete in 199m 16s! D+ }3 j+ [- \! M! N8 h$ v
    Best val Acc: 0.662592$ ?- c: A8 @/ f! y: V6 X7 L- r
    5 G9 F7 P0 n# z
    1) M2 ~& _& L: g* w: U" _
    28 Z0 \+ [) `$ O$ t- {& A+ _+ W  Y+ y
    3
    1 D/ X; V% \/ [$ q4' R( f+ [/ [+ g+ K
    5
    1 U9 Z* i4 f" ^9 S& e( ]6
    # @5 {) a3 y$ R7. Y) Q1 `( {, O! ]5 ~7 X+ N4 v
    86 J0 T! V2 Z" O7 d, w, ^
    98 T# {6 Z5 j! m0 y$ F, x, Z
    10: w, a# H' }6 j0 _8 \. Z
    112 U" R/ g! I0 y, ]0 N( j
    120 P) h, M0 j1 L. Q4 T2 [5 ^" n
    13
    . `; E4 Y7 g% k# V! k14
      B  b" {, A( E" q15
    + r5 D$ h) y8 R, v" j7 P* U  H7 A) A16$ h! L' s% F- {4 T# Y+ h2 l
    17
    / X) w: y& _* ^8 w18
    6 t7 h  \2 L! q8 c+ ~196 ^+ F, p6 [" g, Z& L& ~
    202 S: K8 v! v5 q- ]
    21
    $ S) c. U4 Y& W6 [; l' l; P; R22
    * ?* K* g8 J% a- l; i4 A2 c23
    ; `. s$ p" q3 f: ^7 l1 W24
    . x; k) l( \4 B8 H5 u25. M6 a) \  S) S. C" p
    26+ M; x: u! s4 d7 g3 M9 K& z5 G
    278 u" U, e( o1 ?) R
    28! x; |9 o/ ?7 G  n, v
    29
    8 j. l/ D2 P4 ^- B30
    4 Y3 n  V; A) J5 Z- A' j& F31
    $ E8 O! l3 t7 ]9 l; h8 I  H. c$ I32
    5 h0 P$ V" H1 h. g, P7 v. u1 ~- T338 K3 G% O2 [& a; G; U
    34& [/ X# v. K' x0 L% G0 x
    35
    7 E& Y) j8 ]- C! H* M2 P36
    ) S( t2 i3 X( `9 l. L37# v+ o  }- ~8 z
    380 e5 x6 e" g1 Y# @
    397 y2 b9 H1 }5 k8 w; q% Q
    40
    $ m( \$ e( e, k410 O4 O! M, U  l; z* {
    42  J. n3 O8 H9 f; c  b1 Y/ b  Y/ g* }
    7.3 训练所有层
    0 t7 s. i4 ~8 J  q6 T; ^. p# 将全部网络解锁进行训练
    * J; Z( g& i9 W1 yfor param in model_ft.parameters():9 F0 g: L9 Z( Q1 x0 d( v/ C
        param.requires_grad = True3 l# G( A" Q& C3 d$ U* A% m

    2 b( o% W! ^8 X6 x' w. G# 再继续训练所有的参数,学习率调小一点\' K, w6 N9 {" W  j$ V
    optimizer = optim.Adam(params_to_update, lr = 1e-4)' U5 z8 C/ ^! C1 S2 k
    scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)
    . I& W3 f& O  `- |8 S8 N  Z9 T& \
    # 损失函数, B; c2 s( ]) _+ L
    criterion = nn.NLLLoss()
    7 _* x' i- O5 V& R) b2 a3 f1
    " t$ A2 e- h1 Q) q/ V21 U& a$ _7 E1 r0 w5 m* K
    3
    0 r# I1 f) h* A' L4 F4
    * R/ o4 `+ }/ C' J5) `: I0 y# o) ^5 F5 s& f* Q
    6
    . ?* |( n+ E6 A3 B9 J/ [! ?7: G/ g5 |! ~+ E4 b% H' t" x( j" s
    8
    ) F8 J# Q; T/ h1 y- J9* T1 h: [6 ^8 c* d6 u
    108 f; `) B/ j4 |2 d& e
    # 加载保存的参数
    / u: \. @9 j0 U2 t# 并在原有的模型基础上继续训练7 A$ C% @, j, U; J* k+ Y
    # 下面保存的是刚刚训练效果较好的路径$ b5 Y- c! K/ l: |8 E. _
    checkpoint = torch.load(filename)/ p4 k( v3 N* t* I' ?& p# Q1 s
    best_acc = checkpoint['best_acc']
    : _% K  h! w7 ^, W' y" H0 d  m7 omodel_ft.load_state_dict(checkpoint['state_dict'])
    , I' Y' T8 h' v/ Doptimizer.load_state_dict(checkpoint['optimizer'])
      i5 g: m0 _. Q0 `. i0 B1; {" y! \4 X# H4 R3 {: t
    26 p+ ^8 f$ L3 G. i6 b( o
    3: C- Z; L; S3 V% u+ q
    4
    : [$ y; ?: ]) z+ e: |9 Z5# v  P; w  ^& D5 ]% ~" T
    6
    " t% \$ J* x2 K' Q" V70 u9 W2 k4 x* ?& [$ L' k! g
    开始训练
    ! I% z& @/ m$ J" N9 s2 D( O注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考; L9 C5 s; p) G

    ) A* ]8 d; f" y. \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"))3 F. m- j& K: _  ]
    1# `3 Y* |- R& Y' D& m8 q9 D* D8 k  U6 j
    Epoch 0/1. e; \4 O; \% v8 I
    ----------" l' w0 S( e# S; c* m
    Time elapsed 35m 22s
    3 D4 Z) [3 V1 T7 ltrain Loss: 1.7636 Acc: 0.7346% P. A( X. }" b  P! T# _' l- u6 B
    Time elapsed 38m 42s
    . t1 [: H% `8 X- y3 ovalid Loss: 3.6377 Acc: 0.6455/ v) |; E. Y) s3 |
    Optimizer learning rate : 0.0010000; }# J4 Z3 O6 b6 }- ]2 X
    : A" U& E1 E/ R8 U- n! }$ Q
    Epoch 1/1
    : P- S8 |1 J# T----------% K) ^/ H7 B' v& w( ?5 c, O9 H' ~
    Time elapsed 82m 59s' Y" i: S) u0 \
    train Loss: 1.7543 Acc: 0.7340& C  J/ z% I9 d6 v2 _
    Time elapsed 86m 11s8 `( h2 l! H; B: ]- j
    valid Loss: 3.8275 Acc: 0.6137
    / w! m, k5 g: S) \' o& XOptimizer learning rate : 0.0010000
    : X$ p" Y4 j7 r& l, S+ F8 ], y- m
    # L- V/ j5 d+ F% Y1 U; U7 H) eTraining complete in 86m 11s
    , @3 S9 k% Q$ x% l* dBest val Acc: 0.6454776 _/ Y5 q6 \( P4 J

    ) p" T7 e) H$ v# D! T1
    / R7 N9 ^0 |, B2
    7 j, W* F8 u3 {9 X3 J4 C9 q; C! H1 S3" U* }4 x7 V3 k/ `% Q/ G& j& y, Q3 _! }& L
    4
    6 W& M* _/ k/ i# `56 g* E! N3 C6 n$ a' G
    6
    # h( Q4 |+ N0 T2 _1 Z9 N; i/ j* H7
    , T' u; N+ g+ y85 K4 w+ X  G8 Q4 O
    9
    % V8 Y2 ^& ~: D6 F10
    : v+ Q# n3 M- |6 c! M110 ~- @: m6 V. c" {1 E+ [2 u
    12
    + |0 R  m; a: W9 i( j# n138 q& h, v6 v( z6 Z' M7 y9 N
    14
    ) E$ c5 C( V) o8 j# q15
    2 |# X: v- q: ^3 s3 i. v. A6 }16
    5 ~) V) x+ K( M' \8 P& Q0 s179 z, k+ j5 x. [
    18. ]$ a' {! Q/ L7 A
    8. 加载已经训练的模型
    2 O/ i2 i4 S3 `; f  q相当于做一次简单的前向传播(逻辑推理),不用更新参数) y- p: G0 r3 X% A  T0 Z

    9 A" J0 O! ~* M9 B7 z# g6 xmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
    ' t* h# E% z3 Q6 v% r5 C4 f
    6 m3 @1 f1 q) Y  y) c" I" i0 T# GPU 模式! D# P) y( [- h4 }
    model_ft = model_ft.to(device) # 扔到GPU中/ |' O( p. v: W5 i) [( F

    - n* m4 x2 w+ T7 [% z5 d9 R# 保存文件的名字* j; ]* o' d7 b* N
    filename='checkpoint.pth'+ p5 m& @( c* ^  |. R
    " o* {5 G8 P* H6 c7 Z& C6 s5 G
    # 加载模型
    ( m# t8 L: ^! p/ r5 G& c' Mcheckpoint = torch.load(filename)3 }, j& I3 J6 H4 b
    best_acc = checkpoint['best_acc']. a& a6 g. m6 \- S0 U) S0 r
    model_ft.load_state_dict(checkpoint['state_dict'])
    ! Q9 J2 g% q# [+ h1
      Q4 j4 Z" l* q2& O8 n9 \: }! A" e. ]) ~8 [
    3
    ' v* b1 `8 O, W6 e0 ~4: G8 F7 ]* u$ ^+ P: @7 y8 E
    5
    1 ~0 r. g# g9 G6
    2 U8 g4 K' q7 e) `/ g- T; T7
    # U7 \- R% Z7 {7 A6 u8
    0 X2 s' O2 s$ |% _( `9
    1 r) T) z0 K6 j: v10
    6 }* |. |, b7 A& \' C- E2 ?11
    % @! h& P  F* j. }! c# r. @12
    ; d3 @# N% j5 B<All keys matched successfully>
    : X; K$ u- A, A/ F3 V- ^* o/ ?' }: F1
    3 Z* N" F$ j" a# e  @/ p' qdef process_image(image_path):3 B$ R9 Y% m/ J
        # 读取测试集数据
    ! _% y7 k, A0 u! o' r, e( \    img = Image.open(image_path)( U: u" x# S+ e4 }' w7 w5 I8 d9 l
        # Resize, thumbnail方法只能进行比例缩小,所以进行判断
    , T; x' g3 n+ Q    # 与Resize不同' t7 @- }7 E; ?# C/ b+ B
        # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小& i3 R; n- L, v6 K2 c. h# \
        # 而且对象调用方法会直接改变其大小,返回None
    4 l9 {) i) ~2 q1 x+ P) ~    if img.size[0] > img.size[1]:
    % m0 ]' y0 J+ e6 E: r        img.thumbnail((10000, 256))
    $ w4 [+ J( i" N9 ~    else:; J1 S: m4 t7 P: j$ v& D- d, [' V' c
            img.thumbnail((256, 10000))2 T) [* j* z! o, z7 Z* o

    : \" ?) y: N5 t' _$ e    # crop操作, 将图像再次裁剪为 224 * 224
    " {& j9 V% e3 E+ A+ @    left_margin = (img.width - 224) / 2 # 取中间的部分# A2 Y7 l( d* q% u
        bottom_margin = (img.height - 224) / 2
    1 n3 o  z/ N0 I- I# y+ m    right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度) m* W+ {( U' v, S4 _
        top_margin = bottom_margin + 224
    ; M$ \+ K1 H, m# ^4 R# ~6 Z9 L6 V! u  f% E' \8 _
        img = img.crop((left_margin, bottom_margin, right_margin, top_margin))
    , E+ n4 g3 t: Q7 S- X1 P) q6 \& N' n
        # 相同预处理的方法) O3 H5 e+ o8 \9 n
        # 归一化2 ]( o- M0 X- a  I6 M) y5 m! c
        img = np.array(img) / 2554 `% a! z$ s5 M4 U8 s4 S( N
        mean = np.array([0.485, 0.456, 0.406])1 i: I. r( h& v$ E8 H6 O
        std = np.array([0.229, 0.224, 0.225])9 C' j1 ?( l$ |9 F
        img = (img - mean) / std! p7 h9 J# M1 S4 @0 L# Q, _

    " l4 t& O* W- X3 p/ R, ~, ~* d3 r    # 注意颜色通道和位置
    ( ~& Z; L' G" @: U    img = img.transpose((2, 0, 1))) _) v" a  ^2 c! o+ ?# d: I

    ' {- @  |; F$ I3 P6 E8 Q. B    return img* n4 [  `' j# n
    & a& L$ s4 Z2 ^
    def imshow(image, ax = None, title = None):/ N% b) o" Q* ]. N
        """展示数据"""* M9 v, Q4 j1 w! j0 C- M
        if ax is None:& v: Q. H, _7 R( ?; F
            fig, ax = plt.subplots()
    3 b5 A2 ]; `! {
    " _- H" j3 ~: v$ P0 X/ L    # 颜色通道进行还原
    : j; |& A9 m3 G4 o3 _$ A0 e    image = np.array(image).transpose((1, 2, 0))
    7 P% F( f/ _9 G9 V/ A
    4 r+ b, m  B8 ?4 [' W: R- k: L    # 预处理还原' Y7 |7 a) Z! W5 a, `
        mean = np.array([0.485, 0.456, 0.406])
    4 ^& E0 Y2 s/ W& S* q    std = np.array([0.229, 0.224, 0.225])
    ; q; U( U. E1 Q( R0 t" k; _  y! \    image = std * image + mean
    ( M& z4 s3 I; G# g; h    image = np.clip(image, 0, 1)
    ( x4 s7 H$ o* p& K& Q' X6 Q- |2 K# E) I2 ]! v' Z0 _
        ax.imshow(image)
    , a  W8 o3 Q- s2 o    ax.set_title(title)
    4 {9 Z6 p0 m- w  G- N. h. A9 T3 ^3 E
        return ax  X0 k6 a' j8 K. C
    7 d# z- f' V6 J* `
    image_path = r'./flower_data/valid/3/image_06621.jpg'( v# f2 W1 D, Y* L8 D& @% L
    img = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理
    6 f* {* n( H/ @9 |3 bimshow(img)
    ' C3 A; k; R% m& v( G/ y# K6 B0 P/ i
    15 v  z" }( K( p# d
    23 ?. H3 p9 ], C% X
    3
    5 g. f, ~: u2 v8 t40 K( b9 ]1 w' T4 s7 R: \
    5% l* x7 i( {/ w2 K) \: R1 u
    6
    5 N. a7 k, v9 _5 M/ g7
    * H! L. O$ B! E0 Q& Z+ Z8: I" u& [9 w8 Z
    9
    ) {0 T0 ~* |/ s7 k, a/ r10! A. B/ n, l/ L7 @1 P! R
    11
      v' q, k* Z* O1 m& Z2 r/ X12
    ; \* u) I* x2 K' Z1 X13
    5 T& L4 s) e7 @' @14
    : i- l) }& T' U0 {15
    5 v# }. c$ h: ?16
    ) x8 i. b- \6 O- F* q, ?176 ?  s( T; ~) j! b! e; M
    183 v9 S% i: f3 h% p9 z- q4 h% w
    19
      ^' |' G! o2 y1 E8 e  g20# s" u; o3 s3 f) ~( R
    21  a# }& W9 A* Q) ^" F5 f. K
    22
    0 W, H) I  W7 J7 f23
    0 u* K7 A) g7 j- f7 N, u24
    9 t; x3 e% S7 A) M: c* D& N25+ x( j; l' d9 ]6 J( o( k
    26
    8 E  ^" L% U: z27
    ' x0 }/ c. u3 D* [28
    + t3 o7 m3 O6 S: f29
    ( h, a0 S( P8 c, Z- p2 J" ~30: D$ D% [2 _6 m5 L$ E# j
    317 [, c# K7 I" W7 X# k2 V- j+ A
    32
    7 Y9 r/ |$ {$ |0 \2 k- W' x33
    6 h5 @  v1 P/ G% o1 X7 ^! {# I8 R34: w) M" z+ s: b' D
    35
    8 n1 W9 x" V* y! p3 q& A# c36
    5 @1 ]1 I# n! V% G6 K4 \2 A, _37
    $ l9 D. C5 X# U  |2 I38. i, w" e8 ?9 V8 p1 Z# G
    398 f2 C' P2 r  u
    40" ?. A, n9 M: f* @
    410 T0 O! f' A/ G- B- F
    42. [2 }3 Y- |1 Y9 r0 c8 F7 @
    43
    : X/ }( r3 M5 v2 U44
    % h/ P# Q/ r. R# _* J# d% u7 A459 W  Y4 ]7 b0 ^/ p
    468 w' B2 p9 [& x/ S6 ^& G( n$ s- N
    47. p5 x2 K* u8 T6 Y$ r
    48
    + r) k) N0 @8 i' D# _1 [6 L2 K49
    3 c* O# F; d2 H/ v- J* T50
    ; `  H  h" m1 B4 a! ^- B) w512 _4 v4 B: q$ y1 c9 c/ F
    52
    6 h5 @/ T: [8 J; ^; S* a+ h536 n% O1 |$ ^- D4 k5 o
    542 n7 D! n9 L* \1 `$ o
    <AxesSubplot:>  L$ Q2 l3 ]9 b, W+ k
    1% y  T$ h# i4 F) x

    7 t9 U( T7 e  ?上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确& [& j5 K2 D1 t2 `2 k
    ; b( N( ^5 S1 _/ z: e5 M+ q0 d8 p
    img.shape# J3 v  _% N# V, [2 v) L  m
    18 N- V* y; E5 M
    (3, 224, 224)
    ! B8 I+ Z( q1 D' l3 Q6 Y" C0 k1
    : J% n: P4 j* r" B证明了通道提前了,而且大小没改变
    ) p* _9 }- A" V$ l4 C7 T; T- A
    $ [2 n3 Z! ?% f+ h, |' G9. 推理
    * W- k' ?  y# ?* I9 mimg.shape
    " [2 Y& z0 c* z. P3 v
    . L3 D9 M3 h# s' R* l# 得到一个batch的测试数据
      q( s. Y2 R9 }8 D- M; L6 Vdataiter = iter(dataloaders['valid'])
    - \/ s0 v3 G& |5 N5 T! O9 j2 E. mimages, labels = dataiter.next()+ D8 z4 ]9 e& _: i5 W

    ' |6 v3 q/ R# Gmodel_ft.eval()% D0 D7 v7 _  g9 X! O

    3 T  Y7 d) q3 S* N' [" Xif train_on_gpu:' {+ F4 q2 }& ?
        # 前向传播跑一次会得到output* P/ Q+ y! b$ `. F( P" ^
        output = model_ft(images.cuda())/ ]) n8 A; n4 l
    else:
    : t& X6 E1 d* Z6 ?  I5 Z    output = model_ft(images)
    7 ~  ]3 |1 C4 {. z9 f
    1 c- U1 G2 p# z) p: S# batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值) K6 U; f. D- @6 w  `
    output.shape
    % N! o8 q( @' v- U" r* p/ T& b) P$ t) N. D
    18 g8 `4 j0 f( d
    2
    2 W8 h+ j6 R" n( L; o' a34 s' B. a2 H4 X9 S, e$ x' C0 z
    4, r% N. y5 N+ K
    53 C) Y; x& Q3 b
    6# ?. ~4 N+ s, m. f
    7! k, I. O5 F$ {4 C5 P1 z. v
    8
    * z" k5 S# C4 g- z6 o/ V2 S) X9& }4 _9 N( J0 ^+ V% h$ ]
    10
    ; I& P' J6 n, g5 K11
    + q" n! N% j" J" n' K. z124 O, l! |3 C% l" _, U: L, x
    138 y: L* j0 P6 _, ~
    14
    . @$ }+ D3 j/ ?1 N  M" i" D15
    , E/ i) p: o8 ?1 N  \, F16
    2 L- c- d0 q! F8 S2 Ltorch.Size([8, 102])4 ]0 i6 t5 D1 z" b$ Y
    1
    ; y- M5 G' F; }  H0 J0 e9.1 计算得到最大概率
    2 s8 H- L3 Q2 `- L( @_, preds_tensor = torch.max(output, 1)
    9 m" W. w+ s0 F# O" p0 n& E4 @% d& M& ?6 @0 _
    preds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量3 ]" |: Q8 {9 x
    1
    ) {: o  d% \$ \) l; U! e2" R; @2 u& W. Q5 C; W7 S+ h) a
    32 ^6 G' y" I* }! B9 ^, N
    9.2 展示预测结果
    / K5 f$ C( U2 H( {5 kfig = plt.figure(figsize = (20, 20))
    + q; h$ i( @9 fcolumns = 4
    2 b2 b' d# ]& ?# Yrows = 2
    5 U- n/ h# H5 `
    & Y# s/ o  }: Ufor idx in range(columns * rows):0 Y9 q5 s0 A* P/ l3 C' V
        ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])
    : s1 s0 j' Q) h    plt.imshow(im_convert(images[idx]))  d! O8 M) _- p, R; x; c
        ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), 7 Y1 o  f. p8 I% `! T) C) a) g
                    color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))" M1 n: o$ n5 F9 |, A$ g
    plt.show(); y  g9 _9 G" w) S4 _& j) j7 x
    # 绿色的表示预测是对的,红色表示预测错了
      X% n. s# I. ?5 ?7 e1
    8 Y% l% B- C% m8 Q* }2% c; E4 Y$ s+ H( o2 @7 y  ^2 R; s
    3% L2 j2 P8 P: L3 Q; X' y& Z! |! s; k
    4, D4 L& f' d: e7 ~8 c1 M
    5
    . q  p/ b# B; {) Y6/ }2 c( g+ X% w9 ^" q
    7
    3 w& p+ O3 b) Q9 V, j$ i8
    . }6 ^& |7 w" ^/ V9 y1 Q& w9
    ( d% d! V7 _7 L$ L10) `2 b$ a& P3 T# X. |
    11
    8 l  ~7 Z* Z9 k; z' W( z/ ^
    4 Q; R) e, z3 b& l9 q2 f' z' K5 G# n4 v. K$ l: B3 D
    . `! e9 v' C4 P; s
    ————————————————9 A$ E, v6 i  C: s
    版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。+ ]. b8 p! r7 W* W3 Y( O( ^+ E
    原文链接:https://blog.csdn.net/LeungSr/article/details/126747940, W1 U$ l% U$ e' |

    0 o* g- ~0 ?; C
    " j, T5 }& |' ?; p# x5 ^
    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-8-1 20:51 , Processed in 0.542905 second(s), 51 queries .

    回顶部