Construindo Modelos com nn.Module
nn.Module é a classe base para todos os modelos PyTorch. Subclasseie-a para definir camadas, implementar forward() e registrar parâmetros que os otimizadores atualizam durante o treinamento.
Busque em todas as páginas da documentação
nn.Module é a classe base para todos os modelos PyTorch. Subclasseie-a para definir camadas, implementar forward() e registrar parâmetros que os otimizadores atualizam durante o treinamento.
Cartão de receita de referência rápida - pronto para copiar e colar.
import torch.nn as nn
class MLP(nn.Module):
def __init__(self, in_dim: int, hidden: int, out_dim: int):
super().__init__()
self.net = nn.Sequential(
nn.Linear(in_dim, hidden),
nn.ReLU(),
nn.Linear(hidden, out_dim),
)
def forward(self, x):
return self.net(x)Quando usar isso:
state_dict()."""building_models.py - CNN customizada para classificação de imagens."""
from __future__ import annotations
import torch
import torch.nn as nn
import torch.nn.functional as F
class SmallCNN(nn.Module):
def __init__(self, num_classes: int = 10):
super().__init__()
self.conv1 = nn.Conv2d(1, 32, kernel_size=3, padding=1)
self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1)
self.pool = nn.MaxPool2d(2)
self.fc1 = nn.Linear(64 * 7 * 7, 128)
self.fc2 = nn.Linear(128, num_classes)
self.dropout = nn.Dropout(0.25)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.pool(F.relu(self.conv1(x))) # 28->14
x = self.pool(F.relu(self.conv2(x))) # 14->7
x = x.flatten(1)
x = self.dropout(F.relu(self.fc1(x)))
return self.fc2(x)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = SmallCNN(num_classes=10).to(device)
x = torch.randn(8, 1, 28, 28, device=device)
logits = model(x)
print("logits shape:", logits.shape) # (8, 10)
total_params = sum(p.numel() for p in model.parameters())
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
print(f"params: {total_params:,} trainable: {trainable:,}")O que isso demonstra:
__init__ tornam-se parte de model.parameters().forward define a computação; chame via model(x), não model.forward(x)..to(device) move todos os parâmetros e buffers para a GPU.flatten(1) preserva a dimensão do lote enquanto colapsa as dimensões espaciais.__init__ registra submódulos (nn.Linear, nn.Conv2d) e parâmetros (nn.Parameter).forward executa o grafo de computação; hooks podem interceptar entradas/saídas.nn.Module rastreia o modo training vs eval para comportamento de dropout e batch norm.state_dict() retorna um dicionário de tensores de parâmetros para salvar/carregar.named_modules() e named_children() percorrem a árvore de módulos.| Padrão | Use Quando | Exemplo |
|---|---|---|
nn.Sequential | Cadeia linear de camadas | Classificadores MLP |
Subclasse nn.Module | Lógica forward customizada | Conexões de atalho ResNet |
nn.ModuleList | Número variável de camadas | Redes de profundidade dinâmica |
nn.ModuleDict | Submódulos nomeados | Arquiteturas multi-head |
# Congela o backbone, treina apenas a cabeça
for param in model.conv1.parameters():
param.requires_grad = False
# Taxas de aprendizado diferentes por grupo
optimizer = torch.optim.Adam([
{"params": model.conv1.parameters(), "lr": 1e-5},
{"params": model.fc2.parameters(), "lr": 1e-3},
])forward() diretamente - pula hooks e wrappers nn.Module. Correção: sempre chame model(x).super().__init__() - submódulos não registrados. Correção: chame super().__init__() primeiro em __init__.model.eval() para inferência - dropout e batch norm se comportam incorretamente. Correção: model.eval() antes da inferência; model.train() para treinamento.forward - novos parâmetros a cada passagem, nunca treinados. Correção: defina todas as camadas em __init__.model.to(device) e x.to(device).model.eval() ou nn.GroupNorm para lotes pequenos.| Alternativa | Use Quando | Não Use Quando |
|---|---|---|
Subclasse nn.Module | Arquiteturas customizadas | MLP trivial de 3 camadas (use Sequential) |
nn.Sequential | Pilhas feedforward simples | Necessita de ramificações ou conexões de atalho |
PyTorch Lightning LightningModule | Loops de treinamento estruturados | Aprendendo os fundamentos do PyTorch |
torch.nn.functional | Operações sem estado em forward | Necessita de parâmetros treináveis |
nn.Module contêm parâmetros treináveis (Linear, Conv2d).nn.functional fornece funções sem estado (relu, conv2d com pesos explícitos).__init__; use funcional para operações únicas.sum(p.numel() for p in model.parameters() if p.requires_grad)numel() retorna o número total de elementos em um tensor.state_dict mas não são treinados.self.register_buffer("nome", tensor).forward por conveniência.nn.Module pode conter outros módulos como atributos.model.children() retorna filhos diretos; model.modules() é recursivo.model.eval() para validação e inferência.def init_weights(m):
if isinstance(m, nn.Linear):
nn.init.xavier_uniform_(m.weight)
model.apply(init_weights)register_forward_hook inspeciona ativações intermediárias.print(model)
# ou
from torchinfo import summary
summary(model, input_size=(8, 1, 28, 28))torchinfo mostra contagens de parâmetros por camada.def forward(self, x: torch.Tensor) -> torch.Tensor.def forward(self, x):
residual = x
out = self.block(x)
return F.relu(out + residual).to(device) - movimentos parciais causam erros de dispositivo.Versões do Stack: Esta página foi escrita para Python 3.14.0 (estável 3.14, manutenção 3.13), FastAPI 0.115+, Django 5.2, Flask 3.1, Pydantic 2, PyTorch 2.6+, pandas 2.2+, Polars 1.x, ruff 0.9+, e uv 0.6+.
Revisado por Chris St. John·Última atualização: 19 de jul. de 2026