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.Sequential和nn.ModuleList。 - 可自定义输出:允许用户配置列的可见性、行设置、嵌套层的显示深度以及详细程度。
- 内存估算:计算模型的估计总大小,包括输入大小、参数大小以及前向/反向传播大小。
相关
- 项目
- 项目
- 项目
- 项目
- 项目