import argparse import copy import os import torch import torch.nn as nn import torch.optim as optim from torchvision import datasets, models, transforms from torch.utils.data import DataLoader def build_model(arch): if arch == "mobilenet": model = models.mobilenet_v2(weights='DEFAULT') for param in model.parameters(): param.requires_grad = False num_ftrs = model.classifier[1].in_features model.classifier[1] = nn.Linear(num_ftrs, 2) elif arch == "swin_t": model = models.swin_t(weights='DEFAULT') for param in model.parameters(): param.requires_grad = False num_ftrs = model.head.in_features model.head = nn.Linear(num_ftrs, 2) elif arch == "swin_t_finetune": model = models.swin_t(weights='DEFAULT') for param in model.parameters(): param.requires_grad = False num_ftrs = model.head.in_features model.head = nn.Linear(num_ftrs, 2) for param in model.features[7].parameters(): param.requires_grad = True for param in model.norm.parameters(): param.requires_grad = True for param in model.head.parameters(): param.requires_grad = True elif arch == "swin_s": model = models.swin_s(weights='DEFAULT') for param in model.parameters(): param.requires_grad = False num_ftrs = model.head.in_features model.head = nn.Linear(num_ftrs, 2) for param in model.features[7].parameters(): param.requires_grad = True for param in model.norm.parameters(): param.requires_grad = True for param in model.head.parameters(): param.requires_grad = True else: raise ValueError(f"Unknown arch: {arch}") return model def train(data_dir, model_save_path, arch, num_epochs, batch_size, lr, use_scheduler): train_transforms = transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.RandomVerticalFlip(), transforms.RandomRotation(30), transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.3, hue=0.1), transforms.RandomGrayscale(p=0.1), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) val_transforms = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) image_datasets = { 'train': datasets.ImageFolder(os.path.join(data_dir, 'train'), train_transforms), 'val': datasets.ImageFolder(os.path.join(data_dir, 'val'), val_transforms), } dataloaders = {x: DataLoader(image_datasets[x], batch_size=batch_size, shuffle=True, num_workers=4) for x in ['train', 'val']} dataset_sizes = {x: len(image_datasets[x]) for x in ['train', 'val']} device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") print(f"Using device: {device}") model = build_model(arch).to(device) trainable_params = filter(lambda p: p.requires_grad, model.parameters()) criterion = nn.CrossEntropyLoss() optimizer = optim.AdamW(trainable_params, lr=lr, weight_decay=0.05) scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs) if use_scheduler else None best_model_wts = copy.deepcopy(model.state_dict()) best_acc = 0.0 for epoch in range(num_epochs): print(f'Epoch {epoch}/{num_epochs - 1}') print('-' * 10) for phase in ['train', 'val']: model.train() if phase == 'train' else model.eval() running_loss = 0.0 running_corrects = 0 for inputs, labels in dataloaders[phase]: inputs, labels = inputs.to(device), labels.to(device) optimizer.zero_grad() with torch.set_grad_enabled(phase == 'train'): outputs = model(inputs) _, preds = torch.max(outputs, 1) loss = criterion(outputs, labels) if phase == 'train': loss.backward() optimizer.step() running_loss += loss.item() * inputs.size(0) running_corrects += torch.sum(preds == labels.data) if phase == 'train' and scheduler: scheduler.step() epoch_loss = running_loss / dataset_sizes[phase] epoch_acc = running_corrects.double() / dataset_sizes[phase] print(f'{phase} Loss: {epoch_loss:.4f} Acc: {epoch_acc:.4f}') if phase == 'val' and epoch_acc > best_acc: best_acc = epoch_acc best_model_wts = copy.deepcopy(model.state_dict()) print() print(f'Best val Acc: {best_acc:.4f}') model.load_state_dict(best_model_wts) os.makedirs(os.path.dirname(model_save_path), exist_ok=True) torch.save(model.state_dict(), model_save_path) print(f"Model saved to {model_save_path}") if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("--arch", choices=["mobilenet", "swin_t", "swin_t_finetune", "swin_s"], required=True) parser.add_argument("--data-dir", default="./data/split") parser.add_argument("--output", required=True, help="Path to save model weights") parser.add_argument("--epochs", type=int, default=20) parser.add_argument("--batch-size", type=int, default=32) parser.add_argument("--lr", type=float, default=0.0001) parser.add_argument("--scheduler", action="store_true") args = parser.parse_args() train(args.data_dir, args.output, args.arch, args.epochs, args.batch_size, args.lr, args.scheduler)