一、time库,显示程序运行时间

import time

# 在开头设置开始时间
start = time.perf_counter()  # start = time.clock() python3.8之前可以



# 在程序运行结束的位置添加结束时间
end = time.perf_counter()  # end = time.clock()  python3.8之前可以

# 再将其进行打印,即可显示出程序完成的运行耗时
print(f'运行耗时{(end-start):.6f}s')

cpu测试vgg16的时间

import torch
from torchvision.models import vgg16

import time

# 在开头设置开始时间
start = time.perf_counter()  # start = time.clock() python3.8之前可以

myNet = vgg16() # 实例化网络模型

img = torch.randn(32, 3, 256, 256)

pred = myNet(img)

print(pred.shape)

# 在程序运行结束的位置添加结束时间
end = time.perf_counter()  # end = time.clock()  python3.8之前可以

# 再将其进行打印,即可显示出程序完成的运行耗时
print(f'运行耗时{(end-start):.6f}s')

运行结果
在这里插入图片描述

二、可视化模型参数:

1. 使用自定义函数计算模型参数

import torch
import torchvision.models as models

def count_parameters(model):
    return sum(p.numel() for p in model.parameters() if p.requires_grad)

# 加载 ResNet-18 模型
model = models.resnet18()

# 是否使用GPU计算
device = torch.device("cuda" if torch.cuda.is_available() else 'cpu')

# --------------------------使用自定义函数count_parameters()计算模型参数--------------------------
print('ResNet-18 Model Parameters: ', count_parameters(model)/ 1e6, 'Million')

在这里插入图片描述

2. torchsummary.summary():可视化模型网络结构,计算模型参数

summary(model, input_size, batch_size=-1, device=“cuda”)

  • model:网络模型
  • input_size:输入的图像尺寸(c, h, w)(必须是元组,或者列表)(不能是ndarray或tensor)
  • batch_size:输入数据的批量大小(想要调试的输入形状为[n, c, h, w],必须改成input_size=(c, h, w), batch_size=n,否则提示输入尺寸错误)
  • device:默认将“输入数据”放在GPU上, 此时模型也得放在GPU上(手动); 或者模型不移动,指定device=‘cpu’。
# 方式1:使用CPU计算网络模型结构
import torch
import torchvision.models as models

from torchsummary import summary

# 加载 ResNet-18 模型
model = models.resnet18()

# 是否使用GPU计算
device = torch.device("cuda" if torch.cuda.is_available() else 'cpu')

# --------------------------使用torchsummary.summary()现实模型结构、计算模型参数--------------------------
summary(model, (3, 224, 224), 1, device='cpu')  # 输出模型网络结构   (默认,将数据放在GPU上,模型也要放在GPU上否则会出错)
# 方式2:使用GPU计算网络模型结构
import torch
import torchvision.models as models

from torchsummary import summary

# 加载 ResNet-18 模型
model = models.resnet18()

# --------------------------使用torchsummary.summary()现实模型结构、计算模型参数--------------------------
summary(model.to(device), (3, 224, 224), 1)  # 输出模型网络结构

运行结果

----------------------------------------------------------------
        Layer (type)               Output Shape         Param #
