TylerYep/torchinfo

View model summaries in PyTorch!

解决的问题

PyTorch 的默认 print(model) 仅提供网络架构的有限信息。torchinfo 提供了类似 Keras/TensorFlow 的 model.summary() 的详细、格式化的 PyTorch 模型摘要,这对于调试和验证神经网络结构至关重要。

工作原理

该工具接收一个 PyTorch nn.Module 以及一个样本输入尺寸或实际输入数据。它执行一次前向传播,跟踪张量在各层之间的形状变化,计算参数数量、输出形状以及乘加操作(Mult-Adds)的数量。

适用人群

使用 PyTorch 的深度学习从业者和研究人员,需要可视化模型架构、验证张量形状并估算模型内存占用量。

主要特性

  • 全面的指标:显示层名称、输入/输出形状、参数数量和 Mult-Adds。
  • 灵活的输入:支持提供输入形状(作为元组/列表)或实际输入张量,包括具有不同数据类型的多个输入。
  • 高级架构支持:支持 RNN、LSTM、递归层、nn.Sequentialnn.ModuleList
  • 可自定义输出:允许用户配置列的可见性、行设置、嵌套层的显示深度以及详细程度。
  • 内存估算:计算模型的估计总大小,包括输入大小、参数大小以及前向/反向传播大小。

相关

  • 项目
  • 项目
  • 项目
  • 项目
  • 项目