{"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":"<a href=\"https://www.kaggle.com/code/fummicc1/fummicc1-cassava?scriptVersionId=115383373\" target=\"_blank\"><img align=\"left\" alt=\"Kaggle\" title=\"Open in Kaggle\" src=\"https://kaggle.com/static/images/open-in-kaggle.svg\"></a>","metadata":{}},{"cell_type":"markdown","source":"## Data","metadata":{"papermill":{"duration":0.004271,"end_time":"2023-01-03T13:39:57.223551","exception":false,"start_time":"2023-01-03T13:39:57.21928","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nimport pandas as pd\nfrom torchvision import transforms\nfrom torchvision.io import read_image\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import Dataset\nfrom PIL import Image\nfrom typing import Dict, List, Tuple\nfrom tqdm import tqdm\nimport pandas as pd","metadata":{"papermill":{"duration":3.412991,"end_time":"2023-01-03T13:40:00.640786","exception":false,"start_time":"2023-01-03T13:39:57.227795","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-08T04:24:49.849330Z","iopub.execute_input":"2023-01-08T04:24:49.849714Z","iopub.status.idle":"2023-01-08T04:24:49.858095Z","shell.execute_reply.started":"2023-01-08T04:24:49.849676Z","shell.execute_reply":"2023-01-08T04:24:49.857005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append(\"/kaggle/input/timm-pytorch-image-models/pytorch-image-models-master\")\nimport timm\n\ntimm.list_models(\"efficient*\", pretrained=True)","metadata":{"execution":{"iopub.status.busy":"2023-01-08T04:24:49.860234Z","iopub.execute_input":"2023-01-08T04:24:49.860895Z","iopub.status.idle":"2023-01-08T04:24:49.871668Z","shell.execute_reply.started":"2023-01-08T04:24:49.860840Z","shell.execute_reply":"2023-01-08T04:24:49.870616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def class2dict(f) -> Dict:\n    ans = dict()\n    for name in dir(f):\n        if name.startswith(\"__\"):\n            continue\n        if not _is_primitive(getattr(f, name)):\n            ans[name] = class2dict(getattr(f, name))\n        else:\n            ans[name] = getattr(f, name)\n    return ans\n\n\ndef _is_primitive(value) -> bool:\n    primitive = (int, str, bool, float, List, Dict)\n    return type(value) in primitive","metadata":{"execution":{"iopub.status.busy":"2023-01-08T04:24:49.873338Z","iopub.execute_input":"2023-01-08T04:24:49.874080Z","iopub.status.idle":"2023-01-08T04:24:49.881224Z","shell.execute_reply.started":"2023-01-08T04:24:49.874044Z","shell.execute_reply":"2023-01-08T04:24:49.880339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"class Config:    \n    batch_size = 64\n    num_workers = 8\n    n_epochs = 15\n    lr = 1e-5\n    model_name = \"vit_large_patch16_224\"\n    is_kaggle_notebook = True\n    resized_height = 300    \n    resized_width = 600\n    \n    base_input_path_for_kaggle = \"/kaggle/input/cassava-leaf-disease-classification\"\n    base_input_path_for_local = \"./\"\n    \n    @property\n    def base_input_path(self):\n        return self.base_input_path_for_kaggle if self.is_kaggle_notebook else self.base_input_path_for_local","metadata":{"execution":{"iopub.status.busy":"2023-01-08T04:24:49.884265Z","iopub.execute_input":"2023-01-08T04:24:49.884943Z","iopub.status.idle":"2023-01-08T04:24:49.891385Z","shell.execute_reply.started":"2023-01-08T04:24:49.884907Z","shell.execute_reply":"2023-01-08T04:24:49.890423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = Config()\nconfig.n_epochs = 15\nconfig.batch_size = 64\nconfig.num_workers = 8\nconfig.lr = 5e-5\nconfig.model_name = \"efficientnetv2_rw_t\"\nconfig.is_kaggle_notebook = True\nconfig.resized_height = 600\nconfig.resized_width = 800\n\nconfig.base_input_path","metadata":{"papermill":{"duration":0.018484,"end_time":"2023-01-03T13:40:00.664302","exception":false,"start_time":"2023-01-03T13:40:00.645818","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-08T04:24:49.893774Z","iopub.execute_input":"2023-01-08T04:24:49.894531Z","iopub.status.idle":"2023-01-08T04:24:49.905102Z","shell.execute_reply.started":"2023-01-08T04:24:49.894495Z","shell.execute_reply":"2023-01-08T04:24:49.903968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_path = os.path.join(config.base_input_path, \"train.csv\")\ntrain_img_dir_path = os.path.join(config.base_input_path, \"train_images\")\ntrain_df = pd.read_csv(train_path)\ntrain_df.label.value_counts().plot(kind=\"bar\")","metadata":{"papermill":{"duration":0.33751,"end_time":"2023-01-03T13:40:01.006509","exception":false,"start_time":"2023-01-03T13:40:00.668999","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-08T04:24:49.906697Z","iopub.execute_input":"2023-01-08T04:24:49.907386Z","iopub.status.idle":"2023-01-08T04:24:50.117997Z","shell.execute_reply.started":"2023-01-08T04:24:49.907351Z","shell.execute_reply":"2023-01-08T04:24:50.116783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import Optional\nfrom torch.utils.data import Dataset\nimport torch.nn as nn\nfrom torchvision.transforms import transforms\n\nclass CassavaDataset(Dataset):\n    def __init__(self, annotation_path: Optional[str], img_dir_path: str, img_transforms: transforms.Compose):\n        self.has_label = annotation_path is not None\n        if annotation_path:\n            self.annotation_path = annotation_path\n            self.annotation_data = pd.read_csv(annotation_path)\n        else:\n            self.annotation_path = None\n            self.annotation_data = pd.DataFrame()\n            self.annotation_data[\"image_id\"] = list(os.listdir(img_dir_path))\n        self.img_dir_path = img_dir_path\n        self.img_transforms = img_transforms\n        \n    def __len__(self) -> int:\n        return len(self.annotation_data)\n    \n    def __getitem__(self, index: int):\n        data = self.annotation_data.iloc[index, :]\n        image_id = data[\"image_id\"]\n        if self.has_label:\n            label = data[\"label\"]\n        image_path = os.path.join(self.img_dir_path, image_id)        \n        img = Image.open(image_path).convert(\"RGB\")\n        if self.img_transforms:\n            img = self.img_transforms(img)\n        if self.has_label:\n            label = torch.nn.functional.one_hot(torch.tensor([label]), num_classes=5)\n            label = torch.squeeze(label, dim=0).float()\n        if self.has_label:\n            return img, image_id, label\n        else:\n            return img, image_id","metadata":{"papermill":{"duration":0.015811,"end_time":"2023-01-03T13:40:01.027496","exception":false,"start_time":"2023-01-03T13:40:01.011685","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-08T04:24:50.119425Z","iopub.execute_input":"2023-01-08T04:24:50.119743Z","iopub.status.idle":"2023-01-08T04:24:50.131427Z","shell.execute_reply.started":"2023-01-08T04:24:50.119717Z","shell.execute_reply":"2023-01-08T04:24:50.130276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Network","metadata":{"papermill":{"duration":0.004436,"end_time":"2023-01-03T13:40:01.03669","exception":false,"start_time":"2023-01-03T13:40:01.032254","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import timm\n\n\nclass CassavaNetwork(nn.Module):\n    def __init__(self, config: Config, output_dim: int = 5):\n        super().__init__()\n        self.base_model = timm.create_model(\n            config.model_name,\n            pretrained=not config.is_kaggle_notebook,\n            num_classes=output_dim,\n        )\n    \n    def forward(self, img: torch.Tensor) -> torch.Tensor:\n        return self.base_model(img)\n    \n\nmodel = CassavaNetwork(config=config)","metadata":{"papermill":{"duration":1.262427,"end_time":"2023-01-03T13:40:02.304048","exception":false,"start_time":"2023-01-03T13:40:01.041621","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-08T04:24:50.134218Z","iopub.execute_input":"2023-01-08T04:24:50.135036Z","iopub.status.idle":"2023-01-08T04:24:50.476465Z","shell.execute_reply.started":"2023-01-08T04:24:50.134998Z","shell.execute_reply":"2023-01-08T04:24:50.475357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Trainer","metadata":{"papermill":{"duration":0.005608,"end_time":"2023-01-03T13:40:02.315177","exception":false,"start_time":"2023-01-03T13:40:02.309569","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import numpy as np\nfrom typing import Optional\nfrom torch.cuda.amp.grad_scaler import GradScaler\nfrom torch.cuda.amp.autocast_mode import autocast\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader\nimport torch\n\nfrom tqdm import tqdm\nimport sys\nimport os\n\n\nclass TrainerOutput:\n    loss: float\n\n    def __init__(self, loss: float):\n        self.loss = loss\n\n    @property\n    def score(self) -> float:\n        return self.loss\n\n\nclass Trainer:\n    model: nn.Module\n    config: Config\n    optimizer: optim.Optimizer\n    lr_scheduler: optim.lr_scheduler._LRScheduler\n    loss_function: nn.CrossEntropyLoss\n    dataloader: DataLoader\n    scaler: Optional[GradScaler]\n\n    def __init__(\n        self,\n        model: nn.Module,\n        config: Config,\n        optimizer: optim.Optimizer,\n        lr_scheduler: optim.lr_scheduler._LRScheduler,\n        loss_function: nn.CrossEntropyLoss,\n        dataloader: DataLoader,\n        scaler: Optional[GradScaler] = None,\n    ):\n        self.model = model\n        self.config = config\n        self.optimizer = optimizer\n        self.lr_scheduler = lr_scheduler\n        self.loss_function = loss_function\n        self.dataloader = dataloader\n        self.scaler = scaler\n\n    def advance(self, verbose: bool = False) -> TrainerOutput:\n        scaler = None\n        if self.scaler:\n            scaler = self.scaler\n        self.model = self.model.train()\n        GPU_DEVICE = torch.device(\"cuda\")\n        epoch_loss = 0\n        with autocast(enabled=True):\n            with torch.enable_grad():\n                for batch in tqdm(self.dataloader):\n                    if len(batch) == 3:\n                        imgs, ids, labels = batch\n                    elif len(batch) == 2:\n                        imgs, labels = batch\n                    else:\n                        continue\n                    self.optimizer.zero_grad()\n                    imgs = imgs.to(GPU_DEVICE)\n                    labels = labels.to(GPU_DEVICE)\n                    out: torch.Tensor = self.model(imgs)\n                    out = out.float()\n                    if verbose:\n                        print(\"out-shape\", out.shape)\n                        print(\"out\", out)                        \n                    if verbose:\n                        print(\"labels\", labels)\n                    loss = self.loss_function(out, labels)\n                    if verbose:\n                        print(\"loss\", loss)\n                    if scaler is not None:\n                        scaler.scale(loss).backward()\n                    else:\n                        loss.backward()\n                    epoch_loss += loss.item()\n                    if scaler is not None:\n                        scaler.step(self.optimizer)\n                        scaler.update()\n                    else:\n                        self.optimizer.step()\n                self.lr_scheduler.step()\n        loss = epoch_loss / (len(self.dataloader))\n        return TrainerOutput(loss=loss)\n    \n    def save_model(self):\n        config = self.config\n        out_path = os.path.join(\"./\", \"weights\")\n        os.makedirs(out_path, exist_ok=True)\n        out_path = os.path.join(out_path, f\"epoch_{config.n_epochs}_base_{config.model_name}.pth\")\n        torch.save(self.model.state_dict(), out_path)","metadata":{"papermill":{"duration":0.024654,"end_time":"2023-01-03T13:40:02.34482","exception":false,"start_time":"2023-01-03T13:40:02.320166","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-08T04:24:50.480588Z","iopub.execute_input":"2023-01-08T04:24:50.481031Z","iopub.status.idle":"2023-01-08T04:24:50.498570Z","shell.execute_reply.started":"2023-01-08T04:24:50.480991Z","shell.execute_reply":"2023-01-08T04:24:50.497409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Classifier","metadata":{"papermill":{"duration":0.004534,"end_time":"2023-01-03T13:40:02.354125","exception":false,"start_time":"2023-01-03T13:40:02.349591","status":"completed"},"tags":[]}},{"cell_type":"code","source":"from typing import Callable, Dict, List, Tuple, Union, Optional\nfrom typing_extensions import Self\nimport numpy as np\nfrom torch import Tensor\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\n\nimport sys\nimport os\nimport pandas as pd\nfrom tqdm import tqdm\n\nclass ClassifierOutput:\n    move_to_full_nn_img_ids: List[str]\n    count: int\n    acc: float\n    loss: float\n    df: pd.DataFrame\n\n    def __init__(\n        self,\n        count: int,\n        acc: float,\n        loss: float,\n        df: pd.DataFrame = pd.DataFrame()\n    ):\n        self.count = count        \n        self.acc = acc\n        self.loss = loss\n        self.df = df\n\n    @property\n    def score(self) -> float:\n        return self.acc\n\n\nclass ClassifierTrainInput:\n    def __init__(self):\n        pass\n\nclass ClassifierTestInput:\n    def __init__(self):\n        pass\n\nclass Classifier:\n\n    model: nn.Module\n    activate_function: nn.Softmax\n    loss_function: Optional[nn.CrossEntropyLoss]\n    c_high: float\n    c_low: float\n    dataloader: DataLoader\n    phase: str\n    phase_input: Union[ClassifierTrainInput, ClassifierTestInput]\n    on_classify: Optional[Callable]\n\n    def __init__(\n        self,\n        model: nn.Module,\n        activate_function: nn.Softmax,\n        dataloader: DataLoader,\n        phase: str,\n        phase_input: Union[ClassifierTrainInput, ClassifierTestInput],\n        loss_function: Optional[nn.CrossEntropyLoss] = None,\n        on_classify: Optional[Callable] = None,\n    ):\n        self.model = model\n        self.activate_function = activate_function\n        self.loss_function = loss_function\n        self.dataloader = dataloader\n        self.phase = phase\n        self.phase_input = phase_input\n        self.on_classify = on_classify\n\n    def infer(\n        self, verbose: bool = False, handle_all: bool = False, calc_acc: bool = True\n    ) -> ClassifierOutput:\n        net = self.model\n        loader = self.dataloader\n        loss_function = self.loss_function\n        count = 0\n        correct = 0        \n        GPU_DEVICE = torch.device(\"cuda\")\n        called = False\n        net = net.eval()\n        epoch_loss = 0\n        label_df = pd.DataFrame()\n        with torch.no_grad():\n            for batch in tqdm(loader):\n                if len(batch) == 3:\n                    imgs, ids, labels = batch\n                elif len(batch) == 2:\n                    imgs, ids = batch\n                    labels = None\n                else:\n                    continue\n                imgs = imgs.to(GPU_DEVICE)\n                if labels is not None:\n                    labels = labels.to(GPU_DEVICE)                \n                outputs = net(imgs)\n                if loss_function and labels is not None:\n                    loss = loss_function(outputs, labels)\n                    epoch_loss += loss.item()\n                outputs = self.activate_function(outputs)\n                if verbose:\n                    print(\"output\", outputs[:10])\n                _, pred_indexes = torch.max(outputs, dim=1)\n                pred_indexes: Tensor = pred_indexes\n                if verbose:\n                    print(f\"{self.phase}_predicted_indexes\", pred_indexes)\n                    print(f\"{self.phase}_labels\", labels)\n                if not called and self.on_classify is not None:\n                    if not handle_all:\n                        called = True\n                    self.on_classify(imgs, pred_indexes)\n                \n                if labels is not None:\n                    labels = torch.tensor(list(map(lambda nums: (nums==1).nonzero().item(), labels))).to(GPU_DEVICE)\n                count += pred_indexes.shape[0]\n                batch_df = pd.DataFrame()\n                batch_df[\"image_id\"] = ids\n                batch_df[\"label\"] = pd.Series(pred_indexes.detach().cpu().numpy())\n                if len(label_df) == 0:\n                    label_df = batch_df\n                else:\n                    label_df = pd.concat([label_df, batch_df])              \n                if labels is not None and calc_acc:\n                    correct += int(torch.where(pred_indexes == labels, 1.0, 0.0).sum())                    \n            acc = correct / count\n            epoch_loss = epoch_loss / (len(loader))\n            label_df.reset_index(drop=True, inplace=True)\n            ret = ClassifierOutput(\n                count=count,\n                acc=acc,\n                loss=epoch_loss,\n                df=label_df\n            )            \n            return ret","metadata":{"papermill":{"duration":0.031072,"end_time":"2023-01-03T13:40:02.389767","exception":false,"start_time":"2023-01-03T13:40:02.358695","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-08T04:24:50.500204Z","iopub.execute_input":"2023-01-08T04:24:50.500804Z","iopub.status.idle":"2023-01-08T04:24:50.524536Z","shell.execute_reply.started":"2023-01-08T04:24:50.500768Z","shell.execute_reply":"2023-01-08T04:24:50.523597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Run","metadata":{"papermill":{"duration":0.004264,"end_time":"2023-01-03T13:40:02.398965","exception":false,"start_time":"2023-01-03T13:40:02.394701","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import torch.optim as optim\nfrom sklearn.model_selection import KFold\nfrom torch.utils.data.dataset import Subset\n\nindex = 3\nname = f\"run-{index}\"\nnotes = \"\"\n\ntrain_path = os.path.join(config.base_input_path, \"train.csv\")\ntrain_img_dir_path = os.path.join(config.base_input_path, \"train_images\")\ndf = pd.read_csv(train_path)\n\ndf.label.value_counts().plot(kind=\"bar\")","metadata":{"papermill":{"duration":1.808145,"end_time":"2023-01-03T13:40:04.211843","exception":false,"start_time":"2023-01-03T13:40:02.403698","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-08T04:24:50.527262Z","iopub.execute_input":"2023-01-08T04:24:50.527521Z","iopub.status.idle":"2023-01-08T04:24:50.738510Z","shell.execute_reply.started":"2023-01-08T04:24:50.527497Z","shell.execute_reply":"2023-01-08T04:24:50.737519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport torch\nimport os\n\ndef set_seed(seed):\n    np.random.seed(seed)\n    random_state = np.random.RandomState(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    return random_state","metadata":{"execution":{"iopub.status.busy":"2023-01-08T04:24:50.740085Z","iopub.execute_input":"2023-01-08T04:24:50.740697Z","iopub.status.idle":"2023-01-08T04:24:50.748744Z","shell.execute_reply.started":"2023-01-08T04:24:50.740658Z","shell.execute_reply":"2023-01-08T04:24:50.747657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"set_seed(3407)","metadata":{"papermill":{"duration":0.021793,"end_time":"2023-01-03T13:40:04.239622","exception":false,"start_time":"2023-01-03T13:40:04.217829","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-08T04:24:50.751090Z","iopub.execute_input":"2023-01-08T04:24:50.751971Z","iopub.status.idle":"2023-01-08T04:24:50.759918Z","shell.execute_reply.started":"2023-01-08T04:24:50.751935Z","shell.execute_reply":"2023-01-08T04:24:50.758596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"k_fold = KFold(n_splits=2)\ntrain_dataset = CassavaDataset(\n    annotation_path=train_path,\n    img_dir_path=train_img_dir_path,\n    img_transforms=transforms.Compose([\n        transforms.ToTensor(),    \n        transforms.RandomHorizontalFlip(),\n        transforms.RandomVerticalFlip(),\n        transforms.ColorJitter(brightness=0.5, contrast=0.5, saturation=0.5),\n        transforms.Resize((config.resized_height, config.resized_width)),\n        transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n    ]),\n)\nval_dataset = CassavaDataset(\n    annotation_path=train_path,\n    img_dir_path=train_img_dir_path,\n    img_transforms=transforms.Compose([\n        transforms.ToTensor(),\n        transforms.Resize((config.resized_height, config.resized_width)),\n        transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),        \n    ]),\n)","metadata":{"execution":{"iopub.status.busy":"2023-01-08T04:24:50.761912Z","iopub.execute_input":"2023-01-08T04:24:50.762568Z","iopub.status.idle":"2023-01-08T04:24:50.793494Z","shell.execute_reply.started":"2023-01-08T04:24:50.762534Z","shell.execute_reply":"2023-01-08T04:24:50.792652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.cuda.amp.grad_scaler import GradScaler\n\nmodel = CassavaNetwork(config=config)\nmodel = model.to(torch.device(\"cuda\"))\nmodel = nn.DataParallel(model)\n\nif config.is_kaggle_notebook:\n    pretrained_weight_path = \"/kaggle/input/weights/epoch_15_base_efficientnetv2_rw_t.pth\"\n    model.load_state_dict(torch.load(pretrained_weight_path))\nelse:\n    for fold, (train_index, val_index) in enumerate(k_fold.split(df)):\n        train_subset = Subset(train_dataset, train_index)\n        train_dataloader = DataLoader(\n            train_subset,\n            batch_size=config.batch_size,\n            num_workers=config.num_workers,\n            shuffle=True,\n        )\n        val_subset = Subset(val_dataset, val_index)\n        val_dataloader = DataLoader(\n            val_subset,\n            batch_size=config.batch_size,\n            num_workers=config.num_workers,\n            shuffle=False,\n        )    \n\n        optimizer = optim.Adam(model.parameters(), lr=config.lr)\n        scheduler = optim.lr_scheduler.ExponentialLR(optimizer, gamma=0.95)\n        loss_function = nn.CrossEntropyLoss()\n\n        activate_function = nn.Softmax(dim=1)\n        \n        scaler = GradScaler()\n\n        trainer = Trainer(\n            model,\n            config,\n            optimizer,\n            scheduler,\n            loss_function,\n            train_dataloader,\n            scaler=scaler,\n        )\n        classifier = Classifier(\n            model,\n            activate_function,\n            train_dataloader,\n            phase=\"train\",\n            phase_input=ClassifierTrainInput(),\n        )\n\n        val_classifier = Classifier(\n            model,\n            activate_function,\n            val_dataloader,\n            phase=\"test\",\n            phase_input=ClassifierTestInput(),\n        )\n\n\n        for epoch in tqdm(range(config.n_epochs)):\n            epoch += 1\n            train_out = trainer.advance()\n            print(f\"epoch: {epoch}, train loss: {train_out.loss}\")\n            # train_infer_out = classifier.infer()\n            # print(f\"epoch: {epoch}, train acc: {train_infer_out.acc}\")\n            val_infer_out = val_classifier.infer()\n            print(f\"epoch: {epoch}, val acc: {val_infer_out.acc}\")\n            if epoch == config.n_epochs:\n                trainer.save_model()","metadata":{"papermill":{"duration":7267.6653,"end_time":"2023-01-03T15:41:11.909946","exception":false,"start_time":"2023-01-03T13:40:04.244646","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-08T04:24:50.796446Z","iopub.execute_input":"2023-01-08T04:24:50.796699Z","iopub.status.idle":"2023-01-08T04:24:51.511715Z","shell.execute_reply.started":"2023-01-08T04:24:50.796675Z","shell.execute_reply":"2023-01-08T04:24:51.510773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submission","metadata":{"papermill":{"duration":0.847939,"end_time":"2023-01-03T15:41:13.562158","exception":false,"start_time":"2023-01-03T15:41:12.714219","status":"completed"},"tags":[]}},{"cell_type":"code","source":"test_data_img_dir_path = os.path.join(config.base_input_path, \"test_images\")\ntest_dataset = CassavaDataset(\n    annotation_path=None,\n    img_dir_path=test_data_img_dir_path,\n    img_transforms=transforms.Compose([\n        transforms.ToTensor(),\n        transforms.Resize((config.resized_height, config.resized_width)),\n        transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n    ]),\n)\n\nactivate_function = nn.Softmax(dim=1)\n\ntest_dataloader = DataLoader(\n    test_dataset,\n    batch_size=config.batch_size,\n    num_workers=config.num_workers,\n)\n\ntest_classifier = Classifier(\n    model,\n    activate_function,\n    test_dataloader,\n    phase=\"test\",\n    phase_input=ClassifierTrainInput(),\n)\n\nout = test_classifier.infer(calc_acc=False, verbose=True)\nout.df.to_csv(\"submission.csv\", index=False)","metadata":{"papermill":{"duration":1.48478,"end_time":"2023-01-03T15:41:15.840399","exception":false,"start_time":"2023-01-03T15:41:14.355619","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-08T04:24:51.516184Z","iopub.execute_input":"2023-01-08T04:24:51.518783Z","iopub.status.idle":"2023-01-08T04:24:52.103428Z","shell.execute_reply.started":"2023-01-08T04:24:51.518743Z","shell.execute_reply":"2023-01-08T04:24:52.102344Z"},"trusted":true},"execution_count":null,"outputs":[]}]}