================================================================
            Conv2d-1          [1, 64, 112, 112]           9,408
       BatchNorm2d-2          [1, 64, 112, 112]             128
              ReLU-3          [1, 64, 112, 112]               0
         MaxPool2d-4            [1, 64, 56, 56]               0
            Conv2d-5            [1, 64, 56, 56]          36,864
       BatchNorm2d-6            [1, 64, 56, 56]             128
              ReLU-7            [1, 64, 56, 56]               0
            Conv2d-8            [1, 64, 56, 56]          36,864
       BatchNorm2d-9            [1, 64, 56, 56]             128
             ReLU-10            [1, 64, 56, 56]               0
       BasicBlock-11            [1, 64, 56, 56]               0
           Conv2d-12            [1, 64, 56, 56]          36,864
      BatchNorm2d-13            [1, 64, 56, 56]             128
             ReLU-14            [1, 64, 56, 56]               0
           Conv2d-15            [1, 64, 56, 56]          36,864
      BatchNorm2d-16            [1, 64, 56, 56]             128
             ReLU-17            [1, 64, 56, 56]               0
       BasicBlock-18            [1, 64, 56, 56]               0
           Conv2d-19           [1, 128, 28, 28]          73,728
      BatchNorm2d-20           [1, 128, 28, 28]             256
             ReLU-21           [1, 128, 28, 28]               0
           Conv2d-22           [1, 128, 28, 28]         147,456
      BatchNorm2d-23           [1, 128, 28, 28]             256
           Conv2d-24           [1, 128, 28, 28]           8,192
      BatchNorm2d-25           [1, 128, 28, 28]             256
             ReLU-26           [1, 128, 28, 28]               0
       BasicBlock-27           [1, 128, 28, 28]               0
           Conv2d-28           [1, 128, 28, 28]         147,456
      BatchNorm2d-29           [1, 128, 28, 28]             256
             ReLU-30           [1, 128, 28, 28]               0
           Conv2d-31           [1, 128, 28, 28]         147,456
      BatchNorm2d-32           [1, 128, 28, 28]             256
             ReLU-33           [1, 128, 28, 28]               0
       BasicBlock-34           [1, 128, 28, 28]               0
           Conv2d-35           [1, 256, 14, 14]         294,912
      BatchNorm2d-36           [1, 256, 14, 14]             512
             ReLU-37           [1, 256, 14, 14]               0
           Conv2d-38           [1, 256, 14, 14]         589,824
      BatchNorm2d-39           [1, 256, 14, 14]             512
           Conv2d-40           [1, 256, 14, 14]          32,768
      BatchNorm2d-41           [1, 256, 14, 14]             512
             ReLU-42           [1, 256, 14, 14]               0
       BasicBlock-43           [1, 256, 14, 14]               0
           Conv2d-44           [1, 256, 14, 14]         589,824
      BatchNorm2d-45           [1, 256, 14, 14]             512
             ReLU-46           [1, 256, 14, 14]               0
           Conv2d-47           [1, 256, 14, 14]         589,824
      BatchNorm2d-48           [1, 256, 14, 14]             512
             ReLU-49           [1, 256, 14, 14]               0
       BasicBlock-50           [1, 256, 14, 14]               0
           Conv2d-51             [1, 512, 7, 7]       1,179,648
      BatchNorm2d-52             [1, 512, 7, 7]           1,024
             ReLU-53             [1, 512, 7, 7]               0
           Conv2d-54             [1, 512, 7, 7]       2,359,296
      BatchNorm2d-55             [1, 512, 7, 7]           1,024
           Conv2d-56             [1, 512, 7, 7]         131,072
      BatchNorm2d-57             [1, 512, 7, 7]           1,024
             ReLU-58             [1, 512, 7, 7]               0
       BasicBlock-59             [1, 512, 7, 7]               0
           Conv2d-60             [1, 512, 7, 7]       2,359,296
      BatchNorm2d-61             [1, 512, 7, 7]           1,024
             ReLU-62             [1, 512, 7, 7]               0
           Conv2d-63             [1, 512, 7, 7]       2,359,296
      BatchNorm2d-64             [1, 512, 7, 7]           1,024
             ReLU-65             [1, 512, 7, 7]               0
       BasicBlock-66             [1, 512, 7, 7]               0
AdaptiveAvgPool2d-67             [1, 512, 1, 1]               0
           Linear-68                  [1, 1000]         513,000
================================================================
Total params: 11,689,512
Trainable params: 11,689,512
Non-trainable params: 0
----------------------------------------------------------------
Input size (MB): 0.57
Forward/backward pass size (MB): 62.79
Params size (MB): 44.59
Estimated Total Size (MB): 107.96
----------------------------------------------------------------

3. torchstat.stat()计算模型参数量、内存开销、FLOPs

( torchstat.stat():
输入:

  • 限制输入维度必须是三维)

