QQ登录

只需要一步,快速开始

 注册地址  找回密码
查看: 2823|回复: 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)实战案例! v6 c- Q2 }/ ], h+ G- a+ p( J) A

    9 w% x2 P3 O& T* h6 X. B3 P文章目录% @. W* y. \8 l& H9 v  j7 `4 L' d3 u5 F
    卷积网络实战 对花进行分类: k: h) h* ?: Z8 ^, o0 I
    数据预处理部分& i7 r8 C9 d  O  f7 l4 i4 ]
    网络模块设置0 P3 u) J2 b3 Y1 x0 B3 s
    网络模型的保存与测试
    $ ?) Z" p9 D/ I/ y9 Z数据下载:
    $ [3 _- Y9 ?% U* }6 {1. 导入工具包; z# ^# b# T4 Z' A& u
    2. 数据预处理与操作
    5 j# {& g9 m# l/ a6 e3. 制作好数据源
    / K- w9 h# @0 P. F读取标签对应的实际名字
    0 ^( j' j. D  T4.展示一下数据( H) v8 J( i2 N4 L
    5. 加载models提供的模型,并直接用训练好的权重做初始化参数
    1 U* U: `# A; u# s. @  f3 O6.初始化模型架构
    : n  B! M3 E/ C: ]; b6 }) v7. 设置需要训练的参数$ Y9 T9 m  r! f0 v. [
    7. 训练与预测* f; y6 `) N/ Z
    7.1 优化器设置6 z' Y4 b. P" d7 Z) C- x
    7.2 开始训练模型
    - w) g2 d" B% H3 [. Z# q! `1 @2 e+ G% T7.3 训练所有层
    9 @8 G- {. {0 C0 D开始训练: `4 U5 ]4 p. a% p5 C
    8. 加载已经训练的模型
    - Y; c4 q4 @. e( c! j  v9. 推理1 R/ F5 r7 y, I1 p/ g. y- Y
    9.1 计算得到最大概率; Y4 v- n2 k8 B; x/ B$ r+ _! ]8 z
    9.2 展示预测结果0 E( W$ w$ B- g  k9 ]  ^) W/ I
    写在最后' ]: T& V' j; O
    卷积网络实战 对花进行分类) P* g3 }  x& `' r; f" g6 m
    本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能% Y. T& ~$ p9 S1 e
    & b: I: ]1 V2 r/ `- X
    在文件夹中有102种花,我们主要要对这些花进行分类任务
    / V5 h5 L) u/ ^* W5 k文件夹结构
    2 O2 d% j/ W5 k# [0 s1 B& j" p. p* ^0 H9 \
    flower_data
    3 M; s" c0 w( [7 ]1 v
    ! B8 @; g5 u) f/ n4 V) U( ?' etrain* z/ B% E% t% G' b4 f# j

    3 [5 C( w1 x- R1(类别)
    % X4 K0 r, i. }. o2
    % U7 x& Q4 s! m* _* B4 oxxx.png / xxx.jpg
    " k3 [5 k- j; R4 }8 h( E7 E# h, [5 cvalid+ h! t: r5 \/ p/ O
    7 K2 D$ f/ F" J6 _8 C7 K
    主要分为以下几个大模块2 i' W  P: l) f3 F; C4 `1 q0 P' A
    2 L+ y, W1 ^* N# E+ S1 d
    数据预处理部分
    9 E3 T2 |4 S$ V5 j; }6 V* s0 n数据增强& K7 L3 p+ k2 g: w/ v" @
    数据预处理7 y3 k6 g) f- @# Z
    网络模块设置
    ' G7 [- H! `3 B- ~/ e+ J5 g加载预训练模型,直接调用torchVision的经典网络架构
    9 @9 `* d4 C3 n9 N  g  G因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务
    / ?8 [5 R6 h$ I4 N+ p3 r* ?: i网络模型的保存与测试
    9 i& |& _0 H7 W- \" b* }模型保存可以带有选择性
    . n6 x! t2 M2 R, @( h6 `数据下载:
    7 W: W) L* K# Y2 u/ U; a. P) \$ ~https://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset' X, }# F$ E6 o' z

    + Z, m2 z# Z) R# O+ i) A改一下文件名,然后将它放到同一根目录就可以了
    & g+ C) k+ H- ~' C! `! d: M) P+ k3 G4 ?  w5 H& f% n
    下面是我的数据根目录% N  p+ G7 x1 H% p0 j2 _4 a; f- O

    ; G5 W' t/ V% J; [1 y8 M2 j6 c( i2 a. K- ^# M$ v
    1. 导入工具包
    : w& E. i* R, S" R2 Z& l+ uimport os
    * D9 p& ?6 c/ m( zimport matplotlib.pyplot as plt1 Q6 |+ U; o5 d1 U4 b" J
    # 内嵌入绘图简去show的句柄! @. ~/ m8 i% s+ O+ ~! h  [
    %matplotlib inline
    % ?) k* q/ _& `import numpy as np3 N& E1 x' _" a1 ~2 Z% y$ w
    import torch, [  L+ V1 |- r( x, }
    from torch import nn) o5 G# U7 ^  I5 Q

    0 M$ W, t0 t/ @# n  i* B5 fimport torch.optim as optim
    7 S; }7 O6 X" O7 G5 `import torchvision
    * Y. j3 B& ~7 Y+ b- y# D+ L$ dfrom torchvision import transforms, models, datasets
    . ]! H/ H2 o0 e# p2 W( n# T+ o$ l" _' X) S# q$ _0 \
    import imageio8 q4 i0 R8 L6 q* u" |6 m: z: d9 j
    import time  H% W7 r9 ^) d0 _# A9 {; M5 }
    import warnings$ K7 M5 b6 N% o, i' S
    import random
    ) y9 s& t8 s0 r+ `import sys$ l4 \. K1 x: J6 T7 Z6 ~/ q
    import copy
    # m3 h' N" H' l' Himport json; H7 h, T" r8 L' J, R" e
    from PIL import Image. s0 s! z0 H, p6 q
    - q; O1 ]$ f1 h( e

    6 z# ~2 O; ?7 w0 w* g  V8 b$ Y1+ w' e( B3 \# w6 E* A0 f9 x4 |8 [
    2% ~* }  S2 j7 j: C* F* }& g
    3/ V( {, `6 m: j. g# g
    4
    8 S6 ]' B  Y. }/ ~& _5; z* T" }% U! G0 ?$ E. {* F
    6
      o$ i' K. G5 W$ G5 t' t7, D  U, K+ Y# B* ~  H* X; L6 Q
    8
    ( E* H+ G7 E* v. B1 g- P9
    + f2 A) B& w! C, _. z0 O1 ]10
    9 O6 a! Q6 ^) `- x7 A- ]9 J11$ W1 c9 I. r, a
    12
    ( e/ r/ _( c! F  P13
    * {. S4 {0 R( J14; K* e1 S1 t5 q  E
    158 D. B/ l$ `; E7 |$ R4 E
    16
    0 \' o& L  J2 s* p/ R7 S17
    % C8 Q% m/ x+ U" k1 @& d18
    0 `3 e5 \9 P5 x: a: w19
    4 D" W3 e3 @# g$ z20% d: v% g0 {& u" W/ R/ S: U4 w
    212 i. I# B; [" D$ F9 ]
    2. 数据预处理与操作7 m; C  `5 x' }6 O
    #路径设置$ |, D  _3 j: _$ p  F) n
    data_dir = './flower_data/' # 当前文件夹下的flowerdata目录
    ' _: Q$ {  H# O& [* N& @train_dir = data_dir + '/train'! u, q9 C: N! r9 E( {
    valid_dir = data_dir + '/valid'
    & U  g& e8 h/ {6 ]: k% i1
    ! ]+ `" V' W8 n* g  f27 r6 k+ {" y1 ~+ q
    3* Y9 ^1 Z7 V- o2 ]  \! g
    4# U/ T' E" y6 ?# y; f3 f
    python目录点杠的组合与区别
    5 V+ g( {2 T  T- G注: 里面注明了点杠和斜杠的操作
    ' Q0 ~/ g0 R8 O% k# u  T, i+ L. E" \& v
    3. 制作好数据源( I2 F' F- p( I1 q1 i# [8 X! E
    data_transforms中制定了所有图像预处理的操作
    9 {9 J; b3 g% g7 Q$ @, a& dImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片6 h! C4 T) p( ~. i$ G% K
    data_transforms = {- P3 s8 I/ P" V5 h1 I6 [( Z
        # 分成两部分,一部分是训练/ e) x8 [! |0 `7 S" s
        'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间4 B+ w- ~+ m; B$ J# W! f  u; d
                                     transforms.CenterCrop(224), # 从中心处开始裁剪: d+ }. X8 T) g5 O1 g% y0 W# K
                                     # 以某个随机的概率决定是否翻转 55开+ y/ A$ S* ~' ^  p- {% j) A
                                     transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转
    : t. i5 M- @. g) j; ^$ B0 ~: n8 N                                 transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转
    " I. V" @9 F- ^$ y                                 # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相; [( J/ ~5 P- S8 D9 ~
                                     transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),
    6 S  y1 ?0 o1 r                                 transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB
    0 L, ?9 I- d" x5 P5 k& C$ `/ C; a                                 # 灰度图转换以后也是三个通道,但是只是RGB是一样的
    ! `" t4 U# `2 j! r5 F                                 transforms.ToTensor(),
    ; Z2 i0 U6 [  a1 Z% ^                                 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差
    4 W' F6 s! G& Y3 i1 |- r# N6 z* F                                ]),# T. n) v3 n: J+ V
        # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化  u4 E7 l' j4 d- i4 l/ @4 f4 `: `( Z
        'valid': transforms.Compose([transforms.Resize(256),
    7 V4 b# Y' O" L                                 transforms.CenterCrop(224),
    6 Y; e4 t: N* Z                                 transforms.ToTensor(),
    ' i- U: u9 X: }/ y& S                                 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同
    0 O( s% f7 c( X# o9 m- c! ^# \4 ^8 {                                ]),# g( N) g! q0 V; }7 z2 m
    }
    & M# |& A5 P+ Y' r
    6 d2 U" V( _4 e# O' w& n1" ]6 {6 `; e4 B
    2; x4 Z( D2 U4 `* k: a2 [! M% ~
    3
    / x, u0 A) L3 r# m3 B4 B4( K1 Q$ t; ^9 z, J8 M
    5
    . R/ C* M! x2 }$ c6
    2 B8 P. E5 d4 `, f2 I77 C: A# c. B8 J+ P. l
    8
    5 x0 c+ O' y( u: w- O6 J9
      o* ?" j* M+ P9 Z10
    9 n8 b7 u8 b' K) E  S11
    - V. o/ j& S" w% F! \: r- q12
    8 [* |3 r; T9 d13
    ; m  X) |+ j/ F2 w+ I& v2 Y14
    & O* S% \% U7 n- }* {15
    " q+ s0 N  A/ x! {16
    ' b" n( B. w; N+ O' \17
    ( [; i4 {/ @7 T6 `8 j& D6 h18
    % M. ]7 F7 F* [0 o4 g19" X, H7 J# x# r7 a
    20
    , E. s- x" ^8 U6 n& G6 n- q2 F210 R+ u, a" L0 ^) j( V
    batch_size = 8
    , [) Q" B+ [: x: |+ ]  G5 Q1 cimage_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}6 v' S& Q- w) [; E3 I6 F) k0 w
    dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}
    # T$ d+ s# r/ ?: i) fdataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']}
    6 B, l- ?7 u) H- b3 \$ O8 |class_names = image_datasets['train'].classes
    / X7 }5 j& ]1 q( W0 R. q: f* ?8 [" c$ B8 Y2 W4 s
    #查看数据集合
    ( i3 O+ B! O+ g) J! wimage_datasets: Y7 y# j" f3 z9 u6 C: E
    6 V, R1 s, P/ g! C" D/ F
    1
    0 G8 @5 _2 u" B9 m5 W4 y4 Z23 Q# Q+ x9 \. t
    3
    # I. t( q  B' q6 A; r4 _4) C+ m5 {+ L0 T; i. ~9 G- |  j
    5: l* E8 W# h6 t4 C3 @0 {
    6
    3 R4 E( M2 U2 B7
      I) z  ^# L) o$ X7 P% M/ H; b# C2 U8
    : e7 p9 X+ J# h) s9- h" z# U: f- w! R4 p# y
    {'train': Dataset ImageFolder- Z" {. r& U9 h6 N, m3 e
         Number of datapoints: 6552' c. }+ E1 [% E  \& X2 u
         Root location: ./flower_data/train- d8 Y6 M2 S/ k) Z" y, D7 H: {
         StandardTransform
    % ~' [# Y* a; k! h2 B9 W- s Transform: Compose(
    ' @; w# ]' u, C. j: t& F- p8 F7 x; u                RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0); M; z& P# u% ]2 D3 S/ m1 q1 x9 H; w
                    CenterCrop(size=(224, 224))
    # Y3 q( f" S+ s2 t* o+ ~) j) l7 ]                RandomHorizontalFlip(p=0.5)
    3 b. y. q4 J" ?6 u- ?! J# |% P                RandomVerticalFlip(p=0.5)
    ) J9 q3 C1 r' w0 }: N                ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
    , O% t: h* p4 I8 C+ k) K$ X7 I                RandomGrayscale(p=0.025)
    6 J$ p1 X$ s$ x. E  G  M                ToTensor()( A/ Z; _3 m2 k* t. \6 r+ V
                    Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    : A1 u) m1 l) ^) e3 G  O7 I            ),& e3 y7 k0 o- v
    'valid': Dataset ImageFolder4 V+ E# U# W8 v( }1 G3 |1 r
         Number of datapoints: 818
    0 r/ ~: G8 s8 i, B8 n! W     Root location: ./flower_data/valid
    ; A% d- Q4 j7 ~. X* N! a     StandardTransform
    - O! G$ ^2 W% ] Transform: Compose(
    1 x/ i6 y+ F% T4 f                Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)$ O- T* V0 W; @& B( t
                    CenterCrop(size=(224, 224))
    % C& }( V+ R& j                ToTensor()
    , k& v  u  t- Y                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    3 Z7 j% R1 X6 O4 N0 k            )}' l- B/ @$ _* L6 y
    % v3 R4 D1 J9 L
    1
    ! q" W6 i6 D# m3 E7 \* C2
    3 z- u! L9 K+ y/ O2 q' t9 m7 n  l3
    ; n  m, {2 T5 N4
    ! @: S. K) f; w5' ^1 ^" _+ q) ~& t* _: I
    6
    ! K7 L% i$ r  I7* Y) p; c  I4 {# U
    8
    3 m6 x5 P: {8 @9* y5 s  |& M* S$ v: p0 h
    10
    $ e' w1 u# k3 a; b11
    8 G4 B( X8 H) e# N( m. Q12
    ' Q' `( n* V* _# _- h/ j' i$ n13' K  q& l4 `8 g, g0 k4 B. j* v* s% u
    14
    # b5 t4 B5 ~+ s7 l( A4 f15
    & x" ]$ v5 K  z# H7 K' h1 d4 k* X* e4 k16) x% a) ~, G% t/ e3 X" @
    17
    $ O3 X* A- H2 v" K( Y3 d18
    7 w- U+ K6 X8 R: @$ X+ y& V195 k+ e% Z3 e( p
    201 @, _" P; Q1 b
    21( n$ m3 F" W5 ^4 W, S0 D
    22
    + M* W$ a# B9 {( n7 S9 @237 ~/ a- ^# W) `* Y
    24
    4 i; Y4 r3 H( W3 `2 G# 验证一下数据是否已经被处理完毕6 W/ r6 l+ Q3 s5 a3 M
    dataloaders8 X( F7 V1 P" i0 p; ?- X$ }
    1
    / O, O6 d% D5 [2
    . g# }" t3 o4 [1 |1 |  \8 C; a: M{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,+ k5 W# q) U% D# J0 {# f6 C* Z
    'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}
    * J  t9 k0 A3 a* U) R0 x  v1
    9 ]3 D  A6 g, M1 s2- c6 q0 G9 r" o" m- O1 W
    dataset_sizes
    0 y- v0 v' t2 c15 |1 `- N7 ^8 x. m0 {; u
    {'train': 6552, 'valid': 818}
    9 g4 @4 L2 H) p2 U: I2 m1
    % M& h4 [2 G% f. }6 x5 f读取标签对应的实际名字
    " X7 b% d8 o( Y2 p2 D2 w* F9 F使用同一目录下的json文件,反向映射出花对应的名字
    + z, c, b/ k+ B4 E8 F1 `" S7 l: `" r/ X
    with open('./flower_data/cat_to_name.json', 'r') as f:
    ; b9 W- i$ ]8 ~    cat_to_name = json.load(f)- U6 S8 j7 ~4 l! B/ q7 m' _
    1
    + \% Q' N6 P9 [  N( R2
    / f6 J+ o% ?. z. v' scat_to_name
    # J- U  W0 l+ `6 ^, `1" T4 g- @# O/ G* h3 `2 {- K0 B4 A
    {'21': 'fire lily',0 U4 n4 L% e/ ?' f7 ]
    '3': 'canterbury bells',. i, ~" D* n% J3 j5 Q
    '45': 'bolero deep blue',5 O% `$ W$ d/ m/ \% Y
    '1': 'pink primrose',+ b% J+ R* S/ b
    '34': 'mexican aster',  C" z; n$ @* C7 L* X2 O# j
    '27': 'prince of wales feathers',
    8 r, _- Y5 n% w& b( z '7': 'moon orchid',! ]- m5 F1 g3 k0 |5 f
    '16': 'globe-flower',' m0 k. s% b! u8 s; n
    '25': 'grape hyacinth',
    5 @  \/ e& X. w# r' b- K& V/ X '26': 'corn poppy',# B- A6 x: T' J7 J+ P( Y
    '79': 'toad lily',% \- T  ?6 v9 j' `! @/ [( N- O
    '39': 'siam tulip',' X. u$ ^; |6 y& V
    '24': 'red ginger',
    # a0 R. ~) Q. { '67': 'spring crocus',
    - v, x2 V1 ]- k  V& y0 s+ K, R '35': 'alpine sea holly',! a3 ^6 G0 }1 A
    '32': 'garden phlox',# G! E1 {3 X- {# Y# ?) c
    '10': 'globe thistle',
    : C% [8 @. i% r2 c& n. I. X '6': 'tiger lily',
    + c9 I; c9 z; y/ }) F6 h4 X0 [ '93': 'ball moss',6 S& R, W; K7 {* `
    '33': 'love in the mist',$ G+ a8 ?* U- ?) I3 ?5 j) U' R( p
    '9': 'monkshood',
    4 M/ O2 V# S' j" o+ `. f9 k! N, m '102': 'blackberry lily',- Z  m* F' l: \1 g4 z9 P5 r- e  P* k
    '14': 'spear thistle',% o+ D" w& g  @0 f" u3 ]. Y' T4 z
    '19': 'balloon flower',
    , j; R8 I" P1 z+ S, I. y '100': 'blanket flower',& N1 b5 K- j* c( r6 z4 o, ]& e, `" s
    '13': 'king protea',
    7 F/ n+ s* d7 N7 ~, H* V5 `8 A* } '49': 'oxeye daisy',1 ?6 x- z. m  X( n, x
    '15': 'yellow iris',
    1 Z; V6 e/ K: f '61': 'cautleya spicata',
      m* E+ _. T& G; M, e '31': 'carnation',
    + u* D0 W& o  o4 g) z; u0 _  p$ \ '64': 'silverbush',
    , d( R0 D1 X+ Q( |9 z# i" a2 L; Y1 ] '68': 'bearded iris',
    $ I+ c" d# z' C '63': 'black-eyed susan',
      f2 R, g7 N& M# Z '69': 'windflower',
    6 G+ k6 G7 j0 K# ~ '62': 'japanese anemone',
    1 N* t' ?+ @. d+ v5 P '20': 'giant white arum lily',
    7 c4 ~5 T& J( i" R0 B3 b5 l* L/ [ '38': 'great masterwort',5 d1 a. B1 |0 a+ l1 J/ l
    '4': 'sweet pea',8 J1 E9 R9 j! R
    '86': 'tree mallow',
    ( L$ y. |; r0 o+ S" W6 T '101': 'trumpet creeper',
    " B0 c& e; \8 `- o/ u4 C. t- h# u5 Z '42': 'daffodil',
    4 n$ s  L1 W% z) f. {( A1 C" S1 G! h '22': 'pincushion flower',
    * h4 A  |4 Q& r8 _ '2': 'hard-leaved pocket orchid',& H8 L7 K' l8 ]& e! o2 y! x0 A1 _
    '54': 'sunflower',
    6 \+ v& ~" U) B% f; B( S '66': 'osteospermum',
    7 a  W6 p: s; Q8 K$ r. k '70': 'tree poppy',
    . v3 J4 z) i6 r& i& c' @6 ? '85': 'desert-rose',
    9 m- \* D+ ]" w) r+ E- x8 v$ e- C '99': 'bromelia',9 }3 m" q! G; t- J: E
    '87': 'magnolia',5 [* N4 r0 [* b  I3 t' ]! a, ]  k
    '5': 'english marigold',
    9 F, I9 }% l! j '92': 'bee balm',9 m) F+ q7 g7 @/ K; W( F
    '28': 'stemless gentian',, l. g8 f7 m; F; ^" C
    '97': 'mallow',+ _0 X6 q6 G/ ^1 Y# X
    '57': 'gaura',
    ( j. X2 a# d, z0 V0 Q! X# N '40': 'lenten rose',6 ^/ P5 B: m# _2 q
    '47': 'marigold',, z( b, X/ w& a% |* S
    '59': 'orange dahlia',
    ; @+ f) E" L% A '48': 'buttercup',* p7 e" W) ]* b' T& A4 ^5 i
    '55': 'pelargonium',
    : J; @. C6 a8 z3 c' j9 D6 T '36': 'ruby-lipped cattleya',& |) V' B' j( Y* Q  ~
    '91': 'hippeastrum',
    ) U  w# V5 ^3 p, y '29': 'artichoke',
    9 p" N/ |# \' q! I$ I '71': 'gazania',  R+ f8 O( c+ c  n, b
    '90': 'canna lily',
    + `/ J& |9 Q8 p  `/ N. M+ }/ ] '18': 'peruvian lily'," r4 O2 z8 ?5 H+ _
    '98': 'mexican petunia',* C8 K$ @$ N* Z- J
    '8': 'bird of paradise',
    3 Z1 v+ e" [( @& D2 {* m '30': 'sweet william',, [; |5 N5 ~9 k! Z' O
    '17': 'purple coneflower',
    ; M) G- |/ j& R/ A6 r '52': 'wild pansy',
    3 n+ Z; V: D! T! e. c3 X6 } '84': 'columbine',1 q! ?: Z8 `6 ?+ U2 M  n: L
    '12': "colt's foot",
    % U( N1 C. Q$ [, g$ M4 K '11': 'snapdragon',
    0 i7 B; M' E) ?% D: | '96': 'camellia',/ `7 o; U5 j# f  e; F  S) {
    '23': 'fritillary'," H  O! _' {; G  l' y( I( U; j: @+ T
    '50': 'common dandelion',
    0 M- E6 c) T7 p+ H '44': 'poinsettia',
    8 F! s0 T0 l' l  X* s  d! A '53': 'primula',1 O, |0 G) m! }# ]% y/ I8 F
    '72': 'azalea',9 t; J5 p/ i: t& C+ N& `* j+ c% X
    '65': 'californian poppy',
    ) B0 J8 E) h. D$ b9 l/ ?/ J+ d8 ~: v- A '80': 'anthurium',) c5 J4 Z4 `) Y6 r; `" a4 k( ?
    '76': 'morning glory',8 E/ p  e# ^; b
    '37': 'cape flower',7 s1 h$ J  u) ]& X0 f/ L
    '56': 'bishop of llandaff',
    ; D1 {; G' C1 h '60': 'pink-yellow dahlia',- h5 L4 i7 h+ {* A* M7 ]
    '82': 'clematis',
    3 `- A) k5 V% b) b8 m0 G1 O '58': 'geranium',) X5 \# g+ S% M5 W$ S1 j* D- R
    '75': 'thorn apple',
    3 j0 i3 {( E; \ '41': 'barbeton daisy',6 @3 w- K% E& l! `2 @! l! [
    '95': 'bougainvillea',
    ! {2 f" t! n$ }1 O& ~8 E# G '43': 'sword lily',
    " ?, S* o9 `! R% O9 E8 f '83': 'hibiscus',- D9 {0 c6 F3 d; y' _( m' P+ ~* S
    '78': 'lotus lotus',
    . e' ?1 y1 V4 Z$ v6 Z2 b, ~ '88': 'cyclamen',
    + _$ B. T: M- L" o$ B( _  d '94': 'foxglove',7 Y; B* y# p  {7 `4 e, F
    '81': 'frangipani',
    0 P" R1 H- P9 h5 ]" I. M '74': 'rose',, b/ {  S- ?0 j: ]
    '89': 'watercress',
    , U; f! s% a5 ~- }0 z6 s8 f '73': 'water lily',+ ?9 I5 z7 z2 \' s
    '46': 'wallflower',
    " @7 g# |% c3 y '77': 'passion flower'," t4 d6 ]0 e9 [2 {- {# G& R
    '51': 'petunia'}7 }3 G' ]. b. w) B3 ?

    1 I( J5 d$ B: v( a12 J7 [& U: ~, {% ^' I
    22 E2 b# X8 [' o
    3
      c$ B9 m& }! U5 M1 R. a4) Q2 U) z  v, ]! X& y, y9 l
    5- L4 h/ i4 }. @9 J% n, ]  d9 ?
    6
    * N; I* [$ ]7 f6 L78 C+ ?3 i* ?4 T5 u3 O  b( n1 K
    8
    1 P, W# S# ^6 v3 R0 b/ L1 c5 ^9
      _6 {& S3 f8 z, `' e/ z10
    ! }: ^, _3 ~" E( r. g11% o9 E3 k- b' w6 d
    12- l; ^0 i, C* G7 I) M% L
    138 L( ^, k% k- F: J+ Y- K$ e
    148 X( H) ~, }; h) y& W4 _
    159 S* H) o; {* }( K. t" A
    160 _. m7 _; a8 Z; L0 F8 h( \( m
    17
    , e* F+ m4 O& z18- Q3 ^1 e* n8 P" E8 l
    195 U, b- U5 U/ v9 c* L+ x/ V
    20
    & ?3 h' q# }# \, @; Z  v21
    % r4 _+ K/ `2 r! A/ V( }22( a' E$ ]$ ]6 }0 B) d' u- u! o
    23; y& I3 _/ z: O+ R, C$ J# D
    249 k: e. m6 P: ^5 }3 M% n
    25& \' {- M+ s. Y, p1 g
    26* x' q3 `6 G. u+ Z  b9 `" L( F# X
    27
    * Q4 a3 T3 t& J$ k1 k28
    4 s) l+ x- |4 f  I: A29
    8 P1 B; ?9 \  H. G* R30; y2 z7 |8 K. Z& q: I6 `
    31+ e$ `! r/ }7 c$ P
    32
    & ^: v) Z. S& V- @9 d33
    3 }( F+ k2 G& F; i: U; D34
    / H( q# z: z+ @! ^& J8 X& G358 o1 }% C0 L& p9 \( n/ s
    36& F* B7 G4 E  I* N, h; v
    37% s% \" ^; e: y/ m1 {* M5 ?
    38
      Z% I% i1 }0 {& v1 O39
    2 u) t6 C/ d  P) _40
      l8 K. n4 B% ~8 R41
    9 z+ e4 D7 w! s9 ^7 t% J: ~42! a6 f' H9 ?& q
    43
    1 Y* Q2 k! c: w/ V2 ^: W445 s7 }9 b! |8 Q& z0 y
    45
    . N+ u4 @) q$ Y6 Q462 \$ X1 G2 G  |+ E! h0 z. B
    474 t, z5 w" T# o" [) V# d
    48
    - @# a! D% a0 F! b/ B  v  }49( g( x  e6 ]; Z' M# z
    500 A" x0 @1 p3 l" d, B
    513 M- J0 a6 T0 u
    52& I" R2 `1 W0 ^2 ]% H
    53( x- N0 N  g! C$ ]0 j$ z6 ?
    54# d! Z5 `7 j' I7 t* e  `" q& C
    55- T  p0 W# S% e, ~
    56+ E" j& g, K- m' k# e9 j
    57
    - f1 s* u2 H" }1 C/ ?587 E$ F7 |9 f9 ]/ r: @. d
    59+ G  R# _; d: R/ K$ v3 {
    60
    ' w3 b2 W) [7 Y) ^616 O% K/ ^+ {# x# u/ w9 P' h
    62( L2 _! [5 h. Q9 j6 w
    63; U- q  W1 ~! o" P
    64
    4 s) g) R; M  \4 |) z/ l5 [5 J65
    ' X0 K( a4 D- Y' C2 `( @66' F3 s; C' Z; Y6 N1 H6 |" y
    67
    : h5 D: q' ~: z68! k; T% n1 A. X2 H
    69
    & f9 b4 t' L+ t; F* W70
      ~; m. M, z" i  J5 N71/ @2 r& \4 ]' P; G  F3 I, j! O( k4 D
    72& x8 Q1 A# J  E  [/ Q  L
    73% m; q+ y, l8 c; L6 z4 Y
    749 c5 ]8 o" p9 w* z
    75- m; j( \& _8 r6 {
    76
    . ^  G4 G( o+ y77
    : c3 ]. K6 g5 Z78
    ! s9 `$ q3 D3 ?0 _79
    : l( Y9 D- I& Z3 x% s80( T# T7 R3 y  Y: q+ H3 ?
    81% T7 r7 E) S, q! f) q- i
    823 h8 t+ g! ^2 v4 j& m+ O
    83; K0 Y) q) k6 Z- a5 \, Y' A% w
    84
    5 Y# v# `, E2 \% ~+ z85/ }# V% l, c9 R( z# L& c2 U
    86
    ' D2 {% W+ U( @/ K4 s1 Y87
    / x- K# f9 l( Y: a1 h5 Q3 S88
    ! a3 C) I" f# E1 }0 J, K  t0 U1 `89: r# h$ o1 t' D$ J
    90
    ) K! B; p( R5 E) L91
    / Q" [& h* \2 s! K5 b6 D/ l92
    6 n% E& x+ Y. ?3 w3 U  u" I0 d) W93
      _6 H9 c# K0 p# O6 l& {, t$ h94. v- G( p9 Y6 n/ ~9 ]) I; J6 G
    95  a; r0 P+ a/ R
    96& v/ [5 j) q; r2 I# ^* j2 ]) a/ @
    97
    1 r. G+ w; m/ v- E2 X) p# M98
    - V. }6 }+ e2 _3 v4 q1 K0 q99$ U& B6 L8 V# O' f
    100
    1 s+ X# o+ X1 k! H; b8 u  }101
    ) y8 J7 \- k8 l: {" K$ K! Q102
    ) r" F: z; f7 C9 f2 n3 |4.展示一下数据
    . ?( U' X! R+ y2 {def im_convert(tensor):
    * G8 T0 w$ e1 d# E6 U    """数据展示"""
    ( _; g; T3 K* ]8 o    image = tensor.to("cpu").clone().detach()
    / R4 b& w5 r. i' E+ C6 P    image = image.numpy().squeeze()" N9 G" D# P! F" F- Z& u3 ^
        # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图
    $ K: F7 C' F$ N' D. e2 K2 _# Q5 g    # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)' v" U6 x0 w% t7 ~( k  C
        image = image.transpose(1, 2, 0)
    + ?1 C( ?% k1 Q( m% w* S& a    # 反正则化(反标准化)- d1 C# k6 {  u1 j" y1 u& V" S
        image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))& k3 f" S7 x' d) s& n6 t! a

    & o: M6 M. u4 j$ B    # 将图像中小于0 的都换成0,大于的都变成15 `8 O& G. ?( Q
        image = image.clip(0, 1)9 P6 c6 E3 D8 A6 u

    ! I' V! h  j0 ]# Z    return image9 W7 p" ?% k1 x* C3 n7 N
    1
    3 h4 K4 K: C5 n# L7 q4 n* p9 H2
    $ [/ X5 K6 s6 e5 g. ]' l/ i/ S3; [3 |- m0 m7 ]
    4
    " y& l+ t8 F4 Y" u, x; h4 A& T53 w# Y% E, ~8 R
    6# O; K* W  e# ?6 y; R6 I  _
    7
    6 q# O6 B! z/ L5 N; t+ J9 L& z8- h( T( R7 `" ~) |) R- g
    9
    9 v. s, D2 Z+ f8 J& w10
    ; H) @  I/ N9 K: n( n; X" ^) v11
    ; B# ~! p3 z# z0 ~5 f3 J12
    $ W% m7 A- g3 A( j, h$ f; `7 I13
    0 [( r" r6 G% l14
    . Z. I5 `4 |. H6 a5 o* V; G- V" k# 使用上面定义好的类进行画图5 |1 f- w# {8 \2 o; ?+ F/ S- \, K5 X
    fig = plt.figure(figsize = (20, 12))
    ( L) G# V( o( ~columns = 4- z# G  L( G7 L
    rows = 2
    1 I0 G: w2 z& D/ U
    ; g6 y# Y  i: w0 N! m! g, X% J# iter迭代器
    : Q7 `2 l3 Q% ~( a# 随便找一个Batch数据进行展示
    - }& s" A5 Q" W. idataiter = iter(dataloaders['valid']), O' Q; }; o; {- x/ f! v, P
    inputs, classes = dataiter.next()
    3 z! E* H" o# ]1 G) K# K+ S5 c& Z+ T- m3 M* J
    for idx in range(columns * rows):
    3 P& k1 g0 r$ q+ m( H" p    ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])
    " B/ k+ b, q+ K4 f+ K; ~# S    # 利用json文件将其对应花的类型打印在图片中6 M3 e' e$ K- o6 U+ H) ~; f1 |
        ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])
    1 a2 g# U7 e- C: A0 j    plt.imshow(im_convert(inputs[idx]))3 q/ A5 G2 h: E! `& q
    plt.show()8 Z, }2 q9 B1 r( n! ~1 _' D
    & m- x' f$ I- j* W
    1
    & x0 U/ H5 n; O5 ^2 @. ~! z* t2$ J- a% c9 a! f& Q2 @0 Z8 y* {6 J4 |
    30 O7 d  {( S6 x* k- z; S" H, G5 E
    4# P* M$ |1 B" n$ l3 a* Z: @
    5
      m; `  {, |* }( h6 M- T8 K6( c" ]1 X( ], i/ b! ^0 D
    7% z/ Y' G+ H1 n1 C* d3 S
    8
    7 K+ S. X- K# u% l( y91 f. K3 u3 c$ U' @, X# m0 U# x1 i
    10
    0 f0 A  v( m$ m' H* Q' a11
    ! X) y( p) _6 u3 l12. c: w- i* j' ~) I/ K
    13
      U) O4 F5 y) H- j( i14
    9 {! G- I) e) v- T( I15
    . x) z* m7 q) G% _" R( a9 }* |16
    " M6 G( f9 z" G8 |) Z8 a( H0 `* x  ?( x$ g
    0 v0 q6 S' T( A9 h* T1 A( {
    5. 加载models提供的模型,并直接用训练好的权重做初始化参数/ i* H- {( n2 b1 Y8 `2 R7 A
    model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
    ! ~1 {$ w, V' _- ]8 l# 主要的图像识别用resnet来做
    " A* S) T! U  O! s  v1 j( M# 是否用人家训练好的特征4 O+ G2 z& a+ q: f- u" C
    feature_extract = True
    1 x; ^1 D% S) ?1 g1 P# ]1
    ' \$ L6 n* A+ x# D# V, |2
    ' R' R. ~$ {$ u) i3+ j' @/ w  o" ~2 w. `; Y& r
    4" {# T- B: R( K8 p3 O1 i: T
    # 是否用GPU进行训练
    0 h! j; u+ K( k( @0 utrain_on_gpu = torch.cuda.is_available()8 b6 p! i3 C2 A" J; l. I' v. E

    ( M4 [1 \/ E* a  _. `5 Iif not train_on_gpu:
    5 h! ]% ^8 `# R" W3 c4 t8 m    print('CUDA is not available.   Training on CPU ...')1 j; i( w4 E% D6 I! P
    else:1 f% Z  I& ^3 u" ~$ J
        print('CUDA is available! Training on GPU ...')
    2 v& j# x9 e4 s: D. ~
    9 Q" ?& q% K$ c* xdevice = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')" S/ P7 t7 C# h
    1  C; ?- E) y6 M+ P" W
    2- H( o3 w& j% B% l" _! X0 s8 l
    3
    % N  q6 |$ _6 R, u  @* J6 Z48 U  x1 e; T8 J$ y" N" i9 Z
    5
    * B4 k; r% }  y7 [$ {/ h4 P$ B67 V/ m  }+ x4 a, K) y2 L5 ]% p
    7
    ; T0 Y1 y! k: u& |6 u, E+ y9 V, h8. u9 F! i0 A  s* e0 @1 G0 Q
    97 M# B1 @2 R/ S% a5 W. @( k
    CUDA is not available.   Training on CPU ...
    , @' r9 {: A+ Y3 y1+ s: w( A! X$ |% _# Y8 I* z
    # 将一些层定义为false,使其不自动更新
    8 ], h* Q9 H2 m7 H6 |0 edef set_parameter_requires_grad(model, feature_extracting):) @4 b( o" I8 l/ k! X1 `! @
        if feature_extracting:
    0 V: c. s0 ?1 d' l& h& u3 U        for param in model.parameters():9 T) }+ k# l$ T
                param.requires_grad = False8 S2 [" O# A* t! j. |) u! p
    1, e+ \6 Z  Y* k6 Y5 _  d" f
    2
    7 b% k  `* C& b* E# S' ?" V. u3$ ^8 N+ t" @; C. G
    4" K9 C* V- P, v
    50 o) P/ E$ a% f" x
    # 打印模型架构告知是怎么一步一步去完成的
    % M7 h7 @7 @& a7 D+ e# 主要是为我们提取特征的4 H, ~- E$ c% y9 W$ H0 f

    ! Y, K$ r$ S  m8 C6 S" J: cmodel_ft = models.resnet152()& B$ N9 ~! b% B, _
    model_ft# h/ P$ l* m+ S7 m) n3 h
    13 W0 F$ E1 a, Z& }4 S4 w5 \
    2& }* ^2 J9 X7 ]& t2 Y2 t- \8 \( T
    3
    / c$ j6 W  r2 v1 D  n3 }7 r2 x4
    / T" D# Z  Y! Y- a$ U3 A3 f" x5! O6 M5 n8 G* ~. _: V
    ResNet(
    9 l3 i$ |" s. r' t& W  (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
    : ]: O  ~, x3 {! c7 i$ d5 V! a  (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)2 a- }# @/ e, U" G: w3 b4 t; M4 J" i
      (relu): ReLU(inplace=True)4 Z% z3 r2 m- n1 }
      (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
    . i# a( o; C0 P$ K  (layer1): Sequential(3 M, j8 w" M& s' _
        (0): Bottleneck(2 i3 d1 Y$ ?  [
          (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)6 v- v5 Y! U. v! }9 D$ ~
          (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)  G9 k: G1 P, B
          (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)4 Y  o, y  `7 ~
          (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)+ F+ V, Q3 d% O2 H2 c& S- X
          (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
    9 D; L& i+ P  a% y2 Z$ w+ k- x) ~9 u      (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)5 V; S, x! {" H2 s- d- [. J8 M# ?: I
          (relu): ReLU(inplace=True)
    - S' O( V0 ~! {! z$ E  ]      (downsample): Sequential(7 M  i) l' p' N6 S# P8 ]
            (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False). J* k0 V& L/ d# A% ?
            (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    5 p2 b) ~7 u: O# Y7 p. j      )2 Y9 G- Q6 y8 C$ U9 i( p
        )
    . X. G+ `: P* D/ X. m1 Z# i) @中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。
    4 H4 |6 T, F& ~) J: R% u    (2): Bottleneck(
    2 B. S0 U7 R5 W; d      (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)+ w" w# L. l4 i, B$ D1 x
          (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)% S, m9 n4 y8 |
          (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)$ Z* i% ]! p& ]7 q5 s
          (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)' ]' W+ \4 Q/ D8 v
          (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
    ' _3 m% B  v' {0 ?. b! l$ ?      (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    1 d' j, s' \! x; G# _      (relu): ReLU(inplace=True), Z) R9 z  m9 W2 n# u% {! _
        )
    / G3 R2 U0 M  U  H. _  )0 U! q' K' ~, j5 m- \3 a
      (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))2 S4 r' |+ H3 H
      (fc): Linear(in_features=2048, out_features=1000, bias=True)
    $ ~5 j# _0 h" r& v+ r)
    / F2 h( @- ]3 D5 S
      d8 F5 v* J$ r* y4 {  \# X( \1
    / F8 L, g+ f+ {7 C0 \24 ?4 u2 P' C2 ~. z
    3$ a, e- m) Y3 t
    4& C2 a2 A1 T3 M' N8 s0 m3 D
    54 h% W) r0 ~1 w" Q. `" w
    6
    # O* g% J  M2 r. |7 Z0 m7
    / z$ n$ k& j7 c: Z' N84 ?7 ]6 p' x1 f* N: p( g1 ?. D- @
    9$ s  A+ w( `6 ]
    10" J' x( i" D( y
    117 c4 b/ B/ \- Y& W7 p8 H" h, |
    12# z; B/ v8 ^, M! V  L7 e
    13
    7 m+ T" G" L. c0 D148 ?% N( G9 M9 Z2 j
    15- C% ?  Q" b+ V9 u
    16
    1 `( d. v" B: v% P7 g  ]) C9 E% A17
    0 x, H( K. }* [" p1 q187 q  n* y% }6 c" N. y( d4 O
    19
    1 `+ L$ a7 ^7 X: C9 a20
    ( O+ X& z3 A$ C$ `! n* u21
    & i7 q: I, J- d9 D9 g# f221 M9 O+ @( j+ R5 k! W
    23
    ; n: Y1 O2 ^% u7 _- P24- m( o7 p: K8 h6 ]- u( i: i
    25+ K% v3 D8 t- O" N  T9 C5 ^
    26
    8 w% x1 Q8 o! q" @$ q0 c277 U2 x# p9 t3 q$ |" s1 y8 b/ S
    28; o! h# i7 v  }: G7 s/ @
    29
    1 Q4 o8 ]7 L1 g& u7 f30/ g  s( V/ k6 T# o+ r; h  J& ]
    31' f$ a+ e/ D- |  L5 A  Q2 ^" G3 S
    32* H9 Y8 R! g5 y; m; t3 M) a
    33
    + x, Z: Q: s/ X: Q6 x8 l: i9 L最后是1000分类,2048输入,分为1000个分类3 q  O% G% x* P4 P
    而我们需要将我们的任务进行调整,将1000分类改为102输出
    / ]7 J  r" J% C  G) v: M# r
    # x# U) W4 f" x! U6.初始化模型架构5 \1 e, `* q" z" l
    步骤如下:
    : v- V3 ^$ x8 q' K$ M3 W3 M- _: x, T
    将训练好的模型拿过来,并pre_train = True 得到他人的权重参数
    + g8 m, l% U3 X! ~0 t7 w可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)( W1 i: M, [9 }' |. v6 a6 `
    无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数  @: F, u& c( z+ K' H
    官方文档链接
    : Q/ l$ j  C) N7 f) ^4 `4 B& ], Dhttps://pytorch.org/vision/stable/models.html  D& k8 G7 T4 V$ T
    ( V5 Z6 K: [2 x* }
    # 将他人的模型加载进来
    0 M0 R% ~5 H. ]1 ^) M# e# Udef initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):
    9 ^! C; E" Q/ W+ \4 I    # 选择适合的模型,不同的模型初始化参数不同2 E. q6 R6 e. n
        model_ft = None
    : g% N- R/ {! g# g, s1 v+ V, S$ i    input_size = 0
    ! N6 [% L' `' j% M# ], v1 w' u/ P
        if model_name == "resnet":/ i1 V9 v/ v! d7 e7 @0 I! H
            """% d8 R' b% W( D: Q9 v" ?$ y
            Resnet152
    ( H& ~6 Q0 t! N        """/ F6 Y% F3 Y8 R* T& K& P) ?6 P
    # ^' N! h- A5 v  Z
            # 1. 加载与训练网络
    5 q8 ]7 x- U  G8 S        model_ft = models.resnet152(pretrained = use_pretrained)
    $ X6 L7 p/ Z0 u+ @' s4 ?        # 2. 是否将提取特征的模块冻住,只训练FC层) i$ y; p- s5 Y. m' I
            set_parameter_requires_grad(model_ft, feature_extract)0 i8 K. D$ {1 x9 J. c; w0 ^) _8 U
            # 3. 获得全连接层输入特征
    / H# ~) Q5 b+ c; Z        num_frts = model_ft.fc.in_features
    3 q% M1 O4 T. `: B, s        # 4. 重新加载全连接层,设置输出102: z# p2 F% k/ a7 b$ ^5 o
            model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),, G# H  u- m! O/ b
                                       nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1
    8 R1 a/ S$ L9 M' w9 {        input_size = 224) y! w5 T$ Q8 [9 q, J! n& @0 y" F

    3 u) b/ U7 e9 f" L: f    elif model_name == "alexnet":
    % M; F7 S' R4 a  T        """
    7 l8 f% S9 B0 s, I6 b0 ^/ Q* V        Alexnet
    ! Q8 {' Q7 Z: ~% L2 K$ t( ~3 ^        """/ T  b: n* ]! L5 f
            model_ft = models.alexnet(pretrained = use_pretrained)1 k/ o# \* d  E$ p2 C2 E5 a
            set_parameter_requires_grad(model_ft, feature_extract)
    7 @3 x& |% V. H9 S  T- h* Z/ b+ h* g, V/ M
            # 将最后一个特征输出替换 序号为【6】的分类器2 H- u6 y9 t$ i% |
            num_frts = model_ft.classifier[6].in_features # 获得FC层输入- I9 K4 |' N  k! S, i/ _
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes): p. z. \* j& ^4 }2 F& }; w
            input_size = 2249 r& p0 S2 J& p4 _7 Q( Z- g; t, x  g% D

    . T% H) C& s, S* V% m' p    elif model_name == "vgg":7 P; C: l: |2 m* `8 N6 A
            """$ e+ T9 ?7 @0 [  u! |" _
            VGG11_bn6 @+ [- X4 u' ~' S
            """3 E7 ?% J1 n8 s, y6 }9 G
            model_ft = models.vgg16(pretrained = use_pretrained)
    - ]% }+ o& E2 G6 y+ o7 q        set_parameter_requires_grad(model_ft, feature_extract)
    7 H3 z& ^3 t* W, ]+ W' H: O        num_frts = model_ft.classifier[6].in_features/ }$ W, S1 T0 g# [3 s2 c. L
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes)
    8 P5 t$ h7 H1 x7 n        input_size = 224
    8 ^" f% @% U. o5 O
    8 t7 I# v' p9 B+ n" G' t; W: V    elif model_name == "squeezenet":1 E" N% }5 n  H+ p
            """! x- _) A* x: A+ {8 |
            Squeezenet4 f$ C( R0 D. ~. E0 z! ]- e
            """
    ; F+ ?. j/ u0 G+ E        model_ft = models.squeezenet1_0(pretrained = use_pretrained)
    4 ^5 j2 m9 p9 H3 G6 z- w        set_parameter_requires_grad(model_ft, feature_extract)
    $ y. z0 U( T- O& Y        model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
    , ~$ i( a3 b5 U' L* k        model_ft.num_classes = num_classes
    # l5 l& ^# I3 @4 |5 e  {$ ^        input_size = 224
    0 x: i2 n& H% G9 k6 {% r
    - w; w' ~" f7 }3 S% d    elif model_name == "densenet":$ K- i7 q$ N' e6 n( q
            """
    : a8 ^1 x% D8 `- i, g% h        Densenet
    ( v7 X: K: Q/ c2 @6 x. d& h        """
    / R* V* j% m) p' e; l  j+ ?/ ?        model_ft = models.desenet121(pretrained = use_pretrained)
    + j1 f7 t  H5 o4 F- V8 M/ O        set_parameter_requires_grad(model_ft, feature_extract)! ~( z1 t, f, P: `4 j$ y7 f
            num_frts = model_ft.classifier.in_features
      }) l  u& ?) Z, U1 w8 {        model_ft.classifier = nn.Linear(num_frts, num_classes)
    ! X) q: [( l3 {+ @4 K        input_size = 2243 v; @4 N5 T% V" _# B
    : u/ x$ S9 P$ v% O, a# d, |$ t6 M
        elif model_name == "inception":; b1 X1 X$ ~9 e: X8 `
            """
    7 q' N4 ?& @+ n! n/ L9 R! {        Inception V3
    ( q% V, k; w/ t  ?3 N        """
    4 `/ l5 K' X* ?0 g. Y  o        model_ft = models.inception_V(pretrained = use_pretrained)" j9 a& S9 q; q- E9 b) y$ y% l/ X
            set_parameter_requires_grad(model_ft, feature_extract)' N+ \$ h# h+ ^
    - P0 D9 o4 P8 ~9 y6 D+ |5 ]
            num_frts = model_ft.AuxLogits.fc.in_features2 a6 x; b& p4 Q
            model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)2 j2 @2 q, a1 \

    " N9 E( G, e2 q$ U        num_frts = model_ft.fc.in_features
    ( ^/ o- [, i: j+ i, s        model_ft.fc = nn.Linear(num_frts, num_classes)- u6 b1 p6 v) M* i8 }' s& z
            input_size = 299
    , H' \7 H) z; c9 ~( B  ?- O
    , F0 B8 b8 N4 O7 ?2 K+ n/ N    else:$ S2 Q" A& d3 |/ q0 x
            print("Invalid model name, exiting...")
    $ V) R! t8 b1 Y$ R) F, J7 s! s        exit()
    % {- R6 ]# V( `% h5 q, A/ o, J& Y4 x! I7 j. X& w
        return model_ft, input_size
    0 n4 @/ v: P3 ?# W% l( T" A: [* r) k* k" ?% k3 C
    1
      D' h8 D6 N! y0 |' p' K2
    ! j: s( M: a6 D3 Q6 c8 W32 `' C5 F5 p" y6 Y1 |4 C# }
    48 R! ~( n7 ~7 |1 G# o' V& j
    5+ R6 i! a# c8 h
    63 |1 F, j  F+ {8 Q5 j. J
    7
    . r( ?6 e# I) h( T+ W9 c8
    $ l' f% d$ i4 W, Q; n; _9
    * ]* U) ^. }; T( u1 w& K2 ~, M& I5 ?10
    ! X4 s' H" J* \, C11
    8 o9 I) y9 R4 q6 J, |1 v) v12
    6 v' `  X( X$ l; S132 G+ D( t, E  M$ L
    14
    1 b' S! M( G' }" l' |; ?& c# W15% t# L3 @) C2 O# S
    16
    * @& D' I# s1 {8 R17
    - |1 ~0 t* R8 [. p+ S5 q! _2 j. o182 _9 d) r: E+ s
    199 ?2 ~7 b" o7 N1 Y
    201 g: T: s$ u* ?2 X+ e
    21
    7 }% P9 ~* ?1 n+ i4 W( h' y22, `; }3 v; A3 Y+ s  y: t
    23
    & c* p( r7 S1 {' n1 z  i24
    2 T* j& F5 _- S, U2 G& b25
    / t" c7 |& h% }2 I26
    5 i- A$ Y% s8 F4 q4 D6 @3 l27
    - q# ?+ B' i* E) p3 W* {; `28: [4 O8 o; A- @9 u# c
    293 ?" v7 C9 w9 U5 ]" d2 @
    30
    : Y/ r! z* u# F3 y1 y31
    ( A/ x  P# e- s1 ^; D; S32# b( f; y$ U; v
    33
    / ]% M% V2 p0 n' G6 N; `5 t1 t347 ]" [6 l: v. d' S
    35, N0 y+ K+ z  q6 r/ ]1 I
    36+ G3 {; o, o, l7 \" r" I7 P* A3 ^, f
    37' u2 O' V- j0 X; M; K1 C: B
    38
    - k. m" I9 o% l9 V& z39! R3 v9 l/ W7 T7 h' n
    40
    & T* h+ @) x" C2 w/ Q  b3 d+ X" K41, J5 g4 E: s: K: w6 {& u& M) G( ]9 O
    42
      H7 W( h' S# @! c/ ^8 _% ]; s! G43* S3 @! _  `, R1 S5 B8 r9 w0 L
    440 Y* j2 H2 R5 m/ a% p+ b
    451 Q& @% e! f5 @! w/ e$ N
    46' d4 k1 R5 }* \2 c; q+ M( @$ i
    47
    + ^, E% X: o6 |* [4 ?) a48
    + L, m5 R: }" K0 z7 z8 z1 \; X49
    9 X; L, J  J; E50
    ) F  D# E8 L  @4 D( ^518 K$ M" H7 l, C6 L1 O1 t/ O
    52
    2 B* e+ N0 V( ^: J  K) O( \6 g53
    ! _1 ?# p! \7 v% E7 X) Z54
    ' g5 h- b% s& P/ `; E/ Z557 V( _5 L0 g$ l4 l* ~% M9 O/ b2 J- k; x
    563 l0 \& Q1 o9 ^' T5 \5 D; i
    572 a$ X0 A8 w4 s) m
    58
    " o! d7 J- T0 M) O/ q1 }2 ?59% c6 J/ b- O* l% V% w  l
    601 D+ ?7 G4 i% @( n9 P
    61
    8 Z$ G& P' p9 B8 t+ N62
    1 o: ]9 {9 y! s+ @4 t3 b$ d# K4 M5 n# ^63( y3 X) [2 b$ a" w; x& `- y/ G
    64" D# t* P3 t; h1 j& U, `& A" q
    65% B6 Q( k8 p, e0 G
    66) T6 l: p: n: s$ m$ ]6 t7 [
    67
    5 ]+ t5 m+ v; o68
    0 i% z7 D1 H' t696 q) }6 s& I. E! q6 z( O2 Y
    702 B, |0 ~8 ~1 x1 H
    71
    8 |: }; P1 w* [# S+ Z" j72
    3 Z/ r5 x2 z  y, q$ z73' E. A0 y  _* |# p7 v
    74
      |  o) L7 {- a8 V6 M% x6 `. b75
    $ z, p; _4 f5 X' I) X$ ]+ ~# |76, O- ], w5 ^* Y5 }" Q
    77
    4 j4 C% g& K  Z: {& v78* [; a6 e5 t. I4 F- d
    79
    ( \5 @6 s& T$ c/ ?7 J# A( K* l# c4 {80
    % C2 E* n9 p1 v81
    # C+ a* ?. {! _# l7 s, ?. |82. R( K# M' \5 H" W
    83
    3 M! B5 J7 f5 q" R, f( v7. 设置需要训练的参数5 U8 A! O+ ~9 I: U
    # 设置模型名字、输出分类数5 X: k% A* ?' g
    model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)+ z3 W" A5 K3 Y. A  f- _) m* m( ]
    ( ?% I& [0 k* N! c* V
    # GPU 计算
    ) j3 G5 @( x8 l' z# zmodel_ft = model_ft.to(device)/ Q) T; |1 g# R$ M7 ?: V) ]
    9 X9 Q4 g$ L# ]0 P5 j( S0 }
    # 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取
    1 r* r. l& ]# D/ C4 c  S3 Y& [filename = 'checkpoint.pth'( H# |9 o* C) [7 J3 t1 P" \8 X
    - o  S# D5 {5 j* @
    # 是否训练所有层  J2 @0 F! p% }. R, d
    params_to_update = model_ft.parameters()
    , W' V) l4 T7 T' |% f; z# 打印出需要训练的层
    ( K5 T" T) C" D" w: N, Rprint("Params to learn:")% b$ v+ u' V6 r  b$ T8 W# n
    if feature_extract:
    ' }/ ]9 p' t2 u3 ]7 v3 h    params_to_update = []/ ~: O7 H8 g5 t/ N, b% H
        for name, param in model_ft.named_parameters():5 y: w* h" R; X4 {% D+ h5 v  o
            if param.requires_grad == True:
    + p! Q0 x% _$ {2 \            params_to_update.append(param)+ l5 w8 O! z. d- t* F( g9 ^1 ]6 d  S3 n
                print("\t", name)* h4 Y& V7 l. d+ t# e
    else:
    6 D" _( b$ M! R# s% i! r    for name, param in model_ft.named_parameters():
    + ]+ Q0 [  W- t3 Z6 v9 Q) x& L        if param.requires_grad ==True:
    0 p5 }0 M6 c% `' M; s            print("\t", name)
    # Z) H* _( u8 F% Z+ O% i  O% B& K7 K: K4 n% r8 g
    1
    ' D- F. r+ ^: W' H: r* P8 o$ f  H# D2
    ! s& E9 j6 M! n  O- B! f5 M. v3
    & O# M! _0 r8 h" T5 o+ Z. ]+ e4/ Z9 a/ \5 F4 A4 P
    5
    % y2 e  Q4 M# X2 q  Q6/ ?( q; U; R* C: x# V
    7* n! g, k$ b( |; g
    8: k8 q8 N  x  f- V/ a: {8 m3 t) X
    9
    % R- R/ H5 P- t10
    " M! h$ t& a( C' K% {- z/ j1 y11
    & v& h' [# I# }3 \; ~- D8 ~12
    6 E  U8 I6 f5 i2 k( `9 I9 e13
    # Z0 ?4 [9 f0 N& a14: l: t- @1 R4 n$ }
    15
    0 _* P& f& k$ y: ]" c; W16
    - ]0 I+ c) s+ T1 E& G171 B% h$ ~1 ^$ M
    18
    & y: |* k! w6 ?+ X) x19/ ?( u8 @% a9 g
    20
      n- m1 |, [, R6 N21
    , Y6 O) w* O, {7 }, ~22
    + I3 x5 g2 v3 Y& [23& i% F# _" x: C0 e- z) r# ^: W  e
    Params to learn:
    $ F3 B) j$ J% }5 w5 f& a" @  [         fc.0.weight2 @( U5 w6 @' P- }/ r  W/ [
             fc.0.bias: p4 g# K3 l: u$ T9 b
    1
    & S' B% q% i  i$ C+ q" |& Z2
    & a  }6 L' T( X2 E3$ d4 [' W0 E3 T' H0 ]* {% S
    7. 训练与预测4 t- {% ^* ]! J5 u; N* H  K
    7.1 优化器设置( C0 G, h( t9 l
    # 优化器设置
    , S0 I( m8 D2 M- Q: t" Foptimizer_ft  = optim.Adam(params_to_update, lr = 1e-2)
    " [! V0 F: m% O% s# 学习率衰减策略
    8 E$ A! l) V& k8 b8 ?# s3 Xscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)' \6 q# G5 v" N3 o
    # 学习率每7个epoch衰减为原来的1/10: Q# h" K, J' X2 z% Z
    # 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算2 S3 G. Y. J, C0 L, G: o

    8 P* g5 X- @! P" J0 g# Mcriterion = nn.NLLLoss()& V, B* u6 l/ }$ h# P! ]
    1: f6 t, o% e) ?+ u9 T4 Q4 M
    2! v7 d8 l8 a7 I; d$ r9 Z, O
    3
    . m! a% |. ?9 G4+ m' V! o# w: P
    5
    7 g5 }- g( R, _' A2 w: S! |' ]6- U0 I% M( V8 F, g/ {
    70 M+ [6 y4 k  Z. B: W
    82 f, l& x+ z+ A  S5 V5 y4 M
    # 定义训练函数
    7 J2 s1 V8 q4 H#is_inception:要不要用其他的网络7 p. ?& a: K# K
    def train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):
    3 Z3 {7 e% S. v) U$ X: n7 _2 `: L    since = time.time()
      Q4 ]: D0 V; i, P/ g: Q2 H    #保存最好的准确率5 |5 v8 ]4 r# i5 w; S
        best_acc = 0. J4 ^. C' `% K6 p' y
        """
    ' O1 y( d# \" ~( _# O    checkpoint = torch.load(filename)& D7 `, _, g* I/ F8 Q
        best_acc = checkpoint['best_acc']# c8 `  d+ U' \$ {( G1 R
        model.load_state_dict(checkpoint['state_dict'])# Q, |8 q7 j2 U& C9 W9 V
        optimizer.load_state_dict(checkpoint['optimizer'])8 h) R' ~/ Z# s) n3 C- s2 V0 N
        model.class_to_idx = checkpoint['mapping']
    # V% R( K4 j+ }+ [* [6 U    """
    " G! L0 N" {" A    #指定用GPU还是CPU! v: ^) p# t- B: ^
        model.to(device)
    4 q7 \9 c% y! N( M. ^    #下面是为展示做的9 n/ C+ f8 m( {7 c
        val_acc_history = []
    $ ]. i4 I7 |2 U, k* u    train_acc_history = []
    - a7 b; t+ X' `2 f0 @: c    train_losses = []/ G; t! [) @9 m$ C; [4 H' q
        valid_losses = []
    ( y* N7 |( r  @    LRs = [optimizer.param_groups[0]['lr']]
    : p( z3 ~! i3 J" p# C' v# @! T    #最好的一次存下来
    & `7 y2 l3 w$ Z! S- j  g4 d    best_model_wts = copy.deepcopy(model.state_dict())  k; b2 T* ~/ c

    2 H% z0 A/ O$ b- t/ \" P: A    for epoch in range(num_epochs):0 O( X! N+ m0 d% I, k# y5 ^+ V
            print('Epoch {}/{}'.format(epoch, num_epochs - 1))
    / A( B. P7 r: S$ |9 N, W% u        print('-' * 10)$ D$ E1 c" D, D

    5 L" q5 c; A  B  W        # 训练和验证
    ! x# b! I3 M" q! v! }6 j1 ?& m        for phase in ['train', 'valid']:6 {' F+ b2 U+ f# ?( j7 R9 y  i
                if phase == 'train':* r( b+ D2 G0 M
                    model.train()  # 训练
    2 W& V( T; D! P' t$ o/ p            else:8 f. q4 q& i& M$ S1 o
                    model.eval()   # 验证7 F  Y, K; O( b: Q  h( g8 m; L$ G

    " C: {4 x4 {9 a6 v            running_loss = 0.0- R  n- q6 w6 b$ h& L5 Q
                running_corrects = 0% v" U( `( @( [
    , U8 c% G; E( {' r* e. ]
                # 把数据都取个遍# e; d* t+ I9 d9 O. `5 z5 f. G
                for inputs, labels in dataloaders[phase]:5 x7 t' l* Q7 ^2 W, D' ]- A
                    #下面是将inputs,labels传到GPU! x) N, o  s0 E* w9 z2 Z: F( A
                    inputs = inputs.to(device): {6 p6 G' x8 A$ [+ J: d
                    labels = labels.to(device)" f4 N" P6 ^/ g: K* d$ ^$ {

    ! y; w/ A+ G$ k# `' M                # 清零: x( w* ]- N( y, B
                    optimizer.zero_grad()- t5 W2 r5 T# n5 j
                    # 只有训练的时候计算和更新梯度# s3 V) B0 ^# P9 O
                    with torch.set_grad_enabled(phase == 'train'):
    + M* Q% n: t# H6 t. z; C                    #if这面不需要计算,可忽略
    + N1 ]  `6 o8 F5 S                    if is_inception and phase == 'train':
    7 n. P# @6 V% [2 R* l/ e                        outputs, aux_outputs = model(inputs). q7 H" i9 H$ {& m! v
                            loss1 = criterion(outputs, labels)5 k2 ]7 l: k' M; ]
                            loss2 = criterion(aux_outputs, labels)9 @$ i5 }# m( Y! _) R$ y* e
                            loss = loss1 + 0.4*loss2, a4 X9 w( p+ h8 c0 d4 R
                        else:#resnet执行的是这里7 m) P+ \  O1 M$ w# B
                            outputs = model(inputs)
    9 P9 ]- f+ h. {; I; U                        loss = criterion(outputs, labels)
    2 L2 J8 i, r  _, h7 t0 Y; B- w7 U$ K  [* f- o
                            #概率最大的返回preds
    ( B, p7 A- m9 j: q* X" j                    _, preds = torch.max(outputs, 1)
    . ^5 C1 T7 @  g; D0 J
    4 @- J1 ]7 W' R3 u2 q$ w6 c$ D                    # 训练阶段更新权重
    . Y& r+ a, S- ]: n. r4 G                    if phase == 'train':  {( w6 p1 Q! v  U/ {
                            loss.backward()
    ' T5 B# D: H8 j8 P3 }, i! y  E) A                        optimizer.step()
    , l; o' e1 @3 S7 V
    2 R& k- C/ a( o                # 计算损失
    , ]0 _' }  E( O: ^5 C2 U# d$ ~% u; _9 i                running_loss += loss.item() * inputs.size(0)1 }$ _7 l! R, Q( }* E
                    running_corrects += torch.sum(preds == labels.data)* q) c5 B- V% u3 [+ X" ^- B" f$ `

    0 C" m% k, P0 |8 @0 T- F            #打印操作6 }# @1 H9 H  N. r
                epoch_loss = running_loss / len(dataloaders[phase].dataset)
    5 z, ]9 @" h, D            epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset), ^7 b  [2 ^  r5 L$ k3 p
    # D( ^  [' y6 ?

    7 j  e1 U8 K% _; n            time_elapsed = time.time() - since
    ! B! e' Y; o. C            print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
    , S" O( \! x, ~, e            print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))
    2 ?; U  t& W6 j7 y8 k* V8 q7 I, f
    " G1 y) c2 C/ }# Q0 ]3 m5 R, `9 p. c6 \
                # 得到最好那次的模型
    , I0 P. }/ e( x+ X* |            if phase == 'valid' and epoch_acc > best_acc:# g6 J/ ~3 m# G
                    best_acc = epoch_acc
    & ?7 V* X4 c/ K7 a                #模型保存. P' X7 p$ c- j1 ~  a8 x
                    best_model_wts = copy.deepcopy(model.state_dict())+ J- n- w6 T- O# l$ \4 a
                    state = {* A. g' i3 w. D2 ^
                        #tate_dict变量存放训练过程中需要学习的权重和偏执系数
    1 p( z+ w( k% U& p                  'state_dict': model.state_dict(),
    0 B; V- V6 p& L5 J7 e- Q, `                  'best_acc': best_acc,9 N  _6 Y' S) Q  A
                      'optimizer' : optimizer.state_dict(),
    9 `7 K. ^( k0 |& ?6 D                }
    ! m1 W2 P! N# l9 o) q                torch.save(state, filename)
    + O! s* L2 Q; n6 n            if phase == 'valid':8 s- x3 {1 g* R+ B) H
                    val_acc_history.append(epoch_acc)
    # I3 ?/ h7 a& ]& Y' ]                valid_losses.append(epoch_loss)
    . M2 M, \, t* ^5 L3 s1 D                scheduler.step(epoch_loss)
    9 u2 P5 D. m$ D% u            if phase == 'train':5 ^" r/ j7 A1 k& A7 L/ {: W# S
                    train_acc_history.append(epoch_acc)
    0 T- ^& v% m. k, P. U% A                train_losses.append(epoch_loss)
    & w1 f9 n1 U4 e5 o8 W6 f; ]9 V# Z7 b
    0 {; z* W9 {4 e, h        print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))
    1 n7 ?- [4 m$ p5 @, Z+ X        LRs.append(optimizer.param_groups[0]['lr'])$ E# ]: e: T: X
            print()/ S- B( x; y3 `" ?

    # W" A$ {  M0 o" f; ^0 U    time_elapsed = time.time() - since
    0 Y2 k8 e8 Y9 B  f    print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))$ U# K& L. o: c( r
        print('Best val Acc: {:4f}'.format(best_acc))0 p9 s, X0 z1 {) ?

    9 o$ D7 a, V$ v; J  _    # 保存训练完后用最好的一次当做模型最终的结果9 N& S1 Z  c7 u3 L, Q
        model.load_state_dict(best_model_wts)
    / ^  a) h6 d" B' G9 B) ]0 [    return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs
    ' S( V6 V# H; t2 H; l0 w( d; T: h7 Y7 q9 I
    * b  [+ Q5 e  t* P- L2 X
    1
    & E: `" F6 q2 b; J$ F  N" B2, a/ y: g  P: c/ K
    3( l- Q$ Q' H0 W" {: A8 Y8 |
    4; p& s3 g6 a/ w$ ?- B
    58 f  I/ l& X5 Y! B. |
    6
    9 a1 g/ d" }' V  C" J( @7/ M- ]6 V4 n( B" q  a4 y
    88 o+ R/ {5 ~7 R( e4 s8 t
    9
    / ]! g3 T' q; C9 ]0 a, _10; g! F; D# S: L, b
    11+ E" ?5 d, z' e# s& f6 l& O8 b
    121 J5 h1 O) U8 }' _# Y
    13
    9 _+ C& y# |8 R5 C! Y$ y! E14
    ; }& [  y4 m3 ~3 @) c8 ^15+ Q, d5 h; |3 q8 D% `4 u
    16
    ; T. Q0 R1 E# X* E  f17: E( s3 ?8 _" _- u- {+ w8 ~
    18: C5 J1 p: M+ _* p3 ]/ g
    192 ]0 t! |6 X! `9 M+ C- u/ g( q
    20
    3 B5 ?0 Z+ r' c: D2 F215 p4 x: I4 g3 r; N4 U
    22
    8 y) O7 h5 G8 p. m. I3 j) D4 j23' r. O* }, L: g+ c
    24% D& M6 Z+ m; r6 a& k2 K, f
    25
    , m# O; b, T7 F  Y" _# ~26
    ! M- X5 V' _& V, O, A4 i27- P& g( z1 U" t1 }9 w( R
    28
    ) y  [& Y9 f& y, `' U5 `- j3 G& f& f2 J29' L- h/ f, |8 y/ |+ A
    304 {7 A/ \' T4 O5 s- u
    31
    . Z9 P: C, p+ Y4 S, s7 Q* W32; e2 m( O+ y4 C: e
    33% @" p# g  Z+ e+ i% Z8 E  G8 s8 s
    34! v0 ?% m9 S2 Q! s) R6 |: Z0 T5 Q
    35
    4 @, a: {; R; T$ j  i36
    7 \- o; W7 o4 f+ r1 Q377 k4 o0 v0 K+ |: c& H
    388 ]$ U6 b" A8 @
    39" Y3 |# e% P$ n( Q; d. K3 Y5 J
    40
    : `( E; S) h7 U+ l/ N  ]- |7 B- b1 ?41
    % \- q3 T6 H; y7 d42* h! ]8 u) o/ N! P3 u4 {
    43
    4 S( z; ?# \1 X4 u444 c* g8 b" q+ X* y# t  U7 ~9 c
    45* r( G0 a, u1 N, Q
    46
    " z, j# o/ i. }, ^8 p- P  q" N47
    & Z: N# h) x% W48
    7 ]2 u- d+ W' S6 F# q+ o2 R* l" E49( Y- Z% r1 N/ H
    50
    1 R6 z. T3 M" t4 `1 N9 I51
    , M2 \; g6 u0 g2 V$ U52
    " n- A$ j& N8 ?+ I. c53
    7 f4 b* ?+ _: g. |0 W) u! Q- Y: F547 `) F2 y3 H$ c' w4 Q1 T8 u
    55
    2 O1 U6 L, H8 C# m% ~  a56
    6 {8 P5 C" M& \+ V2 q- M57
    5 `  f+ e! Q# a4 y4 `3 D1 U580 i: g( Z$ c  T0 t7 s+ f+ H; b; h
    59& C  e  G* z" _7 ?' [% h' G. X# `
    606 q0 c  D. }! e7 i8 y, E
    612 l9 v1 m; Q1 r, w0 u
    622 n' ^' l! U7 r
    63  U: Y# M/ w% e- G, n
    64
    5 q% r; w, K6 `) N65
    , q6 o4 L* ]7 P. P9 a66
    7 f; \# R! w" c& h3 a! h67! }$ n- z8 u# D: r/ v
    68
    % U) N; u, d- ]) h69* w1 a% Z3 G! s2 W' k9 b* ^
    70( J& e; G0 T% \) @
    71$ Y7 c' U3 o! }1 b
    72! \9 {' ^$ m' A/ P& P) r
    73
    & Y( x3 Q( _, ]8 c74$ M" J; J( V) V0 u
    75
    ' P0 h! m( k- |76
    ' Z! e8 F2 \, C/ b, g4 g* n9 r8 \+ q77
    * N7 w2 }. U5 v" P, v4 q78
    2 x' ^* u% g+ s8 j79
    2 @4 i+ g- t& {+ f3 Z80
    2 `) f2 j4 _4 u* Y# @81' u( q' l# N! K& `7 Z
    82
    * {/ ~# E! [9 |838 ^; V( B/ l; x* ?" [( w1 }: m
    842 k9 t, N" V" t
    85
    ) V' h. \/ ^0 B1 K6 S4 p86
    ' `0 b, Y& x% b' q2 J1 L2 K( D( \87: Z4 C  r' X" c. s- Z
    88
    # G, a  w1 g8 K: J6 h- j1 g89
    ' k* B* J( l; g90
    $ R& v; Y; f& z% L  K91
    ; L# @8 o2 S6 e/ z( Z# ~923 `! Y1 P) ?9 D3 `% i
    93- |9 Y0 O# X, f  Z6 B" V  P
    94* ^' C, M; A! E
    95$ ~1 R: O8 @! I/ j, q
    96
    & k/ p5 p6 ~) j! ~( J  H! i971 U- d! {7 `% N3 J- q- N
    98
    . ?5 t6 l& x, U' d& D' t7 U99( `$ ?( E% e6 z
    100. B- {- L/ I+ p0 u, R
    1010 g% n# e. G7 D0 i
    102& Q* _) x8 q& [: s6 Z+ ?
    103
    / [4 _, L- ~4 Y$ X5 Y104
    0 N3 X: h. G" w$ _' C105
    ! o$ i2 x; ^  R8 i8 a2 d9 P% ~106
    4 b; Q3 M: b* L7 p3 k- E107) b9 _* }. ]( m% @" [/ x( s
    1085 ?* R+ f0 G7 d+ y" n
    109. O( h5 L3 n! |2 X
    110
    " T' d1 ]9 |9 l& V/ N, l  l) L111
    * u2 ]1 q" C0 Y4 R/ V' M# k% V5 z112" I5 @* b" @6 K1 Z1 M
    7.2 开始训练模型
    & ^7 U4 I8 q2 Q7 U我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次& b: z" R: R& B5 F

    & M5 ]. x# ~7 R0 p8 T#若太慢,把epoch调低,迭代50次可能好些
    7 o. K- ~9 u$ V0 h) u#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合" t% `6 \/ D1 k. c+ V# u3 {0 s
    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"))6 h" T( z$ }- n7 u

    $ ^( Y; B  K  o  r1. G# c, _; R& ]) x/ N
    2
    6 Y& _' f) Y, t" V  O30 U8 u" O# w  y: G8 X
    4
    1 {: |. c9 @& B8 ?+ [Epoch 0/4
    6 n5 v1 m- D* E0 c2 ]5 i----------
    + `& k3 [9 e# {/ P1 KTime elapsed 29m 41s3 w' y5 ?* p2 ~  D
    train Loss: 10.4774 Acc: 0.3147
    ' V" u" p1 @( p# h8 lTime elapsed 32m 54s
    ! S. A* V: A( L0 E" X; M' xvalid Loss: 8.2902 Acc: 0.4719& T( i2 e& R! z3 h
    Optimizer learning rate : 0.0010000
    6 J. [* ?: {3 y0 `* T+ G
    : a- ~# l7 X9 Z% L3 ?, vEpoch 1/4
    + c9 w4 k1 a2 b2 c9 a----------0 [% X5 e+ R; |) x  N, ]3 r* ^
    Time elapsed 60m 11s
    " P7 K7 S8 |% C1 h0 C- p  C5 htrain Loss: 2.3126 Acc: 0.7053
    2 U4 U+ Y; A; {, kTime elapsed 63m 16s: G" ^1 s* Q& b1 q8 _
    valid Loss: 3.2325 Acc: 0.6626$ R4 k' r% U) R1 m$ e" M
    Optimizer learning rate : 0.0100000
    & F( o+ |$ c* @: ^/ G7 ^. f, p2 @* ~; I  V& J2 n
    Epoch 2/4
    1 T9 Q- B, w4 i. E( M+ W----------) h1 Y7 ~' N9 n
    Time elapsed 90m 58s. ?) ~3 k. ^  i; _) F
    train Loss: 9.9720 Acc: 0.4734
    ( O4 s6 c: S; s9 U4 ^) u; W& F1 c' K: mTime elapsed 94m 4s
    / d! ?3 r! x! O, ovalid Loss: 14.0426 Acc: 0.4413
    - o! L. I3 z( v. U' o' |: YOptimizer learning rate : 0.0001000( W3 A# G8 S2 K2 @# ?

    ) j$ F) s! t6 O9 [! ~; uEpoch 3/4
    * V4 N; f: ^! B4 H2 V6 T& y----------/ F  _7 i, F+ L3 ?: R
    Time elapsed 132m 49s
    ! n. [7 L* P1 ]train Loss: 5.4290 Acc: 0.6548
    $ {1 o3 N+ T' XTime elapsed 138m 49s, y& A2 ?' j: G1 o6 @( ^
    valid Loss: 6.4208 Acc: 0.6027
    5 z% a, W( }3 O. J8 WOptimizer learning rate : 0.01000003 \: K- M7 K" I2 T6 ?2 I
    6 u! x" q' t. m/ P/ p* |5 v! u7 N
    Epoch 4/4
    $ c% w: m/ d# ]2 P& X----------
    ' H6 J* g- T, C  DTime elapsed 195m 56s1 d0 [0 j8 K# x5 V+ s
    train Loss: 8.8911 Acc: 0.5519' U1 M1 G' R# x8 q# y5 Q
    Time elapsed 199m 16s) K- M! M! U1 r" Q/ F4 E  B( S# v
    valid Loss: 13.2221 Acc: 0.4914& I: x8 ]+ o( P/ ^+ F
    Optimizer learning rate : 0.0010000; o: Y& ?# }" }' Z; H+ }! t

    3 A! p# M3 @  ~& mTraining complete in 199m 16s
    # k3 J; n% N4 w  X# a3 V, {# \Best val Acc: 0.662592" B- ~% b8 i- Q6 V0 \& u9 o8 z5 v

    # Z0 c3 D( [5 U5 Z- l1
    : M* Y% k1 M6 ]+ C26 c( k9 [" A" P, ?0 K% }
    39 y. _/ m. a) v: F3 `
    4- P$ N* R: C+ I+ v% B3 O3 }* l2 Z# [
    56 V1 ~& C( [0 L4 X) ]  V
    63 w' B: b5 [0 N" O9 P
    7* x& u" F3 `; `. X0 c% R- S
    8( R2 I. q* j! U$ Q/ Q
    9
    & m$ n& I" ]7 q; ~' F1 s# [101 d* e# t% r4 t. ]& c  m3 Z/ v
    11
      ^+ y$ G, H. T1 ~$ j( p7 k12
    & l8 R" V; }9 c  S; f. s* T$ R! o139 s; G1 r2 @. T1 j, d% q  F
    14
    ( p8 @+ d: _1 G. U0 s15
    + ~7 T& q! T* T8 Y16
    7 G9 m2 P6 }6 V( R176 b) W: _5 s3 |) X& a* q2 J# m
    18
    ) V' h) `" H$ Y. R# o) w( Q19
    $ |) E% B/ y/ c8 z. o. P20' N" @1 [" N; a2 n. n
    213 z7 ]3 f* ~& I9 Z
    224 ?% M( _# _6 n' }4 o; V+ e
    23
      Z: z; f( ~  R( z! a& ?24* D% |5 H$ p9 a# \) u3 O
    25
    ! A4 ^1 U$ h6 o/ i; G# N26
    ) P0 P# ?, ^5 A) \- Y27: e3 O  p% q$ \! J: X. B
    28; p1 p7 [; x& O7 s
    29/ L" @: R" ^3 Z1 a
    30
    " @: p, t6 B/ E2 q5 |31
    + |0 d# Y/ F+ Z% y32
    1 p  m4 _& r8 Z4 c- w3 M5 K33
    0 _  p1 i2 [5 Y& l- o: ~34
    4 Q' e5 \1 T& c- ?. J  r  V. K# X35
    , r2 U* C! ?! _36) S, f6 P, G) k4 J" b2 K: G7 C
    37
    2 i! S8 h, U* E" e( G" n38  ?  {' ]- U5 S
    39
    , N) M1 t4 V1 U- _$ I* j! b8 ~40
    ( e" b8 s: W) q2 ]! j0 @41
    " L, _9 L: O4 ^# ?425 q' f/ M. x$ {
    7.3 训练所有层) t+ |0 y" j1 U: L7 a+ K9 P
    # 将全部网络解锁进行训练
    & X; m5 O+ b$ Tfor param in model_ft.parameters():
    9 C- I+ \1 U- I: z" k! D- e  b    param.requires_grad = True
    3 _# m# H- x. t
    % {" A1 h3 a" i4 c, m3 m# 再继续训练所有的参数,学习率调小一点\
    - D5 L, C- ]; v- d# p+ goptimizer = optim.Adam(params_to_update, lr = 1e-4)' t: g( L. p' j1 {" F( }- h
    scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1); U2 U4 A% u# C6 L4 T; N
    ' J2 h. e3 l2 s; ?7 ?
    # 损失函数
    ) `1 t  m: @6 X: Ycriterion = nn.NLLLoss()
    2 r; f- E2 {% m, h8 v1 I1
    ! S# k$ ~+ E' ?2% ~0 I) E# ?+ J, i, h
    3
    ( Q0 c( i& F4 W% k' A6 S9 B9 H42 C- \. [7 \9 t
    5
    $ m. k9 G1 X, ^/ Y64 ?+ l+ v5 F# {
    7$ F! s& {& M  O
    89 w8 s! G3 \$ C
    9  r. Y4 ^/ C2 a5 W
    10$ s+ O! u1 X6 A
    # 加载保存的参数: F" ^0 c0 L, _6 k0 e" s0 C
    # 并在原有的模型基础上继续训练6 V. y, h! T- ]; }2 m/ s: ?$ L. F
    # 下面保存的是刚刚训练效果较好的路径
    5 ?% L( j) t" n1 ?) m& p8 \9 j/ Rcheckpoint = torch.load(filename)
    2 G" M0 L0 s7 r& `best_acc = checkpoint['best_acc']
    ; t  d* e2 C6 G( H7 T% jmodel_ft.load_state_dict(checkpoint['state_dict'])! S; i9 g! B5 R
    optimizer.load_state_dict(checkpoint['optimizer'])! R$ J7 @6 s) O1 a; G: D
    1
    0 Y8 ^& P5 D- w3 N; H5 a0 R$ C! a2
    : E- K/ E7 i& f8 d. e/ q8 m4 x3
      d% S( {. |  U4. M/ g2 P: J+ o9 h! u  h' s- u% r* k" Z
    5
    # W" s  q1 \' x8 o3 q4 e" |( v6
    % J6 U' Y2 ?7 [. k* P; Y7
    + A' M1 |6 ^! N* m5 I1 y) f开始训练5 x0 W4 e' C" c' [2 b0 ]
    注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考
    * E* g! e* o. b5 O' k  m9 |- \% X9 f3 y0 O6 `" k3 o1 A; U
    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"))
    8 i7 i& ~2 O, e7 P$ Z7 D1
    2 i4 P  j3 L2 r) V8 |- D( wEpoch 0/1
    ! I$ V* b7 r- a7 R" D, I----------" L5 t; v' X. a, a, `4 c
    Time elapsed 35m 22s
    ( K+ ^2 @: {. F6 etrain Loss: 1.7636 Acc: 0.7346/ p" R, b1 j( a$ P. T
    Time elapsed 38m 42s
    0 W- }$ ?/ P: ], b" X* Uvalid Loss: 3.6377 Acc: 0.6455
    4 n1 j# I) J4 n  q3 B" POptimizer learning rate : 0.00100002 J7 g# G+ Y% a/ z
    * d4 P3 i% B  P  O9 l2 ?8 i' c& ^: W
    Epoch 1/1
    9 ~8 x. g0 h0 k) G& F/ V9 G3 u' R----------
    ( ^2 {6 w2 C& h- w2 CTime elapsed 82m 59s
    1 j! Y& K/ n  @$ R* q# e& ]train Loss: 1.7543 Acc: 0.7340
    7 x5 K2 I$ @" PTime elapsed 86m 11s2 H% v. e* R, }
    valid Loss: 3.8275 Acc: 0.6137( J7 G) F" y% ~' Q4 V. O$ Y
    Optimizer learning rate : 0.00100007 {! @3 i3 d6 G$ a, A1 F& M0 S- h7 R0 R
    6 G' N4 H2 E. s8 V% m8 @
    Training complete in 86m 11s
      R2 K8 d* Q4 q  h# P9 DBest val Acc: 0.645477
    7 T8 }% \9 {7 {+ |1 z! [6 I' t1 U# g+ w" ~
    1
    / s9 F; B/ A5 Q* D; [2
    4 ]0 f' J  H: T7 W' F( a6 J3* N- ~. E2 r. C% [3 O! D( X
    4
    * n6 O' D7 W: D' o! _5
    3 D' M) K0 F1 t* Q* Y  {8 h5 H6 S6
      u) D$ l3 j2 S0 ^$ R7 R+ F7- }; p4 T. ^: Q/ B6 o4 w
    8
    ( i8 [. a% N; e9 n96 I- K& Q$ w/ T6 z
    10$ j% ~) n) |6 @  K' u
    11
    2 U* _8 S9 s9 u$ p8 z$ T12
    6 B( f1 I3 A' s( [8 c2 H. ^6 h13
    * v' T  P2 _! A! M/ n. n147 M1 b$ o1 S4 K  J9 v
    15& Z2 r/ V( S9 e* g1 E; w
    16
    0 l7 p% y$ o7 ^1 Q, k- R17
    % G5 ^" A% I' w- k/ p0 Q/ U1 }18
    . e- D6 I* W8 }: p9 q1 h" ?# A" t. j8. 加载已经训练的模型: s' h( y$ Q0 m4 h" p
    相当于做一次简单的前向传播(逻辑推理),不用更新参数
    7 P8 o8 }- ]8 N
    : J5 z! ~. f6 ^) H/ rmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)6 n. S! N$ n; O/ H" z! {" F

    - Y6 P8 b0 z; W: ]% }$ y4 I* n7 p: r# GPU 模式& h9 x5 d7 R3 D- l4 i2 K) T4 b
    model_ft = model_ft.to(device) # 扔到GPU中
    6 j) _6 M# _" \$ F) h; S% @* t: }0 }
    # 保存文件的名字
    - l9 h* M6 {7 m6 s" P  Z3 Wfilename='checkpoint.pth'
    * |2 e/ G/ n* x! q+ d+ A5 x& @9 q
    # E  c7 T' Y/ j; }! h# 加载模型
    - f, }" N- P' ]: ~# D* ?checkpoint = torch.load(filename)
    3 V% T5 H( q" b& V  |6 U+ o5 Wbest_acc = checkpoint['best_acc']
    / u( [% [) w8 Y& X0 C5 nmodel_ft.load_state_dict(checkpoint['state_dict'])7 ?1 ]; r  I  u" e) O
    1& X! H4 `7 G) R5 ^
    2
    ) B4 X* D3 d9 t" a) f( `, t  F32 m2 g3 ]( j1 O% Z3 y
    40 h6 X  m  B, {6 u; S; Q8 [- h
    5
    4 M5 [6 Q3 ^$ C+ @* C' I67 l, U, r7 w) f: ?
    7+ H: y0 `, \/ P2 {& z2 [
    8# O! O/ j. `5 Z2 z& W
    9. C2 U- Z& {! B% t
    10% t5 V# e) c  a. w
    11. e4 o7 v( r4 E- n9 m
    12
    . W% q/ g' Q- E$ G<All keys matched successfully>3 X/ K, @6 @# P
    1; ]* S  {6 w9 @0 }5 L- ?
    def process_image(image_path):! M5 v/ N  L* j: ]9 H
        # 读取测试集数据" i% S5 w& U( N% _1 q; T3 P: V
        img = Image.open(image_path)3 |* r" ^: B% n# k" N
        # Resize, thumbnail方法只能进行比例缩小,所以进行判断
    0 J  M) y8 o. j3 F! `5 e    # 与Resize不同
    / I' E3 u# o* j    # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小1 c$ h1 R% `5 h" z, K
        # 而且对象调用方法会直接改变其大小,返回None
    ; G1 n' V: }. ?* K7 O' q  i    if img.size[0] > img.size[1]:
    0 c/ U& g9 Q; Y" @6 I; o) f  R* c        img.thumbnail((10000, 256))
    2 h, w) `: ~  j9 Y4 g    else:
    + L7 x% d) D- g; L        img.thumbnail((256, 10000))" h" z% e8 K# c1 r1 R
    . ]: j8 Z' r" h- J2 {
        # crop操作, 将图像再次裁剪为 224 * 224
    9 ?" i6 c# e) q0 r% Z  {  d0 n    left_margin = (img.width - 224) / 2 # 取中间的部分
      K6 x: b- J; x$ R( D( s. O* n    bottom_margin = (img.height - 224) / 2 / K% W* Z* B8 Q: Q
        right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度
    ; i) g, c/ ]4 B    top_margin = bottom_margin + 2242 `0 i( I- R4 x1 u8 B
    + \0 q$ w, i$ V1 }4 R; i
        img = img.crop((left_margin, bottom_margin, right_margin, top_margin))& N3 X$ U0 J8 ^/ P
    8 F! _' p2 D1 {5 U9 ^+ Z& |
        # 相同预处理的方法
    ; R: G$ w4 z; o% p4 I+ M$ Q    # 归一化- @9 e7 a+ u6 ?" p5 |; W3 V
        img = np.array(img) / 255! r3 v' P/ K: r; I  [: W2 n+ B6 k
        mean = np.array([0.485, 0.456, 0.406])* Y+ {) i/ s: q5 z7 F' ^) x
        std = np.array([0.229, 0.224, 0.225])
    3 K' J' N3 ~5 [1 W5 {! V    img = (img - mean) / std
    ( i) O6 j' q3 L. X& o+ e! N% D% N- o! Z; ?) }# X
        # 注意颜色通道和位置
    . m6 {. \7 V7 H    img = img.transpose((2, 0, 1))
    / j3 z$ y2 j6 o, x" p7 Q
    3 \* R7 Q( d( x, {; i    return img
    ) G1 ^( _4 N: L1 t$ ]
    + G4 y5 ~7 O4 L  N( G+ \def imshow(image, ax = None, title = None):% a3 _4 C: d& `
        """展示数据"""
    $ B9 w& L. ]% D6 r    if ax is None:1 _* J, F7 P, i1 N1 D0 Q
            fig, ax = plt.subplots()
    ) L: @8 r. c  m/ z' i- r8 X% k0 o: ?# s" b$ H5 C
        # 颜色通道进行还原. J$ Z& K/ q8 }/ \
        image = np.array(image).transpose((1, 2, 0))8 O# O( B9 G( K( w8 l, P( f

    - e( C) f( D4 {8 v0 g    # 预处理还原3 w# `9 I- Z4 U# h# u6 y4 |
        mean = np.array([0.485, 0.456, 0.406])
    . }: e4 j5 n3 O% f! P5 S    std = np.array([0.229, 0.224, 0.225])/ A' o6 N5 v' ?' k0 M' b- \2 Q; }. s: v
        image = std * image + mean
    2 V4 {8 k. q$ T/ ]& D; \( {! G$ w    image = np.clip(image, 0, 1)7 ^) q4 E6 B6 Z

    - q. B; f1 ~" r$ [. S/ _( r    ax.imshow(image)3 h. j$ q3 |; }5 S7 e$ G
        ax.set_title(title)
    7 X! d3 B) j6 k
    0 v4 s1 g# v$ `, a6 s    return ax
    % I' ?8 T( m0 b8 b; d( Q
    ' [: j# l, G7 Rimage_path = r'./flower_data/valid/3/image_06621.jpg'
    ; R: z- v; F; J" r/ Mimg = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理4 r- ]1 I. ~9 p& i! L
    imshow(img)# z/ S: C/ z( }1 g( ]9 l/ ]% p5 n+ r. n

    $ S, m; E; p# |1: x* x0 ?. j( u& ~* G. R
    2  G; i8 k4 `" W8 _# q) O
    3) J  O, r- X' d$ s* G$ w% v
    4
    4 d- T4 f8 y( k5
    0 q5 z* s% ^) E  i1 ^1 J$ ?/ c8 ]6
    9 r, r/ k/ Q0 Y+ s8 O. F" X3 w7
    ' I: r1 T# a4 S5 S8
    . B% L- u$ V; q9 m( J9
    5 N; B) j& g( O( B10! A5 V0 E3 |5 f: b
    11* @5 u; j6 d* f+ e2 p! G' E
    12
    / z0 f! Q, A* m, T9 G13) {/ G0 W0 v2 ^  M
    146 W; l+ r# O/ s* H5 k3 e3 A
    15
    , D7 Y, Y6 u( o1 [5 e- J& x8 M( a$ N16
    ; ?# U- d0 O6 o9 @176 j" q. t( p7 m5 D$ C& y0 q1 k
    18( Y7 }; q) Q3 P- g; Y
    19% R/ O" ~# Z, z' @4 L' J$ w
    20' i8 X( `* \$ D6 v' w) k4 O- q4 ?
    215 F/ Z1 K) i; D* B" x/ k# |
    229 E4 W% B- G5 @4 I2 r9 x) H
    23( u- O0 S9 a( ~. r4 R; v8 [
    24
    3 N- D8 a% [  A0 x, u1 R& }25! F5 e5 b. f; C' ]( S' Y1 l
    26+ M* K- u' o6 f: ?$ I
    27+ S8 s' X" }2 U4 S* [
    28
    : _  q2 |5 z. r2 o# f9 E29
    4 `2 }% G; O- c" M! S/ N" C- M308 I$ G. G* W9 ?9 D- ?5 r: h( C) w
    31
    6 U" P7 x- v  ^+ U1 P3 P32/ X; y* ]* |: J' i% h  |! R* J
    33! E5 h  f2 Q8 o2 v
    34
    9 l0 q/ R8 ]+ M- l5 e/ v8 r3 U# N35) g- M2 h6 Q1 Q+ c2 B) V
    36
    1 }1 {3 W- v, B37
    2 S4 O. w: `$ p38
    9 R) g/ v) ?3 @4 q" D39
    5 J$ X0 ]" G  t! l6 z40
    # _1 Y+ a' O6 N% \5 O41
    7 f- D3 W1 f3 p  j42- ?. L3 n5 b( s# o  G7 I
    43
      b' i  j4 p$ `9 ~4 M: e1 I444 `; {! x# ^( Y9 K! `1 X, O
    45
    # v' V- v. T. {. q46
    % O  \5 C+ M5 u) y: ^47
    $ G- N- X# T) Y& M: s48
    , u; d7 p' o# }2 }  a" w& ?; Y& a49
    $ I1 q: t+ x1 E. Q50) n* h. n" ?6 i( _6 P
    51
    2 K! }& g  z7 g52; u+ Q6 y2 q$ n
    53
    # ^9 A- N& d* J2 s  L- S: w9 ^54- ^5 b1 ^8 O$ h9 D. `. p1 A
    <AxesSubplot:>
    2 b0 a9 g. j9 D3 n- b7 ]$ t1# l* _/ E% o: @# y

    3 ]5 Q0 f7 z% s3 g上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确
    ; Y9 A1 M/ ?; Q  K, }& @8 @; S
    / `7 I0 E2 B! @* e% C6 @img.shape
    - o, {0 z  F9 `* a5 C  H1% E9 q1 E9 q! ^+ g7 d) z" L
    (3, 224, 224)
    2 [6 U1 L' l, B$ E15 o& U0 R% X8 ~3 U: h  s
    证明了通道提前了,而且大小没改变) o7 ?4 s3 ~+ M

    6 f  L/ J8 M  a- S5 |: v9. 推理: l' M5 f! e, ~6 m) p
    img.shape, B+ |& u! |* o. w1 b

    % I1 E" Q0 I1 \# 得到一个batch的测试数据0 k, P# t1 r- j. L( l" c
    dataiter = iter(dataloaders['valid'])
    , C9 o, X$ r/ @0 a2 J" ^( zimages, labels = dataiter.next()
    " [7 r; D! v( O8 T. j8 q# _5 k
    # K6 m) E( Z4 S1 Cmodel_ft.eval()5 M) `+ b0 W9 q( I. t( @) B
    / n! R) b& `7 |9 i
    if train_on_gpu:
    ; ?; F1 @2 o. S8 c    # 前向传播跑一次会得到output, l1 q& U) p4 s% N* \
        output = model_ft(images.cuda())$ P4 ?# z, a) y9 H$ Z5 _6 a2 P! ~
    else:
    / l! e- ?' I+ J: U* }    output = model_ft(images)5 G3 _- l, m7 U% ~
    4 x! I. S5 r# E6 }" r2 k. B$ _
    # batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值2 w( T: u( h8 c/ u/ ?6 x) T/ p) h
    output.shape
    ! u" P4 Q& }6 W( Y0 O9 `8 J# N+ x- b' M" Z$ _
    16 @$ Q1 D; r- L
    24 j! t/ Q, B$ w3 @4 n
    37 P2 M5 e' b, m
    4
    . W/ A# G& R+ p2 [+ E- r- K) q5
    4 G& [  z" k! Q. J) k; h, \- j6
    ' _7 I# M! E1 H7
    / L( ?9 m6 p) d! i  U8. r# w1 v' ?+ D
    9% N* ~' [7 P8 k" g
    10
    2 q) M: o- b. s7 _! s11
    8 S3 F" l5 r- e. \12
    / M; {' K9 t0 w- p' v$ P) p3 G! d0 m13  k1 [5 J2 V3 Y& N
    14+ h" D  d% u! L. w. @. R  Q4 F
    15
    + @6 ~9 \4 Z7 M16
    5 C( ]& C0 n& mtorch.Size([8, 102])
    6 N6 m+ m7 J; y, n5 k: R1
    / ?. [6 k8 l+ J- H( k  _  e9.1 计算得到最大概率" D. h% ?6 i1 i9 T# @' x$ U
    _, preds_tensor = torch.max(output, 1)
    5 L5 |% R6 v, t" r6 f* T: C5 u& d3 D9 h  _$ Q1 H! I# E/ H
    preds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量
    4 C; n4 U$ l# W, i" V/ p1
    + L, b0 F) Z! z+ t2" A1 v. H2 ?/ c% W( L1 T: p
    3
    4 `& X% S4 F& S8 j7 t" d# f. ~9.2 展示预测结果. @  X: b2 G3 V, X# S- c3 V# i
    fig = plt.figure(figsize = (20, 20))
    ( ]# q9 p3 |: t. c/ @2 Y1 Dcolumns = 4+ d5 `- U3 O. l' t9 l! {' ^
    rows = 2
    9 k6 Z3 D  @) g9 X: U' u5 A
    ; ~$ R- _" H0 B! ^1 z9 Rfor idx in range(columns * rows):
    3 F" k% b$ Y3 c4 j    ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])2 H" z) I* r  F5 S8 ~3 ]
        plt.imshow(im_convert(images[idx]))- d" T: e0 m( a; x/ K- Y3 N
        ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), : Y4 W" w/ `% H1 p) f" _2 x
                    color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))
    ( ?  p* V& o9 Gplt.show()
    6 v" A5 ]" k8 R  ^6 u# 绿色的表示预测是对的,红色表示预测错了
    ! v6 U  u3 ^0 T: Z' R' N, i* a1
    " L- u0 \8 t7 ?0 v; H4 v1 ?6 g7 W28 q. U" X. \5 M9 S1 o4 g+ o
    3
    # V5 ?5 q: x4 L' K# C4
    ) N  `" e6 C: l; ]3 |/ B1 F  [5
    + X2 c/ W2 E: C( x& u. G6% f. r( I) H0 l# ?6 z3 w( z
    7
    1 C8 F/ {+ g( T: b- P8
    ( |: [: }3 y: r! c2 Y( q9! ^+ q6 w1 G1 I8 h1 I" M: U6 Y
    10
    ; |2 R5 \) U2 O% U11
    ) T, Z' q( q5 W
    0 i5 o; J/ o5 A' C0 {( g% n+ W0 w# b8 k# p# M" R

    0 W) h9 c, b8 d————————————————; Q; y: P% d3 U0 v: m7 d. ~
    版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。- [9 W! c3 C" p+ R( N  z( F* L4 e$ h
    原文链接:https://blog.csdn.net/LeungSr/article/details/126747940. X9 U8 X& e* r9 x9 j

    6 p9 ^7 V* a( d4 A, E$ D
    2 K$ l+ n1 ^. |7 Z0 Q
    zan
    转播转播0 分享淘帖0 分享分享0 收藏收藏0 支持支持0 反对反对0 微信微信
    您需要登录后才可以回帖 登录 | 注册地址

    qq
    收缩
    • 电话咨询

    • 04714969085
    fastpost

    关于我们| 联系我们| 诚征英才| 对外合作| 产品服务| QQ

    手机版|Archiver| |繁體中文 手机客户端  

    蒙公网安备 15010502000194号

    Powered by Discuz! X2.5   © 2001-2013 数学建模网-数学中国 ( 蒙ICP备14002410号-3 蒙BBS备-0002号 )     论坛法律顾问:王兆丰

    GMT+8, 2026-9-22 15:36 , Processed in 0.877310 second(s), 50 queries .

    回顶部