CoolFace
Apppublic

GirlsINSAT/Projet_AIF_HF

sourceHugging Faceupdated 7mo agoView on Hugging Face
1likes
train.py144 linesDownload Raw Back to root
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()