输出:

  • Total params网络的整体参数量
  • Total memory模型进行推理时候所需的内存
  • Total Flops网络完成的浮点运算
  • Total MAdd:网络完成的乘加操作的数量。一次乘加=一次乘法+一次加法,可以大致认为(Flops ≈2*MAdd);
  • MemR+W:MemR+W = MemRead + MemWrite;
    (MemRead:网络运行时,从内存中读取的大小)
    (MemWrite:网络运行时,写入到内存中的大小)
import torch
import torchvision.models as models
from torchstat import stat

def count_parameters(model):
    return sum(p.numel() for p in model.parameters() if p.requires_grad)

# 加载 ResNet-18 模型
model = models.resnet18()

# 是否使用GPU计算
device = torch.device("cuda" if torch.cuda.is_available() else 'cpu')


#  --------------------------使用torchstat.stat()计算模型参数量、内存开销、FLOPS(限制输入维度必须是三维)--------------------------
stat(model, (3, 224, 224))
[MAdd]: AdaptiveAvgPool2d is not supported!
[Flops]: AdaptiveAvgPool2d is not supported!
[Memory]: AdaptiveAvgPool2d is not supported!
[MAdd]: AdaptiveAvgPool2d is not supported!
[Flops]: AdaptiveAvgPool2d is not supported!
[Memory]: AdaptiveAvgPool2d is not supported!
[MAdd]: AdaptiveAvgPool2d is not supported!
[Flops]: AdaptiveAvgPool2d is not supported!
[Memory]: AdaptiveAvgPool2d is not supported!
[MAdd]: AdaptiveAvgPool2d is not supported!
[Flops]: AdaptiveAvgPool2d is not supported!
[Memory]: AdaptiveAvgPool2d is not supported!
                 module name  input shape output shape      params memory(MB)             MAdd            Flops  MemRead(B)  MemWrite(B) duration[%]    MemR+W(B)
