QQ登录

只需要一步,快速开始

 注册地址  找回密码
查看: 2838|回复: 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)实战案例% z7 B! H3 g$ Q8 v" Y- l

    $ t3 E$ r0 l: e文章目录
    % N+ _' I* s0 N! w卷积网络实战 对花进行分类
    7 K5 Q$ j1 F7 `9 x, v数据预处理部分+ e1 U* k+ v0 X1 U) \, N, o
    网络模块设置" [' C& Q" [* t0 e% F) a
    网络模型的保存与测试- i; O" ?4 Q8 _, y. K' B( K
    数据下载:
    ( L$ n3 ^1 r, \; V- q1 F1. 导入工具包2 f1 a3 t" i, W# W3 i# C) }
    2. 数据预处理与操作. w' P, q! a. F
    3. 制作好数据源
    * @& D& _& I- S" @读取标签对应的实际名字: E+ f% t; x& ?3 W. ?; Y
    4.展示一下数据8 @) N) L% w* d3 y
    5. 加载models提供的模型,并直接用训练好的权重做初始化参数- A9 G+ i* W: q
    6.初始化模型架构$ D' A: H9 D& u( g, v
    7. 设置需要训练的参数' J2 ]  k* m( U* [7 Z5 Z# e3 z8 A
    7. 训练与预测
    1 _. W" i4 F- e, P7 _2 f  _7.1 优化器设置# ]& Q- g9 c/ R! _& E
    7.2 开始训练模型
    2 T+ _1 b8 I! Q  O1 n7.3 训练所有层
    9 A7 N# d4 v7 b  Z/ P, `开始训练9 d! |) {9 C& I
    8. 加载已经训练的模型& E* M! Q& H, {) Z
    9. 推理) A* j; Y# O& x$ F- q2 C$ f' @
    9.1 计算得到最大概率% |1 d4 {3 M( F! E- ?2 u' ~
    9.2 展示预测结果
    / @- ?9 y) T% t: `' u9 @4 l) g写在最后
    # U! a9 _" U& U' o5 ^卷积网络实战 对花进行分类
    ! _" G3 S! _% {7 D/ r% K7 X5 [本文主要对牛津大学的花卉数据集flower进行分类任务,写了一个具有普适性的神经网络架构(主要采用ResNet进行实现),结合了pytorch的框架中的一些常用操作,预处理、训练、模型保存、模型加载等功能' M+ d( d* `. Q: Z" E+ A

    ) X% Q# d7 X% l在文件夹中有102种花,我们主要要对这些花进行分类任务
    7 a2 T* l, B9 _% I* O3 C. B文件夹结构/ ?% c5 G% |0 N

    # ?5 I# R- H3 `' H7 D8 ~0 T. Wflower_data
    % ?8 `* p. e# M" S  N/ @! M$ K6 ^  v# G" X; J6 O2 W9 r5 m! g
    train
    ! w, U  @1 h( Y* n
    ) Z8 U# u) c1 a2 D: Z% Y7 X' ?; C1(类别)
    ) W1 n3 U5 G, M1 k- C* i2
    5 c  T# w$ K, W* |# j+ Sxxx.png / xxx.jpg" m6 E, `, I- u2 ]" ~% L
    valid
    ) O, [% ?8 f/ k# T, `
    ' L! n% k6 P$ B3 T, w主要分为以下几个大模块
    * ]5 W9 j2 [% r
    2 c3 [9 n+ e& q! W4 }数据预处理部分7 r8 l1 h/ k8 A- I$ N5 ~7 ~  K
    数据增强/ u- v" e+ p" @1 o. m
    数据预处理$ t- h* G" f! d, F# z! G* e$ t, M
    网络模块设置5 @" p& {9 j) S
    加载预训练模型,直接调用torchVision的经典网络架构# e# I- t8 K) Z+ K8 y
    因为别人的训练任务有可能是1000分类(不一定分类一样),应该将其改为我们自己的任务- m3 k9 B1 I6 @+ ?( o/ f
    网络模型的保存与测试
    9 A, l) I( E0 f0 k. ?' ?; F" n1 K! d模型保存可以带有选择性+ A! n/ C* Z( P% x
    数据下载:* J. k  F  n$ G6 \% X) C9 G
    https://www.kaggle.com/datasets/nunenuh/pytorch-challange-flower-dataset/ h. U+ i3 S0 _5 o: {
    , w8 u" o& }6 F& v( j5 I" j
    改一下文件名,然后将它放到同一根目录就可以了1 j& [4 F  S3 l' Y& O

    ' L3 ^% R# A9 ~. H' Y; A8 ]下面是我的数据根目录
    ! t$ t5 Z) h" }8 o. y  d- ]0 y# @  |; x

    5 b: i- M: h0 U" U1 }# u  E+ l1 R; X1. 导入工具包: i! v: y% V8 A9 Q' ]
    import os
    * T7 T- [7 o+ j6 x4 m) C. |import matplotlib.pyplot as plt
    6 I& S/ M1 P* S$ L. G# 内嵌入绘图简去show的句柄2 f/ y6 [, ^% k
    %matplotlib inline   h! V9 D! A9 I% }$ P) W! R) L) f" f
    import numpy as np4 M6 g7 y3 D; D$ e
    import torch; [- m( S* v- k7 i$ h
    from torch import nn8 L( r3 Y. {3 A( \) B) ?2 _4 e' ]. }
    2 I! t8 v* j9 E
    import torch.optim as optim
    % i) h+ t" G. @7 x4 C" uimport torchvision
    . }2 {/ o1 M2 A& I$ F# l" w) Cfrom torchvision import transforms, models, datasets
    . ?- G, V- C7 Y* G! [& `) {. ?  [8 ~+ c% h1 ^) Y; r2 O( m! U1 i
    import imageio
    - h5 T  y6 w" W% w  Aimport time8 u* {% Q' u- R; N4 i- d
    import warnings
    ( c! @. O" t- c$ o; F' r7 W: }import random# [# h) b5 s; U7 m3 B7 |% P' z
    import sys
    ( G5 [7 N/ X  ]: d# _import copy
    ! S' J5 {! ]+ q1 u& |. q! m1 ]import json$ L' x, I; b0 I! ?
    from PIL import Image
    $ r; d4 j: P) o
    6 F9 n* P& d3 m/ b, _" m( {3 o" U4 n' F0 H3 q! ^
    1" B" ]; E! Z) d
    2
    + V0 ?: A' C4 G3 H4 p: s5 Y3) u/ E& v7 K9 k1 K) j7 O
    4$ t- w2 M$ x' `, d0 \0 ^
    5
    . s3 W. o: T2 |, u8 ~% S/ Y4 O62 }% j+ c6 W6 e3 ]) z9 O$ B
    7. V/ t+ _/ K& i
    8
    ; k- W  f  ]& G0 R0 q3 Q97 M2 ]! f1 n. c. s4 K; T9 P& h
    104 f- ?: F0 ^, Q
    11- D* t% `' `4 P  Z1 K
    12; Z/ {1 T$ `% p5 |; `0 ~8 E. ]' `4 G! s& V
    13& q/ c7 z# S' }
    14
    ! [0 @+ {- n& {15$ c8 A8 _+ @; I$ U& H. l8 ~- C
    16
      q8 y( O+ C! Z$ c17
    : k7 e! u! b+ t1 M181 q' D/ ^2 T4 }+ G8 E% l! U, e
    19% t' G' `) I# r/ ?* J- b
    20
    4 k/ Z$ A3 L7 b$ Z5 B0 a21* F9 G5 V' w' K- x! p
    2. 数据预处理与操作& _. }; p5 |3 t# |* M# M
    #路径设置
    / j* r8 w. r5 ?) g5 P1 ?data_dir = './flower_data/' # 当前文件夹下的flowerdata目录
    , t7 B! w! y" E% \+ u9 O* N% ^! h+ p. mtrain_dir = data_dir + '/train'. ^" @, }" v" Y9 f& g
    valid_dir = data_dir + '/valid'' s( V+ |8 |* L, A% I7 V1 y/ D! M& n
    1
    ' ~% C8 F' n" q: C2
    2 c9 V5 e  F% E0 t, j* Z3
    ( f- `* \0 x; ]# E) @: i/ j4
    2 V! I( U8 x- }. `$ rpython目录点杠的组合与区别: z) F" X7 ^3 B
    注: 里面注明了点杠和斜杠的操作
    7 [* y- e, u6 K% W9 S) r$ i% Q% O4 V; {+ q, F
    3. 制作好数据源' B# r! j' ~6 o( Y
    data_transforms中制定了所有图像预处理的操作
    1 I$ G- b0 F6 ~, L: _. iImageFolder假设所有文件按文件夹保存好,每个文件夹下存储同一类图片3 ^/ g5 }# q6 W3 |0 ?+ H6 B
    data_transforms = {
    ; E4 t1 M' A0 R" b- ?    # 分成两部分,一部分是训练
    + Z3 [* E- F) l; K* @0 c+ o    'train': transforms.Compose([transforms.RandomRotation(45), # 随机旋转 -45度到45度之间8 P; F+ x' Q9 e! o  s& M
                                     transforms.CenterCrop(224), # 从中心处开始裁剪
    ; d* E. l) ?4 ?5 g9 e& z5 `& |                                 # 以某个随机的概率决定是否翻转 55开. E! U( x7 M% m7 ?1 ~1 }
                                     transforms.RandomHorizontalFlip(p = 0.5), # 随机水平翻转
    ) A5 U; d2 i1 D4 Z! x# s8 H9 d                                 transforms.RandomVerticalFlip(p = 0.5), # 随机垂直翻转* D4 z; Y+ T3 z
                                     # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相5 q9 S, ~6 \, n! l# S, A) s
                                     transforms.ColorJitter(brightness = 0.2, contrast = 0.1, saturation = 0.1, hue = 0.1),
    8 V0 Y7 F# i! U, N. s+ {                                 transforms.RandomGrayscale(p = 0.025), # 概率转换为灰度图,三通道RGB
    & T8 d# j- ]" G$ f3 U5 g  N                                 # 灰度图转换以后也是三个通道,但是只是RGB是一样的! g( m: m3 {% q' E% E6 o5 {/ ~; o
                                     transforms.ToTensor(),% o+ P$ r+ a4 f) v# g$ v, e
                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值,标准差
    & `2 }5 b, a  I7 f8 L1 O9 C                                ]),+ {- X3 x4 f# ~2 _& Q8 |0 c' O2 G- D
        # resize成256 * 256 再选取 中心 224 * 224,然后转化为向量,最后正则化) F8 O" M/ w" s
        'valid': transforms.Compose([transforms.Resize(256),
    6 v1 j: I8 s( n: j' x. v. l# ]* T                                 transforms.CenterCrop(224),* h' {  r& B: G" |
                                     transforms.ToTensor(),3 a/ V+ e  g: z9 g- d# o, x
                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # 均值和标准差和训练集相同+ s% S5 p- C; ?+ D0 L
                                    ]),
    3 s% L$ W, U$ p- w9 q4 h}4 P5 ~: ~1 i7 ^( `+ F

    * I8 s- y$ S  S& Q" I% G# A1' S8 I7 J+ x& g' H
    28 A6 D& Z, O2 Y. \) a
    3
    + M1 T, V9 p! D2 V4
    ' m8 e$ S$ A3 b6 l$ p% `5" P) A  ?  X% ]  M; o
    6
    / G/ a( d) L; ]) m8 `* b7, D& V+ m6 [% i0 T8 N1 E8 Q. G
    8' R2 Z7 G7 ]2 k& `
    9
    9 r  c4 f# J! [* q/ Q7 N106 Z. _/ F1 A& m
    11
    $ l3 l+ n" ?! @' f4 _12
      P) V; i9 T' w13* u# B1 C( C& t) _
    14
    3 G! p. H! z; r/ r15
    % [" O4 _/ S, {" C0 N16( w4 q" f) K- a/ M% f
    17
    - K- }7 d1 G& |. D3 q182 C6 ^* n7 q+ F( g" n
    19" y$ g3 [/ l1 A+ g
    20
    5 K( W( r6 I, G( C3 V; r2 g21
    5 y& _! h+ }9 r/ L& n) [batch_size = 8
    5 i! B" q: A  J/ `# B5 g2 j5 Himage_datasets = {x: datasets.ImageFolder(os.path.join(data_dir,x), data_transforms[x]) for x in ['train', 'valid']}
    ( C1 k" B; e- ^7 T- tdataloaders = {x: torch.utils.data.DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True) for x in ['train', 'valid']}
    4 W+ U  f0 t2 `- \, sdataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'valid']} 8 w+ i; e. p( r2 O0 j
    class_names = image_datasets['train'].classes
    1 S4 ]3 ^! Z% F7 `1 M( O! {/ \) I0 k. y* D0 q* ~
    #查看数据集合/ y! F# ^0 P3 S3 A
    image_datasets  O9 e3 h2 `5 @+ C, o

    ( M* k. Y) ^% e8 t2 v1
    9 l  n" y/ C! @1 N; }- @20 u( J: v2 `( r6 f; X
    3& {( _3 M2 }' B
    4
      L5 f. q6 X- G) i5
    : q2 `6 h; s% W' U6 u- l' a6
    ! {9 r3 I4 K+ R7 @7 w$ u1 |5 \75 {* N+ v/ N9 A0 f3 ?$ s" z
    8* H$ V. {4 ]9 e* |0 ~
    9
    : B+ e3 a0 `! I3 _! U{'train': Dataset ImageFolder
    $ `# O  `) E' k     Number of datapoints: 6552
    ' Q& X* ~1 x$ ~2 _1 n  e     Root location: ./flower_data/train
    - _; [# \) E  k) k0 V" [/ p; i+ f     StandardTransform; s5 I$ I/ S, q: q* q
    Transform: Compose(
      ?6 f; m- M5 G/ A                RandomRotation(degrees=[-45.0, 45.0], interpolation=nearest, expand=False, fill=0)
    : \- I2 C3 k" n2 F( [/ ^                CenterCrop(size=(224, 224))
    3 w. Z: |; t' j                RandomHorizontalFlip(p=0.5)% W" H& [! F/ H* p8 A! B6 w' n
                    RandomVerticalFlip(p=0.5)
    0 t% m9 X9 ^6 Q' |% X8 D3 d                ColorJitter(brightness=[0.8, 1.2], contrast=[0.9, 1.1], saturation=[0.9, 1.1], hue=[-0.1, 0.1])
    ! S  l# Q: r* s' D  Y                RandomGrayscale(p=0.025)
    ) C5 }# o# t! r) n  n2 F                ToTensor()
    . ^+ D' \- D0 ^/ l+ c                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    ! [, A6 T6 Q0 T7 J2 W5 j* m            ),6 K$ |* J5 U+ a, m* H. i
    'valid': Dataset ImageFolder
    ( h) ~# n8 z) k3 ~4 r' ^     Number of datapoints: 818
    5 Y6 U& ~- n. c. M) q) R( @) e, W     Root location: ./flower_data/valid
    2 b6 A( t* ?# w, F: c5 d5 j3 ]     StandardTransform$ |7 u$ I$ Q7 L! L0 S
    Transform: Compose(
    7 p$ o, v' }+ h5 J/ P                Resize(size=256, interpolation=bilinear, max_size=None, antialias=None)
    . C2 F  B- @6 U: u3 R. S                CenterCrop(size=(224, 224))
    4 j* t$ w9 m9 b" {1 l                ToTensor()
    ! A4 j4 j4 b$ D3 J7 w6 Q/ c                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    ! o5 t  C/ f9 z            )}2 w+ n+ ^# Z' v+ g0 R
    ( q% R3 J# w1 k5 |/ [2 d
    1
      E% U; d! y$ j3 y$ i( ?2' @4 }8 W- Q& ^! Q, Q# Y- V+ y
    3
    " @1 C4 m( `3 D0 L$ x4& G2 l2 b/ |5 B7 o2 T
    5
    * b. R4 V4 I8 O8 P: N+ v- w6
    2 J; s( P6 o4 y( Y7
    ; Z; O/ t0 Q( J" X8
    5 {4 z6 ]# `. h- d' d3 A9
    ' F4 g# ]4 B* X- a: R' j10
    * _: \  V0 y) v$ I0 v8 `  K* J11$ x; C% y* M+ F
    12
    4 B7 X, X% k/ r) [( q3 ^) _131 y3 h2 ~& w/ p& @3 n7 ~
    14& o* h* _9 |* p7 m
    15
    : O8 @7 @* _4 ~/ Z+ q16
    $ e% u+ l1 L3 Y1 Z, e174 L1 D2 j' V  v! N
    18
    + T! [) S0 [1 q19( {: m1 K5 {% V* j  O' f6 q
    20
    * X4 d& v) s2 ]9 D+ @21
    + \% |3 d+ q$ i! e22
    2 W; q0 L/ S1 z  C: H  S23
    ' `, U; B3 {( o8 u, H24! Q2 F" i. `3 H
    # 验证一下数据是否已经被处理完毕( I8 ~  ]1 G- Y. w% O( ?
    dataloaders" I' X; W1 r3 |" C( T8 b) t
    1
    2 d7 u5 p% Q2 q+ c; Q2
    9 Z3 Y/ e( F9 l4 ~0 I  o8 w{'train': <torch.utils.data.dataloader.DataLoader at 0x2796a9c0940>,# d! z( d9 i  w, V. g
    'valid': <torch.utils.data.dataloader.DataLoader at 0x2796aaca6d8>}7 y3 f# s: L3 R  j7 A3 i. J
    1: F* G/ q9 c) d1 ]" _% N9 m
    2* ~9 s, l3 ^; _- S4 S. v
    dataset_sizes
    6 l; f1 ?- u7 w1& v: K5 H* j; P- s9 B% k
    {'train': 6552, 'valid': 818}! d& q  h. U4 d, V. P  |
    1/ O" _0 `# z: R& l- e! T  ~% {
    读取标签对应的实际名字
    " s2 a* H) u. X$ c. c9 X2 x' Z) l- ^使用同一目录下的json文件,反向映射出花对应的名字
    / ^9 j7 X+ @" r2 a$ s& P4 q: [
    " B1 [/ K+ s6 \+ ?with open('./flower_data/cat_to_name.json', 'r') as f:4 R( n$ B8 h0 a) g& @; L) e  q
        cat_to_name = json.load(f)
    , I. ^9 Q; H8 G6 Z7 e( M# K11 q% j" l% e0 o) T  S
    2( y6 F+ i5 Y0 i  F5 H
    cat_to_name7 O: D' j& i8 d6 V
    1
      t/ _. @& i% @# `% H$ q{'21': 'fire lily',; x- J; a: g) Y3 G& e2 m% c' U5 H
    '3': 'canterbury bells',
    7 W% p' o+ A% p9 D+ E '45': 'bolero deep blue',& e5 l: _7 t! o
    '1': 'pink primrose',
    4 @( o8 H2 W# [; S! {, R. T '34': 'mexican aster'," y; q) V2 u1 F8 j8 A: Y
    '27': 'prince of wales feathers',8 }3 Y0 |, U) j0 ?( z# M) g# N7 ~
    '7': 'moon orchid',
    * m7 X# j6 a/ ? '16': 'globe-flower',
    4 W, D5 `7 e& A6 p: |* j '25': 'grape hyacinth',
    * q4 k& X. v% |5 v& |+ t% L '26': 'corn poppy',
    ! }, s( p0 M) ?: S- }$ g3 }, s '79': 'toad lily',
    & p1 ^! ]$ D0 H) b '39': 'siam tulip',* J4 }. D' R5 A/ ^( O$ W
    '24': 'red ginger',, `# q( S+ {3 G3 }$ `/ _# N) F
    '67': 'spring crocus',
    1 [/ \7 ?3 {) t. r- S* [6 V- w1 m '35': 'alpine sea holly',. Q) ]8 z8 `% r" D
    '32': 'garden phlox',! A. v% a1 Z! Y9 {1 v( E
    '10': 'globe thistle',. k( h7 t7 i, r" M1 a4 X( F, x+ K
    '6': 'tiger lily',. Q, h# n6 w4 A1 @/ p" W1 ^
    '93': 'ball moss',/ }+ s  w4 }. U4 F! @1 l& b% X+ U
    '33': 'love in the mist',
    0 d% J1 P$ C/ t! } '9': 'monkshood',
    + J- n! }; ?7 V& l; K# I+ K6 j '102': 'blackberry lily',
    6 J* J7 P6 B6 ^) E+ ] '14': 'spear thistle',
    9 z6 R9 J3 @8 ^( J '19': 'balloon flower',' Z+ c7 \/ ]/ H7 @9 @6 M) v) L
    '100': 'blanket flower',2 f- r! q5 E8 J- \
    '13': 'king protea',7 O0 F& I. y* q" t) F5 Y
    '49': 'oxeye daisy',
    + H4 u2 M% y  Y '15': 'yellow iris',
    - ^: g7 u0 j2 s8 ]4 r* H2 x* i: ^ '61': 'cautleya spicata',6 P7 \- f) \* _& h4 K2 {
    '31': 'carnation',
    + U) Z2 _- O+ p0 a '64': 'silverbush',* O( R0 g4 h: V+ m8 I" M. V
    '68': 'bearded iris',1 J3 ]: U+ C5 D9 e
    '63': 'black-eyed susan',* y5 q# i  r( \) t& ]+ e7 d
    '69': 'windflower',
    ( d3 C! B- h1 S8 w( N* b '62': 'japanese anemone',5 c1 g/ x+ A, f$ f
    '20': 'giant white arum lily',- c. E5 o# j. v) Z4 {/ m0 U; W
    '38': 'great masterwort',1 p* b% `9 X6 @, p4 ?
    '4': 'sweet pea',
    8 j6 N1 H* O, D4 u9 A '86': 'tree mallow',! w9 {; G. S6 z+ b
    '101': 'trumpet creeper',2 O4 ]; p7 Y& E3 a' Y
    '42': 'daffodil',1 [6 H$ e2 W$ \% l
    '22': 'pincushion flower',- v! H9 [) v5 B1 t+ B6 o
    '2': 'hard-leaved pocket orchid',
      Q) l! E  A$ A8 ]+ i; ^ '54': 'sunflower',1 d  X1 K) X6 D9 q
    '66': 'osteospermum',
    6 C* D- L+ p% |3 o, K* O '70': 'tree poppy',+ u4 u$ Z. p5 c' _- k
    '85': 'desert-rose',5 P& w$ b- s) B) }+ C# M7 M9 k
    '99': 'bromelia',5 m6 Q3 L9 J& x6 S9 M
    '87': 'magnolia',$ o- e  q. S9 ]! ~! |, U. M
    '5': 'english marigold',7 z0 }6 P& O- T  u: J) V
    '92': 'bee balm',+ b# _4 B5 D8 H8 c' l/ W
    '28': 'stemless gentian',2 y  l  ~, u# W! f  I
    '97': 'mallow',; L3 `* b/ L4 M* T( ?6 |' s
    '57': 'gaura',
    ' l5 ~9 U" d! E2 h0 B) @3 s5 W1 W '40': 'lenten rose',8 K( A6 X: [/ W6 x2 V$ n
    '47': 'marigold',6 W! x' X# @: z1 _$ T" k
    '59': 'orange dahlia',+ u5 I7 V* H% i0 S$ h4 R- @
    '48': 'buttercup',. ]3 G: M3 Z. |7 u$ `. }' h
    '55': 'pelargonium',
    7 r9 C$ p2 o. _* R7 M8 F- A/ E$ i' e '36': 'ruby-lipped cattleya',1 a2 W# M/ L8 J) g, l* M
    '91': 'hippeastrum',, E, i- ]5 f- i6 }3 Q: f" z
    '29': 'artichoke',
    ) T2 n3 O2 D+ `5 q: _ '71': 'gazania',
    ! {4 @. g  W9 l& r '90': 'canna lily',* l, j& A: V. Y8 s4 G# m
    '18': 'peruvian lily',
    : X5 w0 M$ K6 M# @  |+ j5 ` '98': 'mexican petunia',4 m1 t& j' g- j; j: e
    '8': 'bird of paradise',0 X; M4 y* I9 P4 G; S
    '30': 'sweet william',+ ~2 J) q- O( z
    '17': 'purple coneflower',% X! U0 H. y3 g9 \9 |. ^% ^
    '52': 'wild pansy',
    0 Q7 s. G9 y  @& |) m2 K '84': 'columbine',6 k: e* D4 z- W9 f: G8 h
    '12': "colt's foot",- U' q+ S/ `0 e
    '11': 'snapdragon',
    / A; P2 l$ T( h; `5 }) w '96': 'camellia',) ]0 Z" \( T# u. s0 C. @6 ]* s
    '23': 'fritillary',
    , W- x3 `$ T8 B5 h" ^- O9 }$ r '50': 'common dandelion',: |3 l+ g+ o, d: N/ v
    '44': 'poinsettia',
    * H. V) a6 m/ |: Y6 v '53': 'primula',5 V# e: J8 m, i
    '72': 'azalea',4 i9 r/ _# a  E( C$ E- Z1 N) d
    '65': 'californian poppy',
    / P* [6 }: r: I# J) M$ Z '80': 'anthurium',
    , {3 V1 ^- m5 c( F; p: k+ m. [4 S/ L '76': 'morning glory',
    4 K0 B1 f. P; r8 Y3 D9 m) H; Y- } '37': 'cape flower',
    & f. ?4 T2 E! \4 u/ b3 ] '56': 'bishop of llandaff',
    7 u+ w/ W. t" C1 r- l& m/ W# j '60': 'pink-yellow dahlia',
    + K3 s* c& n7 I' k '82': 'clematis',
    % k( b6 R! j2 A0 }4 \7 C* d; F '58': 'geranium',
    1 x0 ^) N& e$ r0 D+ y* X '75': 'thorn apple',
    & T& Z& g: V, M" W9 N '41': 'barbeton daisy',
    9 D* C% Z& y) S: q( N, X/ @ '95': 'bougainvillea',5 u# k4 L, o0 ?0 j/ Z
    '43': 'sword lily',
    0 m/ P7 W! W& s) |, s '83': 'hibiscus',
    4 v% K# A% n& L- { '78': 'lotus lotus',0 I' D( ~- _  q) L( a' S5 R6 \
    '88': 'cyclamen',/ E6 t/ \; E* r2 U* V
    '94': 'foxglove',
    ! u4 s" A" Q5 d* `- }  ~8 H '81': 'frangipani',
    6 k/ w5 h: V9 ]$ H( m7 b '74': 'rose',8 [3 i' D: Q' h- s! `
    '89': 'watercress',; v! h/ L" @  [2 l4 ~3 U5 I
    '73': 'water lily',! A# Y4 E# _$ }( _' _2 q
    '46': 'wallflower',
      C, V( y8 g0 a+ F0 r/ S4 Y '77': 'passion flower',
    : w* D  r& Z+ v7 b; H* `  P; J3 k '51': 'petunia'}
    : _4 ^  x3 x7 c# _( p/ j. {( l2 Q1 `) G
    1( L) x1 y" @4 h3 V; z! F
    21 N& V; M8 \; U7 u6 v$ U. r. y  l
    3
    + `4 K" |0 t2 c2 L44 {; x, E9 `2 U' ]; [) t
    58 v3 m+ T0 k9 d/ m9 T! z. H) F6 _
    6
    / q$ W5 h  X$ D( j7! ~3 e( P: k# G
    8
    ) `- b6 z4 k3 X7 i* s& Y1 p+ _8 X, v9
    3 y* C) a! k' I# x10
    6 p+ j3 b0 n' B( P& N  {9 e6 C, g+ ~113 Q2 \/ W& Z) o) w1 n# S0 J6 A' O/ D
    12
    / B! N: s- E& |6 Y$ J* s* l" g; C13& n1 f5 \5 s  P% z: w, i3 N! M
    14& R8 m6 D) I, N# r2 Q% t
    15
    * {; J  _( i/ C16* G1 a1 I% c& |! p$ }" M
    176 }' r& b6 |# \2 E6 m0 p
    18
    ) ?+ v; x7 ~3 l% ~% P" x) Z19
    $ }6 U3 n) m' d! ]4 v: \, i6 Q1 i20
    ! I+ }, A8 V: H- C. @: Y7 @21
    ) t" f* R0 }7 `# \' v  w& v22
    1 @9 `* w$ q* n9 S; n23
    # P) }% v* g! T6 ^' W# T1 J24. V4 `) \4 K, `
    25
    % d4 O2 M, Y( E3 c3 X. P4 G: B/ L& |, H260 h1 d: ~; }" g1 E1 j' S
    27
    + j& ^% R0 E4 k! n( E  o- f# W28+ {1 S" r8 c# J; ]1 X7 j
    29/ c3 z* q9 @) A* K- I- D! {1 A
    308 [0 b( m2 Z# [$ B$ f, D- \
    31
    / j8 x* i4 u4 P% ]32* a- _. p  }9 G! P' h
    33
    . L" K  M; v3 |) A34; f$ ]$ j" z% H
    35
    * B1 G5 B6 m7 q3 G5 z" u36# L) b+ C) Y5 L. I) `% Z( _) `
    37; r0 [3 C* Y! Q  A: G# B( o( Q
    383 S% m" j4 |2 K
    392 F" u( _6 l. u: `7 f( k+ G
    40, N1 j  E% C9 k
    41. L: d  @3 K2 e& K0 @) k
    42
    # y* l; O/ ]- U$ O% M+ u43, @* U* P( f( S' x9 P. s
    44
    * T9 I+ f4 G- b# |$ w( X' {453 ~# V( e3 K  N. Q( `& J* D
    464 o5 `/ B: `/ V" N
    47
    : e5 i7 N# Y4 Q6 a% ]48
    ) @& K/ Y2 @- A9 w+ i: ]49
    2 u5 N. Z& M3 R% M* Z( a7 T1 i50
    9 K" W, J& z$ a511 l8 X2 W: r. n1 X5 P/ Z8 Q
    52
    5 x" n4 Y* Q: j9 _% s( _53
    $ u" S" ?2 M1 N54* t0 X2 H1 T2 F2 a
    551 o, \4 y6 a( A
    56
    ! v. u7 [) l5 D) O/ D57$ D" s1 L# m& }5 i; s
    58
    0 G: p9 R$ h  T+ \$ E4 x: T59  G& p8 c( k) [5 p1 L" K- Y
    60; y5 Z( }) r5 J+ W0 w
    61
    & c6 Y" j, k+ c' @" \2 V. W62! q; I6 X* e! {+ Z9 N1 l8 v% Q
    63
    0 `8 h' e' ?9 O" u" v( s0 ]4 [64' {1 `; Q( @4 m; S% L5 N9 g
    658 l0 l" H$ T8 d4 U
    66
    ; k! s, l7 ]# {" U# P! F67; w' e* h0 b+ C$ |  {
    68
    : U6 g0 Z  ~* m: W% |; `2 x69
    5 ]: x' I  N" h" B# R0 f70
    " f1 G$ Q$ G2 \0 _71
    , B) Z/ ?2 F/ L9 S5 |& v; U729 V! n: l2 R  G1 R0 Y. S
    738 c; `, ]& X* [) c# \
    74
    ! h  t5 R  I$ v! \75" g5 G+ K/ F2 \( ~9 G  G% J2 ^
    76' M' t" m6 _8 A: ^- i6 Z( T+ `
    77+ `: ^, D. p9 |) T& B; M4 \' q
    78
    ( U# z. t& B1 ], `! D  F79
    . j  v0 j! @2 M( ~! X6 I80
      a9 n/ S6 Q' l+ B81
    : ~# A4 n8 A! L# P4 _% S' f824 c' m9 j9 P  {* X* f) ?
    83( d8 k+ W$ y) U( M3 Q5 C9 y2 }" X
    84
    5 _4 B5 }: O# [  O/ S: I859 z2 z! J) T' ]6 v1 p, F2 c! k) J% i
    86
    3 J" Q- t( q7 ^5 r* t2 ]; l0 f87
    + m8 d+ v" t) K2 |0 N88
    8 |2 u  ]9 w2 P2 d898 k8 Y3 a3 Z" \
    90
    3 ~' f  N" C. u- O! S91! q4 ~$ j- o2 n
    92
    " Q* G2 S; |' ?* ]. Y4 B: @& U93
    9 ], ~1 l4 V8 W2 k% W6 K; v/ `) B0 @94+ _/ @) J) N& V- Z' o% M. ~- y/ ]
    95
    : c2 B0 q' w. @. }96
    1 S- n( a- L# N, ]97
    % d$ W% G3 {/ C: ~98
    7 U0 [' U; A8 ~, {7 }99
    ; L" _4 v! C7 E100
    . E' O; ~8 _" }" G7 |9 d5 y101& Y: p$ J' {2 q" z/ z8 a2 |
    1022 D2 ]6 o" i( q1 n4 L/ L( e
    4.展示一下数据
    " m0 \3 ]% {( x* I' A- d2 pdef im_convert(tensor):
    * B5 `; b& n% g* ?+ s; e/ @    """数据展示"""$ \$ Z* B3 \) i- G# m# }
        image = tensor.to("cpu").clone().detach()' Y8 t$ M3 _4 G) X/ p' B- W
        image = image.numpy().squeeze()" k' }  E; L* r+ S# L/ m! b
        # 下面将图像还原,使用squeeze,将函数标识的向量转换为1维度的向量,便于绘图
    + t3 d0 c; `* x* A8 X. f! a    # transpose是调换位置,之前是换成了(c, h, w),需要重新还原为(h, w, c)+ U) v# h0 e0 T
        image = image.transpose(1, 2, 0)
    + F! J+ K1 F" F0 C0 P# J$ P    # 反正则化(反标准化)) C  A  ~; H; \6 Q) b0 O
        image = image * np.array((0.229, 0.224, 0.225)) + np.array((0.485, 0.456, 0.406))
    0 M( ~5 p$ ~: ]( H% S4 e4 z& @! N5 v- j* F
        # 将图像中小于0 的都换成0,大于的都变成1
    * k# Z9 Y' h8 N0 Z2 l& L5 s! q4 m    image = image.clip(0, 1)( _& Q4 C+ r1 R" Y+ J

    7 `2 B' [+ }; I    return image
    ) [' w( ]7 m) C6 F/ V4 Y' o& x1
    3 y+ T. a1 f# `2) ^. V  C  u2 y$ t7 O& Z
    3# _8 ~1 \' f+ H# l4 C& o" k  Y
    43 l& a' W" f) A: _% }  M
    5
    . a+ s  o, k7 t% \# j6" u! Q% ^/ ]: s& r2 t( O2 q$ C$ T
    7; W  U' ?9 O0 d4 B: B6 v5 \: l/ R
    84 `5 O" M) R# K; O: o- e
    9
    . I$ Z' k2 P' C# V2 @$ b; l10
    - a- X  |% u" s' U1 e116 f. {  d9 o+ M$ H+ Z( @: e
    12
    2 B" k, O+ `6 l3 k( z5 s' }; K* O) n3 j13
      X; r9 T* K6 N, @3 R( ?14
    # m4 [, [1 d3 w6 d) `# 使用上面定义好的类进行画图* g/ ^" X3 T" B6 n& t& r
    fig = plt.figure(figsize = (20, 12))5 F8 T" f0 J9 I! m; L: Q' u
    columns = 4# D9 \7 f' i/ e9 j' A, V
    rows = 2# _3 @+ G1 I' Z! `+ x* O

    - N- g9 y! g( T, V1 m& }5 k# iter迭代器
    * ?* e) D' Q7 Q; E! _2 l, U# 随便找一个Batch数据进行展示4 f, i4 u" W( o3 j% q0 q* @0 b
    dataiter = iter(dataloaders['valid'])7 I* V9 c2 s1 u+ H- v
    inputs, classes = dataiter.next()
    1 K% _* A' Q3 s& s7 [
    5 m+ {9 ]! N# R$ mfor idx in range(columns * rows):+ o" H+ ]; {; j& V
        ax = fig.add_subplot(rows, columns, idx + 1, xticks = [], yticks = []). \  Z: G% {$ T6 N4 O; v
        # 利用json文件将其对应花的类型打印在图片中
    3 v7 z. n* H. v    ax.set_title(cat_to_name[str(int(class_names[classes[idx]]))])2 ^6 t' L* Q2 [7 z: s" w& @
        plt.imshow(im_convert(inputs[idx]))
    - q1 y+ u' @5 u0 _plt.show()4 M( k! B3 U0 V* |7 d
    9 _5 W7 P7 P* C+ V
    1: V9 g2 `8 u- o( I; W
    2; p$ l0 z" a  k( {! T$ N
    3
    4 R: y5 l$ C, @48 D0 q; q: [: P4 X/ W
    5
    2 H4 _1 Q" J( Y67 j/ z: z4 U3 c+ ~+ ^
    7
    0 M$ [2 j! N2 m0 V$ t) z& ^; R8
    & \' V: n. |: V% S0 S" `  `: v9
      R, d' ^- U0 U- U+ f. ^4 i" f10
    4 j2 v" O4 p  t1 m5 U0 V' g' [, G9 g11
    8 A! N6 [5 d7 H; M12, V3 n5 [8 S/ A* j- P
    13
    ( q4 g& L( W5 q* d, Y: j14
    $ s5 h$ a9 V+ t. @15' s# Y5 a8 O) B( y, o
    16
    1 b2 R! ?, X, _5 y( B2 D" U1 {) K' ^* K4 O0 ]) x

    6 |2 e. {6 a) O9 b! [% m- ]5. 加载models提供的模型,并直接用训练好的权重做初始化参数5 Y! y- G. U) Q8 M9 M
    model_name = 'resnet' # 可选的模型比较多['resnet', 'alexnet', 'vgg', 'squeezenet', 'densent', 'inception']" n4 F4 e* t5 r# z
    # 主要的图像识别用resnet来做
    & p- V! m+ N. u$ p. d/ D9 H# 是否用人家训练好的特征! [# _1 F% d- T+ |$ S) F' J
    feature_extract = True
    " x& I1 a- X0 @6 Q3 {- [5 Y$ R' P1, v; m0 F* V' _. y. t% N
    25 Q! D* O/ V7 J: L# ~1 H
    35 J. M1 @" v5 n3 j4 F2 m
    4
    & y5 Y' S/ \; q1 S+ m: i# 是否用GPU进行训练9 w5 D0 j3 e, P; r2 r; x! @
    train_on_gpu = torch.cuda.is_available()
    ; i' F1 U8 ~5 |7 {) C1 k
    4 l7 \7 b5 Y# a3 zif not train_on_gpu:; W. d4 d3 g' I' l2 D; s, Y
        print('CUDA is not available.   Training on CPU ...')
    + ~& K" y; ~9 h: z, D/ W" velse:( I$ \4 B' A( F
        print('CUDA is available! Training on GPU ...')
    " A3 d! m* f  Y2 K0 a4 z& b! }
    ) x/ I/ v. [7 u) p6 q; m9 K8 Ddevice = torch.device("cuda:0" if torch.cuda.is_available() else 'cpu')0 I* Q) S' b2 y6 p) r4 b7 q
    1
    3 e9 _" b; ]- N5 {  p) W2
    4 y& H0 q3 H  b3 [- K3
    ( A! L: x. R. s# R# f( W4
    : i1 e0 n2 o9 P! p6 a5
    ! h6 y$ r! o: B9 G6 S( f& @66 z" A' Q& X  I3 M8 e
    7. G5 `/ c0 ~5 h' E* F/ f
    8
    6 O) `) |. H6 }3 k: s9
    6 Q( C8 K/ H; k, O: kCUDA is not available.   Training on CPU ...1 i1 ]7 }) c4 u
    1$ g* V3 E& x# a1 m$ ~
    # 将一些层定义为false,使其不自动更新
    , J4 h7 G" H/ X5 ldef set_parameter_requires_grad(model, feature_extracting):6 j; w+ ^0 d7 i- A/ v$ H
        if feature_extracting:
    2 s7 s$ y6 s/ |/ h        for param in model.parameters():9 P# d4 ]9 x5 M0 y; v2 r& d
                param.requires_grad = False' t3 u4 J% i/ J4 ~& y! _
    1
    & u) d% x' \+ C% f# M2
    + y& R8 ^8 Q! O- R* P: `3. w, v) Q: S! }
    4! Q/ z$ i0 B. y
    5" I0 g$ j+ ]0 b  [
    # 打印模型架构告知是怎么一步一步去完成的2 R0 }+ p% w6 S1 e
    # 主要是为我们提取特征的
    / V2 i9 I3 M5 U3 l; K1 W# X9 R1 B  }" }% O8 l) B
    model_ft = models.resnet152()
    # ?: a( d# D3 T  Q5 Rmodel_ft
    7 M5 d5 Y% `7 ^5 u1
    / y0 j& T0 J6 u8 u0 R2 x2! l9 h  ?" F3 p8 f; ?! R
    35 B+ `6 ]& z. U  H: V
    4+ o" A; ?2 t; T6 Y2 p* S
    5( L+ g6 y) [! {+ O( G# g% {4 D
    ResNet(. k: A2 d/ H; u' A0 C8 x4 \
      (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
    ( P! }6 r$ P6 _5 G0 [+ [! ^  (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    ; y8 Z  t+ A5 Z  (relu): ReLU(inplace=True)( w3 s) K' S7 j6 R( E5 ^
      (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
    % L$ b4 {; ?; E( x6 R  (layer1): Sequential(
    + T- D# @, ?' ]: }' h3 f    (0): Bottleneck(4 W! {( _2 m0 H1 ?4 n- N; Y
          (conv1): Conv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)
    ' _) `% E+ V$ }  \! |- r      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)3 v0 k# J5 s: O" S
          (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
    + A( q7 Q8 O9 X3 y3 t& `& D      (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    : v' Z) e" R, r8 b) R      (conv3): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)3 J; d- [" d- t  m6 k
          (bn3): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)5 N0 |8 c2 _& x% O; L7 G: o* h5 k* x
          (relu): ReLU(inplace=True)& I# X# ?" {% g9 X4 N2 G- n9 I
          (downsample): Sequential(+ H( f9 K& V5 f
            (0): Conv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)' f5 C% U; v8 Z9 B+ w
            (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    9 O( Z. _7 J5 q5 p/ W/ {6 p0 E$ q      )
    2 }4 f0 s0 ^" J0 o1 j1 s( d- d  h! I9 o    )
    7 P9 v3 h6 L  H" h: |! i3 |- g中间还有很多输出结果,我们着重看模型架构的两个层级就完了,缩略。。。! d9 a, u- k  Q4 X: @6 K
        (2): Bottleneck(
    3 R% s, I% d4 D; G      (conv1): Conv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)
    0 d" I, H1 u; D( y, c      (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)- Q9 z* D2 a8 P3 s% |) X" q) }1 |
          (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
    . s( v4 b; F# ~' X      (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)- g0 ], q, `' N8 k
          (conv3): Conv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)
    $ y/ N2 S  L8 W1 I" k$ N      (bn3): BatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)# i$ D6 Z1 U# H0 K. v% d1 b, j
          (relu): ReLU(inplace=True)
    1 q$ B. m0 O- k. m    )" ]; j2 i( f( k1 R
      )' p. d- `9 [7 ?# O7 f
      (avgpool): AdaptiveAvgPool2d(output_size=(1, 1))8 B9 ?, x* y7 T$ \
      (fc): Linear(in_features=2048, out_features=1000, bias=True)! u; b1 d& I6 F; D
    )
    5 A8 I9 u8 P2 C) s1 o3 |# M3 `+ ~6 v: X5 P6 l% c
    13 r: A4 e4 c% F
    2+ B5 |  |. e6 B6 _/ e
    34 N! P1 V8 H9 \' E: t. h
    4
    8 T$ X- o/ _; J$ a3 f# R0 u5) F; T" L1 S) q
    6- `8 q3 n4 P/ X. g
    7
    , b, f0 P& \( `/ f! m9 V( c0 g: d8
    & y6 Z7 T+ x3 Z7 l/ _9 @* r% U9( M9 g) F& n8 L# t- L0 ]
    10
      c- R) b& J3 n0 j5 a8 {11
    - {( q% \# ~+ d$ i12* P; C- J2 n4 j/ I* e
    13% P4 t& r) |8 E
    14; q# P. U/ l; V+ |0 L( n1 U" j
    15
    & E. E# t; B) K16
    % p# K+ h- `7 e17# u) S0 K  A  G* S) Z1 E1 G! ~# D
    188 i9 }5 P+ u3 b4 D- s4 {
    19
    ( V! g( z4 a& d8 X" L20" Y$ z) z+ b% s! w
    21' y  a/ P9 w, S& L7 N
    22
    % q" @/ r1 H1 O1 l23
    : ?& j' b- N! e! [: s0 c24  {. k0 e! ]$ S9 Q
    25
    4 e6 c6 \1 \; D9 U% K8 Q+ M! W26
    $ H# s1 y; @% _! L& E+ B& u27' O* U( K" d9 H0 u1 j% w2 V
    28
    ' j- P( a* s+ c292 d) o7 q1 s0 t$ w& k/ I: J8 _
    309 ^* ]  A6 y* U' Y9 P" u' F
    312 P- L5 ~- u$ ]5 N
    32, u/ t' x! }- P1 t
    330 S3 W' a  ?4 H9 V, W  a0 }
    最后是1000分类,2048输入,分为1000个分类- r5 f! A  ]. J4 ^9 q% g
    而我们需要将我们的任务进行调整,将1000分类改为102输出* v; f8 {! D! s3 A, n& ~8 u8 U, m

    3 |/ U5 N( R6 n  Z+ z/ X6.初始化模型架构, Z2 m* D( |, ^# G  P+ _" s
    步骤如下:
    - z7 O' i8 t7 ^- [* X) H4 F9 @
    1 b  L1 E0 w! _3 l' C+ y将训练好的模型拿过来,并pre_train = True 得到他人的权重参数
    , t9 h  T! B8 Z6 H4 S/ g9 ~可以自己指定一下要不要把某些层给冻住,要冻住的可以指定(将梯度更新改为False)5 r6 b/ F3 t7 }# f2 I
    无论是分类任务还是回归任务,还是将最后的FC层改为相应的参数
    4 r9 D* L. I( C( S: \官方文档链接- \( J! Y. H" m6 c
    https://pytorch.org/vision/stable/models.html
    # a' v, Z; z  K: W& k( Y  B  i5 j4 W$ X+ r
    # 将他人的模型加载进来
      z$ `8 r: x% {8 _% [def initialize_model(model_name, num_classes, feature_extract, use_pretrained = True):
    % t# c( O! {# n7 j! i0 G    # 选择适合的模型,不同的模型初始化参数不同
    : P+ S0 ~$ |  F# w    model_ft = None
    0 g% h8 n0 V/ v4 M    input_size = 00 |6 H4 F- V% c5 _' l" ^
    . q! P2 g. ~2 `( L
        if model_name == "resnet":. v6 Y2 ~4 U3 d( b& Q! M
            """
    3 L2 u% u- C: c3 h( ]. ^* l, J        Resnet152
    3 H6 d$ N' O( s2 v        """
    7 Y, e8 I, S  a& v- a- F! O; K
    - n$ Q1 V, j2 I4 Q        # 1. 加载与训练网络
    7 ?5 S% Z1 D* X1 X$ G5 g2 d        model_ft = models.resnet152(pretrained = use_pretrained); k9 ]( U2 _, x1 b! _5 V
            # 2. 是否将提取特征的模块冻住,只训练FC层" f9 e7 {0 m" t+ P' [9 b  t( I
            set_parameter_requires_grad(model_ft, feature_extract)
    3 Y+ n; j7 V" T0 G- Y& n& |* c        # 3. 获得全连接层输入特征* [4 W! ]1 M$ z4 a2 t! W6 f
            num_frts = model_ft.fc.in_features
    - M, F# t' U# A        # 4. 重新加载全连接层,设置输出102
    5 q9 m* @9 Z3 z! S2 }; E        model_ft.fc = nn.Sequential(nn.Linear(num_frts, 102),- M! Q* z, _0 Y& I" ]
                                       nn.LogSoftmax(dim = 1)) # 默认dim = 0(对列运算),我们将其改为对行运算,且元素和为1
    . S$ a2 X% m% g9 N/ D        input_size = 224
    3 C5 r9 o# G9 a- H: T& j0 \
    3 Y- e' v7 V: P' r* s: ~    elif model_name == "alexnet":
    1 ~6 q" J  H  X; \! W        """; \& _: f; Q! A$ g- V. h6 s1 ^
            Alexnet
    * n) Q4 C+ Q! |5 ?# ]        """6 P; c0 J0 {+ ~8 f' `) R
            model_ft = models.alexnet(pretrained = use_pretrained)
    : q/ a" Y% ?& Y; Z! u% O" f- t* i$ P        set_parameter_requires_grad(model_ft, feature_extract)
    7 I7 Y% J" l3 `6 T1 W# F% ^% n$ N' |) k8 y  L
            # 将最后一个特征输出替换 序号为【6】的分类器
    % [0 J( `% R7 A2 W! X- @# N        num_frts = model_ft.classifier[6].in_features # 获得FC层输入: O" B* Z1 b3 B0 k' \" C$ I* L
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes)3 o. |- \( v- Q' n
            input_size = 224
    + D6 x8 M  m$ t8 G# B' r; k* s, _" v# v5 ^
        elif model_name == "vgg":
    / V  w! o$ P( T4 P3 C        """& j( ~5 E5 X5 Q! `
            VGG11_bn
    " f: p% I2 J- C- h: d        """3 r$ q/ n* U8 }! H) c! X
            model_ft = models.vgg16(pretrained = use_pretrained)/ l! I9 f9 A! A; n6 p
            set_parameter_requires_grad(model_ft, feature_extract). [# J! u& K/ ^. ?( n
            num_frts = model_ft.classifier[6].in_features# \; v5 E4 \' ?" C
            model_ft.classifier[6] = nn.Linear(num_frts, num_classes)0 D; u. t" s3 E; @' f; r5 D
            input_size = 224
    : G2 e1 F3 A0 B5 ?. b2 N
    / t! \& l  U" R3 T. Y; Q    elif model_name == "squeezenet":* S( p% }- O# m
            """
    0 s/ T! w1 a, G9 @- h4 H6 s        Squeezenet0 R3 U% H7 q) Z0 d( x% E! }
            """
    ; r2 h9 a0 b3 D+ T& G! G: W        model_ft = models.squeezenet1_0(pretrained = use_pretrained)2 j+ B1 e6 j7 w( f7 E) E
            set_parameter_requires_grad(model_ft, feature_extract)
    . s+ m. P8 k( N. s        model_ft.classifier[1] = nn.Conv2d(512, num_classes, kernel_size = (1, 1), stride = (1, 1))- K! C' x! z' y" s0 @! k3 c
            model_ft.num_classes = num_classes9 ^2 }0 x( v& _( M: [6 T5 y
            input_size = 224, C7 B$ k8 i! U) G5 I

    % U/ c' j, \' I6 R6 C& O    elif model_name == "densenet":
    # o/ p) q" X7 {6 {& O+ `1 I        """
    . @" _# H# N  B. p        Densenet
    ' E3 {5 V& K! H        """
    7 h8 N5 [+ V! C; F/ B  d- o3 N        model_ft = models.desenet121(pretrained = use_pretrained)9 w, o9 \" ~7 c; v, i6 {! n3 J
            set_parameter_requires_grad(model_ft, feature_extract)
    + R) D/ N3 K; t5 b: u        num_frts = model_ft.classifier.in_features
    & T. a4 |1 _5 v4 m& b        model_ft.classifier = nn.Linear(num_frts, num_classes)
    * \: b; j; B2 e' N1 u7 {2 y# J        input_size = 224
    ; E0 N) B; G- Z( z4 [1 b0 _& T2 a& o( T: V7 _# `2 [+ F4 e7 I
        elif model_name == "inception":
    8 N8 ^; k3 r% j1 W" x        """/ K1 {) U2 U# a3 Z8 @0 ]
            Inception V3
    1 z* ~4 n  L. C1 I8 A        """
    * z7 X6 e2 `  Y& f# k        model_ft = models.inception_V(pretrained = use_pretrained)% D$ l- M9 [; r4 r- l( R8 G
            set_parameter_requires_grad(model_ft, feature_extract)
    $ T/ R6 f& I; t: X) {% T8 x! e! v3 S8 P7 z& B, O* o$ Y
            num_frts = model_ft.AuxLogits.fc.in_features; ]! v5 O( I' [1 o
            model_ft.AuxLogits.fc = nn.Linear(num_frts, num_classes)
    5 P9 m) b7 S% X( D& u; q2 M* o% u0 g% R, h* a
            num_frts = model_ft.fc.in_features
    / O2 p* R% [& T0 C        model_ft.fc = nn.Linear(num_frts, num_classes)
    , y4 C1 y3 t! _7 u# s) O, Z        input_size = 299
    6 ~7 i& b1 ~' d  W  P
    9 X0 Q  j% z6 P( d+ k) j    else:
    * I. V$ D. X, }1 ~  G7 |        print("Invalid model name, exiting...")' @3 |2 g$ k( n$ v: B# W. H8 ^
            exit()
    9 J  d* w) Q1 V, n, b9 d* e( I; M$ G1 p# P* {
        return model_ft, input_size; j. B2 D& v8 D! Z1 i9 {
    2 g1 Q6 p0 C( y6 Z% j. R. J
    1
    7 X/ r( Y# V; j9 i0 d# R9 p9 A0 _2
    ' I, f! Q4 W6 ?& m( m3
    % u. r+ Z- x3 q1 _7 ?1 U41 ?. v& a7 m7 J; y
    5
    : L% u' M/ \% X& Y$ v3 }6" b7 ?# v2 g$ J
    70 S& a$ P# g6 }3 b
    8
      n8 h4 J3 [) w. {8 U- F) M; W9( ~, E& I7 g1 K, p
    10
    . J6 i  g4 ?6 j( P/ k2 U7 d' m! K11
    % x/ V" k  p" @. H9 i5 Z120 X! f% `" U8 |; {
    13
    1 Y4 P9 Y/ T) a) N; U( L6 D$ B$ N' V14
    # e( F% e* ?! m15& |) k+ H" Y: l/ r
    16" D+ t/ z( ]2 n& F4 i8 r. `/ ]
    17
      Y# l1 t, L, t7 p. z18
    % z/ Y: e) q3 ^. a0 _$ x9 \( r6 }0 _19
    * [! _& Z9 [+ z1 j20* u: ?7 s0 E4 [8 a, R
    21+ ]) C  q7 [+ K! n; {; J6 e
    22* N- C3 f0 ]- s
    231 v0 K# \4 M) B8 K
    24. Z0 C, v4 A* A) P" i
    25' _* n, a+ A6 R4 F% U+ {4 W2 W
    268 D& ?+ }3 ?7 Q- J, H- m
    272 Q8 k! h3 x7 X% x
    281 r/ w, j: s  I4 s3 ]+ O0 t
    297 T- c! a3 e9 ]0 n' w
    302 u  E6 n+ \+ \: R8 {
    31
    0 |$ V* F! o- X4 z32
    9 M/ W: I% O, s% P* g8 T33
    4 W* H+ V  ~+ f8 {6 x, G; c34
    5 D7 p8 g" w: W! s) {35
    0 M) O0 N" L# M  T9 ?6 u36
    ' w2 q( U4 A# @6 F6 P) ?37+ k3 e5 ?' I; R
    38, s4 q8 }  ]5 t6 M; g- H
    39
    8 l( D$ s2 i0 V, c. Q* M40
      \7 \& a; Q% }% e" r8 O41
    , P3 p5 l! I' J4 l& R+ M) {  W423 o9 I0 j, x3 F- P6 t
    43
    / H) e- N# k( }44+ r5 }6 _5 y% q+ {9 M
    45) N3 B! K* M3 m% z: {# \5 l. b
    46( o: I+ y+ T/ Z" y
    476 g* \0 V. ?- R& Q5 S; s
    48$ `: b9 a( K" J; K
    49/ h1 Q% [6 }' }
    50! r" ]! b7 `* B" e
    51
    ' l0 q8 K: r2 n2 R52/ j, s  l" R" o. R! c$ F
    53% i3 r* H3 E1 l, B7 L
    54
    , C7 o( \, J3 V& Q555 P3 Y; D. B3 m1 k$ h9 x6 [
    562 V& U0 y/ s' Y" ?( e$ Y; i1 z
    57
    4 M! q/ g( Q; N6 k1 \( X584 v; _( G% x6 F5 C# {) a
    59: y7 |' p$ x, y0 Y8 s( M2 T
    60% j! {' c$ u6 X0 l7 Z2 N
    61
    ; c, r; C/ f4 N$ T# a62. z) X1 F1 f) e9 s' {' U% \# c
    63
    ' w+ s4 \3 Z4 q9 U" g- C64
    ' @7 C/ o$ G$ Z! v# p7 h65
    1 ^/ X: A& B# d0 P2 v: u) U! K66- u5 y2 y; D5 o% P% e! m; b
    67% v6 M+ ^& q! J) n8 H6 x
    68
    4 ^8 ]+ n: M: O- v( J69
    ( z$ o3 ?' ?/ L( w* Y2 l70
    " X( |3 [( w- U. s; `* L71
    # [! W; g, P- q/ s8 [4 z72
    ' P4 M& h# }7 G$ c4 _4 K, p73
    % L* w, F6 L# W& s! t9 X9 M74. R* _  E2 {2 P7 L
    75  |6 [1 _( k' J3 c0 g9 ^8 X- Z  Q
    76
    : Z1 p+ r) Q1 G5 {) J77
    0 P" q( h# s" t5 c+ [9 E78
    ; q# `/ I8 O# L  k" V2 C& K' O79* ?$ r9 i+ h4 J
    80
    4 n; V' N0 S! Y8 a/ {81
    * t& h: s5 W# V820 `8 \. P# w: [1 r" \8 y, `
    83% m+ |0 U6 c8 \, k# L+ z0 l
    7. 设置需要训练的参数
    , y* h' \+ p: J& f1 M# 设置模型名字、输出分类数7 N% j1 }* @3 o& W
    model_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained = True)# y) J; T/ P* d: C" m8 V7 r

    8 Y1 M% j# v/ S, i+ C# GPU 计算
    4 _4 c, J1 x' |. X3 i7 \model_ft = model_ft.to(device)
    . u% `4 t8 i6 [5 F( o3 r4 u
    6 f" r' z( ]2 E: Z! H8 H- R# 模型保存, checkpoints 保存是已经训练好的模型,以后使用可以直接读取
    + |: P; G0 X& G1 r2 w  I$ rfilename = 'checkpoint.pth'6 e, l, T/ p; |+ m) a

    8 D" C' f# g* Z6 K1 Z: H2 w# f# 是否训练所有层
    0 g( E% ~) ~% q( N! d! O# }. gparams_to_update = model_ft.parameters()% P, E5 t: z% O- A& r
    # 打印出需要训练的层' Q: d& B/ f8 f# _: p+ e$ z
    print("Params to learn:")
    ( p/ V8 M4 ?) P4 |' s- L" K, ]6 gif feature_extract:0 {0 \: ]( Z; j
        params_to_update = []
    7 z4 B2 \8 O) {    for name, param in model_ft.named_parameters():
    7 ^3 O9 k' P8 ^6 X        if param.requires_grad == True:
    3 W. e* x$ [% m- y" B( s$ p! U            params_to_update.append(param)
    . ?9 ]+ d, k) M3 p. c0 [% G6 u            print("\t", name)# l  O: s( D4 @! [& r. k% C2 Z
    else:6 {1 G  X6 v8 k8 D
        for name, param in model_ft.named_parameters():
    4 T6 B, n! ^$ ^  m0 T1 L1 X        if param.requires_grad ==True:7 D2 a3 a' H3 Q& ~! _
                print("\t", name)2 v( c/ {) z4 o  r! V
    + q0 ~$ w0 y) f& u" g
    1
    ' E* y" u  Q/ j9 q0 K3 h3 \! ]: J% y2
    3 G/ T9 `. a2 W4 O! f: D& ]3
    7 Y2 g2 L  z$ G- |1 `4
    ' c4 ?$ ]& P5 Q$ T5" ^# c+ E, D) K9 d, {. I
    64 H0 i% g6 V/ d# ^) c* }7 g
    7: h" f) [7 c3 V" T* X( Z) r
    89 L. u4 ?# U! i" t$ y8 s
    9
    ) r3 P& s& T6 x1 M! `- v6 \100 \+ y) _" l& |
    11
    + C. Y4 Z) X+ |# q" B7 S12
    9 ?7 G& [2 e$ D13
    $ O" w0 [7 W* ?; @14+ E8 j1 X. x, S6 w+ _) q% @/ M
    15' e9 {3 x. p$ F" ]* l( t  F% H
    16
    . }1 `  O: N: g9 y7 _17
    8 s% O6 m0 U5 k. j) [: e  K- ^  D* E* r18
    3 D/ Q" P+ D" |+ t19
    8 l/ E9 @; s/ R0 ~4 n/ h20( o* i$ ]" [2 n/ n& v7 L
    21
    + m2 B5 l# @& y228 }3 S/ |  p4 u! F
    23/ y; e1 y  u' D0 _+ N
    Params to learn:
    & d4 K/ s' o& U; P0 r7 T         fc.0.weight- O5 |' i; |# H1 t
             fc.0.bias$ y" o$ n" l& [9 e9 s
    19 w! ^; r) f. `3 T$ K0 L
    2
    : H# J! s. x, C& W  V5 N3
    7 [( C, ?& u0 c7. 训练与预测0 P- N# q* [' C/ C8 {3 u6 R
    7.1 优化器设置! q0 M1 D6 H$ E. J
    # 优化器设置
    6 X+ ]1 J% K8 ?/ G  [( t# J! noptimizer_ft  = optim.Adam(params_to_update, lr = 1e-2)
    4 D6 k" X% e% l  N- M3 S" O# 学习率衰减策略
    5 W! ~9 q. f+ `9 E1 ^scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)
    : ^& \) T8 v& q7 v# 学习率每7个epoch衰减为原来的1/10
    % o- q6 }4 k! s3 M1 d# 最后一层使用LogSoftmax(), 故不能使用nn.CrossEntropyLoss()来计算: I, x! t# ?% V7 J
      ^8 c' v# k. Q; O
    criterion = nn.NLLLoss()4 {1 q8 n/ g% G# T
    19 I" e" ~7 _+ J
    2
    ) W9 @3 z4 V4 w31 r2 z) J2 W' F2 f1 Z
    40 d8 W, H3 n5 P# O2 e/ k4 |" F
    5
    , H( U9 g  X; p/ S) R6
    ! x# P* @% s* ~* i7
    $ p$ n( ?% p' `( \  C8
    . \: ?- }* `- D, e3 ~# 定义训练函数3 j. N; K" z, P9 i5 O! l7 M+ I% @
    #is_inception:要不要用其他的网络
    & a' L2 ]- Q% ?: Ydef train_model(model, dataloaders, criterion, optimizer, num_epochs=10, is_inception=False,filename=filename):. u) `$ M( e" z+ W( i% q9 Y
        since = time.time()* [8 \  s: V8 o( y+ j' N
        #保存最好的准确率8 |( h/ c: ^8 @4 C3 y8 Y
        best_acc = 0  M( q+ g" B8 z/ r8 Y3 J' C) `0 h
        """* X- C3 [& |" H  X5 R0 m+ a
        checkpoint = torch.load(filename)
      y2 `- A" E/ l, O2 _! _6 J    best_acc = checkpoint['best_acc']$ a  K: w2 w# D5 G
        model.load_state_dict(checkpoint['state_dict'])
    9 i6 u6 f1 E, S5 d0 S    optimizer.load_state_dict(checkpoint['optimizer'])
    0 ~$ G% p- Z( n- i( r1 T    model.class_to_idx = checkpoint['mapping']
      ^& d! ^0 K' w$ P7 A( A* }    """# |7 b$ p  _# x, \/ k
        #指定用GPU还是CPU8 D+ E& \1 s8 u% }1 `
        model.to(device)
    1 `& [3 G4 J. l% g( |5 `$ Q    #下面是为展示做的
    + @; {$ w3 r) u! h    val_acc_history = []4 l1 u: _' K7 ?, u3 ?6 x
        train_acc_history = []
    ( f5 E  ]! y* e& ?5 h1 y    train_losses = [], V& {  E$ M$ ]; V$ `0 E
        valid_losses = []/ t! p% H& u# d% N9 k: t( B
        LRs = [optimizer.param_groups[0]['lr']], `0 e; i$ }" r6 C; T; \
        #最好的一次存下来+ @7 c7 i. z. M: L# c$ j+ @3 N
        best_model_wts = copy.deepcopy(model.state_dict())5 n9 q& Y% H7 v$ N

    1 [0 S  O+ x* g$ W: |    for epoch in range(num_epochs):, }6 q! \& H+ S- ?2 B7 I
            print('Epoch {}/{}'.format(epoch, num_epochs - 1))0 z2 w8 V# t* e- I# H. w
            print('-' * 10)
    3 X( ?) w0 c- k& {  W: w4 H8 P8 u8 g
            # 训练和验证
    8 r5 P+ L' g9 |; L1 Y        for phase in ['train', 'valid']:+ ]% p/ j( P/ q4 f9 X
                if phase == 'train':0 V+ ^4 n/ ~/ W3 i" j
                    model.train()  # 训练& a! o8 D9 z3 P9 A8 X5 }
                else:8 z8 Y, ^. O/ f& }
                    model.eval()   # 验证
    8 s1 v+ ^* S2 C" |
    8 w2 R- t/ G9 Z7 s1 c! u            running_loss = 0.0
    5 Z, S% s2 C& F" d* L& M! L2 B3 M            running_corrects = 0
    % F5 Y3 ^# j- ?
    & c* ?3 ?3 z/ `" F4 S  y6 }            # 把数据都取个遍
    ) a# K5 i& Q: a9 X1 f            for inputs, labels in dataloaders[phase]:5 [1 U" Q" R6 {, s) x
                    #下面是将inputs,labels传到GPU4 l* C  Y9 J9 {* J  H
                    inputs = inputs.to(device)' w: z; g/ e) y* Y; O, w# s3 R2 j7 |
                    labels = labels.to(device)
    ; Y5 o$ ^9 X8 C$ `/ j9 G0 ?9 ]& K3 m$ Y6 t
                    # 清零2 ^- d+ @5 G/ j; o5 u! q% }
                    optimizer.zero_grad()6 `$ V2 F% h7 s, [" G
                    # 只有训练的时候计算和更新梯度
    * s+ O# D% K- F( S! [- M                with torch.set_grad_enabled(phase == 'train'):
    * Q0 G& |! \0 l# G) G8 t                    #if这面不需要计算,可忽略- D. h/ d4 l: [, O" N* S; P
                        if is_inception and phase == 'train':
    1 g5 E# R' ~' ~) s, I3 _                        outputs, aux_outputs = model(inputs)
    & y% d7 D% }0 b& k                        loss1 = criterion(outputs, labels), x; D' {* `9 }
                            loss2 = criterion(aux_outputs, labels): c! Q  F9 {1 i, @' ?
                            loss = loss1 + 0.4*loss2
    4 V, @. x* ~5 x, |8 ^% {5 C1 `" X                    else:#resnet执行的是这里
    7 U7 C8 n9 i! y9 L: u5 C! h* _                        outputs = model(inputs)
    & z9 }0 H; D2 b5 T% W6 r! s  ~% J% P                        loss = criterion(outputs, labels)
    ! E( c& Y/ ]3 V
    6 ~6 \! d7 R4 B                        #概率最大的返回preds6 K+ l% z+ G' L# K$ A- Y' m2 H3 X0 ]
                        _, preds = torch.max(outputs, 1)
    0 D! q2 N3 _# ?8 S# u, Y! B; }# [  Y* O$ l
                        # 训练阶段更新权重$ O; e4 J; N. P# H5 q
                        if phase == 'train':
    # t( ]9 Q' Y: t6 X+ ~' P/ b                        loss.backward()- D( h& L7 C2 h5 `. e" U
                            optimizer.step()+ |8 j* P0 @# w3 s7 W
    % M+ v  o. Q, R0 ?! u
                    # 计算损失
    ; M$ A& k- x7 a) M                running_loss += loss.item() * inputs.size(0)) \" x+ m' v6 Z2 v4 V: N
                    running_corrects += torch.sum(preds == labels.data)
    0 Q8 T- r/ @/ Y% P1 M# Y" K- g8 N! R
                #打印操作
    " j, F; C8 V- c* s* a  b            epoch_loss = running_loss / len(dataloaders[phase].dataset)
    ; M1 D" u. _" ^+ v1 ]            epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)
    . |) y4 {4 w" r8 E
    ) b( p1 J2 r9 H" `+ n( r8 |1 J
    ' C: G; L' B5 t            time_elapsed = time.time() - since
    : v5 F- @2 f0 o3 j- l            print('Time elapsed {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
    . u4 c+ e6 M. l/ E+ t            print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))
    7 L6 c6 z- X1 K) z" f8 J7 w! g- |( x. R" ?6 k, F* {6 U8 F

    9 K" [& |. k  f  O& ^, \: G- `            # 得到最好那次的模型& S4 H. c; d; p, ?$ W/ d! J
                if phase == 'valid' and epoch_acc > best_acc:
    5 q' }+ ^5 h( {) y- x1 J) k' F                best_acc = epoch_acc
    ' P" \0 l* v; U) P' K9 q                #模型保存9 n% W' y. P6 v, E0 O, a$ r
                    best_model_wts = copy.deepcopy(model.state_dict())
    / O1 H5 q6 t6 C7 ]4 V                state = {: A. w! U& C) ]
                        #tate_dict变量存放训练过程中需要学习的权重和偏执系数
    - g  ]0 x/ ]- W1 W0 A                  'state_dict': model.state_dict(),. k0 _. _+ v0 [# Q. Q
                      'best_acc': best_acc,( s- T2 s/ I1 y- r$ h5 a7 l
                      'optimizer' : optimizer.state_dict(),
    ) u6 Y: l$ e/ E; P% s7 K1 T  S! g, Z                }
    * K9 J8 t4 c& R9 o4 K  e" |                torch.save(state, filename)
    * S4 g2 ?9 X: b; P            if phase == 'valid':0 ^( C0 L& X  T
                    val_acc_history.append(epoch_acc)0 S3 Q" R, D2 r6 w
                    valid_losses.append(epoch_loss)
    % P) K- N/ c+ Y/ S3 ]- K/ T                scheduler.step(epoch_loss)
    : ]* t: U; }/ Y- a            if phase == 'train':" ]( n* n: H) W8 ?6 L
                    train_acc_history.append(epoch_acc)
    ( e' `2 a0 |! d; I# ^8 o2 j, x0 l                train_losses.append(epoch_loss)
    5 `1 ^9 {5 H# G0 X1 o3 C
    % {. ?6 c7 e5 I  ]# F9 B1 F* `        print('Optimizer learning rate : {:.7f}'.format(optimizer.param_groups[0]['lr']))
    3 z  ~  G# u3 O7 ?        LRs.append(optimizer.param_groups[0]['lr'])( z. L, p6 n; N4 m0 }
            print()0 K' U1 c) d: P: B( j* v! Z% {* [$ ^- B

    4 t- d2 J1 v& w    time_elapsed = time.time() - since* _3 h/ G# V" r7 O5 ~0 p
        print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60)). X7 ]8 c4 H9 F- O. C& k. j/ k( z
        print('Best val Acc: {:4f}'.format(best_acc)). _1 u; V' t' M! j. t; D2 Z0 ?) Y2 H

    ' l! H; I1 V- o5 l0 J1 S, F, p    # 保存训练完后用最好的一次当做模型最终的结果0 Y2 P4 q1 G& ]* ]1 F. v
        model.load_state_dict(best_model_wts)
    ) i4 R/ ?' ?: `( |& r  N5 d; Q    return model, val_acc_history, train_acc_history, valid_losses, train_losses, LRs 7 Y% T) @( Q  x9 \8 [! w

    * y) q( t% N' ]: A1 t6 J6 h( c4 @# M2 o) K# G; R4 C0 _
    1" l1 ^0 [. o7 P$ @( L9 o6 a
    24 ~4 @% X& Q2 @7 r
    3* U& s) y/ A7 h" x, a6 h  H
    4/ ?* ~' s! H% c, a. e4 Q% B
    56 ]4 q/ V3 p5 i7 `* \$ y
    6
    7 d' e. M0 ~  R; u' o76 A% Q* z. {7 a  |
    86 L5 o: k% _8 s: P
    94 k; }! U9 T2 J0 j3 [; G
    10: W; D6 j  m+ e+ ~+ b( I& {
    115 a( {1 |. k7 @! E3 q0 x0 j
    12
    ! j, x: r7 l, I0 V' q( J% z13
    , Z9 S7 v- z( z$ T14
    # l% P: Q5 u7 ~15
    ; E/ N( P5 [4 Q. z5 O16
    2 Z+ U. g# V! |, Y/ u6 A17
    3 ?7 j/ W/ w2 h% W8 u18
    $ Z* [# v! @& ]+ ^6 ^0 T( I19
    1 O, v( I  e7 S' ?. j9 Z# Q1 a200 ~! r5 `4 y' S$ n" e
    21- _. ?8 @0 }! X/ U; o- O
    226 {7 t1 }! g. X3 R# G% K
    23" g9 N% N7 C$ Z  {. y
    24- i% k% v( B' L: R  X2 t' z/ Q
    25: j. [7 q4 D$ F  L- S
    26
    3 G* `4 F9 g4 P% `' x1 D1 a$ t27& Z& N* L% q9 H2 `5 f3 @0 c
    28% h! B" m5 S3 h# S/ g& b
    29- j" M1 |+ m. u/ u" g
    30
    ; _& b( m4 Y0 u9 [31
    / M$ s7 N! g  n( o& v, G32! w- r% [5 V9 M7 y" n- p6 D
    33
    5 v2 d9 E0 j/ \/ j346 o/ F  C9 q0 n3 B9 x3 B& O
    357 L( l5 r( ~: o, S7 N
    360 D* m2 E6 s1 g. K. t
    373 e" i8 k% N7 B4 z' a* }
    38' e5 E# M0 l7 ]6 Z% _" b$ N# R
    39; `% K4 T: B  |. s% p6 ]
    40/ |1 ]0 d7 M2 `: ]0 `
    41+ d9 V: I1 b7 S& ~8 y) t; _
    42( u, o8 l  L9 H- Z# ^6 ~  \0 k& K* n
    43
    0 K! O7 q9 I& i# m44
    . \# h  B! S6 k; s45. @. X8 a$ X# R0 u& a! x. M7 z
    46; [: k0 G3 @  I1 \( }
    47
    / y8 s* D* ]( ^( a$ k/ \8 L/ D48! ^1 \" }* h0 b' ~9 a
    49( I9 s* [# v# \' Y) r3 F1 p9 d. j
    508 C1 C! Y  J" M6 [; ]0 L
    517 B% Y5 P% \, V+ |, L
    529 I* N- n) ^* c* L+ Q: s
    53
    6 }! P2 B5 a8 n# G9 u( m: O545 z3 g7 e; d( M. a3 v
    557 W+ \: n) Z9 V, f
    56# e6 U3 a: A7 C' N' j& @
    57
    4 s+ m8 P1 L! m, c6 B- ^' X; y581 V! ]% H% t1 B* u
    59
    : e: t3 e+ |$ d: `% q; S0 g# g1 d' v60$ V7 S2 @: `9 M) H5 {3 C
    618 X+ r3 ^0 ]$ R
    62) z  |, B. ^4 y
    63
    3 C- {9 t' Z# v0 c, m* ~64+ s; E% K( B4 |/ V" l4 f
    65/ G8 e* V2 T+ E6 q% m
    66# y& |. h9 y) [" q0 Z
    67
    / H  o4 U7 E" x' c5 S  y68
    0 o+ ^/ [* N7 K; I/ H" G696 @' l( W" x8 v
    70# K2 x  D' b9 ]+ {) |8 j  Q
    71' I8 p* K- M0 h' t* ?
    72- E1 Z0 w/ P4 g* Q" `' u
    73
    5 ?  C* U) u8 ~6 c* r2 o74) I6 n7 A" p  Q0 i
    75+ k& r1 n8 e8 i( W" ~, d, I5 R
    767 I( |2 r2 j7 }3 k
    77+ a$ u( K3 Z: W$ N
    78
    ; K2 z- ^- D: b4 I79
    , C! ?4 L* v" r808 f& D  v' K2 x
    81
    , L5 T2 Z  P; D82
    ! j6 E( h* K# W; g83
    6 L4 ]4 E6 \1 L7 k) I848 Y2 x- K2 Z5 P/ D) @2 D- b
    85' x; x5 T! Q; z( n& r# M3 p9 H
    86- l9 H5 e' ?( c
    87, D$ |& B; o! b: A. ~
    88
    0 I/ g4 }# g; F; b2 `: m89( _8 ]. M0 V( B* y9 B
    906 l5 ^! L4 Z/ i, t) d0 o9 K
    91$ X  j+ M8 r+ u; Y- i" }
    920 P& E% o9 i* S- H- R- P
    93# Z1 Y$ I  R" b% K+ u- B
    94
    6 a. {, F$ _  z0 X+ b3 q95- @6 g7 F$ A3 f  h' ^5 c; I
    96
    8 ^; h8 D* Z, |" \+ w97
    2 L& y2 F7 V- ], r6 l; j98
    # e' x1 V# |: Y, E" c0 j99# w- s* t% Q& \" C0 C! h" t5 g: S. W0 I
    100
      _% q6 u0 r! ~+ }$ d' Y101* b3 Q0 w2 H; [8 ]7 o6 k3 m
    1021 W1 S9 U& E% l" o7 p( G/ `
    103! `) H4 w4 s. s  R% b4 i" f0 z: u
    104
      Z# J; _; c' e1 Y: m; p! A105$ R' A& n: N0 k5 j1 C, {6 B, \2 `
    106
    . V, j0 z; U. j+ X2 b1073 N/ _. Y8 k* u* Z5 d
    108
    ( j6 ^; A/ F+ p7 w3 u2 S6 C109
      y. T: z: H1 E1107 S3 ^* X, ?) m+ z6 D
    1110 Z) g' R% F1 X& g+ J
    112
    * h! @  k1 ]3 P7 u2 N; ?1 P7.2 开始训练模型
    % Z' E$ V) V. Y7 J4 W. h  E我这里只训练了4轮(因为训练真的太长了),大家自己玩的时候可以调大训练轮次
    - o' J) w2 Y# ?+ t6 }$ Y
    - @! j8 |3 i* f3 {#若太慢,把epoch调低,迭代50次可能好些
    $ B" g  Q, }: Z) c% J#训练时,损失是否下降,准确是否有上升;验证与训练差距大吗?若差距大,就是过拟合
    5 ]  _2 q4 u* W! bmodel_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"))8 l9 v. N% G8 n* s0 c

    : U' s+ O$ N+ V& d1
    1 A: U2 J# }5 K5 I8 |2; E: K% r2 A$ J/ a1 c+ M6 I
    3
    1 N. b9 ~$ T$ C  H$ c4
    2 L/ u+ m6 G* P1 u! cEpoch 0/4. v  w, ^$ T3 _/ t
    ----------1 f$ F% i! z$ A! S: p* G
    Time elapsed 29m 41s# F) T7 w" k& i8 w$ K5 I
    train Loss: 10.4774 Acc: 0.3147: Q2 R& E) f7 J5 n& l; W6 U
    Time elapsed 32m 54s
    0 j  N: S, b! z+ Z% q! B  \& Xvalid Loss: 8.2902 Acc: 0.4719# M4 M4 H  e/ R; O3 j4 f2 ]! ?0 e
    Optimizer learning rate : 0.0010000
    0 t" t& S% X: s# g2 r2 i
    8 S4 L* r9 W! q6 CEpoch 1/4' B1 r' U& o  g& Q! `/ ^- p# J
    ----------
    9 L5 Z$ O7 [  C7 hTime elapsed 60m 11s" [& C' ~! W+ k, L  r3 j( N6 \3 G
    train Loss: 2.3126 Acc: 0.7053
    / m$ U/ P1 O1 q$ P1 x; j1 YTime elapsed 63m 16s  W! B5 d3 F0 s4 S
    valid Loss: 3.2325 Acc: 0.6626
    & ?/ \, C  h5 L* S% c0 bOptimizer learning rate : 0.01000004 E4 H- _9 `( i8 I0 g6 T

    + r' ^) D, ]5 m% n# uEpoch 2/4( J) Z7 u& B+ ~- N7 u% m
    ----------
    4 d- y; X7 h6 K" k0 E7 bTime elapsed 90m 58s; l! p2 H2 `8 J3 X) K3 _
    train Loss: 9.9720 Acc: 0.4734
    . X9 m+ H  v, T7 f9 FTime elapsed 94m 4s0 j; I. \9 \3 L5 F
    valid Loss: 14.0426 Acc: 0.44135 Z0 P/ F1 e" K; h; w; {  j
    Optimizer learning rate : 0.0001000
    * Z$ f  i( o8 _4 F! ^& _2 K* I) t& \, y" K+ ^! `: z
    Epoch 3/44 |2 ^+ B+ J  e/ \6 z+ u
    ----------
    5 Q% F  h. b) nTime elapsed 132m 49s
    4 O. f5 P  p6 D( z! i, Z8 J; Ftrain Loss: 5.4290 Acc: 0.6548
    7 u0 y: Z8 c" q- [Time elapsed 138m 49s& v# L# ?- ]0 S+ w9 d, X
    valid Loss: 6.4208 Acc: 0.6027% e8 _" e9 s2 i9 ~  z. k- H
    Optimizer learning rate : 0.01000004 ]% e: C9 F; X8 f7 m4 Y
    7 z$ [* n7 w" N% G. _
    Epoch 4/4
    $ W9 L) b& Q  i& \$ k----------
    ; p: i0 H8 k. J" OTime elapsed 195m 56s
    + S& ~7 s; y2 b/ Strain Loss: 8.8911 Acc: 0.5519
    ) i9 E9 ~- M- t: w+ b! K& `; `Time elapsed 199m 16s
    4 f. w, W% w4 ~6 P) R" rvalid Loss: 13.2221 Acc: 0.4914
    7 F- h! X, T, p. E: n! O  r* ROptimizer learning rate : 0.0010000
    , }5 ^" t0 n$ o/ Z" K1 `) E' S# s( ?3 H
    Training complete in 199m 16s0 B' f0 u. N, K& q# ?( {
    Best val Acc: 0.662592  H$ t5 F' i, o

    + ^6 h3 _3 g; m6 v; A8 C/ P1; y# o: N5 w! [' i  X9 ^. n/ R
    2
    & F' t: p) |" B" M$ t3
    & Y( |8 k: D6 N  T6 \45 `- H* H' Y; M3 D+ A- d$ I# ]( V9 _
    5
    9 H3 P+ E9 v4 U5 F6. G& Y. }% k5 I6 R: K6 ]
    71 y: \( L2 y' G' h
    8/ v# f! b3 [8 f, Z+ ]1 H/ C1 r
    9: `1 {( ~& N! ?& a3 J, l
    10, t$ C+ j& k2 H- A7 I( I: P5 {
    11
    2 c1 b! v; v5 g* |12
    / M1 I+ r9 V% L3 i8 L; r13
    ! S1 m/ ^: N* ~4 n! w, ^# O' d, Y14
    ! ^, y: {% }( b2 M6 g15
    % D7 o' J: l- n4 b" Q& F2 s4 M16( E' j, g8 S' p# r. i4 E+ |
    17
    $ v6 r* o# @& G- Y187 y1 g: g0 s9 A9 v/ `' X0 w' R. p3 B
    19/ ^, h9 k+ @3 f3 N/ m% ~- ]
    20
    + [2 W0 S! d* b217 k; i, d" @" |0 O. T
    22
    4 C' y. M. o  C4 p) o& Q6 V' u23
    , i; B7 G" r0 f+ f9 f. C24
    7 @% V7 M. x9 }2 |25
    / o+ U/ g/ c# @+ M) B- [' c6 }26' z$ Z$ b/ ~, P; o7 B: @/ F
    27, F' O; t/ r) J1 {1 H7 V! d2 u6 q
    288 Q- v3 t  K. J1 D# S/ u8 q& F
    29
    8 f: M) ]/ J  P0 g  K; C" G% [5 j30' b/ V( Q) B  `, y2 f  Y: O6 R" l
    31/ t! U" R9 }. O, p( y) n8 r) p5 V
    321 {: P& V# v3 w4 g+ O% ]  \& J( n
    334 C* o" h( @' K2 T( o  A" i
    34
    . H3 D5 O9 b. s" D( A35
    ! s# v/ f/ K% Y3 a36
    5 m) O# a. _3 {. d37
    3 Q9 t/ P/ w1 |5 I# V: a. ?38: W% r# E; g2 p( v" K7 F' ^
    39
    $ u1 G- _5 z; K4 @) a; H3 \8 b40
    6 x# b) m& y2 Q  a+ k41
    . a5 V, ^: J" I7 k% M& ?42
    ' c8 T# ^/ a9 O& y9 h# Y9 J7.3 训练所有层" c5 |8 r: O  E; K8 U
    # 将全部网络解锁进行训练% s8 s/ M' A& A, j# S5 m
    for param in model_ft.parameters():8 w! x/ z5 |0 M# V) z
        param.requires_grad = True
      R! u: ~( R$ }( X2 ?* k7 x/ s- Z4 ~3 A' y" A
    # 再继续训练所有的参数,学习率调小一点\
    ! N0 [) z& P  I$ v% soptimizer = optim.Adam(params_to_update, lr = 1e-4). Q- F" ^5 ^6 `4 B
    scheduler = optim.lr_scheduler.StepLR(optimizer_ft, step_size = 7, gamma = 0.1)$ w! u' p& V( F' A5 G& q! _. X& m
    ) P( [5 k7 Q, G3 P2 i% _9 G" Q
    # 损失函数
    : I# m  |/ S; ]0 d% j+ ]criterion = nn.NLLLoss()
    $ z& q1 E1 T' ~# U- u" _/ D17 ?8 Q0 T) z( \+ B* N. ~
    2
    : Y" v( W9 h$ M: e9 f3( J3 _9 p: T. x; I3 i8 Q* {
    4
    8 z$ X  x$ \# w7 N2 `9 `. P  R5
    : ?5 Y2 V; w2 o& Z0 L! w7 M$ Z" Y6( T5 m8 `- g8 l  v0 j9 {
    7* V4 K$ G/ j$ M9 z" G
    8
    3 a# T" J* O0 O* p$ ^0 P9
    1 z- ~- j4 |* z' S10
    ) _5 ~( n) E% T% Y# 加载保存的参数
    & t( ~9 q6 V" [+ _# b$ C: P# 并在原有的模型基础上继续训练" ?$ t/ r2 w( O; s8 u
    # 下面保存的是刚刚训练效果较好的路径* S! w' I  S. Y( v, H: b8 |
    checkpoint = torch.load(filename)# l$ m! {( f0 E% C( ?2 Q& _( ]
    best_acc = checkpoint['best_acc']
    # L% ]- H4 p' W) Q. [: d4 T9 }& Smodel_ft.load_state_dict(checkpoint['state_dict'])
    # V$ X8 @5 D  `; v* D' d7 q5 J+ s. v& Soptimizer.load_state_dict(checkpoint['optimizer'])
    " N8 f3 t9 z& u. s- F- W1
    + c. t6 D6 Z/ f0 f+ O7 h+ E2: H; Z& V, h* r  Q
    36 N/ |# {0 p$ h
    4
    8 `# c6 k: S6 k6 A' s" s& ?50 g0 v& j3 |7 H
    6
    ( D4 e1 `3 Q' v7 J: J7' L, @5 y; r4 t9 P) X3 S# Y% K
    开始训练
    8 j8 u7 z; m' g. t9 i, `6 T1 z- N2 K/ z注:这里训练时长会变得别慢:我的显卡是1660ti,仅供各位参考
    ! [0 q. U# s0 w  `; d  h4 U
    - ]$ L& b2 g3 k2 O# ]" Nmodel_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"))
      @1 C( |/ q$ O1
    8 u% d% L' y+ k! yEpoch 0/1: N- w& g8 o* F
    ----------
    # y$ ^+ Z& I$ ATime elapsed 35m 22s6 i1 i- v* y- M
    train Loss: 1.7636 Acc: 0.7346
    0 f  I+ o2 m6 g! R; l1 FTime elapsed 38m 42s* _1 H% |; o6 F- x) [+ W
    valid Loss: 3.6377 Acc: 0.6455
    8 O/ G9 b8 N; k( }2 NOptimizer learning rate : 0.0010000
    % v' r6 u. S* D* p4 m6 D
    . a* V; Y+ w5 W) F$ U/ D: s0 }Epoch 1/1
    " H7 a8 j1 V  q- c: O2 k----------% z$ |4 M. e/ j
    Time elapsed 82m 59s  r- Q; T$ w' `8 `: ?& B
    train Loss: 1.7543 Acc: 0.7340
    * c8 @( m; D  ?) MTime elapsed 86m 11s
    + ]; X1 T3 \/ X' uvalid Loss: 3.8275 Acc: 0.6137! ?1 D( W% U- ^: P( R5 ^1 D+ q  k0 Q
    Optimizer learning rate : 0.0010000. a0 s1 x" F$ V

    ; j0 L. `1 S  n6 u6 z; dTraining complete in 86m 11s
    " \- o- N9 @% v+ `* s7 bBest val Acc: 0.645477
    , N) v/ Q0 ?5 j' P, ^
    " E; v% n1 V2 ~. B1# v7 l4 ^1 P) e0 e3 z6 |! u. i
    2
    2 L- W" h; [2 z$ Y0 }9 j3: m, g8 d5 s( I: t6 R
    48 W; g) z/ ^7 I1 A5 i2 V9 r( |
    59 |/ }1 |. w) |
    6
    : r9 u, y& y2 u7
    , R4 J; \( U* L( B8 O' A2 o! I8
    7 |( P1 V' @1 g' h. _92 S5 p# ~( i6 G+ ^2 \, `
    107 u2 X1 K' o0 I$ X
    110 s( Z9 y0 I4 b* `7 f+ v% J( b8 S
    12
    % w: c4 x7 h! K) f8 k. H4 k13
    ( z! s& T1 B+ @/ L: B14
    1 l" c' g; ~: ^+ @6 C- `15; e8 h8 y, c0 n& d$ j; T6 ^
    165 k1 I' Q+ [* ~0 U
    177 {$ K- F1 @. o$ o+ I+ e4 J
    18
    ) k/ K$ A( L. m  d0 T7 o8. 加载已经训练的模型9 q3 U' g# W2 ?0 f( s2 b6 I
    相当于做一次简单的前向传播(逻辑推理),不用更新参数$ p0 x+ R/ m9 R5 ?$ W) e, t

      Z; P; Z) O5 y0 Xmodel_ft, input_size = initialize_model(model_name, 102, feature_extract, use_pretrained=True)
    / Q$ ^  v% B' _1 E4 V  o2 @2 d: }& w7 L3 s" {
    # GPU 模式" t* j+ r$ k+ Y1 C; A" ?# ^
    model_ft = model_ft.to(device) # 扔到GPU中( g& E; C/ S0 L

    5 I- \! k, n/ }5 b. X( n: z# 保存文件的名字
    , t" W( ~# g7 ~5 p! g- Dfilename='checkpoint.pth'( ?1 v' \* ^& Y5 t
    - h# Q- {+ p* E& I$ A- Y7 u
    # 加载模型
    ' y: ~! C6 y; n% x# |6 O5 Icheckpoint = torch.load(filename)
    & E: X: K5 S( c; p! m& o* hbest_acc = checkpoint['best_acc']
    ; H  \; b; x& n( Z) Q; L* ^model_ft.load_state_dict(checkpoint['state_dict'])
    . Z) A, c$ j$ f% l3 b. g2 U17 |/ m; h) Z' [6 G
    2% X' Y8 z5 C. u( ~- [
    3
    9 ?- H3 y* z- h/ Z/ i4
    / v" U- e( D2 r. Q5 l5
    # ?% x5 {  s: D# e$ X8 u6
    % n2 t$ q; |" n3 \  G5 S$ k) z1 w7
    ; P: V7 V- U  Q. V  C86 b6 l; h/ W% z1 q
    9
    $ S, I* z6 K. P8 V7 d7 r6 o, M10! w' S5 |) {( v, b; P
    11
    7 E% J  O" N' P. f) A% P! v12
    3 [) ^" d$ R4 o2 K. J3 t, D<All keys matched successfully>
    # V8 D& ]) w/ W' d2 b1: [4 k1 n8 z  \
    def process_image(image_path):
    7 v2 X) S6 C. I9 Q& q4 F, c    # 读取测试集数据
    , u' _. \/ f6 y# ?# ^# m    img = Image.open(image_path)' S8 x+ `$ y6 \1 d
        # Resize, thumbnail方法只能进行比例缩小,所以进行判断
    $ h, y! O* c; h' S+ k* h+ _/ S3 D    # 与Resize不同, X4 ~7 M1 L8 q( v9 f
        # resize()方法中的size参数直接规定了修改后的大小,而thumbnail()方法按比例缩小
    & V% U5 i4 x8 \    # 而且对象调用方法会直接改变其大小,返回None
    ! B( P# Q! [/ [! S0 _( ], y6 Z    if img.size[0] > img.size[1]:
    ! U4 O) u# q# b  s/ M1 x        img.thumbnail((10000, 256))2 x% l4 C  m  l. c2 F6 d9 [
        else:0 E5 l4 K* ]/ a3 q
            img.thumbnail((256, 10000))* e1 k  e; t5 i3 S/ i) y, T' o2 s

    " }4 S3 Q" y) I+ |7 v' G) K# V    # crop操作, 将图像再次裁剪为 224 * 224+ f& S: Z) s: H: X/ N# O# n
        left_margin = (img.width - 224) / 2 # 取中间的部分) M0 W; h. Y- M% ~/ I* n' r
        bottom_margin = (img.height - 224) / 2
    - F* k7 |# ?6 Y/ ]  x6 T7 }    right_margin = left_margin + 224 # 加上图片的长度224,得到全部长度
    ( w* I/ f& W1 e4 S5 Q* ~2 T    top_margin = bottom_margin + 224
    4 T6 b6 D7 e6 {) N
    ; q# L" A! T1 Y- x    img = img.crop((left_margin, bottom_margin, right_margin, top_margin))4 G: {& E+ I) J" K0 r

    . T$ B& b9 L1 R4 {. s    # 相同预处理的方法0 A8 r" N; a6 X( {! ]
        # 归一化' m& J, w) b) q4 ]1 E/ h
        img = np.array(img) / 255
    $ ~, x0 W( h% _9 e6 Z/ ~2 a! r+ B, j    mean = np.array([0.485, 0.456, 0.406])5 ~( }& o; B5 I3 ?& Y+ \# Z
        std = np.array([0.229, 0.224, 0.225])
    1 n# L1 T4 K0 s4 w) h; h    img = (img - mean) / std- m0 Y  h; U# m2 i! f+ q
    4 a' K7 H& a/ w; n3 |$ u0 s6 g3 T0 W
        # 注意颜色通道和位置
    " R. `+ p; l/ G, }- O6 B    img = img.transpose((2, 0, 1))- L+ a3 W: N' J# A

    % ^- p; X: `- R7 |    return img
    - i& I: h$ k# ^% f3 F+ E3 f# Y/ h  Z* P0 i5 `# j+ K) F5 B+ ~7 x; O
    def imshow(image, ax = None, title = None):( ~' i$ [/ c+ y/ ?. x* W
        """展示数据"""' T9 \8 `2 c7 F, d. M
        if ax is None:8 ?9 E& R( ]2 }* [
            fig, ax = plt.subplots(); }; h) q- {6 L1 J& D& g0 m

    ' \; W% k9 Q8 Q, K; d/ f    # 颜色通道进行还原
    3 A# d) l! p  T4 E+ l- {    image = np.array(image).transpose((1, 2, 0))+ h( G& V" Q4 i1 w/ m) z1 v- O1 ?

    1 i( {. C! |0 B2 A8 L+ G7 c! Y    # 预处理还原
    ) M; E/ [1 y# Y    mean = np.array([0.485, 0.456, 0.406])
    # p& m4 ^4 F* A3 P9 H8 |    std = np.array([0.229, 0.224, 0.225])  V: K% {5 T4 @- a: u. N/ C
        image = std * image + mean  H8 c0 ~7 R/ \6 h
        image = np.clip(image, 0, 1)3 Y# w7 K: V% [) C% y9 \) I
    2 a8 |; h, ~6 A: Y, {/ r) P2 t) x
        ax.imshow(image)7 P6 d+ e. O" z( ^* r3 G
        ax.set_title(title)
    : y% T. m7 d$ X! Y) b
    7 o1 p+ w1 M( ]2 X    return ax& z6 u8 a8 j; v4 p4 E* A( a9 J
    0 T. c" w* v' E' U. o7 p+ r. a; H) J
    image_path = r'./flower_data/valid/3/image_06621.jpg'4 y2 G( R+ _$ }, ~' m" D1 Z
    img = process_image(image_path) # 我们可以通过多次使用该函数对图片完成处理" s2 A* K3 j8 A8 t7 Q
    imshow(img)5 q: _# s$ O" i/ D  H

    5 E; c# D4 |: p; P1' l! J% [/ b. o# R, }
    2
    ' V2 L/ V8 y: B% G5 U30 q; H' S+ W4 ?# V
    4
    " y& h& p4 @8 G4 S! A, A. K5& K( j/ r6 Q3 B! n& s# p
    6
    4 R9 _. U1 y" I4 O+ W, E; c* H& F76 M) d- O; l  ^1 u" G) I. B
    8
    & c: \: g  C. B3 U9
    ' q% C0 D3 m6 I10+ ~% N/ _/ l9 y. J3 y
    11; j0 l5 \; c) J; F6 I5 M5 d
    125 n, `2 Q% X( R$ m  M1 O% j, C2 z
    13
    $ [+ X' ~5 Y4 H! Y+ g! l, f; _9 e14
    3 j) X' v) K5 K15
    + ]! v% G( F$ K3 N16
    1 m7 |- d( _0 C3 O17
    1 [6 U: U/ a" z+ {" a18
    ; v0 o/ Y9 n: J2 @/ C3 c19: X- X% o1 Y5 C% v. J8 [! R
    20: u3 ^$ Y1 y; e! g! ]! B
    216 K% C/ o: B1 Q) b
    223 D, a8 o( N* H" v
    23" e1 v* m  u/ k* C; e* p$ i
    24
    ; ^3 b% d1 r1 Q& O8 ?4 r$ y( Y25# Y  F" O- O7 }
    26
    " f* a8 H& e* ^8 U; c. R! K27
    $ |6 I/ t9 C, M# I+ M286 J/ k: ^- l5 q4 L; E' ^
    29
    8 J8 f+ o. a% v: r1 T5 v4 a$ k' g1 ]30" X- Q0 P7 \! I2 N+ Y: Q* T) E
    31
    % d/ J! Y1 g; L8 y32
    5 {8 N7 R9 m, {9 W* Z33) F7 O8 S7 u- U; B" {2 T: U! ]
    349 i6 A& d4 o4 I( G9 Z( l
    35. `; ?( o5 X* u( _6 G" s( e
    36
    - y4 Y4 C+ g7 e6 W% H* B0 }. X! @37
    & }. C) j% ^5 V0 S! D7 z8 {, ^( ^38
    - O+ _% j- I9 D$ X. W. F3 J0 |39
    " d4 e7 h* T" h9 L/ Z40
    % w6 u, w( ^/ Z$ M6 D* O8 r- }+ o4 ~41
    ) j$ x; m# G6 U423 s( G8 L/ }+ p) C) i. h
    43) p7 {6 f  `+ g* O* S
    44* n, r4 w& f, B8 q1 n
    455 y4 Q5 z0 q+ T. r# O7 [$ b0 S
    46
    , I+ p# G7 ~! S! J470 m( M) E6 S7 V+ a2 I
    487 T: y) P# t" v- }' G
    49$ F% N: a. e' t$ U# w6 H
    50
    0 o( o. t- }. @" ~517 r5 F5 m" `4 M2 m% `1 H; C: D% W0 L
    52) E& q! r, _' M
    53- ^4 L% @0 `, {
    54+ Q3 O" Q# @* m; S8 v
    <AxesSubplot:>
    8 s9 H9 X8 @9 ^) G& P' [1 q# C1: `! b; j' e! g

    + Y/ ~" h6 l# v! b& g( q6 |% h上面是我们对测试集图片进行预处理之后的操作,我们使用shape来查看图片大小,预处理函数是否正确' c3 N: ]% g: r4 ]: _( N( M

    ' v  O1 T- _: @  T8 Oimg.shape/ _. q' J/ ^4 V9 C, l
    1! l2 @6 Z. i, c- e
    (3, 224, 224)
    3 w  O( G- U; f2 a& R1; x4 u7 d! s+ C; d$ e8 d3 z
    证明了通道提前了,而且大小没改变% C. C- s( \/ y* s2 t( Z
    " E' F; Z% f  f8 _9 y, s
    9. 推理
    7 N( p) l; v. o/ z( I: q" N: rimg.shape
    8 n  H0 T! v+ V% z" i1 |
    + i. C$ E" M5 s2 l# 得到一个batch的测试数据
    3 t  J, `! A$ }! _dataiter = iter(dataloaders['valid'])
    # o. J0 g9 \% d' W; T7 {5 Yimages, labels = dataiter.next()- Z% ?% o1 F( I* T# A9 ~& `

    / @& c: J1 B2 f6 @3 Q1 r' dmodel_ft.eval(). y4 k6 s6 w$ `- l

    # P+ B6 W8 ~& X- D' H+ m* }0 Lif train_on_gpu:6 H( Y8 l! _  P
        # 前向传播跑一次会得到output# u0 o8 N$ ^7 V8 R$ \9 \, G0 U
        output = model_ft(images.cuda())
    9 Y" Q+ R% o/ w, w0 B! n& N/ |% @else:7 p, ?5 X5 O2 E. F( u
        output = model_ft(images)- I5 D1 `7 \5 j! u) r
    9 K+ c/ @& s, i; a
    # batch 中有8 个数据,每个数据分为102个结果值, 每个结果是当前的一个概率值( j- g9 ?7 u" ?  r* ^& Z# c0 n. ~
    output.shape& w9 [/ n  V8 r8 \+ b
    & J' ]% K3 k$ S5 g
    1
    / C9 N; d" Z; C3 q5 b+ t* O$ ^2! T# B3 e) V+ A
    3; c8 r8 `: G$ @
    4
    . n5 k. o# K# {59 ~4 m7 a3 S, x$ w  s; o
    6
    ( B0 O8 l8 f- |0 h  U1 t! x7
    ; a* s$ d# ?, e  T2 b8 \8
    & D( W; q0 y' I$ ]9
    8 W- e) b* m2 ^" J10: d/ y- I9 u# P3 Y0 i
    11
    & M# b. A+ t* X4 u6 b12
    4 r/ a% R( y) l7 G6 W137 e+ e3 o8 f5 j" y  o% M
    14
    % j2 h/ S  S# [$ r15, e+ m! D# |( E! u! Y
    16, v! b$ r3 c- G+ ?
    torch.Size([8, 102]); _- R+ `) K  I( A
    1
    - \6 a4 v8 V. O  j. g9.1 计算得到最大概率
    4 Z- |- F  Y) \/ K/ B; q; v_, preds_tensor = torch.max(output, 1)2 x* q0 J: L2 a( {  `7 n2 U
    4 }- e( l: E/ L0 m+ I& F, c# [
    preds = np.squeeze(preds_tensor.numpy()) if not train_on_gpu else np.squeeze(preds_tensor.cpu().numpy())# 将秩为1的数组转为 1 维张量8 L7 \* J0 r2 L1 u$ _- K
    1
    ' I; j# V7 W; ^) m: _2/ t/ c; i1 k) J! r
    3! u7 R, R( S0 X" }9 K2 ?
    9.2 展示预测结果
    2 L; C2 A0 P! L. zfig = plt.figure(figsize = (20, 20))
    " _" F/ \# G* ?& O$ |" Qcolumns = 4+ c9 a" F; l  j5 t
    rows = 2
    / r3 _1 @2 Z- C% S# t  B, d5 M+ a$ n: N: r
    for idx in range(columns * rows):  X( Q2 p& w3 Y' D; h; n
        ax = fig.add_subplot(rows, columns, idx + 1, xticks =[], yticks =[])' [# ?: S# {# j9 ?2 @" O
        plt.imshow(im_convert(images[idx]))
    * E2 t' D, c, l. c3 h  R    ax.set_title("{} ({})".format(cat_to_name[str(preds[idx])], cat_to_name[str(labels[idx].item())]), ! K' I0 x5 B' Z! o. s
                    color = ("green" if cat_to_name[str(preds[idx])]==cat_to_name[str(labels[idx].item())] else "red"))
    1 {4 W% s9 p" y, G+ _plt.show()0 J  n  j. L! Z! U% |6 z% P: Y2 X
    # 绿色的表示预测是对的,红色表示预测错了
    7 I6 [, g; n5 {, c: s% K1 H& t1
      A7 J! m% r9 B) i( T/ D2; t1 H/ o! q7 r
    3# Q: T% w( p' {4 }5 z- h
    4
    & U' B% W& h/ k% g& ~* ?5' n& E' k3 w3 W* _4 U+ w
    6
    # F" A6 P' n9 T* a70 B# y& ^0 r. I4 f" E
    8- b$ Q/ V/ [2 e4 Z; L4 [" J  c
    9; S" p% W" s8 M0 ~0 O! {
    10
    $ {# }) r* t% h  _11$ z$ ^# Z. ?1 X9 H& M% X/ f0 ~
    1 q4 x4 V0 x  O( m2 o. B
    7 Z/ D; W$ Y: k/ \" K, `9 C
    5 v! o6 ^- J+ l0 Q
    ————————————————1 `4 O! e; D# O# r
    版权声明:本文为CSDN博主「FeverTwice」的原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接及本声明。- H# r7 G/ S# N4 n/ D) Y
    原文链接:https://blog.csdn.net/LeungSr/article/details/126747940# e9 r" W$ |; x' o

    2 G' v5 [# b5 m6 f( [: A+ `0 u2 n6 R- [
    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-26 06:47 , Processed in 0.552824 second(s), 51 queries .

    回顶部