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を処理可能。 - カスタマイズ可能な出力: 列の表示・非表示、行の設定、ネストされた層の表示深度、出力の詳細レベルをユーザーが設定可能。
- メモリ推定: 入力サイズ、パラメータサイズ、前方/後方伝播サイズを含むモデルの推定合計サイズを計算。
関連
- プロジェクト
- プロジェクト
- プロジェクト
- プロジェクト
- プロジェクト