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

関連

  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト