from datasets import load_dataset import torch as pt import torchvision.transforms as T from torch.utils.data import Dataset, DataLoader from torch.optim.lr_scheduler import SequentialLR, LinearLR, CosineAnnealingLR from PIL import Image from tqdm.auto import tqdm from configuration_nula import NulaConfig from modeling_nula import NulaForImageClassification BATCH_SIZE_TRAIN = 128 BATCH_SIZE_TEST = 256 NUM_WORKERS = 0 MEAN = [0.5, 0.5, 0.5] STD = [0.5, 0.5, 0.5] train_transform = T.Compose([ T.RandomCrop(32, padding=4), T.RandomHorizontalFlip(p=0.5), T.AutoAugment(policy=T.AutoAugmentPolicy.CIFAR10), T.ToTensor(), T.Normalize(mean=MEAN, std=STD) ]) test_transform = T.Compose([ T.ToTensor(), T.Normalize(mean=MEAN, std=STD) ]) class CIFAR10Wrapper(Dataset): def __init__(self, huggingface_dataset, transform): self.dataset = huggingface_dataset self.transform = transform def __len__(self): return len(self.dataset) def __getitem__(self, i): sample = self.dataset[i] image = sample["img"] if not isinstance(image, Image.Image): image = Image.fromarray(image) return { "pixel_values": self.transform(image), "labels": sample["label"] } def get_loaders(): dataset = load_dataset("uoft-cs/cifar10") train_ds = CIFAR10Wrapper(dataset["train"], train_transform) test_ds = CIFAR10Wrapper(dataset["test"], test_transform) train_loader = DataLoader( train_ds, batch_size=BATCH_SIZE_TRAIN, shuffle=True, num_workers=NUM_WORKERS, pin_memory=True ) test_loader = DataLoader( test_ds, batch_size=BATCH_SIZE_TEST, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True ) return train_loader, test_loader def get_model_and_optimizer(device): cfg = NulaConfig( block_channels=(128, 256, 512), classifier_hidden_dim=512, use_se=True ) model = NulaForImageClassification(cfg).to(device) optimizer = pt.optim.AdamW( model.parameters(), lr=1e-3, weight_decay=0.01 ) warmup = LinearLR( optimizer, start_factor=0.1, end_factor=1.0, total_iters=5 ) cosine = CosineAnnealingLR( optimizer, T_max=45 ) scheduler = SequentialLR( optimizer, schedulers=[warmup, cosine], milestones=[5] ) return model, optimizer, scheduler def get_device(): if pt.cuda.is_available(): return "cuda" elif pt.backends.mps.is_available(): return "mps" return "cpu" if __name__ == "__main__": DEVICE = get_device() train_loader, test_loader = get_loaders() model, optimizer, scheduler = get_model_and_optimizer(DEVICE) print(f"Ready to train on {DEVICE}")