{"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":"#### Reference:\n\n1. https://www.kaggle.com/code/itsuki9180/a-segmentation-is-in-reverse-order\n1. https://www.kaggle.com/code/samuelcortinhas/rnsa-3d-model-train-pytorch\n1. https://www.kaggle.com/code/andradaolteanu/rsna-fracture-detection-dicom-images-explore#2.-Image-Data-%5B.dcm%5D\n1. https://www.kaggle.com/code/weixinxu/submit-baseline\n1. https://www.kaggle.com/code/mlwhiz/bilstm-pytorch-and-keras\n\n#### Modeling Reference:\n1. http://dx.doi.org/10.3174/ajnr.A7094\n1. https://blog.devgenius.io/resnet50-6b42934db431\n1. https://www.kaggle.com/code/samuelcortinhas/rnsa-3d-model-train-pytorch#Torch-dataloaders","metadata":{"id":"Ual0RqSlW4F0"}},{"cell_type":"markdown","source":"## Downloads","metadata":{"id":"_qtBuug5W4F3"}},{"cell_type":"code","source":"INTERNET = True\n\nif INTERNET == True:\n    !python --version\n\n    !pip install monai\n    !pip install -q segmentation_models_pytorch","metadata":{"id":"yu3ndgUrW4F4","outputId":"8deeeda2-d82e-45e3-e279-287271b7ba53","execution":{"iopub.status.busy":"2023-05-16T12:33:04.985086Z","iopub.execute_input":"2023-05-16T12:33:04.985793Z","iopub.status.idle":"2023-05-16T12:33:38.13167Z","shell.execute_reply.started":"2023-05-16T12:33:04.98575Z","shell.execute_reply":"2023-05-16T12:33:38.130254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Import Libraries","metadata":{"id":"LA_kCUmHW4F5"}},{"cell_type":"code","source":"import os\nimport sys\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nfrom PIL import Image\nimport cv2\nimport re\nimport gc\nfrom tqdm import tqdm\nimport math\n\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Rectangle\n\nimport skimage.transform as skTrans\nfrom skimage import exposure\n\nimport albumentations as alb\nfrom albumentations.pytorch import ToTensorV2\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torch.optim.lr_scheduler import OneCycleLR, CosineAnnealingWarmRestarts\n\nimport tensorflow as tf\n\nfrom monai.transforms import Resize\nimport monai.transforms as transforms\nfrom monai.networks.nets import resnet18, resnet101\n\nimport segmentation_models_pytorch as smp\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.model_selection import KFold, StratifiedKFold\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"id":"q54eI0HrZ8p1","execution":{"iopub.status.busy":"2023-05-16T12:33:38.135082Z","iopub.execute_input":"2023-05-16T12:33:38.135588Z","iopub.status.idle":"2023-05-16T12:33:55.004536Z","shell.execute_reply.started":"2023-05-16T12:33:38.135535Z","shell.execute_reply":"2023-05-16T12:33:55.003346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config","metadata":{"id":"wxZqi55oW4F7"}},{"cell_type":"code","source":"SEED = 1927550\nIMG_SIZE = 352\nBATCH = 5\nEPOCH = 50\nCLASS = 7\nhidden1 = 128\nhidden2 = 64\nKAGGLE = True\nchannel_3 = True\ntot_slice = 30\nbest_loss = 1\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\nencoder_backbone = 'resnet101'\n\ntrainlosslog = []\ntrainacclog = []\nvalidlosslog = []\nvalidacclog = []","metadata":{"id":"OZCsaDR3W4F7","execution":{"iopub.status.busy":"2023-05-16T12:33:55.01001Z","iopub.execute_input":"2023-05-16T12:33:55.016792Z","iopub.status.idle":"2023-05-16T12:33:55.094002Z","shell.execute_reply.started":"2023-05-16T12:33:55.016756Z","shell.execute_reply":"2023-05-16T12:33:55.092956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if KAGGLE:\n    work_path = '/kaggle/input'\nelse:\n    work_path = '/content/drive/MyDrive/Colab_Notebooks'\n\nbase_path = f'{work_path}/stage-2-preprocessed-zip'\ntrain_df = pd.read_csv(f'{work_path}/rsna-2022-cervical-spine-fracture-detection/train.csv')\n\n#stage2_prep_list = os.listdir(f'{base_path}/stage2-prep')\nstage2_prep_list = os.listdir(base_path)\nstage2_prep_list = [os.path.splitext(os.path.basename(prep_path))[0] for prep_path in stage2_prep_list]\nstage2_prep_uid = [prep_path.split('_')[0] for prep_path in stage2_prep_list]\nvoxel_df = pd.DataFrame(list(zip(stage2_prep_list, stage2_prep_uid)), columns=['id', 'StudyInstanceUID'])\n\nkf = KFold(5)\nfor fold, (train_idx, valid_idx) in enumerate(kf.split(voxel_df, voxel_df)):\n    voxel_df.loc[valid_idx, 'fold'] = fold\n\ndf_train = voxel_df[voxel_df['fold'] != fold].reset_index(drop=True)\ndf_valid = voxel_df[voxel_df['fold'] == fold].reset_index(drop=True)","metadata":{"id":"ZG2rrXStW4F8","execution":{"iopub.status.busy":"2023-05-16T12:33:55.100392Z","iopub.execute_input":"2023-05-16T12:33:55.104197Z","iopub.status.idle":"2023-05-16T12:33:55.329907Z","shell.execute_reply.started":"2023-05-16T12:33:55.104102Z","shell.execute_reply":"2023-05-16T12:33:55.328647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seg_revert = [\n    '1.2.826.0.1.3680043.1363',\n    '1.2.826.0.1.3680043.20120',\n    '1.2.826.0.1.3680043.2243',\n    '1.2.826.0.1.3680043.24606',\n    '1.2.826.0.1.3680043.32071'\n]","metadata":{"execution":{"iopub.status.busy":"2023-05-16T12:33:55.332841Z","iopub.execute_input":"2023-05-16T12:33:55.333759Z","iopub.status.idle":"2023-05-16T12:33:55.339198Z","shell.execute_reply.started":"2023-05-16T12:33:55.333698Z","shell.execute_reply":"2023-05-16T12:33:55.337863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## EDA","metadata":{}},{"cell_type":"code","source":"voxel_df","metadata":{"execution":{"iopub.status.busy":"2023-05-16T12:33:55.340946Z","iopub.execute_input":"2023-05-16T12:33:55.341741Z","iopub.status.idle":"2023-05-16T12:33:55.365755Z","shell.execute_reply.started":"2023-05-16T12:33:55.341692Z","shell.execute_reply":"2023-05-16T12:33:55.364657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# check the image size\n# find the most sized image or the average\nimg_size = []\nfor idx in range(len(voxel_df)):\n    loc = voxel_df.loc[idx]\n    id_ = loc['id']\n    path = f'{base_path}/{id_}.npz'\n    \n    img = np.load(path)['arr_0']\n    img_size.append(img.shape[1]) # equivalent image size for both width and height","metadata":{"execution":{"iopub.status.busy":"2023-05-16T12:33:55.367393Z","iopub.execute_input":"2023-05-16T12:33:55.367891Z","iopub.status.idle":"2023-05-16T12:35:43.725785Z","shell.execute_reply.started":"2023-05-16T12:33:55.367848Z","shell.execute_reply":"2023-05-16T12:35:43.724721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"appr_img_size = [round(size/10)*10 for size in img_size]\nunique, count = np.unique(appr_img_size, return_counts=True)\n\nplt.bar(unique, count)\n\ntotal = 0\ntotal_count = 0\nfor idx in range(len(unique)):\n    total += unique[idx] * count[idx]\n    total_count += count[idx]\n\nmean = total / total_count\nprint('Mean value:',mean)","metadata":{"execution":{"iopub.status.busy":"2023-05-16T12:35:43.72711Z","iopub.execute_input":"2023-05-16T12:35:43.7282Z","iopub.status.idle":"2023-05-16T12:35:44.029333Z","shell.execute_reply.started":"2023-05-16T12:35:43.728159Z","shell.execute_reply":"2023-05-16T12:35:44.02828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Useful Functions","metadata":{"id":"r7lLzzm6W4F8"}},{"cell_type":"code","source":"# apply seed\ndef seed_everything(seed):\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)","metadata":{"id":"a_KZzLgBW4F9","execution":{"iopub.status.busy":"2023-05-16T12:35:44.031078Z","iopub.execute_input":"2023-05-16T12:35:44.03147Z","iopub.status.idle":"2023-05-16T12:35:44.039513Z","shell.execute_reply.started":"2023-05-16T12:35:44.031431Z","shell.execute_reply":"2023-05-16T12:35:44.038276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_transforms = {\n    'train': alb.Compose([\n                alb.Resize(IMG_SIZE, IMG_SIZE),\n                alb.HorizontalFlip(p=0.5),\n                alb.VerticalFlip(p=0.5),\n                alb.Transpose(p=0.5),\n                alb.RandomBrightness(limit=0.1, p=0.7),\n                alb.ShiftScaleRotate(shift_limit=0.3, scale_limit=0.3, rotate_limit=45, border_mode=4, p=0.7),\n                alb.OneOf([\n                    alb.MotionBlur(blur_limit=3),\n                    alb.MedianBlur(blur_limit=3),\n                    alb.GaussianBlur(blur_limit=3),\n                    alb.GaussNoise(var_limit=(3.0, 9.0))\n                ], p=0.5),\n                alb.OneOf([\n                    alb.GridDistortion(num_steps=5, distort_limit=1.),\n                    alb.OpticalDistortion(distort_limit=1.)\n                ], p=0.5),\n                alb.Sharpen(alpha=(0.3, 0.5), p=0.7)\n             ]),\n    'valid': alb.Compose([alb.Resize(IMG_SIZE, IMG_SIZE)])\n}","metadata":{"id":"CTOZzHuWW4F9","execution":{"iopub.status.busy":"2023-05-16T12:35:44.045169Z","iopub.execute_input":"2023-05-16T12:35:44.045444Z","iopub.status.idle":"2023-05-16T12:35:44.056193Z","shell.execute_reply.started":"2023-05-16T12:35:44.045418Z","shell.execute_reply":"2023-05-16T12:35:44.055148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Drop Bad Scans","metadata":{"id":"e2KLHNYVW4F-"}},{"cell_type":"markdown","source":"https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/344862\nhttps://www.kaggle.com/code/itsuki9180/a-segmentation-is-in-reverse-order\n\n\n1. 1.2.826.0.1.3680043.20574: does not include a full cervical spine and should be ignored.\n1. 1.2.826.0.1.3680043.29952: the slices are duplicated, meaning that there are 2 scans stiched to each other.\n1. 1.2.826.0.1.3680043.1363: the scans are reversed, if cannot reverse it again, ignore.","metadata":{"id":"shQpqqF0W4F-"}},{"cell_type":"code","source":"bad_scans = ['1.2.826.0.1.3680043.20574','1.2.826.0.1.3680043.29952', '1.2.826.0.1.368004.23904']\n\nfor uid in bad_scans:\n    train_df.drop(train_df[train_df['StudyInstanceUID']==uid].index, axis=0, inplace=True)","metadata":{"id":"_y90qKzXW4F-","execution":{"iopub.status.busy":"2023-05-16T12:35:44.057674Z","iopub.execute_input":"2023-05-16T12:35:44.05845Z","iopub.status.idle":"2023-05-16T12:35:44.0779Z","shell.execute_reply.started":"2023-05-16T12:35:44.058413Z","shell.execute_reply":"2023-05-16T12:35:44.076861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for uid in bad_scans:\n    voxel_df.drop(voxel_df[voxel_df['StudyInstanceUID']==uid].index, axis=0, inplace=True)","metadata":{"execution":{"iopub.status.busy":"2023-05-16T12:35:44.07932Z","iopub.execute_input":"2023-05-16T12:35:44.079766Z","iopub.status.idle":"2023-05-16T12:35:44.091692Z","shell.execute_reply.started":"2023-05-16T12:35:44.079729Z","shell.execute_reply":"2023-05-16T12:35:44.090638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset Class","metadata":{"id":"OhCycBpMW4F-"}},{"cell_type":"code","source":"# extract image and label of the given data\nclass CustomDataset(torch.utils.data.Dataset):\n    # Initialize\n    def __init__(self, voxel_df=df_train, train_df=train_df, transform=None, test=False):\n        super().__init__()\n        self.voxel_df = voxel_df\n        self.train_df = train_df\n        self.transform = transform\n        self.test = test\n    \n    def __getitem__(self, index):\n        voxel_UID = self.voxel_df.iloc[index]\n        uid = voxel_UID['id'].split('_')[0]\n        label_df = self.train_df \n\n        #print(voxel_UID['id'])\n        \n        #imgs = np.load(f'{base_path}/stage2-prep/{voxel_UID.id}.npz')['arr_0']\n        imgs = np.load(f'{base_path}/{voxel_UID.id}.npz')['arr_0'] # 80 X img_size X img_size\n        imgs = imgs.transpose(1, 2, 0) # img_size X img_size X 80\n        \n        if self.transform is not None:\n            trans = self.transform(image=imgs)\n            imgs = trans['image']\n        \n        imgs = imgs.transpose(2, 0, 1) # 80 X 512 X 512\n        vertebrae = voxel_UID['id'].split('_')[1] # C1; C2; C3; C4; C5; C6; C7\n        if self.test == True:\n            return torch.tensor(np.array(imgs/255.0, dtype=np.float32)).float(), [uid, vertebrae]\n        \n        label = label_df[label_df['StudyInstanceUID'] == uid][vertebrae].values.astype('float32') # either 0 or 1\n        return torch.tensor(np.array(imgs/255.0, dtype=np.float32)).float(), torch.tensor(label).float()\n       \n    # Length of dataset\n    def __len__(self):\n        return len(self.voxel_df)","metadata":{"id":"3AA2ZCh7W4F_","execution":{"iopub.status.busy":"2023-05-16T12:35:44.093384Z","iopub.execute_input":"2023-05-16T12:35:44.094829Z","iopub.status.idle":"2023-05-16T12:35:44.107549Z","shell.execute_reply.started":"2023-05-16T12:35:44.094791Z","shell.execute_reply":"2023-05-16T12:35:44.106325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Test Dataset","metadata":{}},{"cell_type":"code","source":"def plot_batch(imgs, size=40):\n    #plt.figure(figsize=(5*5, 5*(40//5)))\n    fig, axs = plt.subplots(size//5, 5, figsize=(5*5, size))\n    \n    for idx in range(size):\n        #plt.subplot(idx//5+1, idx%5+1, (idx//5+1, idx%5+1))\n        img = imgs[idx,].numpy()*255.0\n        img = img.astype('uint8')\n        axs[idx//5, idx%5].imshow(img, cmap='bone')\n        #plt.imshow(img, cmap='bone')\n        \n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-05-16T12:35:44.108921Z","iopub.execute_input":"2023-05-16T12:35:44.109341Z","iopub.status.idle":"2023-05-16T12:35:44.121444Z","shell.execute_reply.started":"2023-05-16T12:35:44.109301Z","shell.execute_reply":"2023-05-16T12:35:44.120422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = CustomDataset(voxel_df=voxel_df, train_df=train_df, transform=data_transforms['valid'], test=True)\ntest_loader = DataLoader(test_dataset, batch_size=1, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2023-05-16T12:35:44.123361Z","iopub.execute_input":"2023-05-16T12:35:44.123788Z","iopub.status.idle":"2023-05-16T12:35:44.131727Z","shell.execute_reply.started":"2023-05-16T12:35:44.123751Z","shell.execute_reply":"2023-05-16T12:35:44.130639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs, labels = next(iter(test_loader))\nprint(labels)\nplot_batch(imgs[0], size=30)","metadata":{"execution":{"iopub.status.busy":"2023-05-16T12:35:44.133353Z","iopub.execute_input":"2023-05-16T12:35:44.133798Z","iopub.status.idle":"2023-05-16T12:35:51.596088Z","shell.execute_reply.started":"2023-05-16T12:35:44.133759Z","shell.execute_reply":"2023-05-16T12:35:51.594682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs, labels = next(iter(test_loader))\nprint(labels)\nplot_batch(imgs[0], size=30)","metadata":{"execution":{"iopub.status.busy":"2023-05-16T12:35:51.597641Z","iopub.execute_input":"2023-05-16T12:35:51.598122Z","iopub.status.idle":"2023-05-16T12:35:59.325082Z","shell.execute_reply.started":"2023-05-16T12:35:51.598068Z","shell.execute_reply":"2023-05-16T12:35:59.323684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## LSTM Model","metadata":{"id":"_TctoXGlW4F_"}},{"cell_type":"code","source":"# important to have a bigger backbone\n# ResNet101 encoder backbone + UNet decoder => CrackNet\nclass ClassificationModel(nn.Module):\n    def __init__(self):\n        super(ClassificationModel, self).__init__()\n        self.segmodel = smp.Unet(              # cnn\n            encoder_backbone,                  # resnet101 backbone\n            encoder_weights='imagenet',        # pretrained-weight = imagenet\n            in_channels=tot_slice,             # in channels of 40*2\n            classes=CLASS,                     # output channel be 7 from C1 to C7\n            activation=None,\n        )\n        self.lstm = nn.LSTM( # output = (batch, sequence length, dimension * )\n            input_size=(IMG_SIZE*IMG_SIZE), # 512 * 512 * 40(channel)\n            hidden_size=hidden1, \n            num_layers=2, # num_layer = 2 if bidirectional = True. 1 otherwise\n            dropout=0., \n            bidirectional=True, \n            batch_first=True # (batch_dimension, segment_dimension, feature_dimension)\n        )\n        self.head = nn.Sequential(\n            nn.Linear(CLASS*(hidden1*2), IMG_SIZE), # dimensions, input image\n            #nn.Linear(CLASS*512*512, 512),\n            nn.BatchNorm1d(IMG_SIZE),\n            nn.Dropout(0.3),\n            nn.LeakyReLU(0.1),\n            nn.Linear(IMG_SIZE, 1),\n        )\n    \n    # input: (batch_size, tot_slice, image_size, image_size)\n    def forward(self, x):\n        batch_size = x.size(0)\n        img_size = x.size(2) # or x.size(3); equal value\n        \n        x = self.segmodel(x)\n        x = x.view(batch_size, CLASS, -1)\n        # input: batch size, sequence length, input image size\n        x,_ = self.lstm(x)\n        x = x.contiguous().view(batch_size, -1)\n        x = self.head(x)\n        \n        return x","metadata":{"execution":{"iopub.status.busy":"2023-05-16T12:35:59.327347Z","iopub.execute_input":"2023-05-16T12:35:59.328207Z","iopub.status.idle":"2023-05-16T12:35:59.344148Z","shell.execute_reply.started":"2023-05-16T12:35:59.328154Z","shell.execute_reply":"2023-05-16T12:35:59.343002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Settings","metadata":{"id":"fcCRIwyTW4GA"}},{"cell_type":"code","source":"train_dataset = CustomDataset(voxel_df=df_train, train_df=train_df, transform=data_transforms['train'], test=False)\nvalid_dataset = CustomDataset(voxel_df=df_valid, train_df=train_df, transform=data_transforms['valid'], test=False)\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH, shuffle=True)\nvalid_loader = DataLoader(valid_dataset, batch_size=BATCH, shuffle=True)","metadata":{"id":"nKXQngCaW4GB","execution":{"iopub.status.busy":"2023-05-16T12:35:59.345724Z","iopub.execute_input":"2023-05-16T12:35:59.34635Z","iopub.status.idle":"2023-05-16T12:35:59.359372Z","shell.execute_reply.started":"2023-05-16T12:35:59.346313Z","shell.execute_reply":"2023-05-16T12:35:59.358384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model setting\nmodel = ClassificationModel()\nmodel.to(device)\n#model.load_state_dict(torch.load('/kaggle/input/stage2-cracknet-lstm-yolo-window/stage2_cracknet_best.ckpt'))\n\n# Adam optimizer\noptimizer = optim.AdamW(params=model.parameters(), lr=5e-5, weight_decay=1e-5)","metadata":{"id":"HupXNvGTW4GA","execution":{"iopub.status.busy":"2023-05-16T12:35:59.360926Z","iopub.execute_input":"2023-05-16T12:35:59.361542Z","iopub.status.idle":"2023-05-16T12:36:11.563078Z","shell.execute_reply.started":"2023-05-16T12:35:59.361507Z","shell.execute_reply":"2023-05-16T12:36:11.561567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Competition Loss","metadata":{"id":"h99jYAB2W4GA"}},{"cell_type":"code","source":"# Replicate competition metric (https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/341854)\nloss_fn = F.binary_cross_entropy_with_logits\n\n# The logits created from model training is either 0 or 1 for only one vertebrae\ncompetition_weights = {\n    '-' : torch.tensor(1, dtype=torch.float, device=device),\n    '+' : torch.tensor(2, dtype=torch.float, device=device),\n}\n\n# with row-wise weights normalization (https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/344565)\ndef competiton_loss_row_norm(y_hat, y):\n    loss = loss_fn(y_hat.view(-1), y.to(y_hat.dtype).view(-1))\n    weights = y * competition_weights['+'] + (1 - y) * competition_weights['-']\n    loss = (loss * weights).sum(axis=1)\n    w_sum = weights.sum(axis=1)\n    loss = torch.div(loss, w_sum)\n    return loss.mean()","metadata":{"id":"TzQmF9qlW4GA","execution":{"iopub.status.busy":"2023-05-16T12:36:11.566605Z","iopub.execute_input":"2023-05-16T12:36:11.570868Z","iopub.status.idle":"2023-05-16T12:36:11.586415Z","shell.execute_reply.started":"2023-05-16T12:36:11.570824Z","shell.execute_reply":"2023-05-16T12:36:11.58542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training\n","metadata":{"id":"qY4KiVMkW4GA"}},{"cell_type":"code","source":"def train(model, dataloader, optimizer):\n    model.train()\n    scaler = GradScaler()\n    #scheduler = OneCycleLR(optimizer, max_lr=0.0005, epochs=1, steps_per_epoch=len(df_train), pct_start=0.3)\n    \n    train_loss = []\n    train_acc = []\n    \n    for imgs, label in tqdm(dataloader):\n        # set the gradient to 0 at initial\n        optimizer.zero_grad()\n\n        # forward data, making sure the data and model are on the same device\n        with autocast(enabled=True):\n            logits = model(imgs.to(device))\n            loss = competiton_loss_row_norm(logits, label.to(device))\n        \n        #print('argmax:',logits.detach().cpu().numpy().argmax(axis=1))\n        #print('max:',logits.detach().cpu().numpy().max(axis=1))\n        #print('truth:',label.numpy())\n        acc = (logits.argmax(dim=-1) == label.to(device)).float().mean()\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        #scheduler.step()\n        \n        train_loss.append(loss.item())\n        train_acc.append(acc)\n        \n    train_loss = sum(train_loss) / len(train_loss)\n    train_acc = sum(train_acc) / len(train_acc)\n    \n    return train_loss, train_acc","metadata":{"id":"UeR07_B_W4GB","execution":{"iopub.status.busy":"2023-05-16T12:36:11.587835Z","iopub.execute_input":"2023-05-16T12:36:11.592107Z","iopub.status.idle":"2023-05-16T12:36:11.605501Z","shell.execute_reply.started":"2023-05-16T12:36:11.592058Z","shell.execute_reply":"2023-05-16T12:36:11.604164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def validation(model, dataloader, optimizer):\n    model.eval()\n    valid_loss = []\n    valid_acc = []\n\n    for imgs, label in tqdm(dataloader):\n        # no need gradient in validation\n        # use torch.no_grad() accelerates the forward process\n        with torch.no_grad():\n            logits = model(imgs.to(device))\n        \n        loss = competiton_loss_row_norm(logits, label.to(device))\n        acc = (logits.argmax(dim=-1) == label.to(device)).float().mean()\n        \n        valid_loss.append(loss.item())\n        valid_acc.append(acc)\n    \n    valid_loss = sum(valid_loss) / len(valid_loss)\n    valid_acc = sum(valid_acc) / len(valid_acc)\n    \n    return valid_loss, valid_acc","metadata":{"id":"Rf1dVk7rW4GB","execution":{"iopub.status.busy":"2023-05-16T12:36:11.607311Z","iopub.execute_input":"2023-05-16T12:36:11.607717Z","iopub.status.idle":"2023-05-16T12:36:11.619608Z","shell.execute_reply.started":"2023-05-16T12:36:11.607679Z","shell.execute_reply":"2023-05-16T12:36:11.618559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed_everything(SEED)\n\n# initialize the best values to save\nbest_train_loss = 0\nbest_train_acc = 0\nbest_valid_loss = 0\nbest_valid_acc = 0\n\nbest_epoch = 0\nearly_stop_count = 0\n\nscheduler = CosineAnnealingWarmRestarts(optimizer, EPOCH, eta_min=5e-6)\n\n# start training\nfor epoch in range(EPOCH):\n    print(f'### Epoch: {epoch+1} ###')\n    train_loss, train_acc = train(model, train_loader, optimizer)\n    print(f'[ Train | {epoch + 1:03d}/{EPOCH:03d} ] loss = {train_loss:.5f}, acc = {train_acc:.5f}')\n    valid_loss, valid_acc = validation(model, valid_loader, optimizer)\n    print(f'[ Valid | {epoch + 1:03d}/{EPOCH:03d} ] loss = {valid_loss:.5f}, acc = {valid_acc:.5f}')\n    print()\n    \n    scheduler.step()\n    \n    # save train and valid logs\n    trainlosslog.append(train_loss)\n    trainacclog.append(train_acc.cpu().data.numpy())\n    validlosslog.append(valid_loss)\n    validacclog.append(valid_acc.cpu().data.numpy())\n    \n    # save models\n    if valid_loss < best_loss:\n        # save highest values\n        best_train_loss, best_train_acc = train_loss, train_acc\n        best_valid_loss, best_valid_acc = valid_loss, valid_acc\n        best_epoch = epoch\n        \n        # save model\n        torch.save(model.state_dict(), \"stage2_cracknet_best.ckpt\") # only save best to prevent output memory exceed error\n        # reset values\n        best_loss = valid_loss\n        early_stop_count = 0\n    \n    \"\"\"if early_stop_count > 10:\n        print('Accuracy not increasing. Early Stopping...')\n        break\"\"\"\n    \n    early_stop_count = early_stop_count + 1\n\nprint()\nprint(f\"[ Best Train | {best_epoch+1:03d} / {EPOCH:03d} ] loss = {best_train_loss:.8f}, acc = {best_train_acc:.8f}\")\nprint(f\"[ Best Valid | {best_epoch+1:03d} / {EPOCH:03d} ] loss = {best_valid_loss:.8f}, acc = {best_valid_acc:.8f}\")\n","metadata":{"id":"uXGxSt9lW4GB","_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-05-16T12:36:11.621333Z","iopub.execute_input":"2023-05-16T12:36:11.621785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"log_df = pd.DataFrame(\n    {\n        'trainLoss': trainlosslog,\n        'trainAcc': trainacclog,\n        'validLoss': validlosslog,\n        'validAcc': validacclog\n    })\n\nlog_df.to_csv('Logs.csv')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}