GirlsINSAT/Projet_AIF_HF
1
1import argparse2from statistics import mean3 4import torch5import torchvision6import torchvision.transforms as transforms7import torch.nn as nn8import torch.nn.functional as F9import torch.optim as optim10from tqdm import tqdm11from torch.utils.tensorboard import SummaryWriter12from torchvision import datasets13from torch.utils.data import DataLoader, random_split14 15from model import MovieposterNet16 17 # setting device on GPU if available, else CPU18device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')19 20def train(net, optimizer, loader, writer,epochs=10):21 criterion = nn.CrossEntropyLoss()22 for epoch in range(epochs):23 running_loss = []24 t = tqdm(loader)25 for x, y in t:26 x, y = x.to(device), y.to(device)27 outputs = net(x)28 loss = criterion(outputs, y)29 running_loss.append(loss.item())30 optimizer.zero_grad()31 loss.backward()32 optimizer.step()33 t.set_description(f'training loss: {mean(running_loss)}')34 writer.add_scalar('training loss', mean(running_loss), epoch)35 36 37def test(model, dataloader):38 test_corrects = 039 total = 040 with torch.no_grad():41 for x, y in dataloader:42 x = x.to(device)43 y = y.to(device)44 y_hat = model(x).argmax(1)45 test_corrects += y_hat.eq(y).sum().item()46 total += y.size(0)47 return test_corrects / total48 49if __name__=='__main__':50 51 parser = argparse.ArgumentParser()52 53 parser.add_argument('--exp_name', type=str, default = 'Movieposter', help='experiment name')54 parser.add_argument('--epochs', type=int, default = int(10), help='nb of epochs')55 parser.add_argument('--batch_size', type=int, default = int(64), help='batch size')56 parser.add_argument('--lr', type=float, default = float(1e-3), help='learning rate')57 58 59 args = parser.parse_args()60 print(args.exp_name)61 exp_name = args.exp_name62 epochs = args.epochs63 batch_size = args.batch_size64 lr = args.lr65 66 writer = SummaryWriter(f'runs/Movieposter')67 68 # 1. Définition des transformations69 # Les posters sont en couleur (3 canaux) et de tailles variées, contrairement à MNIST.70 transform = transforms.Compose([71 transforms.Resize((224, 224)), # Redimensionnement standard pour les modèles de vision72 transforms.ToTensor(),73 transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)) # Normalisation sur 3 canaux (RGB)74 ])75 76 # 2. Chargement du dataset complet77 # Le chemin '../' permet de remonter d'un niveau par rapport au dossier 'projet_AIF'78 data_dir = '../sorted_movie_posters_paligema'79 full_dataset = datasets.ImageFolder(root=data_dir, transform=transform)80 81 # 3. Division en train/test (ex: 80% train, 20% test)82 train_size = int(0.8 * len(full_dataset))83 test_size = len(full_dataset) - train_size84 trainset, testset = random_split(full_dataset, [train_size, test_size])85 86 # 4. Création des DataLoaders87 trainloader = DataLoader(trainset, batch_size=batch_size, shuffle=True, num_workers=2)88 testloader = DataLoader(testset, batch_size=batch_size, shuffle=False, num_workers=2)89 90 # Accès aux classes (genres)91 classes = full_dataset.classes92 print(f"Classes détectées : {classes}")93 94 95 net =MovieposterNet().to(device)96 97 # setting net on device(GPU if available, else CPU)98 net = net.to(device)99 optimizer = optim.Adam(net.parameters(),weight_decay=1e-4, lr=lr)100 101 train(net, optimizer,trainloader, writer, epochs)102 test_acc = test(net,testloader)103 print(f'Test accuracy: {test_acc}') 104 105 # 1. Gestion du dossier de sauvegarde des poids106 import os107 if not os.path.exists('weights'):108 os.makedirs('weights')109 110 torch.save(net.state_dict(), 'weights/movieposter_net.pth')111 112 # 2. Récupération d'un échantillon de données pour TensorBoard113 # On utilise le loader pour obtenir des tenseurs déjà transformés114 dataiter = iter(trainloader)115 images, labels = next(dataiter) 116 117 # On limite à 64 images pour la visualisation et on envoie sur le device118 images = images[:64].to(device)119 labels = labels[:64].to(device)120 121 # 3. Enregistrement du graphe du modèle122 # Vérifiez que les dimensions d'entrée du modèle correspondent (ex: 3, 224, 224)123 writer.add_graph(net, images)124 125 # 4. Enregistrement d'une grille d'images126 img_grid = torchvision.utils.make_grid(images)127 writer.add_image('movieposter_samples', img_grid)128 129 # 5. Projecteur d'embeddings130 # get_features() doit être définie dans MovieposterNet pour retourner l'avant-dernière couche131 with torch.no_grad():132 try:133 embeddings = net.get_features(images)134 # Conversion des indices en noms de classes pour la lisibilité135 metadata = [classes[l] for l in labels]136 writer.add_embedding(embeddings,137 metadata=metadata,138 label_img=images, 139 global_step=epochs)140 except AttributeError:141 print("Erreur : La méthode get_features n'est pas définie dans MovieposterNet.")142 143 # 6. Fermeture du SummaryWriter144 writer.close()