File size: 2,902 Bytes
65cec23
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
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}")