0                      conv1    3 224 224   64 112 112      9408.0       3.06    235,225,088.0    118,013,952.0    639744.0    3211264.0       0.00%    3851008.0
1                        bn1   64 112 112   64 112 112       128.0       3.06      3,211,264.0      1,605,632.0   3211776.0    3211264.0       0.00%    6423040.0
2                       relu   64 112 112   64 112 112         0.0       3.06        802,816.0        802,816.0   3211264.0    3211264.0       8.38%    6422528.0
3                    maxpool   64 112 112   64  56  56         0.0       0.77      1,605,632.0        802,816.0   3211264.0     802816.0       0.00%    4014080.0
4             layer1.0.conv1   64  56  56   64  56  56     36864.0       0.77    231,010,304.0    115,605,504.0    950272.0     802816.0       8.38%    1753088.0
5               layer1.0.bn1   64  56  56   64  56  56       128.0       0.77        802,816.0        401,408.0    803328.0     802816.0       0.00%    1606144.0
6              layer1.0.relu   64  56  56   64  56  56         0.0       0.77        200,704.0        200,704.0    802816.0     802816.0       0.00%    1605632.0
7             layer1.0.conv2   64  56  56   64  56  56     36864.0       0.77    231,010,304.0    115,605,504.0    950272.0     802816.0       0.00%    1753088.0
8               layer1.0.bn2   64  56  56   64  56  56       128.0       0.77        802,816.0        401,408.0    803328.0     802816.0       0.00%    1606144.0
9             layer1.1.conv1   64  56  56   64  56  56     36864.0       0.77    231,010,304.0    115,605,504.0    950272.0     802816.0       0.00%    1753088.0
10              layer1.1.bn1   64  56  56   64  56  56       128.0       0.77        802,816.0        401,408.0    803328.0     802816.0       0.00%    1606144.0
11             layer1.1.relu   64  56  56   64  56  56         0.0       0.77        200,704.0        200,704.0    802816.0     802816.0       0.00%    1605632.0
12            layer1.1.conv2   64  56  56   64  56  56     36864.0       0.77    231,010,304.0    115,605,504.0    950272.0     802816.0       4.26%    1753088.0
13              layer1.1.bn2   64  56  56   64  56  56       128.0       0.77        802,816.0        401,408.0    803328.0     802816.0       0.59%    1606144.0
14            layer2.0.conv1   64  56  56  128  28  28     73728.0       0.38    115,505,152.0     57,802,752.0   1097728.0     401408.0       2.97%    1499136.0
15              layer2.0.bn1  128  28  28  128  28  28       256.0       0.38        401,408.0        200,704.0    402432.0     401408.0       0.00%     803840.0
16             layer2.0.relu  128  28  28  128  28  28         0.0       0.38        100,352.0        100,352.0    401408.0     401408.0       0.00%     802816.0
17            layer2.0.conv2  128  28  28  128  28  28    147456.0       0.38    231,110,656.0    115,605,504.0    991232.0     401408.0       8.44%    1392640.0
18              layer2.0.bn2  128  28  28  128  28  28       256.0       0.38        401,408.0        200,704.0    402432.0     401408.0       0.00%     803840.0
19     layer2.0.downsample.0   64  56  56  128  28  28      8192.0       0.38     12,744,704.0      6,422,528.0    835584.0     401408.0       0.00%    1236992.0
20     layer2.0.downsample.1  128  28  28  128  28  28       256.0       0.38        401,408.0        200,704.0    402432.0     401408.0       0.00%     803840.0
21            layer2.1.conv1  128  28  28  128  28  28    147456.0       0.38    231,110,656.0    115,605,504.0    991232.0     401408.0       8.37%    1392640.0
22              layer2.1.bn1  128  28  28  128  28  28       256.0       0.38        401,408.0        200,704.0    402432.0     401408.0       0.00%     803840.0
23             layer2.1.relu  128  28  28  128  28  28         0.0       0.38        100,352.0        100,352.0    401408.0     401408.0       0.00%     802816.0
24            layer2.1.conv2  128  28  28  128  28  28    147456.0       0.38    231,110,656.0    115,605,504.0    991232.0     401408.0       0.00%    1392640.0
25              layer2.1.bn2  128  28  28  128  28  28       256.0       0.38        401,408.0        200,704.0    402432.0     401408.0       0.00%     803840.0
26            layer3.0.conv1  128  28  28  256  14  14    294912.0       0.19    115,555,328.0     57,802,752.0   1581056.0     200704.0       8.37%    1781760.0
27              layer3.0.bn1  256  14  14  256  14  14       512.0       0.19        200,704.0        100,352.0    202752.0     200704.0       0.00%     403456.0
28             layer3.0.relu  256  14  14  256  14  14         0.0       0.19         50,176.0         50,176.0    200704.0     200704.0       0.00%     401408.0
29            layer3.0.conv2  256  14  14  256  14  14    589824.0       0.19    231,160,832.0    115,605,504.0   2560000.0     200704.0       0.00%    2760704.0
30              layer3.0.bn2  256  14  14  256  14  14       512.0       0.19        200,704.0        100,352.0    202752.0     200704.0       8.37%     403456.0
31     layer3.0.downsample.0  128  28  28  256  14  14     32768.0       0.19     12,794,880.0      6,422,528.0    532480.0     200704.0       0.00%     733184.0
32     layer3.0.downsample.1  256  14  14  256  14  14       512.0       0.19        200,704.0        100,352.0    202752.0     200704.0       0.00%     403456.0
33            layer3.1.conv1  256  14  14  256  14  14    589824.0       0.19    231,160,832.0    115,605,504.0   2560000.0     200704.0       0.00%    2760704.0
34              layer3.1.bn1  256  14  14  256  14  14       512.0       0.19        200,704.0        100,352.0    202752.0     200704.0       8.37%     403456.0
35             layer3.1.relu  256  14  14  256  14  14         0.0       0.19         50,176.0         50,176.0    200704.0     200704.0       0.00%     401408.0
36            layer3.1.conv2  256  14  14  256  14  14    589824.0       0.19    231,160,832.0    115,605,504.0   2560000.0     200704.0       0.00%    2760704.0
37              layer3.1.bn2  256  14  14  256  14  14       512.0       0.19        200,704.0        100,352.0    202752.0     200704.0       0.00%     403456.0
38            layer4.0.conv1  256  14  14  512   7   7   1179648.0       0.10    115,580,416.0     57,802,752.0   4919296.0     100352.0       8.38%    5019648.0
39              layer4.0.bn1  512   7   7  512   7   7      1024.0       0.10        100,352.0         50,176.0    104448.0     100352.0       0.00%     204800.0
40             layer4.0.relu  512   7   7  512   7   7         0.0       0.10         25,088.0         25,088.0    100352.0     100352.0       0.00%     200704.0
41            layer4.0.conv2  512   7   7  512   7   7   2359296.0       0.10    231,185,920.0    115,605,504.0   9537536.0     100352.0       8.37%    9637888.0
42              layer4.0.bn2  512   7   7  512   7   7      1024.0       0.10        100,352.0         50,176.0    104448.0     100352.0       0.00%     204800.0
43     layer4.0.downsample.0  256  14  14  512   7   7    131072.0       0.10     12,819,968.0      6,422,528.0    724992.0     100352.0       0.00%     825344.0
44     layer4.0.downsample.1  512   7   7  512   7   7      1024.0       0.10        100,352.0         50,176.0    104448.0     100352.0       0.00%     204800.0
45            layer4.1.conv1  512   7   7  512   7   7   2359296.0       0.10    231,185,920.0    115,605,504.0   9537536.0     100352.0       8.38%    9637888.0
46              layer4.1.bn1  512   7   7  512   7   7      1024.0       0.10        100,352.0         50,176.0    104448.0     100352.0       0.00%     204800.0
47             layer4.1.relu  512   7   7  512   7   7         0.0       0.10         25,088.0         25,088.0    100352.0     100352.0       8.37%     200704.0
48            layer4.1.conv2  512   7   7  512   7   7   2359296.0       0.10    231,185,920.0    115,605,504.0   9537536.0     100352.0       0.00%    9637888.0
49              layer4.1.bn2  512   7   7  512   7   7      1024.0       0.10        100,352.0         50,176.0    104448.0     100352.0       0.00%     204800.0
50                   avgpool  512   7   7  512   1   1         0.0       0.00              0.0              0.0         0.0          0.0       0.00%          0.0
51                        fc          512         1000    513000.0       0.00      1,023,000.0        512,000.0   2054048.0       4000.0       0.00%    2058048.0
total                                                   11689512.0      25.65  3,638,757,912.0  1,821,399,040.0   2054048.0       4000.0     100.00%  101756992.0
=================================================================================================================================================================
Total params: 11,689,512
-----------------------------------------------------------------------------------------------------------------------------------------------------------------
Total memory: 25.65MB
Total MAdd: 3.64GMAdd
Total Flops: 1.82GFlops
Total MemR+W: 97.04MB

