QQ登录

只需要一步,快速开始

 注册地址  找回密码
查看: 2782|回复: 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)实战案例4 [( k; C5 B0 P3 O

    2 U5 x$ Q; `0 |# @/ A' u文章目录' ?. q7 ], k4 H' ~
    卷积网络实战 对花进行分类
    # \( D, {8 W# G- O- C! C3 k数据预处理部分! o) ~7 v' b8 v) s: w# p8 h
    网络模块设置# c& Z- b; p; h+ h! b2 X
    网络模型的保存与测试
    % m2 i* m( l  I7 `) [8 i数据下载:7 U3 K* e' b# P, I* _$ n3 J6 x
    1. 导入工具包
    , ^0 ~8 I7 h; `' f2. 数据预处理与操作
    # }; e. X5 ~: d" T4 ?5 q0 _( j3. 制作好数据源, x9 ]0 w, }" K/ Z9 y2 g7 x$ A6 }
    读取标签对应的实际名字
    $ j. R& K; b1 I8 M4 D, b4.展示一下数据
    % ]6 [6 G- [/ \# [5. 加载models提供的模型,并直接用训练好的权重做初始化参数1 {1 W1 ~( s$ u0 |% v' _6 u4 D
    6.初始化模型架构! t7 d+ o6 ^; d3 _7 `. ]
    7. 设置需要训练的参数
    ' b* _( I. l" \9 E7. 训练与预测  _& ~1 ~; E9 P5 |& M
    7.1 优化器设置
    9 P/ J9 M" r/ `7.2 开始训练模型
    " R1 a9 c& [0 j" f* G6 m! C7.3 训练所有层
    5 @2 F) F, t7 L) v  ^2 a, Y3 X开始训练2 a3 _% T2 P8 l' K3 G$ M4 Q* A7 b
    8. 加载已经训练的模型) i) F/ ^& X4 q; A( U1 K' I' l
    9. 推理0 X  }! X1 U8 K3 [5 |( j
    9.1 计算得到最大概率
    - a; J! k" X. {1 ^: J9.2 展示预测结果% s3 S2 T0 g% e4 `* g
    写在最后
    6 M# L! ^: ~. [% I9 T: p! m4 k/ |卷积网络实战 对花进行分类4 [! k9 l  i& Y; a0 n9 g
    本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能
      g! a$ R! h/ ]1 c8 {0 {+ r3 q. l9 z" \
    在文件夹中有102种花,我们主要要对这些花进行分类任务
    + T" a% Q" E) C! |8 T文件夹结构# s& W. W0 Z6 `, W) Q' t8 z

    % q$ g4 O( }& W' K6 h. Mflower_data
    4 [7 E. _- `0 b7 H% Q# {: t) }8 F) |/ B" i7 e; g7 r8 R+ k
    train
    8 B. F0 p% q) p8 q7 O9 v
    + f# j2 B' _. G# R1(类别)3 v( x- t- k# M% l. Z7 d! D* v
    25 e& G7 S& J( u* V4 q
    xxx.png / xxx.jpg
    : {4 z9 W' j$ Z4 r) ^. M" lvalid+ d& X9 b3 N. X& C! V8 G0 }
    & {0 [' |6 ?* Y8 D# ]9 m3 y
    主要分为以下几个大模块
    ( F5 O+ r- z) Z; O; G; F" Y; Q) C0 C7 n- B" U) N# W2 u
    数据预处理部分
    * h9 E9 S: q/ L; c! E3 L数据增强) l/ A! C4 p. q# C
    数据预处理
    - s3 \- X" q0 H; D" |4 r" L# t网络模块设置
    7 r! w' |6 P7 K/ n" J8 T6 D' H: r加载预训练模型,直接调用torchVision的经典网络架构
    / c8 S  M3 D5 Q2 S4 Q) W因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务& H/ h0 I0 X2 S1 D, d! c
    网络模型的保存与测试
    / S0 c1 [) v/ a( p. b模型保存可以带有选择性
    0 T1 Z  o% V4 B2 C& d6 u2 f数据下载:
    / t4 t+ c* |# u0 S, khttps://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset$ {/ l; s! W; t8 J! d

    7 r( W# g0 X1 d* n- Y6 P3 S- S改一下文件名,然后将它放到同一根目录就可以了
    0 m! S9 k+ h% }1 y# M
    . b9 ?9 E  z$ z下面是我的数据根目录; K- T2 L( B0 X* `2 _
    1 {* _) V  D' R6 c( }9 O7 g

    0 o1 V( c3 R6 R7 x$ l3 b5 k# k1. 导入工具包
    0 t# N, F2 o' U/ {: y; w+ Y2 X9 j/ zimport os
    0 {2 N7 l. o" [3 N: Iimport matplotlib.pyplot as plt* j! R2 ]7 W% T) t3 g& r
    # 内嵌入绘图简去show的句柄
    1 f; o1 G- }1 _. m) u  n9 i6 s2 x! ]%matplotlib inline * A5 Y: g( @+ c/ ~3 ]/ T
    import numpy as np6 |! W& J0 r& G: I% N, l' r
    import torch0 e+ ]: C- C! O
    from torch import nn$ x- \: P5 ?2 p: J1 q: Y- L# v# P2 R
    , T; y: V% M' Q! d8 ]! y
    import torch.optim as optim+ h4 ?& i( J4 A: s! N
    import torchvision# k4 @! Q& j( F+ S/ V
    from torchvision import transforms, models, datasets( H+ j: q$ p* Y* ]
    6 R; ^* G% P( }
    import imageio" Y, U9 i; u, A5 D' N% a3 L
    import time. s# i$ E; Y1 V# h7 U& w  b, G
    import warnings; V' G0 a( ^+ ^3 t5 a; Q2 ?& O* p
    import random' ?& m2 `  {% Y
    import sys& X4 \+ b3 e( Z) ?
    import copy2 a0 \! X$ l* ^8 h1 J( W- G" G4 ^
    import json
    6 Y' ?8 u( b  |from PIL import Image
    9 C7 `9 f' w" r) u: ~8 z3 I8 W
    * J7 m, D' b3 I& F  {- U7 L1 {- n2 ]# \, J, r% f
    1
    3 I. g, b" [+ _+ E# q7 {2, E9 M9 |2 z2 d3 {$ L8 m
    3
    4 ~/ P' L2 F, ^% ]6 g0 l4
    7 a0 _' E' e* Z& k! ^% e+ ]5 M8 F5% m1 @) Y$ W) ~( s" t
    6
    3 I6 h! s7 c  P4 ?6 L/ ^4 j3 ]7
    1 U+ q+ P8 R- e* O$ x0 [8
    7 w+ F$ a' V# c% v+ F  V/ V* w9
    & G5 }' Z$ X" y; u# e10
    - L& |6 d7 {' e( c11
    - a! s0 p- @+ O2 x$ h1 g12
    6 h( d) J* G, j/ f% C* a13/ d3 S: f4 ?  w
    14! N3 x. E5 F) U3 X2 g& h
    15
    : j: O$ D' ]$ U/ `7 X: \+ m16& ^; j+ v5 D. w4 C3 N
    17
    ! W* [( W; W4 B8 k) c4 @18
    - V! }% [8 h' t; b+ d7 ]19) G* ]5 s+ t/ T" O. h, M" `
    208 x8 `: A0 S  G$ w' A- x
    21
    6 K+ F% `1 z4 R/ D0 r6 X' j9 Z1 o2. 数据预处理与操作
    4 A+ r2 x  D7 t% M' L#路径设置1 ^3 u/ J. X8 @; a5 r: ?
    data_dir = './flower_data/' # 当前文件夹下的flowerdata目录
    6 x; w9 E8 F8 _! F( g" Ctrain_dir = data_dir + '/train'
    3 s. w+ ~7 r8 v2 Mvalid_dir = data_dir + '/valid'
      K6 u6 s$ |+ O1
    ) ]" ^4 P0 d! L2
    ' i. |3 t4 H% M3
    2 b- K. x- X/ q! F4. N7 F% O' I- m) [' x2 o6 I1 D
    python目录点杠的组合与区别( j+ F: ]5 l% }. g% l' n, F
    注: 里面注明了点杠和斜杠的操作
    - M# u  J8 |$ E: @7 J/ j9 Q  ]- ^
    ; q+ [3 s+ ^; _- }- x2 m3. 制作好数据源
    ! g$ o1 Z# A6 n5 s: h% N( pdata_transforms中制定了所有图像预处理的操作
    2 _0 j$ p& a1 Z$ l$ nImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片8 Q( x3 }+ |1 T: o& n- `
    data_transforms = {  O; B8 g& G$ O: Z7 W3 P  {  x
        # 分成两部分,一部分是训练5 y6 l) ]  L9 P" Y0 ^
        'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间% `; r- l9 }8 K) ^0 h
                                     transforms.CenterCrop(224), # 从中心处开始裁剪0 T7 v4 C1 h9 \# a! H0 Y# s
                                     # 以某个随机的概率决定是否翻转 55开: r1 t( l$ u: q- {1 [4 ^4 n  e1 T/ g
                                     transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转
    8 l3 p+ F: `/ e- r: m& o; N                                 transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转
    ! \+ H% \- y4 E; @7 W+ t                                 # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相, W/ r+ m8 k# a, F, ?- g0 Z
                                     transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),
    7 c: `' D  l) h9 e& Z: @                                 transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB& O2 `+ k2 ]- r. D% F' f
                                     # 灰度图转换以后也是三个通道,但是只是RGB是一样的
    * P2 g# F- y8 w" ^; Y5 @                                 transforms.ToTensor(),5 O6 q% p; Q/ r% \* U
                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差& \. e+ ~, W* U$ Y
                                    ]),
    ) d* F/ s' h$ ^( `+ F+ Z    # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化
    2 ^) Z* ~& d, Q, S3 L: {    'valid': transforms.Compose([transforms.Resize(256)," C) _1 d+ E; U& d0 I. m3 E3 n6 _
                                     transforms.CenterCrop(224),
    , [6 T/ @' r' b8 X7 c8 A                                 transforms.ToTensor(),
    4 a2 g' m  y8 r8 \. o& U# `                                 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同  N5 l1 w. F$ F* R
                                    ]),+ l% j  X4 L+ ~  F
    }
    2 S% W1 B' T$ v( l. y& u5 T+ j- D- ^, v* D, g$ x
    1" d* Q* m/ }3 [( L4 f+ Y
    2% R- _; b; N& a! d. a/ Z2 T
    30 P# ?8 V' R- ?# s
    4: T) k. y& M5 |3 s
    5
    7 f; }' J7 N7 J7 [6
    - B& B) m; L  z; T) d7, y5 x( p8 E/ u" M
    8+ b/ D, G" a! V1 U0 f
    9) \/ g+ F7 W3 X
    10; J, `9 r5 F. w2 k3 c
    11
    " g# h+ l8 ]# S( K6 S" z12! A; @+ U* l8 O/ z. i
    13% z0 D" l; Y' E; x. L
    143 ]& x( Q' g( ?( ?* [( m
    15
    1 @% E/ d7 \4 }* E9 r; W7 D1 S16/ P4 S* f) D8 B
    17* i- D( j- F8 `! K, E) M$ T
    18  q- ?1 u1 q7 v. f. L* T. C
    19
    " R. @/ D! ~, i/ P- v* \) O203 P& d; k( c  {  b2 ?
    21! z- W  ]8 Z5 C' ~' a
    batch_size = 8! h$ d# @- F0 _+ S. G* y* Z
    image_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}3 D" ?1 q# z2 S  f7 b
    dataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}
    6 g5 _# l' K/ a/ F. q' @6 I, V0 ndataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']}
    ! i+ H! j: p$ \3 H4 F+ ?class_names = image_datasets['train'].classes. k/ N$ O8 h+ R5 q- L

    4 _7 R9 }# e' d" k  ?3 t" W#查看数据集合
    * v# `+ n7 A3 S& K. qimage_datasets% J6 [/ o: J( v
    - m' X( h. S( n- s; H
    1! [2 c8 m1 _+ z/ M5 G4 s* [5 Q
    2/ u; ~6 i; R# P* v# r
    3
    & u  L) h" v) e' S- D* ]0 l! j; O4
    " k* E- i: s: w& e50 ]' f5 C3 k9 i8 ^1 d: }
    6
    * D! ~6 b) H# E" r' A5 l( H' T74 g: O8 n% i3 Y' z1 j( y" Q
    8
    , o$ f1 E" J  `% [5 K( i97 L# @4 c3 M% \7 S/ q$ ?* b
    {'train': Dataset ImageFolder
    ' `5 O# y/ w7 _1 {' N2 s     Number of datapoints: 6552
    6 D' |( v4 M4 m7 ~  K: F/ J) u     Root location: ./flower_data/train
    : ~3 C; N8 w$ e  x2 S( k$ J" `     StandardTransform, o* Q+ h" Y. V; W
    Transform: Compose(
    1 e* L3 {8 R- }8 T! k( r" }                RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)
    % O3 d6 d/ K4 `/ \; S                CenterCrop(size=(224, 224))* J3 P' |& t6 f. A0 s, C3 x/ c6 ^) Y
                    RandomHorizontalFlip(p=0.5)
    & J1 r. l6 m+ o9 U  _                RandomVerticalFlip(p=0.5)
    9 X8 m2 M; \# N                ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
    * |; _0 {7 h* c4 o+ r! v- N                RandomGrayscale(p=0.025)0 g0 @# {# u$ _1 j+ A, D
                    ToTensor()
    6 U3 [0 S3 ~& M( T( y# ]5 u                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    ( q) [) d  F7 y( U; m            ),7 n" n, x: n( ?* k( A: ~
    'valid': Dataset ImageFolder. M% G1 D2 ~6 S0 _% R
         Number of datapoints: 818! y+ ~4 O+ `- W: z" F
         Root location: ./flower_data/valid
    ( h' b7 `1 j. s7 G6 H, s0 C3 d% I& N     StandardTransform# v- ]& J2 r+ A- K* }/ l
    Transform: Compose(
    2 a% ~9 @; R% c- `7 d. c; q                Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)# ]+ r; E5 a1 L5 M
                    CenterCrop(size=(224, 224))* Q7 e$ j2 {+ Z7 R+ ]
                    ToTensor()
    ) L) l4 @+ y& A3 X                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    . e1 C9 @- Y8 d+ s9 W            )}$ `9 N' r& J( q# m' e6 R
    ; @: N; i, f* m1 E& l. j1 M
    1; O, h3 a4 \' P  V) u5 O
    2
    8 \7 i( ~0 N4 D9 h: I- ]3
    2 c; o1 Z3 r, x4
    ; h" c: Z) ?3 r) U- G( Q* I) R54 a4 X' z, y' S6 H
    6
    ; z7 {* ~. q7 N7 }. A( ]& R# E7
    * Z- u2 s' Y1 _# ?  m. I* p5 U8
    0 C/ H+ q2 V3 ^$ u) e. T5 j99 b6 y, ~6 \. u) O" Q' j9 a# P5 z8 a2 J
    10
    . ?. Q; q  J  v1 L& O0 [11
    / T: Q& f$ f: i( t# F; A% _12
    9 d! @8 Z& I4 V( P: O+ \13
    ( H4 p* V1 M& j7 n6 K8 N4 J- r8 f14
    0 R% F- o' t( h5 G. m4 i! P0 U& S3 G( o15
    + k2 S7 t% m6 v- k$ B3 s16
    5 ^6 r- Q% E; Y# ?* y17
    . M1 m9 \1 @9 w( e2 v0 \18
    + W7 J9 x: }* ?8 n8 z) D2 w19
    6 Z9 n1 }/ s$ W( F  ^2 S# |. Q* ?20
    2 H+ ?2 z2 h1 B4 @( I' a21
    ' ^' ]/ {) k7 q7 N* C  X5 l22( J2 b  U8 a( A. }. J5 C9 ~, j
    232 u& j1 Z& S1 Q7 Q9 y
    247 s* h! x/ ^$ r5 v2 g$ N
    # 验证一下数据是否已经被处理完毕0 Q# H* C: {* A  d; i
    dataloaders
    # ]0 ^. G# R$ s5 J" d1
    ! C7 N# _) e$ V; J2. ?. u, w* s8 n9 }  E
    {'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,* u3 O8 G% P9 L* Z( n
    'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}+ g6 q! P$ z; r+ _) P
    1
    + J1 J) T  `! I4 k2 |7 _3 R2
    * h' c& ~3 k+ d. y# Y# D2 Odataset_sizes  W" I9 z" b2 ]* h
    1
    5 r& s; U7 Z( g. }. l{'train': 6552, 'valid': 818}
    8 x. L8 f, f) k" V4 z/ @6 ^# ~& L& |1
    6 _! y; S" i" E: T读取标签对应的实际名字
    - {3 @1 N6 S( z: a使用同一目录下的json文件,反向映射出花对应的名字
    1 d; u1 j9 k' b! b3 K/ A: B2 |0 Y3 E
    with open('./flower_data/cat_to_name.json', 'r') as f:
    4 i! [! n+ |! u# c+ o1 o0 ^6 C+ u    cat_to_name = json.load(f)/ D/ B; f1 U8 P$ k
    1
    % @' z4 O+ f" s0 O& q! D24 I+ u! Z5 H. ]! o
    cat_to_name
    - e* P" K1 [% o! b; K+ }/ [1* W5 k4 \. j6 Y# `, p
    {'21': 'fire lily',* T3 F6 w5 _4 I# F- N, x4 ^
    '3': 'canterbury bells',
    ; w! k9 h0 D/ _) ]6 M! O: B '45': 'bolero deep blue',5 K7 |( V- m& j& t
    '1': 'pink primrose',8 C* w7 V/ ]2 @# O
    '34': 'mexican aster',
    2 E" o$ L. m% j, `3 G '27': 'prince of wales feathers',
    1 L$ N! k2 ]& Q; q6 \0 n '7': 'moon orchid',
    5 ]+ l, ]5 g, I& C( T '16': 'globe-flower',6 X: y) r! s) V5 G" Q' B
    '25': 'grape hyacinth',0 o( H3 u9 x( |& ?, d
    '26': 'corn poppy',; G6 n5 m4 y2 F0 X+ B9 b' h
    '79': 'toad lily',
    - d& t2 k( a" I; q& n' r. G7 _& e '39': 'siam tulip',
    / `9 S! @; b2 k+ ~ '24': 'red ginger',
    . i* I- `$ f2 I' g$ c '67': 'spring crocus',
    7 q1 f0 K2 K- `- x5 ^& J* R '35': 'alpine sea holly',
    3 f& `9 s3 R$ j' P1 Y8 H '32': 'garden phlox',. L; c  H, Z& g+ }. H' Q
    '10': 'globe thistle',5 P  t; _$ i* i8 R
    '6': 'tiger lily',3 A+ J8 b5 b  K% Y4 P# W
    '93': 'ball moss',# t' e* K: @3 v6 L
    '33': 'love in the mist',( A+ G, L+ _6 G1 H$ h6 H
    '9': 'monkshood',
    / o5 }- K0 X2 n! i, | '102': 'blackberry lily',
    4 L# r9 S/ {. b$ O, {& y '14': 'spear thistle',+ E$ A2 T# P0 o+ k
    '19': 'balloon flower',
    - t* m" b) e- b; {: n7 f8 u, O '100': 'blanket flower',2 c, ~! Z- d) @- i$ A
    '13': 'king protea',
    3 L5 P* G' `% ]( |3 {% [: i  K '49': 'oxeye daisy',3 X% `) f+ q2 ]$ ~: P- D& h% N
    '15': 'yellow iris',
    & _6 R& ^+ q# J( ~9 V2 d" T3 f '61': 'cautleya spicata',
    + J1 ^4 `% E" v7 T '31': 'carnation',1 B6 j$ d* h" b& N8 w
    '64': 'silverbush',
    ; U* H& Z" E. M" {6 | '68': 'bearded iris',
    4 w* _- Q. A) N '63': 'black-eyed susan',$ P9 i* d/ u, y1 }
    '69': 'windflower',
    0 U( X- e' h4 V' o% _/ o '62': 'japanese anemone',; I2 f2 P% f0 \( J) a
    '20': 'giant white arum lily',: M  t2 ~3 i/ s
    '38': 'great masterwort',9 E; Y0 ~0 S  k
    '4': 'sweet pea'," W" `6 D, N6 ~$ [6 v+ a2 }
    '86': 'tree mallow',
    ; [% w1 M6 S/ k' m, d '101': 'trumpet creeper',  k6 v) c# f: q& J  R
    '42': 'daffodil',
    . T: e' n* L# x '22': 'pincushion flower',
    " f+ G0 [! ]% {. Z! Y2 M2 u" m '2': 'hard-leaved pocket orchid',& G3 a! o. a& x8 f) g
    '54': 'sunflower',% B! ]8 v* y2 M# N0 S# u
    '66': 'osteospermum',
    ! `1 R2 c7 N& T. y4 o& G/ Z '70': 'tree poppy',3 h7 O# h) \$ w& x
    '85': 'desert-rose',
    ; j, j+ q( E& D* B! M '99': 'bromelia',1 b" K4 m, R" i8 Z# j* q! N/ a  Q
    '87': 'magnolia',$ I1 e9 p8 i8 p0 y5 ^1 V
    '5': 'english marigold',
    , X; e) }5 E3 e* E8 | '92': 'bee balm',
    + ]+ i: n9 k! t. W1 Y5 B* Z6 w '28': 'stemless gentian',* Q' y# v- @# ]8 o" H. k
    '97': 'mallow',0 ]* x; U  t2 W7 j# R  g* W
    '57': 'gaura',
    7 x# h9 y6 J% Z9 g '40': 'lenten rose',$ D: e5 L. B9 s- y: T
    '47': 'marigold',
    ) u$ T$ n+ Z  ] '59': 'orange dahlia',
    ! A2 v5 r+ ?  I) [ '48': 'buttercup',( G, `2 y6 U" W
    '55': 'pelargonium',6 e6 G) I" V' `0 B, R
    '36': 'ruby-lipped cattleya',+ d& j" Y5 j6 ~7 t' I: D
    '91': 'hippeastrum',
    & V  d) _& I. J* [0 ?* I% m! a '29': 'artichoke',2 h, S6 Z1 v3 x# d) J6 t
    '71': 'gazania',
    ! z6 e* L( y0 P' t! D: N0 h '90': 'canna lily',: Z2 h8 {/ n, _5 M  Z0 I1 f# b5 l
    '18': 'peruvian lily',# p# d# t& Q$ R! ?! Q% f
    '98': 'mexican petunia',6 _% M9 l$ [& o7 d7 c
    '8': 'bird of paradise',' n' c# g$ j# y; ~+ d; {* L
    '30': 'sweet william',# M3 r- S( M: t! o1 |: s5 U
    '17': 'purple coneflower',4 G0 U, [, ~5 ?5 j2 q  c+ J/ e
    '52': 'wild pansy',; H; H7 a* T: _# T, G
    '84': 'columbine',3 s$ y% E5 S9 p1 X: N' f. E8 T& M
    '12': "colt's foot",3 s0 e' U4 V$ Q3 r5 J) r/ j, N
    '11': 'snapdragon',2 d- Z+ P5 K+ W2 _+ b1 W- d/ {
    '96': 'camellia',; O& }% g6 t4 v$ ]
    '23': 'fritillary',
    ) i' m6 y8 t& ^. E '50': 'common dandelion',
    4 w3 I: W2 _) _8 \, X '44': 'poinsettia',
    9 ~/ _5 L- t# I# j/ ~3 O/ o; b '53': 'primula',; B1 h1 j8 H3 y/ a1 _& a5 u
    '72': 'azalea',
    % y  r9 m0 @: B% i" R! f1 j! G' i '65': 'californian poppy'," |$ ~2 p9 g( P
    '80': 'anthurium',
    % m; h' K( q! `! B: |( L$ ]2 a! F1 _ '76': 'morning glory',
    # t6 M5 V! P& c '37': 'cape flower',( W/ N$ f9 _; a, q7 X
    '56': 'bishop of llandaff',
    5 j: ^5 j7 Q" D9 }4 y1 @' s6 Q '60': 'pink-yellow dahlia',
      A7 }) z( N! ] '82': 'clematis',$ @0 p" q) ?3 n5 a- h- R
    '58': 'geranium',
    % m3 z  q+ ~* ~* C) y '75': 'thorn apple',% }. y* d- a! X. ~
    '41': 'barbeton daisy',
    " `9 {/ N3 l  g/ U/ w; S0 P. X" V% w '95': 'bougainvillea',
    # c0 ]3 k" E. d. N+ p '43': 'sword lily',! K  o  |  x. A% j, g  ]
    '83': 'hibiscus',9 W& [: V& ]. ~: A4 H( O
    '78': 'lotus lotus',+ ~! ^+ v' _5 q5 m
    '88': 'cyclamen',
    / ?7 z: u/ R- y# ~ '94': 'foxglove',
    . s. }7 H" c5 L; i) e7 ` '81': 'frangipani',
    9 I! K  f6 B% p. B5 B '74': 'rose',
    2 T; E' T! v* B) N% |: \ '89': 'watercress',
    6 t  `3 A! y" p0 H9 n5 b '73': 'water lily',
    & F7 W4 r! C" | '46': 'wallflower',
    $ x7 V, q( ^0 R& S" ^1 B '77': 'passion flower',
    ) @# m$ `7 n% G( Q '51': 'petunia'}
    + c7 N( T: W. l+ Z' I2 ~
    + S5 x7 ~  \9 I3 I! J: k  v15 H# Q' E2 t/ H* A. j
    2
    3 _) R( F: C9 R7 S( u3
    * x8 W: i1 P, o$ i  m4
    1 \% N7 O* a7 t# v5- U1 m- B; e0 u8 U* L; B, m
    6- B2 D, f, l2 z; {$ N. I" A8 V
    7
    / f5 {$ c9 Z5 E8, W. d$ O; u  N$ t" J9 L# V% ^
    95 Y! P$ X% W9 W
    107 x0 y  m( D: p" a) W
    11
    , E3 u% W. O6 n3 y& @" o" ~12
    ' k" y. r5 @7 p/ P/ T13
    6 Q: N6 w4 X: \$ g! Y7 j5 I14
    # |2 Y$ F/ a$ W/ i9 L! k+ S* n4 e15
    $ v! K1 M. C4 a; n. D( l16
    & T/ z. e8 A; |+ u; H176 D% P- w9 ]) [
    18
    9 p" v' c/ X* \) G/ F19$ W* b& ^9 t1 d8 t# @
    20( f) L/ q, m9 Z0 g; j1 l' G8 d
    21
    2 X0 d$ ~# p- k- |  S$ l22$ [" m: a4 D2 h' N; l+ \
    23
    ; o4 V5 x4 V2 O! K1 O* l24
    3 m4 }$ k0 T/ _: h25
    7 ]: j2 b/ r$ y5 Q: b2 H26  P7 e  q  w6 S7 y
    27
    5 G, W: h; n! ]! s: w28
    , B# F- p2 s- ?4 u29/ ]0 S* l% i8 ]8 L2 S
    30
    ( V7 Z$ g/ J" t5 Z31
    8 q6 |" m& V* n0 G8 X) h3 d& U  `32" E" h- n2 h2 Z+ r; N
    33
    / v2 y8 Q2 T% v34
      F1 a& |9 ~, n' h# U9 Z! P1 P1 @354 r2 c9 }, v  B1 @% J5 T% V. |
    36
    , _" }: N$ A! x# C: x0 X; N37
    6 ]+ e* A  d& B! U+ N2 H38
    ) j9 o' k9 [( w! v6 X7 Q39
    5 \. H. z; t, J( V% }! ~) Y) J40
    6 v* A0 x+ `! ~% x41
    + D4 D, }. E& p& W. }4 x, l9 m42
    # ~3 i2 y' p% h* I5 o& b& A+ r43
    7 \3 m1 n$ w! H44
    4 O# [- H5 R; p1 C& n6 U- F. g45
    3 ]9 A+ ^7 t' i. g; }46: f4 A8 ]% l. Z! s) q0 l
    47
    " l8 l8 ?+ d5 x1 }& H  `48# m1 `; ]" e5 {( Y
    49
    % }* F) W/ S! o50( o4 N- b/ \9 W, P, {4 M
    51
    , r2 n' C* R' [% e52; [, [" i9 U. j
    53
    ; ^+ r1 W9 n% a9 Z54: a5 P- N1 w& Z  b) F) a
    55
    5 ~7 k' w$ d$ n. M8 j56
    * {2 j# P3 C8 ]1 G5 z, x57
    2 p6 Z9 u2 [! u8 [# ]( h58
    3 P' i2 o9 j. k+ i! Q59
    4 I1 G( W9 Y' K5 G; \5 r60& W3 `, G$ k5 S% u7 l1 a
    617 k9 B/ I/ v$ s: a: H# I' \$ |
    62# j6 O# b( R: j- C, Z
    63
    $ a7 M3 G, s+ D, t. y# ^4 z0 [64
    " Z( Q" X0 @9 |; d6 o% M; U65
    0 y$ ]9 m6 I: B66# j" c0 H+ Q& o( B
    67
    - g( `' q, W) W+ y9 E$ }! T8 W68
    5 p5 u% C" g8 W8 o" X( L69
    - S( c/ z1 I# ]3 p70$ Y$ U, J! L( k2 Q  j4 ]
    71/ l4 i1 b, I; R2 W
    72
    $ F9 ~3 G; F8 D" d! U: w1 Q73
    / h! n1 f# f' K9 ]74
    7 \8 {& M( l* X; p" F; Y( A753 i  U5 f( M: \* I
    76
    , V* a8 o8 v3 f& ?! _; n772 F3 u4 O( ^9 c* S9 A
    78
    1 H8 Y- j3 l3 ~" Z9 J, K79
    0 H, S+ m# q0 v  s9 j- m80
    ' n4 V  B$ e0 J+ y3 x  j817 X! L7 t+ r9 y' T% a
    82
    % h3 ]% p, u( P, y& S5 X83
    1 x8 r3 x- t( Y+ ^2 L. E84
    ! S- W: W; }$ n4 x; G4 ]/ w85! X' G$ @: f0 D! R, U
    86+ G' I& {4 R( I$ z
    87
    ( R' H" {6 h& u2 q3 h4 V# G; F889 V3 r; w, n8 z5 V% C
    89( ^' q% D! V  k: V! m
    909 a' E6 v* o6 D1 L. g& P$ h
    91
    4 d; a; c  q, C/ H92! X6 c. v: X3 m/ z, `5 x' X
    934 S4 x  l8 t4 q9 B: G
    94
    $ `. D9 z' y( y2 o95
    / w% c# Q) S. j2 l2 \965 F! G& l: X% r: f* D4 L
    97  V' c3 W0 |6 x* c& u! I$ ]
    98
    * ^8 |5 ^+ Q/ R5 i99) l, x9 n- B. }
    100; Q8 U# g; _# W4 D+ \
    101
    9 u) D0 D. m. A$ m102
      t# X: W! H5 v0 v: T! m9 ?% i4.展示一下数据5 E9 ^4 t/ F+ H# Y4 p
    def im_convert(tensor):
    $ m& ?& t5 z* b( u2 I    """数据展示"""5 u1 {0 p/ p& y5 ?  ^
        image = tensor.to("cpu").clone().detach()# n6 O/ W6 d& D: N; Y5 D: q; m
        image = image.numpy().squeeze(). ^: Y6 y1 p6 a( \; R* g1 y6 O
        # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图
    " V' v. v( O+ {  A    # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)+ G$ o+ C; ]3 ?$ j- i; |
        image = image.transpose(1, 2, 0)4 u) C/ K: h- [. H
        # 反正则化(反标准化)
    7 q/ F: E! x' m( D9 f) m& ?    image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))
    . P0 H8 w  l8 @4 W7 A1 i
      U- F# w" F; Z+ J& j    # 将图像中小于0 的都换成0,大于的都变成1- C3 M  K$ u# _. U6 e0 ^
        image = image.clip(0, 1)/ @- f! @5 v6 b4 i
    ( A2 O8 v0 s+ X5 _; O9 N) j
        return image
    - E9 Z6 ?* D* G* B4 j1
    " k9 A1 f. Q$ r2 e; b2
    ( G# u- q6 [) ?- m5 Q3
    ' ~2 r; s5 W* C. E1 u% M4
    * b. ^' b7 I; {9 v5
    6 O- J  W3 s1 i( d- ~6* H0 K* R6 e% l$ w
    72 |1 S; Z7 x% i# f+ k$ h
    8) P8 |( I& N$ Z' t2 M0 g; u
    90 q# m! R( M& }$ }1 U
    102 g9 Y# g' M# ^: @5 s+ m
    11
    # T2 [5 S) |  M3 H12
    9 @) s& |0 O2 Y9 B13
    6 R: I! ]3 r; Z' E* F( z/ }144 `  |( }0 u# C7 G/ p
    # 使用上面定义好的类进行画图. B* M: t$ V2 k
    fig = plt.figure(figsize = (20, 12))
    0 |; B. k7 e3 }4 H. v4 pcolumns = 4
    : ?0 w1 Y7 h$ ^1 u2 h, ?. hrows = 2
    4 S0 p/ m2 [4 g& \  m  }' q" e  S2 \$ E7 j% ]  C
    # iter迭代器8 V( L/ s1 J! n* G) |
    # 随便找一个Batch数据进行展示
    / ^6 O- j& Z7 n& Jdataiter = iter(dataloaders['valid'])5 B6 \2 P" R( X' y5 x
    inputs, classes = dataiter.next()
    8 H- {; p: C; B3 z- j" P9 ?1 Q8 R. }% v- Z
    for idx in range(columns * rows):
    " ~5 l9 Y, h* A& {# _    ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = [])
    / W7 r5 n+ A9 r  Q* c' A/ ]    # 利用json文件将其对应花的类型打印在图片中
    2 i( L- P0 u( Z2 Y) q    ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])& @  o6 q# _5 K: F) s9 j* C5 y0 n2 d
        plt.imshow(im_convert(inputs[idx]))
    2 e7 m$ f, ~% D5 K* ~plt.show()
    ' L' Q3 e, }9 b& r4 S) I; S6 e3 b4 i( [" N
    10 z/ {* O% v: y6 F7 e
    2* W7 t$ m$ h1 P% w
    3% {' {2 S/ C4 r/ e' W3 h: v
    44 a4 M- A% u  Q" W, u1 ]- F+ }
    5
    5 j9 C: K* a% `0 q1 L6
    , G# b6 I* h* ]0 n7+ w4 \+ z+ n; P+ A/ R5 a
    8
    ' h/ e& g& @0 V4 ^; w9 S6 r) n9  Z6 r2 c, h4 o1 i  `4 i* m8 ~
    101 D8 D- H" e2 B9 M) }( L
    11) i- T/ v# R0 [+ V  O
    12
    9 v6 K* o8 \4 Q/ m: ]. n13% ?7 c9 |" Z7 b  {5 R
    14
    / C' D1 `. w8 y( s; R15
    0 s4 m- N$ o& X3 u& ^16" s; j; ?9 p& n# p4 {) y# T$ Q

    * A8 ^( {4 R3 q9 a- B8 d- n, h) s/ w$ o4 ~  i8 N- i# U7 W; [
    5. 加载models提供的模型,并直接用训练好的权重做初始化参数3 r4 ]  a/ y: ^. B- q. o/ X
    model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']
    % B# W3 }- w# J) q# 主要的图像识别用resnet来做1 H1 S0 y% w: r8 R- U! f' V2 ]
    # 是否用人家训练好的特征. f4 c/ i' I1 S: f- j* h3 l6 q1 Z
    feature_extract = True5 g4 B6 t, j' ?* d8 |+ M
    13 e2 o  Z. h; T0 m# b1 Z+ T6 N
    2
    ; o3 {4 g$ ~# n9 b4 N+ p/ B/ O. _0 d3
    9 {; U7 d1 D: L- I& z4
    % u' }, [5 g1 u# p/ F& D5 V# 是否用GPU进行训练
      N2 }9 o* f7 M: t- u" ^5 Ztrain_on_gpu = torch.cuda.is_available(): b, H+ R. M" _& E) W
    - F2 A' Q% _6 l! v/ c1 e
    if not train_on_gpu:
    % I8 x. V1 K/ x0 p% B; h8 ~6 w5 S4 n    print('CUDA is not available.   Training on CPU ...')
    ( T' w/ O2 Y& ielse:
    % L& _# t& K/ X( }+ G% J    print('CUDA is available! Training on GPU ...')
    % ~6 u. }  R8 W4 Y' }5 J- \; L
    1 W7 s: _3 R0 e" q7 J% v" fdevice = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')* F8 @0 x  j1 U6 M% ~
    1
    6 V  |, |; j/ w6 I2
    ( u! I$ O4 h2 @( }" G3
    3 Z+ J+ X1 Z- s! I0 V9 [4
      O  k+ P; w' {6 g6 S. Q5
    9 I$ Q# m0 Y2 g& U& t0 m% U, h6
    $ `/ |+ y5 ]% j* H: S7( j0 d! ]5 n+ x+ `, m, O
    8
    ' Q! z: m( Z5 C2 z+ z9 }. C1 d9
    # o; I3 u8 i. l& _CUDA is not available.   Training on CPU ...- H* H# `! a2 U
    1& [6 x) ?+ m: w3 @0 W8 Z
    # 将一些层定义为false,使其不自动更新+ U/ M: f. k$ C' [& P
    def set_parameter_requires_grad(model, feature_extracting):
    $ ]9 g' q$ O* `* a8 ^( I# A1 p5 o& w    if feature_extracting:
    : F' L! Z7 T3 n# x# s' Q5 z* i        for param in model.parameters():
    " z6 D! C5 z6 A1 M+ N# o4 n, W            param.requires_grad = False
    3 a* y2 ^- {, z0 j/ U/ }" j1
    & H4 y" f6 t3 ^. Q% ]4 d2$ p1 y1 q3 l4 Q
    3
    ' @. J6 k& y1 K4
    ( E7 y7 E9 I7 P1 t- ^$ I! V5
    + i% K' P; @8 a9 k9 C# 打印模型架构告知是怎么一步一步去完成的
    : z  ~1 E  k$ G! x# 主要是为我们提取特征的
    9 P! E; f! Q; C" v+ w0 d4 [2 L% w% m* ]1 Y4 C/ |8 }9 s& z7 E
    model_ft = models.resnet152()
      h1 x, z1 |5 b- K4 f, R  hmodel_ft
    7 B  E( Y# p- Y18 c, [9 c3 [  Q$ U2 ~  q' y! w- j3 _( \
    2- C0 w- r0 c1 ^+ [
    3
    " ^. I) Q" W. k% ?9 p, n4: Z" x) J' i7 f  C% M
    5# V+ H' ^: d. j) K5 e3 v
    ResNet(5 T. H: W  _8 J5 s
      (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
    3 k* g3 i  B1 x! r  (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)/ \9 s' v. `' g8 q8 f1 z
      (relu): ReLU(inplace=True)1 m4 h5 i& n- u5 k. k. L3 o
      (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
    - E6 v  D5 D/ R  (layer1): Sequential(
    1 ?9 S; n0 g6 R( _4 f# V3 j$ J. w    (0): Bottleneck(: c2 L- P. ]6 }6 G6 I: O
          (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)4 ]. \  b7 F( s2 g2 x- x7 {
          (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    ' B& U' g" i: U& \      (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
    * ^4 ]3 H$ n- }( @      (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    4 _8 q: h* E- L7 \" L3 P      (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)
    ! ]9 K1 C+ t( M2 I, w7 N      (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)3 R8 j( @% _& ~
          (relu): ReLU(inplace=True)
    9 N3 M5 t4 Q  j4 l$ a4 B4 T$ T      (downsample): Sequential(
    ( r" D; n  z! ?& h" J' }- \        (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)7 N' U5 w; }6 Q: ?3 x5 z
            (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)7 |% z2 y  @' m! B  O
          )
    4 m6 P. m$ W& o6 r6 V8 f    )6 u, Y1 Y0 l" U1 e  F2 Z( `
    中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。
    8 {% _1 F2 R# j; |0 [9 b0 r( w    (2): Bottleneck($ @# J. ?, B5 Z! f* l
          (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)& I% J$ m: X& F+ _' [- b
          (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    ! N# A4 z, N- ?) k+ [% X) |0 S      (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)( T( i3 C/ {( G
          (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
      U1 s- [/ T6 K! {$ ]& s  r& ~$ X      (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
    : H; h- o! {$ q! J# D/ [      (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    $ i8 k) _% ^( A6 J1 s$ ^( f      (relu): ReLU(inplace=True)
    " m9 j# h! \( n; V6 p( B    )
    6 N/ ~" F5 H8 ^" ~  )
    0 P0 L* u* e- M  (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))- u% s% a. p  y& @; t
      (fc): Linear(in_features=2048, out_features=1000, bias=True)
    8 N" K1 p4 l6 H; y5 ^, U1 A6 L) h)- a0 M# }1 \: V6 _, p2 y$ |
    1 D1 X1 E4 M# n2 s* b
    1
    $ o9 c. i# R" y' F25 M1 n8 O2 v, g+ \
    3
    ! V3 u5 m* M4 s1 d4% }2 L: A& |3 n! t% |: [
    5
    / G+ [* E3 }  L5 e7 |2 g) L61 X6 C7 |6 ~8 _6 g- x
    7
    - N4 N) t1 c) ?1 c; u8
    6 P) @) g& v  ^* \9
    6 I& s  d  d7 K10( ?2 a2 I( e" w* v) ?
    11+ r, |3 x: T& K9 {- O; F. N# Z
    12* k/ c8 c- P3 ~: B6 T; X
    13  U' I' e, o0 B; }
    14
    0 G; a, O$ m: X( p, [; B  E15
    * |- i! N: y$ ^' }16' V$ D8 `6 N* @: V5 k$ C
    17
    9 d& N+ e4 p- m+ g9 O/ K9 a7 v+ Q18
    8 n7 ]: h& l. S: p! B19& u8 D1 V. H3 g; w( Z
    20
    + x$ P8 y! w* ~6 r21( W2 L2 m5 Q9 V0 _  h  a
    22# \2 d" O5 ?* L4 B% B" U$ u1 W% b7 v
    237 D! c( ]# X' H$ [, o+ t) l8 A" D& O
    242 U+ G' `" A+ V7 f4 ^- z! x0 n& R
    25( z1 j. u: C2 }. Y) _5 @: Y. A
    265 \) E. n$ T  s# L
    27
    0 V' z; |( x: l* ^28
    # J7 j! g% p$ [" T. t& `8 m3 B7 P298 ~+ i$ a) \# v6 D0 X9 g& g
    30
    * L3 M2 k. b+ o* h, Z" z+ Y31
    ; ]) [1 p7 s+ y' K$ k5 ^( P8 a32
    6 `: z7 I9 O! R( B( E! ]4 |338 ]/ _7 ?3 m# e0 P1 [  t, t* V! q
    最后是1000分类,2048输入,分为1000个分类
    * i* e6 D4 @# b. X6 G0 a* K而我们需要将我们的任务进行调整,将1000分类改为102输出
    6 u6 w1 A2 X( y- R
    - Y( i6 x" U( Z( \6.初始化模型架构
    0 N5 ~6 }, Y2 T0 M" X8 Z: `* q步骤如下:
    / f, H# [0 l1 m. S- p* g+ E9 [! t* l6 h
    将训练好的模型拿过来,并pre_train = True 得到他人的权重参数
    - K& W% m- W' K& F3 B2 e9 V) f' N; l5 J. A可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False); \$ F* K! u: l; U
    无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数
    & Q5 n7 R% c  r官方文档链接, p, @/ S* ?1 @7 }7 ?2 o1 Z" ?
    https://pytorch.org/vision/stable/models.html
    5 m' p, j4 ]1 Z! M  \+ R3 s* \. h+ }! v. A# Q
    # 将他人的模型加载进来
    6 D$ c- W  M4 k4 e; _def initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):
    4 t$ v+ V/ v8 |) n$ `    # 选择适合的模型,不同的模型初始化参数不同# h6 i1 S5 s( ^; G( }- K
        model_ft = None3 w2 m+ X( ~$ @* C8 q9 q
        input_size = 08 g/ h; \# t- }/ t& E5 t

    % B1 p$ s2 Z7 ~    if model_name == "resnet":! `) ?0 S# C6 \/ o
            """
    $ E7 I9 E7 ?4 `5 |8 o, t/ n* W6 E        Resnet152
    1 b# d6 K* b& @7 ]) p" ?        """# r# J$ P/ r( u

    2 S$ m- t% x2 P8 x        # 1. 加载与训练网络+ M5 h" x4 N" b8 ^. I) ?
            model_ft = models.resnet152(pretrained = use_pretrained)0 y3 t& V! N8 J* Z5 i# u9 S
            # 2. 是否将提取特征的模块冻住,只训练FC层1 e6 d' m' l- O( i; f) U+ K- |9 \
            set_parameter_requires_grad(model_ft, feature_extract)
    % w# s) ?0 {; f' V        # 3. 获得全连接层输入特征
      s( Y, s% I5 `( I        num_frts = model_ft.fc.in_features* I" i4 \2 n7 b6 ~/ Q5 y3 n2 T
            # 4. 重新加载全连接层,设置输出102
    " ]+ \" g* V1 Z  f% H/ Q        model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),3 ~5 Z1 U/ e( f5 d9 [& B1 [
                                       nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为13 y- h+ y: i8 I1 b' P* r
            input_size = 224% u3 T" [& p: _8 z
    ' l* C0 G# e' s2 `
        elif model_name == "alexnet":! C: M& Y9 E5 \. a& A
            """
    ! J( p- T/ s; J6 r; J9 F        Alexnet
    % d9 O& R" v8 P, \        """
    2 [: \& U. o7 h0 P( ^" }8 C        model_ft = models.alexnet(pretrained = use_pretrained)! a3 ~9 u9 X' G- G  Y& n, ^4 Y
            set_parameter_requires_grad(model_ft, feature_extract)( U* O( c- |! U( z" t$ [0 q/ O
    2 u. x) e+ [3 d  W4 b6 ~# S
            # 将最后一个特征输出替换 序号为【6】的分类器5 y" e9 z: |2 Y+ W7 i  g
            num_frts = model_ft.classifier[6].in_features # 获得FC层输入, _6 c- D! g. Z! }
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes)1 Q' `& T) b- r% P
            input_size = 224
    6 C' K6 |2 T7 o9 ~6 g9 c& V! t% B' N2 ~* t9 {  {
        elif model_name == "vgg":
    9 |; X; ]# V" M: p0 i/ u! `        """& E( O* |4 T: [; X  b1 D" r6 G/ N  p
            VGG11_bn5 B- d9 x. G/ M: L2 _
            """
    ' p$ S! v. k2 w* Q4 Q3 a; k9 H, R        model_ft = models.vgg16(pretrained = use_pretrained)
    7 i2 I2 k  f9 t4 r        set_parameter_requires_grad(model_ft, feature_extract)
    & X: M( r9 v0 V: W        num_frts = model_ft.classifier[6].in_features
    6 E! M. o* |. X" z$ E) v        model_ft.classifier[6] = nn.Linear(num_frts, num_classes)' l2 W- x6 S& w
            input_size = 224. {( C4 D( N% D
    . Q6 D+ M# _( f3 F; E
        elif model_name == "squeezenet":7 j) q' o2 z* \9 J5 C
            """2 _! T) x2 I( o0 e
            Squeezenet
    : \5 [5 W- {9 K( f+ J) E/ A/ S- B, E0 N        """
    ( @4 r  E7 S" V1 p5 C3 j: \        model_ft = models.squeezenet1_0(pretrained = use_pretrained)
    - G) k) K: I7 Q4 @# |        set_parameter_requires_grad(model_ft, feature_extract)0 n# H  V6 ]( x+ q6 w' {( h
            model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))
    5 H. p+ M0 W+ g1 i; O: p        model_ft.num_classes = num_classes
    5 w+ l8 l8 n( O. J+ S* F6 U        input_size = 224
    8 Q$ \0 o* v1 M' g' k
    / f" A2 i& {  q% p1 U    elif model_name == "densenet":
    / X( n6 g' ^+ y        """; M$ @2 ]2 _( U0 Q: Z
            Densenet
    6 V9 i) Z4 u7 o: l6 U0 l7 o        """
    & J0 o: O2 R+ N) i        model_ft = models.desenet121(pretrained = use_pretrained)
    % g- z, h4 w! e+ v2 O) V+ I        set_parameter_requires_grad(model_ft, feature_extract)2 Q5 [+ S  b6 i& A- P3 v" M; o
            num_frts = model_ft.classifier.in_features
    * s* L6 [' i5 D* v+ {        model_ft.classifier = nn.Linear(num_frts, num_classes)
    4 |% f7 g" C9 P1 c        input_size = 224) n! R* I7 f2 g+ }$ l
    % ]5 L8 C: U4 ]7 F0 X& o
        elif model_name == "inception":
    ) T$ I4 @1 x9 a( d6 ]: h        """  e" W7 I3 C) H3 g5 b
            Inception V3
    ! h# E. u: d0 i( ]' s        """
    2 Y; r$ n: |3 x& Z        model_ft = models.inception_V(pretrained = use_pretrained)
    5 t9 Z9 C( K2 i6 M$ s# r        set_parameter_requires_grad(model_ft, feature_extract)
    , ]9 i; K$ k, O; D! _" {  V$ _( Y- @2 `2 s
            num_frts = model_ft.AuxLogits.fc.in_features
    ) _* f" C1 [7 B        model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)
    5 [4 C" ^9 O% }% A
    ' B1 a3 l; D3 Y9 ~        num_frts = model_ft.fc.in_features0 V- `( {+ q" F# q
            model_ft.fc = nn.Linear(num_frts, num_classes): @* E, j  }+ L: V. h
            input_size = 299
    : {3 t5 `6 [( G$ O  {" x. z" T* E0 m$ ~- n
        else:
    # g+ k8 f# M3 m, G; M" g/ ~  g  ?( v        print("Invalid model name, exiting..."), a2 i' l& z: h" v
            exit()! |: B2 Q! U8 t+ ]6 Y. @' \, C7 n
    $ _9 L2 W* w8 N# R( z
        return model_ft, input_size0 p1 h- ^9 j* r; p: n& J

    ( {$ I+ M: Y. j) Y1$ t5 V, Y4 z+ G! y  Y
    2& o% u9 ~( M% V; q1 B  O* q
    3  f3 @; A% _* n! l2 M: m, ]
    4
    0 q; C& |3 k" h# D( p5- |" J# r4 c# r+ n8 I$ z& X6 |
    6
    7 ~9 Z3 o+ w) @$ w6 Z0 H+ d! u- X5 L7: C: u, V# f7 `; r" d$ f# j
    86 c5 x9 e0 o$ |
    97 I4 v$ S! l4 w3 I" d# I6 r
    10+ W1 v' N8 V1 C: e" j
    11) @" a9 j* A1 X5 X- {$ P! ~
    12
    * P' I* O3 L1 G0 q0 B! E( k13
    ; n" h: q5 D# f' G4 a14
    : v) g& o# c: F  x" D+ T6 m158 P3 l( g7 A' p% V! Z% \: z) Q
    16
    : J+ n& f8 W# ^  E17
    + u/ B+ _5 m/ n" R) V) M18% l( T2 A. v# {" I5 `
    19
    # _/ e- {0 e( s& l( T# h20
    & c& D4 q0 t0 m: G21. K+ ^2 J) d. W0 Y8 @
    22
    , H9 b) V; R$ _6 ^% x3 i- R23
    8 ?3 ^+ N( g7 J+ Z24; N8 A: S* l) V2 M: R
    25. ^: I3 v9 s* q, i& W- Y
    264 c% r0 C: }% M! v* |
    27
    ) G) N, Q/ W+ c2 E( Q* ^; K28) R% {- u, n% B
    29# l# k+ b# S1 k3 U$ ]; j% V/ d
    30
    6 A% Z8 f. i/ I, G  \317 z- Z0 ~4 y& d0 k
    32" h! `( p  L  L! I
    33% |6 p& h( N* |% @) a
    34
    9 A% Y8 C3 ^. x7 E* z" D359 c! M% k( f- J9 x7 D
    36- b1 c/ M3 ?+ P' r
    37# S8 L5 _' M7 |& l
    38) K* t# g( U2 V
    39
    % `+ h" p( `1 q: s6 F1 g409 Y# V$ V, t- ^' v7 v6 J9 X
    410 M" y: T8 S5 H
    42% p- a" b/ s$ o) S" p
    43
      r; q/ D  I7 H: j44
    1 m# G4 a& M" a: B, p$ H45" E" p' P  H: H3 R& h( Z+ f
    46
    & _9 i; |2 r5 i7 u5 _1 _) F47
    + M9 f' Q6 S4 C* [' Y$ p48
    3 y+ `5 t7 U9 q1 n& m49. G. t, w% U( o2 b- y1 ?% L
    507 l) @) N* V+ K
    51
    7 g5 m" K9 X3 Z6 K5 _! N' G6 q52
    2 T9 ^  v- Y$ D/ Z53
    8 [. U( U$ A/ w+ t7 {+ n, q5 L& A54
    $ T2 R3 A4 @9 W% h; q' M55
    $ _" |2 `0 O. c9 j( b5 U$ g- q56
    ; q. B/ J& ]# H% V4 n57
    , _2 W2 j: V- m8 x* W) d58) ^* |, G* y2 T8 B- R- k  ?! R! A
    593 u$ q  j( e2 U4 ]- D# ^6 R
    60
    ' K) z2 H0 b1 H8 m  r61
    ! @1 y, p+ y3 D3 f& W- S62
    " u3 N4 [5 ?! x# A3 g3 V63
    / T  V8 L- h' g5 N& w, `3 ?0 ~6 d64' D) W1 a; O1 u
    659 w5 o$ T/ F( N4 L6 {
    66) R0 J) n4 O1 s: `; ]6 D
    67, ]2 r7 A$ e+ p% t& L
    689 G9 m! E# {5 ?4 M6 a; T
    695 t2 X) }: h2 M
    70& X6 O+ C9 K, o% Y5 B$ W4 z+ _
    71
    " ~1 x# Q9 e. P% A/ V4 L- l72
    . {" A9 @; D3 M8 A  i4 }73* p; C" R6 J, R- t  j
    74+ B4 w/ T" ]) M! K
    75
    3 y7 p$ e, C4 [# `% c0 n# B76
    * |) K& _! ^. N/ t2 q778 }1 ?( F/ u; j/ k* ^* F
    78+ G% U3 V1 A, s0 I+ ?: j
    79
    % v. `8 h) o, M3 u0 M; D0 S80' n- `  F$ A9 |8 A  H% X. v
    818 n( W/ u' Z7 i. W1 }9 o
    82
    ) @, Q, A( @( c5 [% Z2 x" s83
    5 v7 t7 [; X% D0 B. m% L+ ^$ x4 Y7. 设置需要训练的参数
    / G3 l5 P4 O! n7 l9 O# 设置模型名字、输出分类数* O& e; E( @# P/ r' i
    model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)
    6 E/ q* s( d4 E$ b" K  d
    ' z4 u' u# w$ m5 S# GPU 计算$ ]; e9 w6 r  q' l* ^6 U9 Z9 }3 B
    model_ft = model_ft.to(device)
    * F! |/ Y: Z6 [. f" s; X& Q7 l. D+ ?) G' a" s
    # 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取" ^6 x: g. a0 y' u9 R: w! h
    filename = 'checkpoint.pth'
    4 X/ |  `& I; g9 x! W9 P
    + v$ p0 h' u' j, h# 是否训练所有层
    1 l  O9 u: a9 ^0 F- ?" d6 a. Nparams_to_update = model_ft.parameters()9 a6 j" ~; R; z" \3 P  P( K
    # 打印出需要训练的层
    9 q: L$ z8 r6 _; h& ?" Y! ~" Iprint("Params to learn:")
    $ N" ^0 B! x. q1 O3 o; B: Iif feature_extract:( J/ ^$ F" p5 n6 j
        params_to_update = []4 a2 b7 e: _3 E6 v5 d; Y
        for name, param in model_ft.named_parameters():8 B5 u3 N1 C1 F
            if param.requires_grad == True:
    * A3 Z5 |# W0 w# ^% S            params_to_update.append(param)0 V" M# O& t4 x( W: f
                print("\t", name)- }! a, ], S' [! \  b
    else:  p) K9 j# k3 O- H- I3 {" ~6 f8 R
        for name, param in model_ft.named_parameters():
    % J% X+ j: R+ y& ~+ ]  D" ~        if param.requires_grad ==True:
    5 n8 s, L% l7 ~* d0 N+ E7 B            print("\t", name)
    8 t1 Y! `& S& u7 |4 z/ F, H" M0 `2 L! }
    1
    7 g5 A' l  v, ]3 T7 x3 V. u  R4 k2( ^- k) O! V7 C3 v, K
    30 a+ I. U4 K/ ?) t+ a
    4
    , n: r+ V- G+ ^; o" ]- u0 f: u5. b/ {, A* o5 @5 X1 F
    6; ]4 z0 u! ~5 d9 o) ^) N+ J' Q
    7+ M) D- Y# z- W1 R; g
    8
    , R8 F! l$ }$ e4 f) z( {9
    1 c1 S  U* a& a; `' ?/ z10
    - R' r+ e9 V2 Z  S% q. x11  m( b4 [! M8 X. a. J; i* X0 D
    12
    & }. w; r# O1 U, {2 X2 }/ ~9 q: s13! w1 b& I4 T# B) @
    14
      ^. \: f; f( a4 e" r7 }15
    & M3 F$ W% l* r* ?# a" n* \16" y% K9 }% @' B' f; j1 \# P- V
    170 r5 @! F7 m. Z( U  R* r+ @
    18
    3 P! ?$ d  x, s- j* m) y19( z, M7 A7 W4 e/ R" ^; n2 e- [0 a
    20' B2 \* z- f. J0 R. O% }8 N7 H
    21
    ; y2 |3 m; }" a2 |9 z6 [: d22
    7 P5 J* [7 A" f, z: P, y: L1 T$ k237 R9 n4 C# K9 {+ J
    Params to learn:7 F+ J8 O  Z) L# p7 x! c
             fc.0.weight/ j% s* o5 V0 B+ U( x6 z0 j
             fc.0.bias
    5 U2 k- g  K# Q" o8 c' j1
    5 n* U" C8 g) V4 X2
    , R/ k8 k8 j7 c; H: E3
      j- R: ]; w3 m7. 训练与预测
    + G9 k+ T" n7 h7 N2 u( D7.1 优化器设置
    ' \& V8 ?/ V- k- h# 优化器设置! I" `/ W9 i1 f- E; N' E; @6 r
    optimizer_ft  = optim.Adam(params_to_update, lr = 1e-2)
    2 J) t( Q1 F+ B; R- z4 j# 学习率衰减策略
    4 R5 b/ k7 U. U* l7 tscheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)
    5 B; m: k( g: E2 ]# 学习率每7个epoch衰减为原来的1/108 _: W  L8 F8 a: g
    # 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算; Z" a6 ]) w; J
    # L7 \9 k3 k. P: ^
    criterion = nn.NLLLoss()
      X8 C6 t7 K0 v  x1# R7 P3 U% b* b/ P
    2
    ' Y% `7 [8 B. }% D7 c3 \' q31 t5 h! l( b  M
    44 N' J6 o2 l& h- R2 Z$ g
    5
    ) L8 Z  W# w+ q) ?9 F& F4 w' m6
    2 O; W' a" X- G8 g75 O# t# ]: a8 D* ]  Y8 v0 B. y
    81 r. L  l/ Y% v- Z
    # 定义训练函数
    % \4 K; c; x6 R* C#is_inception:要不要用其他的网络
    - s& A  ]+ _+ X7 B/ B2 C2 z2 }. ]def train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):
    " {) z! g) _2 f/ I+ h    since = time.time()
    9 g/ ~$ [  b5 o2 P; c& o    #保存最好的准确率
    8 W! C! W% `& W9 r    best_acc = 0# y1 M& y1 n' L7 x' T, {8 q# J
        """: d" [- J& h$ @$ K1 X
        checkpoint = torch.load(filename)5 z0 a+ p, o/ \
        best_acc = checkpoint['best_acc']6 b& X2 [" q# y3 A2 W. p( |' W1 l
        model.load_state_dict(checkpoint['state_dict']). Y2 I( X! w/ k- L4 n& F& b* p
        optimizer.load_state_dict(checkpoint['optimizer'])( b+ q1 \6 |5 d8 m) S8 r
        model.class_to_idx = checkpoint['mapping']* s0 V  e+ d- m+ K7 N0 S. \7 l- s
        """$ u& J2 W, @0 Z1 D0 @2 [# A& G: F" {
        #指定用GPU还是CPU
    - i: d* s5 @8 q8 E# u6 X, x    model.to(device)7 a- X, a# R; s+ L0 K
        #下面是为展示做的
    , E; p6 N! O9 C' s    val_acc_history = []
    # ^6 K/ y8 d5 z) F2 [& Q/ _    train_acc_history = []. [  b" \, W- M4 m+ s
        train_losses = []; d5 R. x/ O" y$ _2 @
        valid_losses = []
    $ S) D- w# A- {1 c& P    LRs = [optimizer.param_groups[0]['lr']]
    1 X; Y# Z0 V  I9 P    #最好的一次存下来
    3 A6 W$ z$ K- B5 m( F7 r7 z    best_model_wts = copy.deepcopy(model.state_dict())5 M! a5 G8 X5 ^

    5 a4 o0 L6 A0 K5 l0 p1 w5 w    for epoch in range(num_epochs):
    6 J! ^0 X( }0 ?6 Q1 D1 V2 r, c$ s5 Z        print('Epoch {}/{}'.format(epoch, num_epochs - 1))& D$ i& T0 y4 N% h+ }6 m) z
            print('-' * 10)
    ; g, b) M, N9 a6 P' y3 s$ E. L' S/ A4 m/ {8 \6 n3 Q. M# N0 R
            # 训练和验证/ a+ ?2 s/ _* R$ L* V6 x& j5 Z
            for phase in ['train', 'valid']:: \. e+ h' R4 J5 \! ]) Y/ B* z
                if phase == 'train':! H" Y0 F# n: W" o' |! l
                    model.train()  # 训练
    1 `1 M( p; O( k' m3 ]7 L            else:* K8 G1 N5 P$ ?6 u
                    model.eval()   # 验证! P) \% v! A* o
    / r4 K- q; h" I
                running_loss = 0.0
    0 ]6 _" \# R) s$ X0 t* _  r            running_corrects = 0' U  X1 ?0 F9 S- T  s$ @
    ! r/ F. A5 Y2 `* v3 X2 O# d
                # 把数据都取个遍/ v  j6 u3 \0 q* M) b( I4 O
                for inputs, labels in dataloaders[phase]:
      G) ^( z. v/ @; U/ X4 \, f                #下面是将inputs,labels传到GPU
    1 i' Q& _* X% ~) s8 c  i                inputs = inputs.to(device)! d1 N8 a" r3 t
                    labels = labels.to(device)
      Y! h* |, g# m' D% d* [9 y7 C4 K! Y3 S6 a( c' }5 G7 G; Y3 j
                    # 清零% @" y3 r% o2 v  A. s# N
                    optimizer.zero_grad(); h# c. Z' @4 [- v, H! [' I
                    # 只有训练的时候计算和更新梯度# Z2 E6 J8 n# i/ X6 S5 n: U
                    with torch.set_grad_enabled(phase == 'train'):' i2 K( F* H6 a. w
                        #if这面不需要计算,可忽略
    $ M/ {4 X: X: @/ M2 v; ~1 `4 \( D                    if is_inception and phase == 'train':9 [7 ?4 Q; q5 W4 g6 t% i
                            outputs, aux_outputs = model(inputs)) Q4 Z0 C3 E& L5 l* N
                            loss1 = criterion(outputs, labels)8 j+ X% n7 G- r9 C& Z
                            loss2 = criterion(aux_outputs, labels)  ], y' @. O. n( R0 s
                            loss = loss1 + 0.4*loss2
    3 x% L& @8 y$ e# J/ c; S                    else:#resnet执行的是这里
    . m3 v) _" ?* _2 L' ^( ?7 Z                        outputs = model(inputs)
    / l5 W& d3 m; w" L+ X5 S2 |$ M                        loss = criterion(outputs, labels); s) h8 X+ ]$ j7 _
    7 ^6 T5 C0 t1 S% {
                            #概率最大的返回preds& Y' {3 t) `; |0 @( H
                        _, preds = torch.max(outputs, 1)
    , b/ ^& {  x" G4 o0 b- B( X
    ; c. d6 P7 t; @& h& I                    # 训练阶段更新权重
    # K' t% ^! W" V! w, _7 _                    if phase == 'train':  e( d3 p7 @3 D3 H. B( X7 v5 x. e
                            loss.backward()1 t6 w% ^  o7 }# [, F: I8 ^4 E
                            optimizer.step()
    " _. H0 C$ t/ x, X4 g8 }% }/ I
    5 X6 b5 X# E9 C% l' e  b                # 计算损失2 K% m( x; U1 J7 y+ w
                    running_loss += loss.item() * inputs.size(0)
    # a" [) f! c* n5 N- v" |+ M% x                running_corrects += torch.sum(preds == labels.data), m! V( @! T! A/ J
    3 g+ ?# b/ A5 }& C
                #打印操作( [( G7 ?) |, }. H
                epoch_loss = running_loss / len(dataloaders[phase].dataset)
    $ S2 d6 Q3 L. n# E            epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)
    1 w; h- }! q! ]( V4 Z' J& I% [# C4 U5 a0 K& Z
    . u; U9 o) z& |4 z
                time_elapsed = time.time() - since+ ~9 A% L7 i  U" o# Q1 K
                print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))+ s4 N/ ~; ~  X/ J& I- b, C
                print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))2 y' o6 s5 q4 a) b4 }
    5 {! P; Y( S* F/ x( {
    ! m1 W' g9 _2 d/ `
                # 得到最好那次的模型
    " W- e; [2 D/ U! M6 M: {7 i            if phase == 'valid' and epoch_acc > best_acc:/ S$ k* r# d2 O: r8 f, A
                    best_acc = epoch_acc1 y( H: O0 r( N6 q7 K. ]
                    #模型保存
    6 R7 p7 E# y, D- T  v                best_model_wts = copy.deepcopy(model.state_dict())0 h" q: f% `4 s& y4 P* l
                    state = {
    / g$ R: r/ s! i" W) {8 n' K' ~8 x                    #tate_dict变量存放训练过程中需要学习的权重和偏执系数8 G. [& K& _8 v* i: c5 u# z( z
                      'state_dict': model.state_dict(),
    , c, i% m$ r. m  \/ E                  'best_acc': best_acc,
    # Y0 A% l- u4 X5 B+ l- b                  'optimizer' : optimizer.state_dict(),1 h5 \* K1 b' e; c. V7 A
                    }
    / r7 m! y# }  _6 V5 M. X$ }                torch.save(state, filename)0 E8 Z6 E$ n' M9 x
                if phase == 'valid':& \. e6 r8 V1 `( ^" u
                    val_acc_history.append(epoch_acc)3 N3 x9 ?6 G8 e
                    valid_losses.append(epoch_loss)8 n( ~5 u% o. q+ M7 p: S  d& `
                    scheduler.step(epoch_loss)) Y3 g/ Z& [; o; [2 o
                if phase == 'train':& s; g4 [( Z+ y$ K" p" }7 ~& Q
                    train_acc_history.append(epoch_acc)$ M# m6 x4 C1 u. p: G! L3 ~0 R5 k4 _
                    train_losses.append(epoch_loss)
    # V0 ^* e; P4 X
    - U! U4 F" ?4 [4 H" @        print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))6 W! j' _5 \# k7 Q' Q' I2 p6 t
            LRs.append(optimizer.param_groups[0]['lr'])
    8 u+ B5 e6 W9 c; m) k# d        print()& {- E5 g5 D( v% A

    . L2 z- \( j! s# H' f4 i( w    time_elapsed = time.time() - since
    & U& X" K1 d- a4 @5 g    print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
    ) X2 m: P( q3 O* @( e0 d7 K3 U    print('Best val Acc: {:4f}'.format(best_acc))$ k; J: D6 g# W  n7 d8 C$ e8 p

    $ W* n$ D" X: d8 h; H& k: d7 o    # 保存训练完后用最好的一次当做模型最终的结果
    3 S. _# z  J8 g  |' ~) k    model.load_state_dict(best_model_wts)
    , |" }! Q. r& S    return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs
    : ?- i1 d$ B& {# C# g& K
    7 U2 E, I3 A7 N. y9 Q' O, X# T( T. Y+ ]& H3 N& Q/ e  ]/ l  d
    1
    6 J. F3 O% o, X  |, `. ^' g2
    ( Q0 z. A, X9 c, L9 p/ t6 ]' q39 h' _1 E4 e7 u* S
    4
    - m" c9 b- d" u57 Y  Z4 z; L( J4 n0 V. f
    6
    / K4 S. Y7 L. h4 g: A* @3 s9 i7
    " L8 s: ?7 p+ g8! i5 m0 _; F+ ?: f) C
    9
    - ]7 s& A. b( j, s% ^# X0 g10
    1 ^$ [* ~& Z0 G2 X  d% R" B0 m+ _11
    ( j# R& x5 P& w( b( R12
    & w9 t6 e8 l8 X9 _6 ~& q13
    $ H, _" Q/ L& X: B- {' D3 F* O14
    4 u% c) F) ^, x) Q0 A15
    8 a$ [0 N! F* ]$ w161 A$ q9 ^$ R% \" D$ m8 O
    17$ y4 d0 P& }: g6 U% H! @
    187 J' ^' l# G# ?% k+ a0 i% `: `
    19  E+ q/ ?3 Z7 T( q$ l& y# v
    20
    & F; ]% Y# P1 m) @- L# q0 |! W21) S# o+ ^7 k2 U' \' Q
    22' ]: x9 {" V; f' B( R9 S0 e- ^
    23
    4 C: m3 u# B* S" F  f0 R$ n24
    , M% X; V+ i1 Z+ W$ q; K25. g7 V6 V( s7 U
    26
    " b( ?' r% }1 g! p27
    0 o. s5 N# u, M  A28
    ! A( K( U# B  u4 V# Z* ?- R29
      {" u5 d& L( v# B; t2 A+ d30
    , n4 V/ m# N2 \. r. t8 T31
    ( Z0 Z1 c& a4 \0 _32
    / P: R5 T3 Y  L5 I: W331 |2 w4 l  \& v4 L. R& }" I% U
    34
    * ~$ g; Y$ R$ E5 i# {3 A$ t35
    ' Y, @# H# s/ D/ E: X36
    - e8 a+ z' y$ ^379 f) N4 X$ H, G6 m. \! G
    38
    + Z5 h# ~" M: g39
    # m* p+ W* W$ J: q40
    ) U1 Y5 l- }; V/ T& L8 G: G# E5 q41
    ( t' x' x7 b+ M# }# U42) P) m8 G( Y5 \+ \. B
    437 y- [! a9 }4 g2 t+ d1 ?
    44
    & L2 S6 a9 r: s1 n8 M' X45- w. |; U) N" G& g1 {
    46
    . q% z& t3 `/ m; N( q47
    ( m' b* V4 k- ^48( r* D9 b% O$ _
    49
    , W; l7 j$ _0 Z/ G% G50) I- t  t) I8 U  m( E
    518 i6 q; `" j: x6 A! w# {% w
    52
    9 U* `* P; d) I; r- V9 c53
    / h) l, }& q2 u4 S54
    ( h/ h* i: I% o6 M8 i2 \- X; T55! a9 b4 B8 q9 V, a: |& K+ a
    56  I2 b9 P, k" Y5 y
    57- {1 {5 a0 a+ W4 |0 f
    58
    $ M& ~1 C' H& s9 ^! o592 Z# g& M: y: E( `  l* ~" A
    603 Y4 Z' R! E) k5 l1 h6 D
    61- d: t; M; ^2 O9 `3 X2 N  ~; |* Q( L
    62
    ! a0 [" Z5 p( W7 g: X# B2 V# g63
      X6 h% p: u1 ?; \; n' M64
    + P2 q- z& c0 P8 g5 A1 _65
    & p" ^, P7 I. h1 ?) I3 }662 @1 x4 W. Q' |4 O1 i0 G( J
    67
    $ o; N2 N% v8 B+ Q6 S. F687 M; H1 b1 ~: L3 Z2 O" ?' L' h
    69" D+ _' a5 ^% [
    70
    ) c) f! A0 y/ b" b, j71
    9 z* I& g/ r' C1 Q9 z72
    , Q9 {: G4 I' L73
    0 |) ]4 H6 T: {74( o' {# J" C$ n! F! j
    75
    $ u/ e3 \+ S: L- f76  g3 V, K- M9 j# |+ w
    77# o( X# J, I1 b7 t0 a, R
    780 m' j' v$ _, ~9 B& I, Z. w
    79
    # X: f2 l; {9 p! L0 `4 \9 B! _3 q80$ ^' E1 s4 b- G2 Z  d
    81
    ) A1 L9 p4 L. F) j0 }82" L6 s" h# j" T$ c9 ]
    83
    % [  M1 K3 H7 P( i84. N& {5 H9 I( N$ V! ^# h/ K4 `
    857 o% I3 |7 Y3 H" ]/ o" L
    86
    1 O8 ?6 R9 a. `  i6 _& x% h  W3 i87
    + T! @. s( Y( h+ C88
    8 Y. `  E* o% }6 C897 S  s8 g/ d' x2 `( w! Z& @
    90
    ) J8 n' v. t: b8 `91
    9 W! l$ {2 }6 v. v$ g" ]92
    $ n  V! q* p+ O. [93
    ' ?  ]" w! e2 j3 }942 g# [$ ?) x9 S
    952 I, y9 w" w/ ~( |$ d7 E0 @1 s
    967 C  Q; J8 R5 w( d) ?+ S' z; C
    97" U- a. _) |! S, g
    98
    * b- s0 X6 _& S- D99
    4 M' [# @: N" o  y# R1004 B' m8 X$ e/ h2 R: N
    101' E& e$ j$ n+ _
    102
    2 N- @/ s- c6 C4 B+ J: F9 h  ~, X103
    ( b* A3 v) t5 {- e4 X104
    ) [9 U6 ^7 S( g4 Q( p5 D105
    , o, }; S( }+ j9 R; C6 Q106' m: ?9 L# `& K& y% [
    107* G* P7 h- i) ?
    108) q& r  `* x6 Q( T( O2 g- \  U2 j
    109" j+ Q% N5 P& V" W9 d
    110
    ( U" C$ [! F' B8 ~3 ]0 B7 W5 E0 N1119 O! @7 C3 M1 x, K! S
    112. R  F. d/ i3 C# s! p7 {
    7.2 开始训练模型
    ( u, _$ Y- D  u# K我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次
    0 K# e/ ^7 p$ w4 u5 m7 J4 \* Y" ?0 Y. z3 U7 j: d
    #若太慢,把epoch调低,迭代50次可能好些
    + ]1 e3 X7 M  h#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合) v8 Q9 R& I6 V* }! k$ v( M
    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"))
    + r+ A: j4 K( o9 i
    % P! a% y8 M' }7 h' R13 ?& S9 Z3 H, ~. G: \7 V& B
    2
    7 [: Q  D6 f8 I0 ~  H0 A3
    6 _- @% a: I# N- C9 v6 r4# t( \9 ^* @- x9 j. C4 m, K
    Epoch 0/4
    6 `& e% Q- `& K7 h9 E1 b( I----------+ |+ J; o7 e3 z# x
    Time elapsed 29m 41s
    ' y! a; |  [9 y" D, L# Strain Loss: 10.4774 Acc: 0.3147
    8 D* `" [8 B8 L) d7 a" u5 [Time elapsed 32m 54s# D9 ~5 X- ^( c* O6 {4 m) L+ h9 V
    valid Loss: 8.2902 Acc: 0.4719
    ; R$ {  x2 l* r! qOptimizer learning rate : 0.0010000: ^! `3 V& t" Y1 k

    / |1 N% r4 [+ X  w. f" L( c. q( S5 f3 \Epoch 1/4
    & ], z  F1 @. V6 k, n4 [2 ]0 T----------
    ' f) |# F6 m9 P: RTime elapsed 60m 11s/ t  R  \% G( Y9 |
    train Loss: 2.3126 Acc: 0.70535 }$ U! c2 n! g
    Time elapsed 63m 16s, {7 \: L# {) ~1 O) O$ \
    valid Loss: 3.2325 Acc: 0.6626
    # a. b7 G! T: i0 qOptimizer learning rate : 0.0100000
    ; e8 C# ~4 N7 s3 ]
    1 r% J- V' ~! g7 g* O8 T/ x% r$ xEpoch 2/4  A0 Q* p; U" K+ u1 A& f
    ----------
    5 D3 @- h& N3 Y+ _5 f3 V4 y/ a# E5 nTime elapsed 90m 58s; ?" k# l& n2 z7 v  ]
    train Loss: 9.9720 Acc: 0.4734
    " C% E2 Z  W# N) t6 D2 ATime elapsed 94m 4s" L! t$ H3 i2 \- u) _
    valid Loss: 14.0426 Acc: 0.4413, n, h3 y- a* p
    Optimizer learning rate : 0.0001000
    , G5 [; c/ P$ R  Q9 K
    ) Y, a  z3 m. r! f4 \Epoch 3/4* w9 |: X: e( h4 `+ }/ r2 J& V
    ----------
    4 i; k! \, F/ ~Time elapsed 132m 49s& w! _. _* C2 J& b7 D2 I5 H
    train Loss: 5.4290 Acc: 0.6548
    9 {9 ^: Y% C6 C' M( `Time elapsed 138m 49s
    . a! p# T6 x/ Q0 Ovalid Loss: 6.4208 Acc: 0.6027
    + q' H: N9 e- q/ tOptimizer learning rate : 0.0100000
    5 g- {, m& L7 [: F+ _, y* {- Z
    1 @+ U4 K& i' {% {6 z( g% [Epoch 4/46 N# b! T. n( j& g/ x( M+ u
    ----------) D* k0 d' F' S0 e5 s$ D
    Time elapsed 195m 56s$ R6 A0 R& a& @2 b: i: O
    train Loss: 8.8911 Acc: 0.5519: o9 {0 [- Q% F) V# J  s
    Time elapsed 199m 16s6 J1 E( I* F( a( p5 b5 g
    valid Loss: 13.2221 Acc: 0.49141 C7 E" Y) S# \- W. H7 m
    Optimizer learning rate : 0.0010000
    % h9 n2 d- n, w7 i
    & N' q6 R' f9 m' {3 ZTraining complete in 199m 16s
    8 x5 `0 L' m# Y1 eBest val Acc: 0.662592& J$ U5 r! K6 w# \# y( J

    $ D4 n4 q" S. Y* R. \7 `1
    7 Y4 O- g& {4 }" g; J( l7 K2
    $ M7 {0 `8 Z0 D/ u3
    ( h/ h, Q3 ~) V4
    , w. G' m  n& n50 q3 V( n( C" i% a4 c- x; I1 g
    6
    3 l, h, N9 `3 v* E, L7 q" Q" ^7
    7 e- \' J1 {2 i' T) k. [- ?8
    6 r/ E; w8 m1 T9
    & a' e$ {8 ]" Y$ ~10
    : B5 O: s1 g$ Y% a11
    1 ]0 M" _9 b3 T7 r6 D: R# S/ a5 G7 i12
    . S2 G4 r/ D6 {  w6 a/ H13- j, C+ E! B! i
    14
    % v% c; H& k3 P% c* `8 Y15
    7 _1 ]* B- f3 ]. Y16
    ; w. M" i5 X3 D17
    3 a! X" v# L4 Z1 C7 i18* O0 b1 G2 |1 S6 |( J
    19
    + y) ~4 j2 K' P. k; D3 t209 [) R; g; B. h" |
    215 C! S3 y* c1 W3 ]# d+ z* E
    224 J" M2 ^6 s9 D3 q% _8 ^3 k; Y1 e
    23' t9 S, @* ?9 ~$ d8 U+ z0 J) E
    24
    0 H5 ~" t) ?+ ?& F25
    * H; C$ C4 S0 }: m# i26
    : E1 y- ^% `4 I( [27/ [; M! L4 W; w7 p9 Y% _' g' v
    285 n' c. z. D0 a! D! [9 V
    29
    1 x4 P5 _+ p9 Y$ l+ Z& Q6 X% {30
    3 H4 ~" g9 U) `: f/ S  _31; h7 F+ R9 q) a3 a1 _; X& C$ N
    32' L+ O7 f& G4 Z
    33  q& z0 @+ x3 L% ^! e- u  X$ r
    34# }: u* |- i' j: q
    35
    " x1 K1 s. i6 H/ e- I% @0 D. E4 v0 |36! m' m# [: _. Q1 z0 S
    37. u% L; ~7 ^$ p' q0 Y
    380 e" s  d6 Y' C9 P8 _
    39+ ~% X2 L% I. N$ [' t7 U( N
    40% T! l7 T. h3 Z) c& u
    417 }+ L0 D* e, I: s5 D3 Q
    421 u: O8 t' t+ t/ f; }* W1 }" ]
    7.3 训练所有层/ A5 @' Q8 X4 c0 q
    # 将全部网络解锁进行训练3 @4 C1 u7 N+ H) T$ g  W
    for param in model_ft.parameters():, x* o4 y" H* `5 b
        param.requires_grad = True0 ^. R5 ^2 T0 S: F3 q8 M
    ) H3 w3 I: i* a
    # 再继续训练所有的参数,学习率调小一点\5 e# {+ z1 C$ N' L4 f0 l7 o
    optimizer = optim.Adam(params_to_update, lr = 1e-4); B  C7 |- K  g7 W& i; ?
    scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)$ K8 F7 `; N0 p9 [# I

    $ y7 v7 _7 b/ C- \: W# 损失函数
    - _& Y4 ^9 V+ w9 n& Ocriterion = nn.NLLLoss()
    % f0 y* E& ?0 o' U1
    $ p+ B: \. ]3 k4 }& y+ p/ w23 b" l' D$ \2 Z! v9 B
    3
      E7 G. I! y& ]0 }/ w4 u5 i4
    . w% c1 t4 C) U5 w, T. [5& x: U/ {& P- e. _7 ~, Y& E
    60 u' ^- p& a  G) |1 e
    7
    & O; y4 A/ ?/ t$ v( M7 `83 G+ L/ F9 c! f! F2 u
    9
    ( u2 [* _  G/ J4 D  s- q10+ |, O# w- w+ [) u
    # 加载保存的参数
    $ ~( P6 `" h3 c. j/ T5 a+ l# 并在原有的模型基础上继续训练" O1 Q) c; e& G
    # 下面保存的是刚刚训练效果较好的路径
    . A5 Q8 I2 Q& Gcheckpoint = torch.load(filename)! ~& z) F6 z/ I+ t# A
    best_acc = checkpoint['best_acc']
    & A8 @- u8 Z+ ?5 cmodel_ft.load_state_dict(checkpoint['state_dict'])
    " I& A; i* A/ J% y% Q) Y( h0 [! U6 toptimizer.load_state_dict(checkpoint['optimizer'])
    ) L5 A' u  S, J1
    ; |6 x. g1 `; q: t$ u9 j23 h0 x: f) {* C  d
    3/ r/ |0 {& T, [( d0 d/ z, ~' D" ^
    4
    8 F2 q! M* t0 }5
    * T. s1 K/ h0 Q3 O- Z68 U% r* E4 r# K
    7
    1 _! h/ I' v' I; h% S9 R1 r开始训练
    + E+ r9 m5 B& W! L注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考
    ! k4 I% t: U" V+ A% R$ q! b% p& R; s, G
    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"))5 I, R$ V& |9 v- p; x7 j( r6 `- `
    1: d# n; X* p: C
    Epoch 0/14 c; w4 d9 u/ V
    ----------
    # k$ @6 w) E7 e0 STime elapsed 35m 22s
    9 V2 j9 z! r( O" C/ Wtrain Loss: 1.7636 Acc: 0.7346
    ; ^- L: D: _  V0 ZTime elapsed 38m 42s' a" o4 _9 G. B, ]8 K, ^
    valid Loss: 3.6377 Acc: 0.6455
    : D& y) X, g/ M& P7 p* rOptimizer learning rate : 0.0010000" x" e2 q3 t" N3 K) _

    $ r6 M4 f: Y; tEpoch 1/1
    , ?3 g; v- |6 f----------* o$ E8 t0 T* H& x
    Time elapsed 82m 59s
    + O$ s6 M) }' z  I0 V% u7 e4 ttrain Loss: 1.7543 Acc: 0.7340
    * H* k8 ]. `0 K5 dTime elapsed 86m 11s1 }' k; V) f) c, e7 [
    valid Loss: 3.8275 Acc: 0.6137
    4 }1 H* Q( z! Z- B3 L% H3 X9 p' nOptimizer learning rate : 0.0010000
    ( P* m0 l/ C4 E5 [# ~# n; r
    7 s/ o0 g& L% H* c) E2 a$ P8 v/ K1 ~Training complete in 86m 11s( L: r9 T6 [4 w# p1 L3 V
    Best val Acc: 0.645477
    ) [- O# }2 H3 e7 h: U3 @/ h. H" r: l* [
    1
    5 X/ b" D( @) v' H1 F2
    - w4 h- T# M# s# R# C3
    , F5 B8 Z/ B$ p/ `8 |9 ?4( c  E3 D" J9 B# Q; y' F& T( G! Y) s
    5, `3 @/ D. E" ]- T. i
    6
    5 i9 z8 E! o$ M8 _# j- [7& l& A0 N! x+ n% D# G/ n
    8# y& ~( R) x: |: g9 ]' n
    9) m% j- l3 z0 L! p  l
    10
    / S! s! e$ e$ I* j9 S11% v" _) D( b8 F2 r
    122 s! k% l; K* s) v6 G9 _
    135 h8 z* w/ r/ D6 |. @" M7 \. m
    14
    7 W4 L8 ~. K8 D+ K( c, L7 A: j153 H! l8 ?) y; _, G9 O  w+ l
    16& r' i$ |1 _- _( F! S1 `  K. O
    17
    ! }3 t; }& y/ b, e& i0 \18+ `# Q) U7 n( S6 w% s; ?  m" m
    8. 加载已经训练的模型
    # Z# Y. j4 y# G' {8 X; w+ R2 d2 N相当于做一次简单的前向传播(逻辑推理),不用更新参数
    * B# r- N/ p1 a6 Z4 `
    ; t8 f. s! S& U, n( E2 mmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)6 Z7 u% F; v3 s1 h& b

    & Z6 C- [4 o& j. ^) n4 F# GPU 模式& g8 B$ c% I5 q
    model_ft = model_ft.to(device) # 扔到GPU中
    2 e: f  e& k5 r: U
    : w! M- v; o2 Z- Q- l# H8 I) G# 保存文件的名字! h% _$ I2 U7 r
    filename='checkpoint.pth'0 L9 `+ j- d. B( C* t) u& @0 H) ]3 `/ o

    5 [* H+ k- [, m) I# 加载模型3 c9 Z/ J9 n% i4 F5 L
    checkpoint = torch.load(filename)& z! i$ \# D) u: F0 D* G; T+ o6 E0 Y2 m
    best_acc = checkpoint['best_acc']; {* O  ?( C3 s& A) s! t: l
    model_ft.load_state_dict(checkpoint['state_dict'])+ n$ ~. P4 S4 h9 A8 ]; Y/ l
    1
    7 Z! {  h% t3 G  @% d% A# q2
    2 f/ v4 z4 o. j- {' f. P' |3
    ! ~6 ~/ V& ?1 S# T9 c1 m40 g9 h6 _( g( U4 u2 j$ H
    5( b6 i) {* [* A% c4 a# Z
    6
    ' K4 m2 Q6 x; `, `9 ~7
    ; N/ \' \( t9 ]# z8
    : V( F6 m# |6 t7 o  g( U  _9
    * o4 s7 i. ?: K1 h( k7 x3 T102 i1 o, t: {6 f( H& v; O
    11; a, s3 k# g: w( s5 N) D
    12
    7 j* d7 j) S* a$ `% n0 u<All keys matched successfully>
    " o6 ^- S6 l/ Y/ n8 w6 j. |9 t- g1) u! c& p6 o( W* Y- v/ {
    def process_image(image_path):. Q/ v5 X& D. a2 o2 G- @# f
        # 读取测试集数据& `5 g, B  }( F
        img = Image.open(image_path)
    7 y7 \  v) [  O" o. u, \    # Resize, thumbnail方法只能进行比例缩小,所以进行判断* J- e( |+ [" a3 o2 s
        # 与Resize不同3 A+ t* A8 e' ?4 d! H
        # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小0 _5 N; k8 z5 f3 i0 W' }
        # 而且对象调用方法会直接改变其大小,返回None
    8 o! d  D! Z* w% k' Y" K    if img.size[0] > img.size[1]:: ]- j/ I6 q2 q
            img.thumbnail((10000, 256))
    ! i8 z2 I/ o! A+ @3 Z    else:5 r- B; H( K0 m+ C# p2 b+ g/ W
            img.thumbnail((256, 10000))
    - E- M# O6 M  M  s2 R2 L: z, w3 q1 }9 I8 d
        # crop操作, 将图像再次裁剪为 224 * 224$ i+ M/ J4 J+ x3 y, S
        left_margin = (img.width - 224) / 2 # 取中间的部分, z5 s& y. o/ z. m, O- Q
        bottom_margin = (img.height - 224) / 2
    4 |8 o) r) e, D% Z4 ?' n. Z    right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度
    + A% w" S! [# C6 [6 k- o) q  W    top_margin = bottom_margin + 224
    - b0 ?. L4 r+ q/ \) ^% S! F! }
    7 Z3 K& U) N0 S0 P- u- {1 E# V' l3 i7 E    img = img.crop((left_margin, bottom_margin, right_margin, top_margin))
    9 F4 ]6 ~7 Y9 z. g* t, Q- t6 W, a: h& u' m9 p6 S8 O+ Y
        # 相同预处理的方法+ P9 U4 @3 Z9 m# w0 W
        # 归一化
    ; r: f  s+ ], S' b    img = np.array(img) / 255; M" U4 R6 X% h" L7 T0 M' p% r/ x& r
        mean = np.array([0.485, 0.456, 0.406])8 m  H8 q. x0 W' C9 C
        std = np.array([0.229, 0.224, 0.225])
    ! m4 z9 l, l) a# A    img = (img - mean) / std
    . c. v+ F4 c4 J
    . U: ?6 ^$ o, h( U& L    # 注意颜色通道和位置
    + ^8 K8 m% R$ ^: `    img = img.transpose((2, 0, 1))
      y5 ^) e! w2 @( m) V1 y! ~0 W+ |1 O* _+ D* N) Q
        return img  X/ ^  |6 H* S0 d9 L

    ( e3 B# e: o4 O6 U: u0 Qdef imshow(image, ax = None, title = None):
    " D, d6 Z5 I$ ^) a4 ^    """展示数据"""
    , I0 r1 ^* i. M8 @5 y    if ax is None:3 S% g' t- b9 z( a
            fig, ax = plt.subplots()  |8 P( I5 O- @2 ~' z. y
    ! D# D3 q9 g$ R' @! d
        # 颜色通道进行还原# e$ G& l5 A8 J* T* W
        image = np.array(image).transpose((1, 2, 0))
    4 C& I* J+ ^/ o% c' U5 F  ]" u
    8 Y" C2 Y5 s9 {6 C    # 预处理还原
    ( V9 {3 k0 w9 N9 p% }% Q. I    mean = np.array([0.485, 0.456, 0.406])
    $ z0 J! _  p5 F3 W; |+ T) o% ^    std = np.array([0.229, 0.224, 0.225])% \( d$ A( v' u9 z$ u& e' d/ N" L
        image = std * image + mean
    ! F6 B3 h  B8 N8 U! a1 ~. h    image = np.clip(image, 0, 1)6 d6 P% S- P- M* f4 M) Z  {
    7 T5 c% D, h+ @) `7 r; R) m* x: [& n
        ax.imshow(image), S! h% V" W  c1 T
        ax.set_title(title)
    5 [6 Q2 x' `8 q. C) u4 m
    % O+ O, F' p8 k. E1 F    return ax3 w9 S* ]0 p- }" _( \7 q' s

    7 S6 E1 B: h! e  Z, V! Y, Uimage_path = r'./flower_data/valid/3/image_06621.jpg'
    * N! ~- `! v1 x' }5 Aimg = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理
    ( [, Y0 D. I/ a; ^; {1 X3 c8 i5 vimshow(img)* H: Z0 g' F. w" V

    5 w) T. `: R' b1! x/ X% ]; Y  i+ ^6 W* x2 `7 u' n
    2
    7 ~) ]! Z3 R" B6 k" N, l3
      c6 V1 w0 n- o; G1 U3 S4
    7 n. ^; S' b2 Q6 X* G3 ^" Q1 F& k5' j+ V  v0 W- ]* T+ ~: l
    62 i2 X* `) h. j% |5 j7 R
    7, `4 s" {/ F2 t& T
    8
    , b% d3 ~& U& i0 B1 J, N9
    % B, c: t4 u: s& {' z0 }103 A+ C6 @- @/ `- f8 R3 S
    11. i1 }4 ]# a+ f* \7 D
    124 x; g. x, r+ S& Y5 A+ B+ h
    13$ V5 S5 s3 t+ q5 n  I
    149 F+ l" [' @4 o& a/ u' E- U
    15& U- u1 L" g# F( |1 h" F
    16
      C4 F% G+ X" B17
    # T* I# c, [. f$ X. S: o" u( S18( Y8 y. P" ~# ]7 ?: \% I2 [/ u
    19
    3 ^% E7 B- n& I20% c4 W$ V# f9 U+ j% |
    214 k% d8 `, _5 b  g
    22/ i  d4 c( r/ X$ `
    23. R9 h( }4 ^& N1 p
    24
    5 u" @  J, ]2 t, e25
    $ }7 i9 ]# i' v" C3 ?264 H! B. L& j& r2 H: _# P5 N: ^7 }
    278 R2 B3 U5 h* O3 b7 M+ w: a
    28
    - M1 m: U/ ], X( y1 h2 ^29- [1 l. t" {% z( ]
    30: N) s6 K+ A. q
    31
    * @4 n$ _' m2 z( L. i; f32' B  k: L+ n6 z
    33. j0 R  s$ C; m$ Z% g: x' e% Q
    34& X+ C- W& F* e5 k1 W/ C) P: E
    35  T! r9 R+ p* T. g$ J6 ?! Z
    36  e& P* b/ d& `7 R7 ]/ H$ m/ D
    37
    ) i: Z# o$ k( m& J9 U  c7 D% l38
    6 v" q$ o+ _7 [; K39  n: N" d* x  k& _8 H
    40
    & h7 k3 ^. `  q1 X! G414 C* F$ q7 z8 C- F! {+ S
    42- H' ?( y! h" |( _
    43
    1 T0 k) d4 Z& l- l( z# N" O44
    0 K1 R: R6 [3 P1 y+ A* Q4 H45
    : Q, Y/ Y5 p+ v6 c( H46
    # g' F  R6 M4 E/ A2 W0 X47" j( n% ~4 `8 L. N8 d
    48! q1 \* t' N0 E; B0 V, ?
    49
    9 E! ]# M, F& L( G6 }3 i4 ?( l9 q( p, U50
    * f+ [2 u5 u; @: M, j! ~51
    ; N0 H& a- c5 K% [523 w8 b9 H0 j8 ?5 Q' S% J
    53
    ' H6 L4 g& ^. ]54# k8 k" ?5 I5 U! f
    <AxesSubplot:>
    & Q1 O6 K  Y( Y7 `4 k7 s! M1; ^9 V; D( U: w/ K, Z8 ?
    ; S/ V# ?; }! Z9 T1 x6 O0 t0 \
    上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确
    + G/ F- h6 \' v* c6 O* N# D7 O5 L6 l. h1 ^% O7 T
    img.shape& {( }! w9 f- u
    10 K$ e' s+ D/ \3 c6 n# G6 ]
    (3, 224, 224)
    ; J2 T& k1 _+ X0 P9 q2 L1& n2 S$ p7 `; j; b# w
    证明了通道提前了,而且大小没改变
    " h( D7 Z7 m7 U5 ^
    6 r; k9 `& q" }8 }- \$ k9. 推理
    7 R$ [7 q2 J  x" v) Zimg.shape& d* \2 I% |, b8 ?

    / B9 A8 i. J) _* K+ v# 得到一个batch的测试数据7 l1 [* Q( @3 i5 v
    dataiter = iter(dataloaders['valid'])
    ! L2 q0 a- N$ z' d6 ^images, labels = dataiter.next()
    , W4 Z* I  F0 A" }  H1 O& J) o$ f: |+ _6 l9 x
    model_ft.eval(). [7 z1 L$ d+ @' [* b# s
    . e9 p1 o, L" r7 S
    if train_on_gpu:
    + `( |/ \6 Z, L2 H5 C    # 前向传播跑一次会得到output
    # Q, s2 w) }7 z8 Z- E    output = model_ft(images.cuda())$ a6 ^4 o' `. _4 i- x$ J: Z
    else:& u* \1 [8 d0 P* O5 c/ H) V* o
        output = model_ft(images)
    % j$ P5 q' B# I7 f  t* E
    8 F6 e+ W$ d" K% J; J% E* ]# batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值
    ' M9 U9 M5 D: v0 Z$ }; H5 Zoutput.shape5 }+ B* o) C, o" r8 |. k
    " [% l1 u- q7 W/ F! M* o0 C% a0 Z
    1
    1 t3 i4 p* H# i6 y/ U* l25 x1 [. y6 @6 \5 p; }
    3
    8 s# j3 H4 y, B2 ?) ^: n% b4
      b5 p, a* F* j, p# e; |/ [5
    - W( z1 [  x* z) L  K2 R: V/ i6
    6 I( w7 i9 Y! s9 X9 {4 e# L7$ u5 e5 O# ~' T1 Q% F$ n; j
    80 ~  R6 N0 Q  ~
    9
    7 ]# @$ D. P) |4 N& e# x109 q" b6 J8 `% v- o) l7 |$ e3 A$ J9 T
    11
    8 p, R7 ]; e6 U6 z4 N: b& a9 i12# K2 |- n0 x  ]% q( w% V: Z
    13
    ) {4 L3 C& U6 S/ G+ w) ^' C0 b14
    1 n+ ~1 T) J! @3 k6 Z159 c7 u5 F( z+ R8 }
    16
    % s0 d8 M, M8 v/ G9 l3 w2 @torch.Size([8, 102])9 o4 _4 n) V  h8 R0 v
    17 I- _# p  V  n) I$ t
    9.1 计算得到最大概率. ~/ E2 b3 z  o  s. D* n; S
    _, preds_tensor = torch.max(output, 1)6 Y8 g  i% `" v8 E8 F
    & w: d9 m2 I# b4 a5 k9 F9 Z
    preds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量* L  c, u' `3 U2 M% D8 ?0 N5 C$ ?
    1- M3 ]" m0 P1 L1 C! e6 {) x5 ?# _+ M
    2
    7 n2 ~2 X3 \; i+ W0 e2 d& C3
    % S0 L/ K0 Q) o$ b% ?9.2 展示预测结果
    # w% P. E9 I/ e4 Z$ rfig = plt.figure(figsize = (20, 20))
    ( l4 f( @1 W; D& M1 a, Rcolumns = 4
    ' d) ~; [" _8 rrows = 2$ W- S& w* I" J" `9 Y- _5 V4 u
    8 Q& z. T! q; E4 B4 N( B  r
    for idx in range(columns * rows):( Y# G4 M7 b: @. {; A- E
        ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])+ z- e, L, [% \; f
        plt.imshow(im_convert(images[idx])). z, W0 h  Q5 O3 `
        ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), 8 ]! h7 O8 w1 H& M. h
                    color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))
    - |! s- h4 D. T* h9 o' N- ^plt.show()4 n# ]7 P0 p1 m8 r7 \2 H, W
    # 绿色的表示预测是对的,红色表示预测错了# s9 E& q! Z0 d' L
    1
    $ `5 v& C. g3 E- ?5 l  _" S5 N2% g: u7 s, K/ o' {: \$ F
    3' n( V/ t' Z: f% V* e" i
    4
    9 c5 t& ^- j+ U. K% G) w) F% _5
    " X' Y& R6 ]: d- i' {5 {6" k/ H7 i) D! Z$ _0 z2 V8 s
    7
    8 g- K* z  M0 C$ l8. C! Q0 v" ?6 s; e) A' B, d
    9
    ' ~6 @% W. B' \102 i" f5 K( i3 h+ E: C$ G
    11$ O" J3 o' b  W8 {3 i& w

    9 Z0 p9 Z  x3 t+ q. [. R' W, B& Y3 y- |+ H

    " {+ x+ o+ k( q* ]; P————————————————
    + s- f' a+ E* Q# l( e; n! h版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。, L# v+ P2 [) \; E$ |+ y/ _% @% l
    原文链接:https://blog.csdn.net/LeungSr/article/details/126747940
    7 L- V2 @& |  ^4 U% k; |6 N" `% ]# W+ Z
    4 i4 Y4 s% M& _% @; [! G  K/ V! N. U$ w  }( \6 ~+ X/ 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-8-4 13:07 , Processed in 0.449205 second(s), 51 queries .

    回顶部