crtal's picture
refactor: consolidate 8 duplicated train/eval scripts into 2 parameterized scripts
e556bb4
Raw
History Blame Contribute Delete
5.88 kB
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)