4. thop.profile()计算模型的FLOPS、模型参数

import torch
import torchvision.models as models
from thop import profile

# 加载 ResNet-18 模型
model = models.resnet18()

# 准备示例输入
input_data = torch.randn(1, 3, 224, 224)

#  --------------------------使用thop.profile()计算模型的FLOPS、模型参数--------------------------
flops, params = profile(model, inputs=(input_data,))
print(f"ResNet-18 FLOPs: {flops / 1e9} GFLOPs")  # 将结果转换为GFLOPs
print(f"ResNet-18 Parameters: {params / 1e6} Million")
[INFO] Register count_convNd() for <class 'torch.nn.modules.conv.Conv2d'>.
[INFO] Register count_normalization() for <class 'torch.nn.modules.batchnorm.BatchNorm2d'>.
[INFO] Register zero_ops() for <class 'torch.nn.modules.activation.ReLU'>.
[INFO] Register zero_ops() for <class 'torch.nn.modules.pooling.MaxPool2d'>.
[INFO] Register zero_ops() for <class 'torch.nn.modules.container.Sequential'>.
[INFO] Register count_adap_avgpool() for <class 'torch.nn.modules.pooling.AdaptiveAvgPool2d'>.
[INFO] Register count_linear() for <class 'torch.nn.modules.linear.Linear'>.
[MAdd]: AdaptiveAvgPool2d is not supported!
[Flops]: AdaptiveAvgPool2d is not supported!
[Memory]: AdaptiveAvgPool2d is not supported!
[MAdd]: AdaptiveAvgPool2d is not supported!
[Flops]: AdaptiveAvgPool2d is not supported!
[Memory]: AdaptiveAvgPool2d is not supported!
[MAdd]: AdaptiveAvgPool2d is not supported!
[Flops]: AdaptiveAvgPool2d is not supported!
[Memory]: AdaptiveAvgPool2d is not supported!
[MAdd]: AdaptiveAvgPool2d is not supported!
[Flops]: AdaptiveAvgPool2d is not supported!
[Memory]: AdaptiveAvgPool2d is not supported!
ResNet-18 FLOPs: 1.824033792 GFLOPs
ResNet-18 Parameters: 11.689512 Million

