Pytorch代码:打印模型每层的参数数量和总参数量

这个代码片段定义了一个函数 print_model_parameters,它的作用是打印每层的参数数量以及模型的总参数量。下面是对这个函数的详细解释,重点解释 named_parametersrequires_gradnumel 参数的含义:

python 复制代码
# 打印每层的参数数量和总参数量
def print_model_parameters(model):
    total_params = 0
    for name, param in model.named_parameters():
        if param.requires_grad:
            print(f"{name}: {param.numel()} parameters")
            total_params += param.numel()
            print(f"For now parameters: {total_params}")
    print(f"Total parameters: {total_params}")

具体步骤和解释

  1. 定义和初始化

    python 复制代码
    def print_model_parameters(model):
        total_params = 0

    这个函数接收一个模型对象 model,并初始化一个变量 total_params 用于累积总参数量。

  2. 遍历模型参数

    python 复制代码
    for name, param in model.named_parameters():

    这里使用了 model.named_parameters() 方法,该方法返回一个生成器,生成模型中所有参数的名称和参数张量。它返回的是 (name, parameter) 形式的元组。

    • named_parameters:这是一个PyTorch模型的方法,它返回模型中所有参数的名称和参数本身。参数的名称是字符串类型,而参数是一个 torch.Tensor 对象。
  3. 判断参数是否需要梯度更新

    python 复制代码
    if param.requires_grad:

    每个参数张量都有一个 requires_grad 属性,这个属性是一个布尔值。如果 requires_gradTrue,表示这个参数在训练过程中需要计算梯度并进行更新。

    • requires_grad:这是一个布尔值属性,表示该参数是否需要在训练过程中计算梯度。如果是 True,则该参数会在反向传播时计算并存储梯度。
  4. 打印参数数量并累加

    python 复制代码
    print(f"{name}: {param.numel()} parameters")
    total_params += param.numel()
    print(f"For now parameters: {total_params}")

    对于需要梯度的参数,打印其名称和参数数量,并将该参数的数量累加到 total_params 中。

    • numel:这是一个方法,返回张量中所有元素的数量。例如,一个形状为 (3, 4) 的张量调用 numel() 方法会返回 12,因为这个张量有12个元素。
  5. 打印总参数量

    python 复制代码
    print(f"Total parameters: {total_params}")

    最后,打印模型的总参数数量。

总结

这个函数通过 model.named_parameters() 遍历模型的所有参数,检查每个参数的 requires_grad 属性,只有在 requires_gradTrue 时才计算并打印参数数量,同时累加总参数量。 numel() 方法用于获取每个参数张量的元素数量,从而帮助统计参数数量。最后打印总参数量,提供了对模型规模的一个直观了解。

相关推荐
栀椩8 小时前
CODrone 无人机航拍车辆检测
pytorch·python·yolo
盼小辉丶8 小时前
PyTorch计算机视觉(9)——变分自编码器(VAE)详解与实现
人工智能·pytorch·深度学习·计算机视觉
clorinda8 小时前
PyTorch 手写数字识别学习笔记
pytorch·笔记·学习
dadanhuang1 天前
PyTorch深度学习与实践【04】【迭代周期、autograd、构建计算图、*params参数解包、.grad属性】
人工智能·pytorch·深度学习
basketball6161 天前
AI Infra 配置 Conda + CUDA + LibTorch + PyTorch 开发环境:解决版本漂移、编译报错的完整指南
人工智能·pytorch·conda·libtorch
LlmCraft|大模型工程实践1 天前
【PyTorch NLP实战】从零实现 LSTM 酒店评论情感分析(附完整代码逐行讲解)
pytorch·自然语言处理·lstm
大雷神2 天前
HarmonyOS ArkGraphics 2D 自定义字体实操:注册字体并验证中文回退
pytorch·华为·harmonyos
盼小辉丶3 天前
PyTorch强化学习实战——融合人类示范数据的高效强化学习
人工智能·pytorch·python·深度学习·强化学习
李妍.3 天前
PyTorch GPU 版安装记录(RTX 5060 + Python 3.14)
人工智能·pytorch·python
zx_741484813 天前
【机器学习入门】PyTorch 神经网络、卷积神经网络 CNN 实现矿物分类
pytorch·神经网络·机器学习