【深度学习】参数量和GFLOPs的计算

目录

GFLOPs和参数量的计算

未完待续,有空再写

https://zhuanlan.zhihu.com/p/376925457

模型显存占用

未完待续,有空再写

第三方计算库的使用

模型参数量GFLOPs 的计算库一般有两个,简洁一点的用thop,复杂详细一点的用torchinfo,方便起见,我就全部写在一个代码里了,大家学的时候就都一起学一下,韩信点兵多多益善。

注: 说明文档还是建议大家去看一下github的示例,早点看可以少走很多弯路,博主直接问AI导致一知半解,AI还是没有文档权威,血泪教训

python 复制代码
import torch
from torch import nn
import torch.nn.functional as F
from thop import profile
from torchinfo import summary

class ConvLinearNet(nn.Module):
    def __init__(self, num_classes=10):
        super().__init__()

        self.conv1 = nn.Conv2d(in_channels=3, out_channels=32, kernel_size=3, padding=1)
        self.bn1 = nn.BatchNorm2d(32)  # 批归一化层
        self.conv2 = nn.Conv2d(in_channels=32, out_channels=64, kernel_size=3, padding=1)
        self.bn2 = nn.BatchNorm2d(64)
        self.conv3 = nn.Conv2d(in_channels=64, out_channels=128, kernel_size=3, padding=1)
        self.bn3 = nn.BatchNorm2d(128)
        
        self.pool1 = nn.MaxPool2d(kernel_size=2, stride=2)
        self.pool2 = nn.MaxPool2d(kernel_size=2, stride=2)
        self.pool3 = nn.MaxPool2d(kernel_size=2, stride=2)

        self.fc1 = nn.Linear(128 * 4 * 4, 512)  # 假设输入图像为32x32,经过3次池化后为4x4
        self.bn_fc1 = nn.BatchNorm1d(512)
        self.fc2 = nn.Linear(512, 256)
        self.bn_fc2 = nn.BatchNorm1d(256)
        self.fc3 = nn.Linear(256, num_classes)
    
        self.dropout = nn.Dropout(0.3)
        
    def forward(self, x,task,prompt):
        # 卷积块1: Conv -> BN -> ReLU -> Pool
        x+=prompt
        x = self.conv1(x)
        x = self.bn1(x)
        x = F.relu(x)
        x = self.pool1(x)
        
        # 卷积块2: Conv -> BN -> ReLU -> Pool
        x = self.conv2(x)
        x = self.bn2(x)
        x = F.relu(x)
        x = self.pool2(x)
        
        # 卷积块3: Conv -> BN -> ReLU -> Pool
        x = self.conv3(x)
        x = self.bn3(x)
        x = F.relu(x)
        x = self.pool3(x)
        
        # 展平特征图
        x = x.view(x.size(0), -1)
        
        # 全连接块1: Linear -> BN -> ReLU -> Dropout
        x = self.fc1(x)
        x = self.bn_fc1(x)
        x = F.relu(x)
        x = self.dropout(x)
        
        # 全连接块2: Linear -> BN -> ReLU -> Dropout
        x = self.fc2(x)
        x = self.bn_fc2(x)
        x = F.relu(x)
        x = self.dropout(x)
        
        # 输出层
        x = self.fc3(x)
        return x
          
# 创建模型实例并测试
if __name__ == "__main__":
    # 创建模型
    model = ConvLinearNet(num_classes=10)
    
    # 计算参数量
    total_params = sum(p.numel() for p in model.parameters())
    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
    print(f"Total Params: {total_params:,}")
    print(f"Total Trainable Params: {trainable_params:,}")
    print("="*50)
    
    
    dummy_input = torch.randn(1, 3, 32, 32)  #模型输入
    prompt=torch.randn(1, 3, 32, 32)  #视觉提示
    task="deblurring"
    macs, params = profile(model, inputs=(dummy_input, task,prompt), verbose=False)
    print(f'Total Params: {params/1e6:.4f} M')
    print(f'Total MACs: {macs/1e9:.4f} G')
    print(f'Total FLOPs: {macs*2/1e9:.4f} GFLOPs')
    print("="*50)

    # 注意:forward 签名为 (x, task, prompt)
    # 所有 tensor 输入用 dict 按参数名传入(x、prompt),非 tensor 参数 task 用 kwargs 传入
    stats = summary(model, input_data={"x": dummy_input, "prompt": prompt}, task="deblurring", verbose=1, depth=4,
                        col_names=("input_size", "output_size", "num_params", "params_percent","mult_adds", "kernel_size"),
                        row_settings=("var_names", "depth"))
    print(f'Total Params: {stats.total_params / 1e6:.4f} M')
    print(f'Total MACs: {stats.total_mult_adds / 1e9:.4f} GMACs')
    print(f'Total FLOPs: {stats.total_mult_adds * 2 / 1e9:.4f} GFLOPs')

参考文献

相关推荐
韩师傅1 小时前
重生之我成为模型 番外 · 饲养员厨房下——大厨古法餐
深度学习·机器学习·计算机视觉
幻影123!12 小时前
从零训练一个会下五子棋的AI
python·深度学习·神经网络·强化学习·五子棋·alpha zero·mokugo
mingo_敏13 小时前
DeepAgents : 检索(Retrieval)
人工智能·深度学习·langchain
LaughingZhu15 小时前
Product Hunt 每日热榜 | 2026-08-01
人工智能·深度学习·神经网络·搜索引擎·百度
卡梅德生物科技小能手17 小时前
卡梅德生物科普 TNFSF4(肿瘤坏死因子超家族成员 4)
经验分享·深度学习·生活
逻辑君18 小时前
ANNA 认知引擎 · Humanoid 机器人训练白皮书
人工智能·深度学习·机器学习·机器人
hans汉斯18 小时前
计算机科学与应用|改进MeanShift算法在智能监控视频中的应用研究
图像处理·人工智能·功能测试·深度学习·算法·音视频
余俊晖21 小时前
多模态大模型细粒度视觉理解:Vision-OPD在线策略自蒸馏技术方案概述
人工智能·深度学习·算法·多模态·opd
大伟先生21 小时前
OpenClaw 数据采集实战入门
人工智能·深度学习