Spaces:
Runtime error
Runtime error
| import os | |
| import numpy | |
| import torch | |
| from pytorch_lightning import LightningModule, Trainer, tuner, seed_everything | |
| from pytorch_lightning.callbacks import ModelSummary | |
| from torch import nn | |
| from torch.nn import functional as F | |
| from torch.utils.data import DataLoader, random_split | |
| from torchmetrics import Accuracy | |
| from torchvision import transforms | |
| from torchvision.datasets import CIFAR10 | |
| from torch.optim.lr_scheduler import OneCycleLR | |
| import albumentations as A | |
| from albumentations import * | |
| from albumentations.pytorch.transforms import ToTensor, ToTensorV2 | |
| from dataset import * | |
| BATCH_SIZE = 256 | |
| class LitResBlock(LightningModule): | |
| def __init__(self, in_channels, out_channels, kernel_size, padding): | |
| super().__init__() | |
| self.in_channels = in_channels | |
| self.out_channels = out_channels | |
| self.kernel_size = kernel_size | |
| self.padding = padding | |
| self.convblock1 = nn.Sequential( | |
| nn.Conv2d(in_channels=self.in_channels, out_channels=self.out_channels, kernel_size=self.kernel_size, padding=self.padding, bias=False), | |
| nn.BatchNorm2d(self.out_channels), | |
| nn.ReLU() | |
| ) | |
| self.convblock2 = nn.Sequential( | |
| nn.Conv2d(in_channels=self.out_channels, out_channels=self.out_channels, kernel_size=self.kernel_size, padding=self.padding, bias=False), | |
| nn.BatchNorm2d(self.out_channels), | |
| nn.ReLU() | |
| ) | |
| def forward(self, x): | |
| y = self.convblock1(x) | |
| y = self.convblock2(y) | |
| return y | |
| class LitCIFAR10CustomResidualNet(LightningModule): | |
| def __init__(self, dropout_value=0, num_of_inp_channels=3, num_of_op_channels=10, data_dir=".", learning_rate=2e-4): | |
| super().__init__() | |
| # Set our init args as class attributes | |
| self.data_dir = data_dir | |
| self.learning_rate = learning_rate | |
| # Hardcode some dataset specific attributes | |
| self.num_classes = 10 | |
| self.class_labels = ("airplane", "automobile", "bird", "cat", "deer", "dog", "frog", "horse", "ship", "truck") | |
| self.means = numpy.array((0.4914, 0.4822, 0.4465)) | |
| self.stddev = numpy.array((0.2023, 0.1994, 0.2010)) | |
| self.transform = A.Compose([ | |
| A.PadIfNeeded(min_height=40, min_width=40), | |
| A.RandomSizedCrop((32,32), 32,32), | |
| A.HorizontalFlip(p = 0.5), | |
| A.Cutout(num_holes=1, max_h_size=8, max_w_size=8, fill_value=self.means*255.0, p=0.75), | |
| A.Normalize(mean=self.means, std=self.stddev), | |
| ToTensorV2() | |
| ]) | |
| self.accuracy = Accuracy(task='multiclass', num_classes=self.num_classes) | |
| self.dropout_value = dropout_value | |
| self.num_of_channels = num_of_inp_channels | |
| self.num_of_op_channels = num_of_op_channels | |
| self.number_of_kernels = [64, 128, 128, 256, 512, 512] | |
| # Input Block | |
| self.preplayer = nn.Sequential( | |
| nn.Conv2d(in_channels=self.num_of_channels, out_channels=self.number_of_kernels[0], kernel_size=(3, 3), padding=1, bias=False), | |
| nn.BatchNorm2d(self.number_of_kernels[0]), | |
| nn.ReLU() | |
| ) # input_size = 32x32x3, output_size = 32x32x64, RF = 3x3 | |
| # LAYER 1 | |
| self.layer1_x = nn.Sequential( | |
| nn.Conv2d(in_channels=self.number_of_kernels[0], out_channels=self.number_of_kernels[1], kernel_size=(3, 3), padding=1, bias=False), | |
| nn.MaxPool2d(2, 2), | |
| nn.BatchNorm2d(self.number_of_kernels[1]), | |
| nn.ReLU() | |
| ) # input_size = 32x32x64, output_size = 32x32x128, RF = 5x5 | |
| # RESIDUAL BLOCK 1 | |
| self.resblock1 = LitResBlock(in_channels=self.number_of_kernels[1], out_channels=self.number_of_kernels[2], kernel_size=(3,3), padding=1) | |
| # input_size = 32x32x128, output_size = 32x32x128, RF = 5x5, 9x9 | |
| # LAYER 2 | |
| self.layer2 = nn.Sequential( | |
| nn.Conv2d(in_channels=self.number_of_kernels[2], out_channels=self.number_of_kernels[3], kernel_size=(3, 3), padding=1, bias=False), | |
| nn.MaxPool2d(2, 2), | |
| nn.BatchNorm2d(self.number_of_kernels[3]), | |
| nn.ReLU() | |
| ) # input_size = 32x32x128, output_size = 16x16x256, RF = 8x8, 12x12 | |
| # LAYER 3 | |
| self.layer3_x = nn.Sequential( | |
| nn.Conv2d(in_channels=self.number_of_kernels[3], out_channels=self.number_of_kernels[4], kernel_size=(3, 3), padding=1, bias=False), | |
| nn.MaxPool2d(2, 2), | |
| nn.BatchNorm2d(self.number_of_kernels[4]), | |
| nn.ReLU() | |
| ) # input_size = 16x16x256, output_size = 8x8x512, RF = | |
| # RESIDUAL BLOCK 1 | |
| self.resblock2 = LitResBlock(in_channels=self.number_of_kernels[4], out_channels=self.number_of_kernels[5], kernel_size=(3,3), padding=1) | |
| # input_size = 8x8x512, output_size = 8x8x512, RF = | |
| # OUTPUT LAYER | |
| self.max_pool = nn.MaxPool2d(4, 2) # input_size = 8x8x512, output_size = 1x1x512, RF = | |
| self.fc_layer = nn.Sequential( | |
| nn.Conv2d(in_channels=self.number_of_kernels[5], out_channels=self.num_of_op_channels, kernel_size=(1, 1), padding=0, bias=False) | |
| ) # input_size = 1x1x512, output_size = 1x1x10, RF = | |
| self.rb1 = nn.Sequential() | |
| self.rb2 = nn.Sequential() | |
| def forward(self, inp): | |
| x0 = self.preplayer(inp) | |
| x = self.layer1_x(x0) | |
| r1 = self.resblock1(x) | |
| y1 = r1 + x | |
| y1 = self.rb1(y1) | |
| y2 = self.layer2(y1) | |
| x3 = self.layer3_x(y2) | |
| r2 = self.resblock2(x3) | |
| y3 = r2 + x3 | |
| y3 = self.rb2(y3) | |
| y4 = self.max_pool(y3) | |
| y5 = self.fc_layer(y4) | |
| y5 = y5.view(-1, 10) | |
| y5 = nn.Softmax(dim=-1)(y5) | |
| return y5 | |
| def training_step(self, batch, batch_idx): | |
| x, y = batch | |
| output = self(x) | |
| loss = nn.CrossEntropyLoss()(output, y) | |
| return loss | |
| def validation_step(self, batch, batch_idx): | |
| x, y = batch | |
| output = self(x) | |
| loss = nn.CrossEntropyLoss()(output, y) | |
| preds = torch.argmax(output, dim=1) | |
| self.accuracy(preds, y) | |
| # Calling self.log will surface up scalars for you in TensorBoard | |
| self.log("val_loss", loss, prog_bar=True) | |
| self.log("val_acc", self.accuracy, prog_bar=True) | |
| return loss | |
| def test_step(self, batch, batch_idx): | |
| # Here we just reuse the validation_step for testing | |
| return self.validation_step(batch, batch_idx) | |
| def configure_optimizers(self): | |
| optimizer = torch.optim.Adam(self.parameters(), lr=self.learning_rate) | |
| # final_div_factor = div_factor for no annhilation | |
| DIV_FACTOR = 100 | |
| FINAL_DIV_FACTOR = 100 | |
| EPOCHS = 24 | |
| MAX_LR_EPOCH = 5 | |
| NUM_OF_BATCHES = len(self.train_dataloader()) | |
| PCT_START = MAX_LR_EPOCH/EPOCHS | |
| # Based on above found maximum LR, initialize LRMAX and LRMIN | |
| LRMAX = self.learning_rate * DIV_FACTOR #best_lr | |
| #LRMIN = LRMAX/100 | |
| scheduler_params = {"max_lr": LRMAX, | |
| "steps_per_epoch": NUM_OF_BATCHES, | |
| "epochs": EPOCHS, | |
| "pct_start": PCT_START, | |
| "anneal_strategy":"linear", | |
| "div_factor": DIV_FACTOR, | |
| "final_div_factor": FINAL_DIV_FACTOR, | |
| "three_phase":False} | |
| scheduler_dict = { | |
| "scheduler": OneCycleLR( | |
| optimizer, | |
| **scheduler_params | |
| ), | |
| "interval": "step", | |
| } | |
| return {"optimizer": optimizer, "lr_scheduler": scheduler_dict} | |
| #################### | |
| # DATA RELATED HOOKS | |
| #################### | |
| def prepare_data(self): | |
| # download | |
| Cifar10AlbumDataset(self.data_dir, train=True, download=True) | |
| Cifar10AlbumDataset(self.data_dir, train=False, download=True) | |
| def setup(self, stage=None): | |
| # Assign train/val datasets for use in dataloaders | |
| if stage == "fit" or stage is None: | |
| cifar10_full = Cifar10AlbumDataset(self.data_dir, train=True, transform=self.transform) | |
| self.cifar10_train, self.cifar10_val = random_split(cifar10_full, [45000, 5000]) | |
| # Assign test dataset for use in dataloader(s) | |
| if stage == "test" or stage is None: | |
| self.cifar10_test = Cifar10AlbumDataset(self.data_dir, train=False, transform=self.transform) | |
| def train_dataloader(self): | |
| return DataLoader(self.cifar10_train, batch_size=BATCH_SIZE, num_workers=os.cpu_count()) | |
| def val_dataloader(self): | |
| return DataLoader(self.cifar10_val, batch_size=BATCH_SIZE, num_workers=os.cpu_count()) | |
| def test_dataloader(self): | |
| return DataLoader(self.cifar10_test, batch_size=BATCH_SIZE, num_workers=os.cpu_count()) |