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를 처리할 수 있습니다. - 사용자 정의 가능한 출력: 열의 표시 여부, 행 설정, 중첩 레이어의 표시 깊이, 출력의 자세함 수준을 사용자가 설정할 수 있습니다.
- 메모리 추정: 입력 크기, 파라미터 크기, 순전파/역전파 크기를 포함한 모델의 추정 총 크기를 계산합니다.
관련
- 프로젝트
- 프로젝트
- 프로젝트
- 프로젝트
- 프로젝트