Comment Pytorch construit-il le graphe de calcul
Voici un exemple de code pytorch du site Web:
class Net(nn.Module):
def __init__(self):
super(Net, self).__init__()
# 1 input image channel, 6 output channels, 3x3 square convolution
# kernel
self.conv1 = nn.Conv2d(1, 6, 3)
self.conv2 = nn.Conv2d(6, 16, 3)
# an affine operation: y = Wx + b
self.fc1 = nn.Linear(16 * 6 * 6, 120) # 6*6 from image dimension
self.fc2 = nn.Linear(120, 84)
self.fc3 = nn.Linear(84, 10)
def forward(self, x):
# Max pooling over a (2, 2) window
x = F.max_pool2d(F.relu(self.conv1(x)), (2, 2))
# If the size is a square you can only specify a single number
x = F.max_pool2d(F.relu(self.conv2(x)), 2)
x = x.view(-1, self.num_flat_features(x))
x = F.relu(self.fc1(x))
x = F.relu(self.fc2(x))
x = self.fc3(x)
return x
Dans la fonction avant, nous appliquons simplement une série de transformations à x, mais nous ne définissons jamais explicitement quels objets font partie de cette transformation. Pourtant, lors du calcul du gradient et de la mise à jour des poids, Pytorch sait «par magie» quels poids mettre à jour et comment le gradient doit être calculé.
Comment fonctionne ce processus? Y a-t-il une analyse de code en cours ou quelque chose d'autre qui me manque?
Réponses
Oui, il y a une analyse implicite sur la passe avant. Examinez le tenseur du résultat, il y a un truc comme grad_fn= <CatBackward>, c'est un lien, vous permettant de dérouler tout le graphe de calcul. Et il est construit pendant le processus de calcul direct réel, quelle que soit la façon dont vous avez défini votre module réseau, orienté objet avec une manière «nn» ou «fonctionnelle».
Vous pouvez exploiter ce graphique pour l'analyse du réseau, comme torchvizici:https://github.com/szagoruyko/pytorchviz/blob/master/torchviz/dot.py