{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Train Model for Distracted Driver Detection","metadata":{}},{"cell_type":"markdown","source":"## Import Libraries","metadata":{}},{"cell_type":"code","source":"# Install dependencies\n\n# Commented because already present on Kaggle\n# ! pip install torch\n# ! pip install pytorch-lightning\n# ! pip install torchmetrics\n# ! pip install timm\n# ! pip install albumentations\n# ! pip install opencv-python\n# ! pip install matplotlib\n# ! pip install pandas\n# ! pip install scikit-learn\n# ! pip install tqdm\n# ! pip install plotly\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:31.868022Z","iopub.execute_input":"2023-06-16T08:21:31.868440Z","iopub.status.idle":"2023-06-16T08:21:31.873618Z","shell.execute_reply.started":"2023-06-16T08:21:31.868408Z","shell.execute_reply":"2023-06-16T08:21:31.872572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport os.path as osp\nfrom glob import glob\nimport random\nimport time\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport plotly.express as px\nfrom enum import Enum\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import accuracy_score\n\nimport timm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.utils.data as data\nimport torch.optim as optim\n\nimport torchvision\nimport torchvision.transforms as transforms\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport pytorch_lightning as pl\nfrom pytorch_lightning import LightningModule, Trainer, seed_everything\nfrom pytorch_lightning.callbacks import LearningRateMonitor\nfrom pytorch_lightning.callbacks.progress import TQDMProgressBar\nfrom pytorch_lightning.loggers import CSVLogger\n\nimport torchmetrics\nfrom torchmetrics.functional import accuracy\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:31.878671Z","iopub.execute_input":"2023-06-16T08:21:31.879114Z","iopub.status.idle":"2023-06-16T08:21:31.888824Z","shell.execute_reply.started":"2023-06-16T08:21:31.879069Z","shell.execute_reply":"2023-06-16T08:21:31.887896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.is_available()","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:31.891370Z","iopub.execute_input":"2023-06-16T08:21:31.892279Z","iopub.status.idle":"2023-06-16T08:21:31.904349Z","shell.execute_reply.started":"2023-06-16T08:21:31.892253Z","shell.execute_reply":"2023-06-16T08:21:31.903388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Static Variables","metadata":{}},{"cell_type":"code","source":"class StaticDataset(Enum):\n    \"\"\"Class for static training configuration constants\"\"\"\n    ACTIVITY_MAP = {\n        'c0': 'Safe driving',\n        'c1': 'Texting - right',\n        'c2': 'Talking on the phone - right',\n        'c3': 'Texting - left',\n        'c4': 'Talking on the phone - left',\n        'c5': 'Operating the radio',\n        'c6': 'Drinking',\n        'c7': 'Reaching behind',\n        'c8': 'Hair and makeup',\n        'c9': 'Talking to passenger'\n    }\n    DATA_DIR = '/kaggle/input/state-farm-distracted-driver-detection/'\n    CSV_FILE_PATH = osp.join(DATA_DIR, 'driver_imgs_list.csv')\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:31.906247Z","iopub.execute_input":"2023-06-16T08:21:31.906915Z","iopub.status.idle":"2023-06-16T08:21:31.914289Z","shell.execute_reply.started":"2023-06-16T08:21:31.906883Z","shell.execute_reply":"2023-06-16T08:21:31.913389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class StaticLearningParameter(Enum):\n    MODEL_NAME_0 = 'efficientnet_b0'\n    MODEL_NAME_3 = 'efficientnet_b3'\n    COLOR_MEAN = (0.485, 0.456, 0.406)\n    COLOR_STD = (0.229, 0.224, 0.225)\n    INPUT_SIZE = 256\n    NUM_CLASSES = 10\n    BATCH_SIZE = 64\n    EPOCHS = 3 #10\n    FOLDS = 5\n    LR = 1e-3\n    GAMMA = 0.98\n    DEBUG = True\n    TRAIN = False\n    SEED = 42\n    USE_ALBUMENTATIONS = True\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:32.885314Z","iopub.execute_input":"2023-06-16T08:21:32.885672Z","iopub.status.idle":"2023-06-16T08:21:32.892375Z","shell.execute_reply.started":"2023-06-16T08:21:32.885642Z","shell.execute_reply":"2023-06-16T08:21:32.891060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"PLOT = False  # Set True if you want to plot the image","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:32.894264Z","iopub.execute_input":"2023-06-16T08:21:32.895209Z","iopub.status.idle":"2023-06-16T08:21:32.903442Z","shell.execute_reply.started":"2023-06-16T08:21:32.895175Z","shell.execute_reply":"2023-06-16T08:21:32.902340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fix_seed(seed):\n    \"\"\"Fix the seed of the random number generator for reproducibility.\"\"\"\n    # random\n    random.seed(seed)\n    # Numpy\n    np.random.seed(seed)\n    # Pytorch\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:33.644682Z","iopub.execute_input":"2023-06-16T08:21:33.645829Z","iopub.status.idle":"2023-06-16T08:21:33.652035Z","shell.execute_reply.started":"2023-06-16T08:21:33.645784Z","shell.execute_reply":"2023-06-16T08:21:33.650652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset Exploration","metadata":{}},{"cell_type":"code","source":"# Read csv file\ndf = pd.read_csv(StaticDataset.CSV_FILE_PATH.value)\n# Show first 5 lines\ndf.head(5)","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:33.656433Z","iopub.execute_input":"2023-06-16T08:21:33.656720Z","iopub.status.idle":"2023-06-16T08:21:33.689068Z","shell.execute_reply.started":"2023-06-16T08:21:33.656696Z","shell.execute_reply":"2023-06-16T08:21:33.688223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Number of Drivers","metadata":{}},{"cell_type":"code","source":"# Group by driver\nby_drivers = df.groupby('subject')\n# list of driver names\nunique_drivers = by_drivers.groups.keys()\n\n# Number of drivers in dataset\nprint('unique drivers: ', len(unique_drivers))\n# Average number of images per driver\nprint('mean of images: ', round(df.groupby(\n    'subject').count()['classname'].mean()))\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:33.690543Z","iopub.execute_input":"2023-06-16T08:21:33.690863Z","iopub.status.idle":"2023-06-16T08:21:33.720701Z","shell.execute_reply.started":"2023-06-16T08:21:33.690833Z","shell.execute_reply":"2023-06-16T08:21:33.719619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# number of training data\ntrain_file_num = len(\n    glob(osp.join(StaticDataset.DATA_DIR.value, 'imgs/train/*/*.jpg')))\n# number of test data\ntest_file_num = len(\n    glob(osp.join(StaticDataset.DATA_DIR.value, 'imgs/test/*.jpg')))\n# number of categories\ncategory_num = len(df['classname'].unique())\nprint('train_file_num: ', train_file_num)\nprint('test_file_num: ', test_file_num)\nprint('category_num: ', category_num)\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:33.722003Z","iopub.execute_input":"2023-06-16T08:21:33.722382Z","iopub.status.idle":"2023-06-16T08:21:34.118005Z","shell.execute_reply.started":"2023-06-16T08:21:33.722348Z","shell.execute_reply":"2023-06-16T08:21:34.116867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Number of Images per Class","metadata":{}},{"cell_type":"code","source":"# Number of data per class\npx.histogram(df, x=\"classname\", color=\"classname\",\n             title=\"Number of images by categories \")\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:34.121302Z","iopub.execute_input":"2023-06-16T08:21:34.121596Z","iopub.status.idle":"2023-06-16T08:21:34.330112Z","shell.execute_reply.started":"2023-06-16T08:21:34.121572Z","shell.execute_reply":"2023-06-16T08:21:34.329276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Number of Images per Driver","metadata":{}},{"cell_type":"markdown","source":"Number of images per driver sorted by number of images.","metadata":{}},{"cell_type":"code","source":"drivers_id = pd.DataFrame((df['subject'].value_counts()).reset_index())\ndrivers_id.columns = ['driver_id', 'Counts']\npx.histogram(drivers_id, x=\"driver_id\",y=\"Counts\" ,color=\"driver_id\", title=\"Number of images by subjects \")","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:34.331460Z","iopub.execute_input":"2023-06-16T08:21:34.332012Z","iopub.status.idle":"2023-06-16T08:21:34.977640Z","shell.execute_reply.started":"2023-06-16T08:21:34.331972Z","shell.execute_reply":"2023-06-16T08:21:34.976685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Number of images per driver sorted by driver id.","metadata":{}},{"cell_type":"code","source":"# Histogram of number of images per driver\npx.histogram(df, x='subject', color='subject',\n             title='Number of images by subjects')\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:34.979335Z","iopub.execute_input":"2023-06-16T08:21:34.979694Z","iopub.status.idle":"2023-06-16T08:21:35.237812Z","shell.execute_reply.started":"2023-06-16T08:21:34.979662Z","shell.execute_reply":"2023-06-16T08:21:35.236796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Image for each Class","metadata":{}},{"cell_type":"code","source":"# Draw data for each class\nplt.figure(figsize=(12, 20))\nfor i, (key, value) in enumerate(StaticDataset.ACTIVITY_MAP.value.items()):\n    image_dir = osp.join(StaticDataset.DATA_DIR.value, 'imgs/train', key, '*.jpg')\n    image_path = glob(image_dir)[0]\n    image = cv2.imread(image_path)[:, :, (2, 1, 0)]\n    plt.subplot(5, 2, i+1)\n    plt.imshow(image)\n    plt.title(value)\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:35.238953Z","iopub.execute_input":"2023-06-16T08:21:35.240320Z","iopub.status.idle":"2023-06-16T08:21:38.615558Z","shell.execute_reply.started":"2023-06-16T08:21:35.240288Z","shell.execute_reply":"2023-06-16T08:21:38.614355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preprocessing","metadata":{}},{"cell_type":"code","source":"# Add file path column\ndf['file_path'] = df.apply(lambda x: osp.join(\n    StaticDataset.DATA_DIR.value, 'imgs/train', x.classname, x.img), axis=1)\n\n# Add Column by Converting Correct Answer Labels to Numbers\ndf['class_num'] = df['classname'].map(lambda x: int(x[1]))\ndf.head(5)\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:38.617228Z","iopub.execute_input":"2023-06-16T08:21:38.617844Z","iopub.status.idle":"2023-06-16T08:21:40.099240Z","shell.execute_reply.started":"2023-06-16T08:21:38.617810Z","shell.execute_reply":"2023-06-16T08:21:40.098193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"markdown","source":"### Custom Dataset","metadata":{}},{"cell_type":"code","source":"class Dataset(data.Dataset):\n    \"\"\"\n    Attributes\n    ----------\n    df : DataFrame\n        class_num, dataframe with column file_path\n    phase : 'train' or 'val'\n        Set learning or training.\n    transform : object\n        an instance of the preprocessing class\n    \"\"\"\n\n    def __init__(self, df, phase, transform):\n        self.df = df\n        self.phase = phase\n        self.transform = transform\n\n    def __len__(self):\n        '''returns the number of images'''\n        return len(self.df)\n\n    def __getitem__(self, index):\n        '''Get Tensor format data of preprocessed image'''\n        image = self.pull_item(index)\n        return image, self.df.iloc[index]['class_num']\n\n    def pull_item(self, index):\n        '''Get Tensor format data of image'''\n        # 1. Image loading\n        image_path = self.df.iloc[index]['file_path']\n        image = cv2.imread(image_path)[:, :, (2, 1, 0)]\n        # 2. Perform pretreatment\n        return self.transform(self.phase, image)\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:40.101151Z","iopub.execute_input":"2023-06-16T08:21:40.101879Z","iopub.status.idle":"2023-06-16T08:21:40.111351Z","shell.execute_reply.started":"2023-06-16T08:21:40.101845Z","shell.execute_reply":"2023-06-16T08:21:40.110351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Data Transformations","metadata":{}},{"cell_type":"code","source":"class DataTransform():\n    \"\"\"\n    Image and annotation preprocessing classes. \n    It behaves differently during training and during validation.\n    Set the image size to input_size x input_size.\n    Data augmentation during training.\n\n\n    Attributes\n    ----------\n    input_size : int\n        The size of the resized image.\n    color_mean : (R, G, B)\n        Average value for each color channel.\n    color_std : (R, G, B)\n        Standard deviation for each color channel.\n    \"\"\"\n\n    def __init__(\n            self,\n            input_size,\n            color_mean,\n            color_std,\n            use_albumentations=StaticLearningParameter.USE_ALBUMENTATIONS.value):\n        if use_albumentations:\n            # Albumentations Transformations\n            self.data_transform = {\n                # Implement only train\n                'train': A.Compose([\n                    A.HorizontalFlip(p=0.5),\n                    A.Rotate(-10, 10),\n                    A.Resize(input_size, input_size),  # resize(input_size)\n                    # Standardization of color information\n                    A.Normalize(color_mean, color_std),\n                    ToTensorV2()\n                ]),\n                'val': A.Compose([\n                    A.Resize(input_size, input_size),  # resize(input_size)\n                    # Standardization of color information\n                    A.Normalize(color_mean, color_std),\n                    ToTensorV2()\n                ])\n            }\n        else:\n            # PyTorch Transformations\n            self.data_transform = {\n                # Implement only train\n                'train': transforms.Compose([\n                    transforms.RandomHorizontalFlip(p=0.5),\n                    transforms.RandomRotation((-10, 10)),\n                    transforms.Resize(input_size),  # resize(input_size)\n                    transforms.Normalize(color_mean, color_std),\n                    transforms.ToTensor(),\n                ]),\n                'val': transforms.Compose([\n                    transforms.Resize(input_size),  # resize(input_size)\n                    transforms.Normalize(color_mean, color_std),\n                    transforms.ToTensor(),\n                ])\n            }\n\n    def __call__(self, phase, image):\n        \"\"\"\n        Parameters\n        ----------\n        phase : 'train' or 'val'\n            Specifies the preprocessing mode.\n        \"\"\"\n        transformed = self.data_transform[phase](image=image)\n        return transformed['image']\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:40.115966Z","iopub.execute_input":"2023-06-16T08:21:40.118415Z","iopub.status.idle":"2023-06-16T08:21:40.133726Z","shell.execute_reply.started":"2023-06-16T08:21:40.118382Z","shell.execute_reply":"2023-06-16T08:21:40.132667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Split Dataset","metadata":{}},{"cell_type":"code","source":"def create_dataset(\n        df,\n        seed=StaticLearningParameter.SEED.value,\n        input_size=StaticLearningParameter.INPUT_SIZE.value,\n        color_mean=StaticLearningParameter.COLOR_MEAN.value,\n        color_std=StaticLearningParameter.COLOR_STD.value):\n    \"\"\"Create dataset for training and validation.\n\n    Args:\n        seed (int, optional): Random seed. \n            Defaults to StaticTrain.SEED.value.\n        input_size (int, optional): Input size. \n            Defaults to StaticDataTransform.INPUT_SIZE.value.\n        color_mean (tuple, optional): Color mean. \n            Defaults to StaticDataTransform.COLOR_MEAN.value.\n        color_std (tuple, optional): Color standard deviation. \n            Defaults to StaticDataTransform.COLOR_STD.value.\n        \n    Returns:\n        train_dataset, val_dataset: Dataset for training and validation.\n    \"\"\"\n    # data division\n    df_train, df_val = train_test_split(\n        df,\n        stratify=df['subject'],\n        random_state=seed\n    )\n\n    # dataset creation\n    train_dataset = Dataset(\n        df_train,\n        phase=\"train\",\n        transform=DataTransform(\n            input_size=input_size,\n            color_mean=color_mean,\n            color_std=color_std\n        )\n    )\n\n    val_dataset = Dataset(\n        df_val,\n        phase=\"val\",\n        transform=DataTransform(\n            input_size=input_size,\n            color_mean=color_mean,\n            color_std=color_std\n        )\n    )\n\n    if PLOT:\n        # Data retrieval example\n        image, label = train_dataset[0]\n        plt.imshow(image.permute(1, 2, 0))\n        plt.title(label)\n        plt.show()\n\n    return train_dataset, val_dataset","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:40.137917Z","iopub.execute_input":"2023-06-16T08:21:40.140609Z","iopub.status.idle":"2023-06-16T08:21:40.152321Z","shell.execute_reply.started":"2023-06-16T08:21:40.140577Z","shell.execute_reply":"2023-06-16T08:21:40.151374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Create Data Loaders","metadata":{}},{"cell_type":"code","source":"def create_datamodule(\n        train_dataset,\n        val_dataset,\n        batch_size=StaticLearningParameter.BATCH_SIZE.value):\n    \"\"\"\n    Create dataloader for training and validation.\n\n    Args:\n        train_dataset: Dataset for training.\n        val_dataset: Dataset for validation.\n        batch_size (int, optional): Batch size. \n            Defaults to StaticDataLoader.BATCH_SIZE.value.\n    \n    Returns:\n        datamodule: Dataloader for training and validation.\n    \"\"\"\n    datamodule = pl.LightningDataModule.from_datasets(\n        train_dataset=train_dataset,\n        val_dataset=val_dataset,\n        batch_size=batch_size if torch.cuda.is_available() else 4,\n        num_workers=int(os.cpu_count() / 2),\n    )\n    return datamodule","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:40.157375Z","iopub.execute_input":"2023-06-16T08:21:40.159912Z","iopub.status.idle":"2023-06-16T08:21:40.169123Z","shell.execute_reply.started":"2023-06-16T08:21:40.159844Z","shell.execute_reply":"2023-06-16T08:21:40.168216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"class LitEfficientNet(LightningModule):\n    \"\"\"LightningModule for EfficientNet.\"\"\"\n\n    def __init__(\n        self, \n        model, \n        lr=StaticLearningParameter.LR.value, \n        gamma=StaticLearningParameter.GAMMA.value):\n        super().__init__()\n        self.save_hyperparameters()\n        self.model = model\n        self.lr = lr\n        self.gamma = gamma\n        num_classes=len(StaticDataset.ACTIVITY_MAP.value.items())\n        self.accuracy = torchmetrics.Accuracy(task='multiclass', num_classes=num_classes)\n\n    def forward(self, x):\n        \"\"\"Forward propagation.\"\"\"\n        out = self.model(x)\n        return F.log_softmax(out, dim=1)\n\n    def training_step(self, batch, batch_idx):\n        \"\"\"\n        Training step.\n        \n        Args:\n            batch: batch data\n            batch_idx: batch index\n        \"\"\"\n        x, y = batch\n        logits = self(x)\n        loss = F.nll_loss(logits, y)\n        self.log(\"train_loss\", loss)\n        return loss\n\n    def evaluate(self, batch, stage=None):\n        \"\"\"\n        Evaluation step.\n\n        Args:\n            batch: batch data\n            stage: stage name\n        \"\"\"\n        x, y = batch\n        logits = self(x)\n        loss = F.nll_loss(logits, y)\n        preds = torch.argmax(logits, dim=1)\n        acc = self.accuracy(preds, y)\n\n        if stage:\n            self.log(f\"{stage}_loss\", loss, prog_bar=True)\n            self.log(f\"{stage}_acc\", acc, prog_bar=True)\n\n    def validation_step(self, batch, batch_idx):\n        \"\"\"\n        Validation step.\n\n        Args:\n            batch: batch data\n            batch_idx: batch index\n        \"\"\"\n        self.evaluate(batch, \"val\")\n\n    def test_step(self, batch, batch_idx):\n        \"\"\"\n        Test step.\n        \n        Args:\n            batch: batch data\n            batch_idx: batch index\n        \"\"\"\n        self.evaluate(batch, \"test\")\n\n    def configure_optimizers(self):\n        \"\"\"Configure optimizers.\"\"\"\n\n        \"\"\"\n        optimizer = torch.optim.SGD(\n            self.parameters(),\n            lr=self.hparams.lr,\n            momentum=0.9,\n            # weight_decay=5e-4,\n        )\n        \"\"\"\n        optimizer = torch.optim.Adam(\n            self.parameters(),\n            lr=self.lr\n        )\n\n        scheduler = torch.optim.lr_scheduler.ExponentialLR(\n            optimizer,\n            gamma=self.gamma\n        )\n\n        # criterion = nn.CrossEntropyLoss()  # loss function\n\n        return {\"optimizer\": optimizer, \"lr_scheduler\": scheduler}\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:40.176461Z","iopub.execute_input":"2023-06-16T08:21:40.178929Z","iopub.status.idle":"2023-06-16T08:21:40.197773Z","shell.execute_reply.started":"2023-06-16T08:21:40.178897Z","shell.execute_reply":"2023-06-16T08:21:40.196983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{}},{"cell_type":"code","source":"def train(\n        datamodule,\n        class_names=StaticDataset.ACTIVITY_MAP.value.items()):\n    \"\"\"Train the model with Pytorch Lightning.\n    \n    Args:\n        datamodule: Dataloader for training and validation.\n        class_names (list, optional): List of class names.\n        \n    \"\"\"\n    # Create model\n    # NOTE: MODEL_NAME_0 = \"efficientnet_b0\"\n    # NOTE: MODEL_NAME_3 = \"efficientnet_b3\" this is used in the reference code\n    efficient_model = timm.create_model(\n        StaticLearningParameter.MODEL_NAME_0.value,\n        pretrained=True,\n        num_classes=len(class_names)\n    )\n\n    model = LitEfficientNet(efficient_model)\n\n    trainer = Trainer(\n        max_epochs=StaticLearningParameter.EPOCHS.value,\n        accelerator=\"gpu\",\n        # accelerator=\"auto\",\n        #devices=1 if torch.cuda.is_available() else None,\n        logger=CSVLogger(save_dir=\"logs/\"),\n        callbacks=[\n            LearningRateMonitor(logging_interval=\"step\"),\n            TQDMProgressBar(refresh_rate=10),\n        ],\n    )\n    trainer.fit(model, datamodule=datamodule)\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:40.202016Z","iopub.execute_input":"2023-06-16T08:21:40.204994Z","iopub.status.idle":"2023-06-16T08:21:40.215537Z","shell.execute_reply.started":"2023-06-16T08:21:40.204794Z","shell.execute_reply":"2023-06-16T08:21:40.214574Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### Operation check ###\ntrain_dataset, val_dataset = create_dataset(df)\n\n### Create DataLoader ###\ndatamodule = create_datamodule(train_dataset, val_dataset)\n\n### Train the model ###\ntrain(datamodule)","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:21:40.220100Z","iopub.execute_input":"2023-06-16T08:21:40.222685Z","iopub.status.idle":"2023-06-16T09:24:56.943453Z","shell.execute_reply.started":"2023-06-16T08:21:40.222653Z","shell.execute_reply":"2023-06-16T09:24:56.942384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!zip -r file.zip \"/kaggle/working/\"","metadata":{"execution":{"iopub.status.busy":"2023-06-16T09:25:27.237567Z","iopub.execute_input":"2023-06-16T09:25:27.238356Z","iopub.status.idle":"2023-06-16T09:25:38.553217Z","shell.execute_reply.started":"2023-06-16T09:25:27.238318Z","shell.execute_reply":"2023-06-16T09:25:38.552006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls","metadata":{"execution":{"iopub.status.busy":"2023-06-16T09:27:39.963522Z","iopub.execute_input":"2023-06-16T09:27:39.963989Z","iopub.status.idle":"2023-06-16T09:27:41.242159Z","shell.execute_reply.started":"2023-06-16T09:27:39.963952Z","shell.execute_reply":"2023-06-16T09:27:41.240816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import FileLink\nFileLink(r'file.zip')","metadata":{"execution":{"iopub.status.busy":"2023-06-16T09:27:53.067351Z","iopub.execute_input":"2023-06-16T09:27:53.067775Z","iopub.status.idle":"2023-06-16T09:27:53.075827Z","shell.execute_reply.started":"2023-06-16T09:27:53.067740Z","shell.execute_reply":"2023-06-16T09:27:53.074720Z"},"trusted":true},"execution_count":null,"outputs":[]}]}