Architecture diagramimport torch
import torch.nn as nn
from pytorch_graph import generate_architecture_diagram
model = nn.Sequential(nn.Linear(784, 128), nn.ReLU(), nn.Linear(128, 10))
generate_architecture_diagram(
model=model,
input_shape=(1, 784),
output_path="model_architecture.png",
title="MNIST MLP",
style="flowchart",
)
Full graph trackingfrom pytorch_graph import ComputationalGraphTracker
import torch
tracker = ComputationalGraphTracker(model=model, track_memory=True, track_timing=True)
tracker.start_tracking()
output = model(torch.randn(1, 784))
tracker.stop_tracking()
tracker.save_graph_png("complete_graph.png", width=1800, height=1200, dpi=300)
Model analysisfrom pytorch_graph import analyze_model
analysis = analyze_model(model=model, input_shape=(1, 784), detailed=True)
print(analysis["summary"]["total_params"])
print(analysis["summary"]["trainable_params"])