三、collections.OrderedDict()函数:创建按照有序插入顺序存储的有序字典

方式一:使用collections.OrderedDict()函数

from collections import OrderedDict

test_results = OrderedDict()  # 按照有序插入顺序存储 的有序字典
test_results['psnr'] = []
test_results['ssim'] = []
test_results['niqe'] = []
psnr, ssim, niqe = 0, 0, 0

test_results['psnr'].append(psnr)
test_results['ssim'] = ssim
print(test_results)

在这里插入图片描述

for key, value in test_results.items():
    print(key, value)

在这里插入图片描述

方式二:直接创建集合

test_results = {'psnr': 0, 'ssim': 0, 'niqe': 0}

test_results['psnr'] += 1
print(test_results)

在这里插入图片描述

四、torchviz库,可视化模型计算图

安装相关库:graphviz 和 torchviz与相应的软件graphviz.exe

import torch
import torchvision.models as models

# 是否使用GPU计算
device = torch.device("cuda" if torch.cuda.is_available() else 'cpu')

# 加载 ResNet-18 模型
model = models.resnet18()

# 准备示例输入
input_data = torch.randn(1, 3, 224, 224)

# 前向传播
pred = model(input_data)

# ---------------------可视化模型计算图---------------------
from torchviz import make_dot

# 构造图对象(3种方式)
# g = make_dot(y)  # 只显示了输出节点的计算图
# g = make_dot(y, params=dict(model.named_parameters()))  # 生成一个包含模型参数的完整计算图
pic_1 = make_dot(pred, params=dict(list(model.named_parameters()) + [('x', input_data)]))  # 生成的计算图将显示模型参数以及额外的输入参数。

# 指定文件类型与文件路径
# pic_1.format = "png"  # 定义pic为模型可视化后的输出,这里输出为png格式
# 指定文件生成的文件夹pic里面
# pic_1.directory = "pic"

# 生成文件(2种方式)
# pic_1.view()  # 直接在当前路径下保存 pdf 并打开
pic_1.render(filename='pic/resnet18', view=False)  # 在pic目录下,保存为 resnet18.pdf,参数view表示是否打开pdf

运行结果:

在这里插入图片描述

五、argparse库,用于命令项选项与参数解析

import argparse

parser = argparse.ArgumentParser(description='Train Super Resolution Models')

parser.add_argument('--crop_size', default=88, type=int, help='training images crop size')
parser.add_argument('--upscale_factor', default=4, type=int, choices=[2, 4, 8],
                    help='super resolution upscale factor')
parser.add_argument('--num_epochs', default=2, type=int, help='train epoch number')

opt = parser.parse_args()
Logo

魔乐社区(Modelers.cn) 是一个中立、公益的人工智能社区,提供人工智能工具、模型、数据的托管、展示与应用协同服务,为人工智能开发及爱好者搭建开放的学习交流平台。社区通过理事会方式运作,由全产业链共同建设、共同运营、共同享有,推动国产AI生态繁荣发展。

更多推荐