{"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":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Import Required Libraries 📚</h1></span>","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport math\nimport copy\nimport time\nimport random\nimport glob\nfrom PIL import Image\nfrom matplotlib import pyplot as plt\n\n# For data manipulation\nimport numpy as np\nimport pandas as pd\n\n# Pytorch Imports\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda import amp\nimport torchvision\n\n# Utils\nimport joblib\nfrom tqdm import tqdm\nfrom collections import defaultdict\n\n# Sklearn Imports\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import StratifiedKFold\n\n# For Image Models\nimport timm\n\n# Albumentations for augmentations\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# For colored terminal text\nfrom colorama import Fore, Back, Style\nb_ = Fore.BLUE\nsr_ = Style.RESET_ALL\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n# For descriptive error messages\nos.environ['CUDA_LAUNCH_BLOCKING'] = \"1\"","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:58:44.363163Z","iopub.execute_input":"2023-11-14T05:58:44.363988Z","iopub.status.idle":"2023-11-14T05:58:44.373475Z","shell.execute_reply.started":"2023-11-14T05:58:44.36395Z","shell.execute_reply":"2023-11-14T05:58:44.372366Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Training Configuration ⚙️</h1></span>","metadata":{}},{"cell_type":"code","source":"CONFIG = {\n    \"seed\": 42,\n    \"img_size\": 2048,\n    \"model_name\": \"tf_efficientnet_b0_ns\",\n    \"num_classes\": 5,\n    \"valid_batch_size\": 1,\n    \"device\": torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\"),\n}","metadata":{"execution":{"iopub.status.busy":"2023-11-14T06:01:50.156704Z","iopub.execute_input":"2023-11-14T06:01:50.157019Z","iopub.status.idle":"2023-11-14T06:01:50.165528Z","shell.execute_reply.started":"2023-11-14T06:01:50.156969Z","shell.execute_reply":"2023-11-14T06:01:50.164557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Set Seed for Reproducibility</h1></span>","metadata":{}},{"cell_type":"code","source":"def set_seed(seed=42):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed(CONFIG['seed'])","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:59:32.750087Z","iopub.execute_input":"2023-11-14T05:59:32.751083Z","iopub.status.idle":"2023-11-14T05:59:32.757391Z","shell.execute_reply.started":"2023-11-14T05:59:32.751027Z","shell.execute_reply":"2023-11-14T05:59:32.756366Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT_DIR = '/kaggle/input/UBC-OCEAN'\nTEST_DIR = '/kaggle/input/UBC-OCEAN/test_thumbnails'\nTRAIN_DIR = '/kaggle/input/UBC-OCEAN/train_thumbnails'\nALT_TRAIN_DIR = '/kaggle/input/UBC-OCEAN/train_images'\n\nLABEL_ENCODER_BIN = \"/kaggle/input/ubc-efficienetnetb0-fold1of10-2048pix-thumbnails/label_encoder.pkl\"\nBEST_WEIGHT = \"/kaggle/input/ubc-efficienetnetb0-fold1of10-2048pix-thumbnails/Recall0.9178_Acc0.9437_Loss0.1685_epoch9.bin\"","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:59:35.043291Z","iopub.execute_input":"2023-11-14T05:59:35.044026Z","iopub.status.idle":"2023-11-14T05:59:35.048763Z","shell.execute_reply.started":"2023-11-14T05:59:35.043969Z","shell.execute_reply":"2023-11-14T05:59:35.047726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_test_file_path(image_id):\n    return f\"{TEST_DIR}/{image_id}_thumbnail.png\"","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:59:38.315313Z","iopub.execute_input":"2023-11-14T05:59:38.315684Z","iopub.status.idle":"2023-11-14T05:59:38.320534Z","shell.execute_reply.started":"2023-11-14T05:59:38.315654Z","shell.execute_reply":"2023-11-14T05:59:38.319454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_file_path(image_id):\n    if os.path.exists(f\"{TRAIN_DIR}/{image_id}_thumbnail.png\"):\n        return f\"{TRAIN_DIR}/{image_id}_thumbnail.png\"\n    else:\n        return f\"{ALT_TRAIN_DIR}/{image_id}.png\"","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:59:40.076619Z","iopub.execute_input":"2023-11-14T05:59:40.077506Z","iopub.status.idle":"2023-11-14T05:59:40.082215Z","shell.execute_reply.started":"2023-11-14T05:59:40.077468Z","shell.execute_reply":"2023-11-14T05:59:40.081262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"d = {'EC':1,'HGSC':2,'LGSC':3,'MC':4,'CC':0}\ndef get_labels(label_name):\n    return d[label_name]\n    ","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:59:41.783637Z","iopub.execute_input":"2023-11-14T05:59:41.784447Z","iopub.status.idle":"2023-11-14T05:59:41.788842Z","shell.execute_reply.started":"2023-11-14T05:59:41.78441Z","shell.execute_reply":"2023-11-14T05:59:41.78797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Read the Data 📖</h1>","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(f\"{ROOT_DIR}/test.csv\")\ndf['file_path'] = df['image_id'].apply(get_test_file_path)\ndf['label'] = 0 # dummy\ndf","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:59:44.764608Z","iopub.execute_input":"2023-11-14T05:59:44.76533Z","iopub.status.idle":"2023-11-14T05:59:44.785158Z","shell.execute_reply.started":"2023-11-14T05:59:44.765294Z","shell.execute_reply":"2023-11-14T05:59:44.784239Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_t = pd.read_csv(f\"{ROOT_DIR}/train.csv\")\nindex = [i for i,a in enumerate(df_t[\"is_tma\"].values) if a == True]\ndf_t\ndf_t = df_t.drop(index)\ndf_t['file_path'] = df_t['image_id'].apply(get_train_file_path)\ndf_t['label'] = df_t['label'].apply(get_labels)\n","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:59:47.406959Z","iopub.execute_input":"2023-11-14T05:59:47.407895Z","iopub.status.idle":"2023-11-14T05:59:47.998469Z","shell.execute_reply.started":"2023-11-14T05:59:47.407856Z","shell.execute_reply":"2023-11-14T05:59:47.99736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub = pd.read_csv(f\"{ROOT_DIR}/sample_submission.csv\")\ndf_sub","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:59:50.766532Z","iopub.execute_input":"2023-11-14T05:59:50.7669Z","iopub.status.idle":"2023-11-14T05:59:50.779179Z","shell.execute_reply.started":"2023-11-14T05:59:50.766868Z","shell.execute_reply":"2023-11-14T05:59:50.778226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder = joblib.load( LABEL_ENCODER_BIN )\n","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:59:52.984706Z","iopub.execute_input":"2023-11-14T05:59:52.985642Z","iopub.status.idle":"2023-11-14T05:59:52.992894Z","shell.execute_reply.started":"2023-11-14T05:59:52.985603Z","shell.execute_reply":"2023-11-14T05:59:52.991997Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_cropped_images(file_path, image_id, label_t,th_area = 1000):\n    image = Image.open(file_path)\n    # Aspect ratio\n    as_ratio = image.size[0] / image.size[1]\n    \n    sxs, exs, sys, eys = [],[],[],[]\n    if as_ratio >= 1.5:\n        # Crop\n        mask = np.max( np.array(image) > 0, axis=-1 ).astype(np.uint8)\n        retval, labels = cv2.connectedComponents(mask)\n        if retval >= as_ratio:\n            x, y = np.meshgrid( np.arange(image.size[0]), np.arange(image.size[1]) )\n            for label in range(1, retval):\n                area = np.sum(labels == label)\n                if area < th_area:\n                    continue\n                xs, ys= x[ labels == label ], y[ labels == label ]\n                sx, ex = np.min(xs), np.max(xs)\n                cx = (sx + ex) // 2\n                crop_size = image.size[1]\n                sx = max(0, cx-crop_size//2)\n                ex = min(sx + crop_size - 1, image.size[0]-1)\n                sx = ex - crop_size + 1\n                sy, ey = 0, image.size[1]-1\n                sxs.append(sx)\n                exs.append(ex)\n                sys.append(sy)\n                eys.append(ey)\n        else:\n            crop_size = image.size[1]\n            for i in range(int(as_ratio)):\n                sxs.append( i * crop_size )\n                exs.append( (i+1) * crop_size - 1 )\n                sys.append( 0 )\n                eys.append( crop_size - 1 )\n    else:\n        # Not Crop (entire image)\n        sxs, exs, sys, eys = [0,],[image.size[0]-1],[0,],[image.size[1]-1]\n\n    df_crop = pd.DataFrame()\n    df_crop[\"image_id\"] = [image_id] * len(sxs)\n    df_crop[\"file_path\"] = [file_path] * len(sxs)\n    df_crop[\"sx\"] = sxs\n    df_crop[\"ex\"] = exs\n    df_crop[\"sy\"] = sys\n    df_crop[\"ey\"] = eys\n    df_crop[\"label\"] = [label_t] * len(sxs)\n    return df_crop","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:59:56.379914Z","iopub.execute_input":"2023-11-14T05:59:56.380335Z","iopub.status.idle":"2023-11-14T05:59:56.397459Z","shell.execute_reply.started":"2023-11-14T05:59:56.380302Z","shell.execute_reply":"2023-11-14T05:59:56.396438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 7\ndfs = []\nfor (file_path, image_id,label) in zip(df_t[\"file_path\"], df_t[\"image_id\"],df_t['label']):\n#     print([label]*3)\n    dfs.append( get_cropped_images(file_path, image_id,label) )\n\ndf_train = pd.concat(dfs)\ndf_train = df_train.drop_duplicates(subset=[\"image_id\", \"sx\", \"ex\", \"sy\", \"ey\"]).reset_index(drop=True)\ndf_train","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:59:59.181063Z","iopub.execute_input":"2023-11-14T05:59:59.181417Z","iopub.status.idle":"2023-11-14T06:01:49.665185Z","shell.execute_reply.started":"2023-11-14T05:59:59.181392Z","shell.execute_reply":"2023-11-14T06:01:49.664065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndfs = []\nfor (file_path, image_id,label) in zip(df[\"file_path\"], df[\"image_id\"],df['label']):\n    dfs.append( get_cropped_images(file_path, image_id,label) )\n\ndf_crop = pd.concat(dfs)\ndf_crop[\"label\"] = 0 # dummy\ndf_crop","metadata":{"execution":{"iopub.status.busy":"2023-11-14T06:01:49.667694Z","iopub.execute_input":"2023-11-14T06:01:49.668275Z","iopub.status.idle":"2023-11-14T06:01:50.14133Z","shell.execute_reply.started":"2023-11-14T06:01:49.668231Z","shell.execute_reply":"2023-11-14T06:01:50.140399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_crop = df_crop.drop_duplicates(subset=[\"image_id\", \"sx\", \"ex\", \"sy\", \"ey\"]).reset_index(drop=True)\ndf_crop","metadata":{"execution":{"iopub.status.busy":"2023-11-14T06:01:50.142574Z","iopub.execute_input":"2023-11-14T06:01:50.142891Z","iopub.status.idle":"2023-11-14T06:01:50.154916Z","shell.execute_reply.started":"2023-11-14T06:01:50.142849Z","shell.execute_reply":"2023-11-14T06:01:50.153798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Dataset Class</h1></span>","metadata":{}},{"cell_type":"code","source":"class UBCDataset(Dataset):\n    def __init__(self, df, transforms=None):\n        self.df = df\n        self.file_names = df['file_path'].values\n        self.labels = df['label'].values\n        self.transforms = transforms\n        self.sxs = df[\"sx\"].values\n        self.exs = df[\"ex\"].values\n        self.sys = df[\"sy\"].values\n        self.eys = df[\"ey\"].values\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        img_path = self.file_names[index]\n        sx = self.sxs[index]\n        ex = self.exs[index]\n        sy = self.sys[index]\n        ey = self.eys[index]\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        label = self.labels[index]\n        \n        img = img[ sy:ey, sx:ex, : ]\n        \n        if self.transforms:\n            img = self.transforms(image=img)[\"image\"]\n            \n        return {\n            'image': img,\n            'label': torch.tensor(label, dtype=torch.long)\n        }","metadata":{"execution":{"iopub.status.busy":"2023-11-14T06:03:53.006405Z","iopub.execute_input":"2023-11-14T06:03:53.007327Z","iopub.status.idle":"2023-11-14T06:03:53.016862Z","shell.execute_reply.started":"2023-11-14T06:03:53.007288Z","shell.execute_reply":"2023-11-14T06:03:53.01585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Augmentations</h1></span>","metadata":{}},{"cell_type":"code","source":"data_transforms = {\n    \"valid\": A.Compose([\n        A.Resize(CONFIG['img_size'], CONFIG['img_size']),\n        A.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225], \n                max_pixel_value=255.0, \n                p=1.0\n            ),\n        ToTensorV2()], p=1.)\n}","metadata":{"execution":{"iopub.status.busy":"2023-11-14T06:03:57.297835Z","iopub.execute_input":"2023-11-14T06:03:57.298703Z","iopub.status.idle":"2023-11-14T06:03:57.303974Z","shell.execute_reply.started":"2023-11-14T06:03:57.298666Z","shell.execute_reply":"2023-11-14T06:03:57.303119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">GeM Pooling</h1></span>","metadata":{}},{"cell_type":"code","source":"class GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super(GeM, self).__init__()\n        self.p = nn.Parameter(torch.ones(1)*p)\n        self.eps = eps\n\n    def forward(self, x):\n        return self.gem(x, p=self.p, eps=self.eps)\n        \n    def gem(self, x, p=3, eps=1e-6):\n        return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(1./p)\n        \n    def __repr__(self):\n        return self.__class__.__name__ + \\\n                '(' + 'p=' + '{:.4f}'.format(self.p.data.tolist()[0]) + \\\n                ', ' + 'eps=' + str(self.eps) + ')'","metadata":{"execution":{"iopub.status.busy":"2023-11-14T06:03:59.447283Z","iopub.execute_input":"2023-11-14T06:03:59.447656Z","iopub.status.idle":"2023-11-14T06:03:59.455978Z","shell.execute_reply.started":"2023-11-14T06:03:59.447624Z","shell.execute_reply":"2023-11-14T06:03:59.454915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Create Model</h1></span>","metadata":{}},{"cell_type":"code","source":"class UBCModel(nn.Module):\n    def __init__(self, model_name, num_classes, pretrained=False, checkpoint_path=None):\n        super(UBCModel, self).__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n\n        in_features = self.model.classifier.in_features\n        self.model.classifier = nn.Identity()\n        self.model.global_pool = nn.Identity()\n        self.pooling = GeM()\n        self.linear = nn.Linear(in_features, num_classes)\n        self.softmax = nn.Softmax(dim=1)\n\n    def forward(self, images):\n        features = self.model(images)\n        pooled_features = self.pooling(features).flatten(1)\n        output = self.linear(pooled_features)\n        return output\n\n    \nmodel = UBCModel(CONFIG['model_name'], CONFIG['num_classes'])\nmodel.load_state_dict(torch.load( BEST_WEIGHT ))\nmodel.to(CONFIG['device']);","metadata":{"execution":{"iopub.status.busy":"2023-11-14T06:04:01.414718Z","iopub.execute_input":"2023-11-14T06:04:01.415129Z","iopub.status.idle":"2023-11-14T06:04:01.654093Z","shell.execute_reply.started":"2023-11-14T06:04:01.415095Z","shell.execute_reply":"2023-11-14T06:04:01.653243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Prepare Dataloaders</span>","metadata":{}},{"cell_type":"code","source":"test_dataset = UBCDataset(df_crop, transforms=data_transforms[\"valid\"])\ntest_loader = DataLoader(test_dataset, batch_size=CONFIG['valid_batch_size'], \n                          num_workers=2, shuffle=False, pin_memory=False)","metadata":{"execution":{"iopub.status.busy":"2023-11-14T06:04:03.60665Z","iopub.execute_input":"2023-11-14T06:04:03.607043Z","iopub.status.idle":"2023-11-14T06:04:03.613035Z","shell.execute_reply.started":"2023-11-14T06:04:03.607009Z","shell.execute_reply":"2023-11-14T06:04:03.612037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#df_train\nfrom torch.optim import Adam\n \n# Define the loss function with Classification Cross-Entropy loss and an optimizer with Adam optimizer\nloss_fn = nn.CrossEntropyLoss()\noptimizer = Adam(model.parameters(), lr=0.001, weight_decay=0.0001)","metadata":{"execution":{"iopub.status.busy":"2023-11-14T06:04:06.097921Z","iopub.execute_input":"2023-11-14T06:04:06.098847Z","iopub.status.idle":"2023-11-14T06:04:06.107948Z","shell.execute_reply.started":"2023-11-14T06:04:06.09881Z","shell.execute_reply":"2023-11-14T06:04:06.106861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 搭建训练数据集\ntrain_dataset = UBCDataset(df_train, transforms=data_transforms[\"valid\"])\ntrain_loader = DataLoader(train_dataset, batch_size=CONFIG['valid_batch_size'], \n                          num_workers=2, shuffle=False, pin_memory=False)","metadata":{"execution":{"iopub.status.busy":"2023-11-14T06:04:09.150829Z","iopub.execute_input":"2023-11-14T06:04:09.151576Z","iopub.status.idle":"2023-11-14T06:04:09.157033Z","shell.execute_reply.started":"2023-11-14T06:04:09.151537Z","shell.execute_reply":"2023-11-14T06:04:09.155959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 训练模型\nnum_epochs = 10\nfrom torch.autograd import Variable\n\n\nbar = tqdm(enumerate(train_loader), total=len(train_loader))\n    \nfor epoch in range(num_epochs):  # loop over the dataset multiple times\n    running_loss = 0.0\n    running_acc = 0.0\n\n    for step, data in bar:        \n        images = data['image'].to(CONFIG[\"device\"], dtype=torch.float) \n\n        # get the inputs\n        labels = data['label'].to(CONFIG[\"device\"],) \n        \n        # zero the parameter gradients\n        optimizer.zero_grad()\n        # predict classes using images from the training set\n        print(images.shape)\n        outputs = model(images)\n        # compute the loss based on model output and real labels\n        loss = loss_fn(outputs, labels)\n        # backpropagate the loss\n        loss.backward()\n        # adjust parameters based on the calculated gradients\n        optimizer.step()\n        \n        del images\n        del labels\n        # Let's print statistics for every 1,000 images\n#         running_loss += loss.item()     # extract the loss value\n#         if i % 1000 == 999:    \n#             # print every 1000 (twice per epoch) \n#             print('[%d, %5d] loss: %.3f' %\n#                   (epoch + 1, i + 1, running_loss / 1000))\n#             # zero the loss\n#             running_loss = 0.0","metadata":{"execution":{"iopub.status.busy":"2023-11-14T06:04:13.711906Z","iopub.execute_input":"2023-11-14T06:04:13.71281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Start Inference</span>","metadata":{}},{"cell_type":"code","source":"preds = []\nwith torch.no_grad():\n    bar = tqdm(enumerate(test_loader), total=len(test_loader))\n    for step, data in bar:        \n        images = data['image'].to(CONFIG[\"device\"], dtype=torch.float)        \n        batch_size = images.size(0)\n        outputs = model(images)\n        outputs = model.softmax(outputs)\n        preds.append( outputs.detach().cpu().numpy() )\n\npreds = np.vstack(preds)\nprint(preds.shape)\nprint(preds)","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:25:31.221722Z","iopub.execute_input":"2023-11-14T05:25:31.222213Z","iopub.status.idle":"2023-11-14T05:25:31.603693Z","shell.execute_reply.started":"2023-11-14T05:25:31.222165Z","shell.execute_reply":"2023-11-14T05:25:31.602324Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i in range(preds.shape[-1]):\n    df_crop[f\"cat{i}\"] = preds[:, i]\n\ndict_label = {}\nfor image_id, gdf in df_crop.groupby(\"image_id\"):\n    dict_label[image_id] = np.argmax( gdf[ [f\"cat{i}\" for i in range(preds.shape[-1])] ].values.max(axis=0) )\n    #dict_label[image_id] = np.argmax( gdf[ [f\"cat{i}\" for i in range(preds.shape[-1])] ].values.mean(axis=0) )\npreds = np.array( [ dict_label[image_id] for image_id in df[\"image_id\"].values ] )","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:25:31.605716Z","iopub.execute_input":"2023-11-14T05:25:31.606215Z","iopub.status.idle":"2023-11-14T05:25:31.627779Z","shell.execute_reply.started":"2023-11-14T05:25:31.606168Z","shell.execute_reply":"2023-11-14T05:25:31.626458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_labels = encoder.inverse_transform(preds)\ndf_sub[\"label\"] = pred_labels\ndf_sub.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:25:31.629222Z","iopub.execute_input":"2023-11-14T05:25:31.629678Z","iopub.status.idle":"2023-11-14T05:25:31.638153Z","shell.execute_reply.started":"2023-11-14T05:25:31.629639Z","shell.execute_reply":"2023-11-14T05:25:31.637239Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub","metadata":{"execution":{"iopub.status.busy":"2023-11-14T05:25:31.63952Z","iopub.execute_input":"2023-11-14T05:25:31.639829Z","iopub.status.idle":"2023-11-14T05:25:31.653003Z","shell.execute_reply.started":"2023-11-14T05:25:31.639801Z","shell.execute_reply":"2023-11-14T05:25:31.652046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}}]}