Analyse du code de base de pytorch_forward_forward : de la classe Layer à la construction des échantillons positifs et négatifs

La rétropropagtaion traditionnelle nécesssite une propagation en avant suivie d'une propagation en arrière de l'erreur. En revanche, l'algorithme Forward-Forward (FF) effectue l'entraînement entièrement par propagation en avant. Dans l'algorithme FF, chaque couche a sa propre fonction de perte et son optimiseur, permettant un entraînement par couches, contrairement à la rétropropagation où une seule fonction de perte est partagée par toutes les couches.

Classe Layer : Bloc de construction central de l'algorithme FF

La classe Layer est le cœur de l'implémentation. Elle hérite de nn.Linear de PyTorch et ajoute la logique spécifique à l'entraînement FF.

class Couche(nn.Linear):
    def __init__(self, in_features, out_features, bias=True, device=None, dtype=None):
        super().__init__(in_features, out_features, bias, device, dtype)
        self.relu = torch.nn.ReLU()
        self.optim = Adam(self.parameters(), lr=0.03)
        self.seuil = 2.0
        self.epochs = 1000

Méthode forward : Normalisation directionnelle et activation

La méthode forward normalise les entrées en utilisant la norme L2, applique une transformation linéaire, puis une fonction d'activation ReLU. Cette normalisation directionnelle permet de se concentrer sur la direction plutôt que sur l'amplitude des entrées.

def forward(self, x):
    x_direction = x / (x.norm(2, 1, keepdim=True) + 1e-4)
    return self.relu(torch.mm(x_direction, self.weight.T) + self.bias.unsqueeze(0))

Méthode train : Entraînement par couches et calcul de la "bonté"

La méthode train implémente la logique d'entraînement FF. Chaque couche calcule indépendamment la "bonté" (moyenne des carrés des valeurs d'activation) pour les échantillons positifs et négatifs, et utilise une fonction de perte pour pousser la "bonté" des échantillons positifs au-dessus d'un seuil et celle des échentillons négatifs en dessous.

def train(self, x_pos, x_neg):
    for _ in tqdm(range(self.epochs)):
        g_pos = self.forward(x_pos).pow(2).mean(1)
        g_neg = self.forward(x_neg).pow(2).mean(1)
        loss = torch.log(1 + torch.exp(torch.cat([-g_pos + self.seuil, g_neg - self.seuil]))).mean()
        self.optim.zero_grad()
        loss.backward()
        self.optim.step()
    return self.forward(x_pos).detach(), self.forward(x_neg).detach()

Construction des échantillons positifs et négatifs : Préparation des données d'entraînement

L'algorithme FF nécessite des échantillons positifs (correspondance exacte entre l'entrée et l'étiquette) et négatifs (non-correspondance). La fonction overlay_y_on_x insère l'information d'étiquette dans les premiers pixels des données d'entrée.

def overlay_y_on_x(x, y):
    x_ = x.clone()
    x_[:, :10] *= 0.0
    x_[range(x.shape[0]), y] = x.max()
    return x_

x_pos = overlay_y_on_x(x, y)
rnd = torch.randperm(x.size(0))
x_neg = overlay_y_on_x(x, y[rnd])

Classe Net : Assemblage de plusieurs couches et prédiction

La classe Net assemble plusieurs couches Couche en un réseau complet. La méthode predict teste toutes les étiquettes possibles (0-9) pour chaque entrée et choisit l'étiquette avec la plus grande "bonté".

class Reseau(torch.nn.Module):
    def __init__(self, dimensions):
        super().__init__()
        self.couches = [Couche(dimensions[d], dimensions[d + 1]).cuda() for d in range(len(dimensions) - 1)]
    
    def predict(self, x):
        bonte_par_etiquette = []
        for label in range(10):
            h = overlay_y_on_x(x, label)
            bonte = []
            for couche in self.couches:
                h = couche(h)
                bonte.append(h.pow(2).mean(1))
            bonte_par_etiquette.append(sum(bonte).unsqueeze(1))
        bonte_par_etiquette = torch.cat(bonte_par_etiquette, 1)
        return bonte_par_etiquette.argmax(1)

Flux d'entraînement complet

Le flux d'entraînement complet est géré dans la fonction main. Il charge les données, construit le réseau, crée les échantillons positifs et négatifs, et entraîne le réseau.

if __name__ == "__main__":
    torch.manual_seed(1234)
    train_loader, test_loader = MNIST_loaders()
    
    reseau = Reseau([784, 500, 500])
    x, y = next(iter(train_loader))
    x, y = x.cuda(), y.cuda()
    x_pos = overlay_y_on_x(x, y)
    rnd = torch.randperm(x.size(0))
    x_neg = overlay_y_on_x(x, y[rnd])
    
    reseau.train(x_pos, x_neg)
    
    print('Erreur d\'entraînement:', 1.0 - reseau.predict(x).eq(y).float().mean().item())
    
    x_te, y_te = next(iter(test_loader))
    x_te, y_te = x_te.cuda(), y_te.cuda()
    
    print('Erreur de test:', 1.0 - reseau.predict(x_te).eq(y_te).float().mean().item())

Étiquettes: PyTorch forward-forward neural-networks adam-optimizer mnist

Publié le 7 août à 22h54