{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9187072,"sourceType":"datasetVersion","datasetId":5504483}],"dockerImageVersionId":30761,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport glob\nimport random \nimport math\n\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport matplotlib.gridspec as gridspec\nfrom tqdm import tqdm\nfrom types import SimpleNamespace\nimport albumentations as A\nimport cv2\n\nimport torch\nimport torch.nn as nn\nimport torch.functional as F\nimport timm\nimport pydicom\n","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:15.592442Z","iopub.execute_input":"2024-09-12T04:42:15.592853Z","iopub.status.idle":"2024-09-12T04:42:24.543484Z","shell.execute_reply.started":"2024-09-12T04:42:15.59281Z","shell.execute_reply":"2024-09-12T04:42:24.542328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=1234):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = False\n    torch.backends.cudnn.benchmark = True","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:24.545524Z","iopub.execute_input":"2024-09-12T04:42:24.545963Z","iopub.status.idle":"2024-09-12T04:42:24.557461Z","shell.execute_reply.started":"2024-09-12T04:42:24.545911Z","shell.execute_reply":"2024-09-12T04:42:24.556282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def convert_to_8bit(x):\n    lower, upper = np.percentile(x, (1, 99))\n    x = np.clip(x, lower, upper)\n    x = x - np.min(x)\n    x = x / np.max(x)\n    return (x * 255).astype('uint8')","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:24.559222Z","iopub.execute_input":"2024-09-12T04:42:24.55963Z","iopub.status.idle":"2024-09-12T04:42:24.572092Z","shell.execute_reply.started":"2024-09-12T04:42:24.559589Z","shell.execute_reply":"2024-09-12T04:42:24.570208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dicom_stack(dicom_folder, plane, reverse_sort=True):\n    \n    dicom_files = glob.glob(os.path.join(dicom_folder, '*.dcm'))\n    dicoms = [pydicom.dcmread(f) for f in dicom_files]\n    print(len(dicom_files))\n    plane = {'sagittal' : 0, 'coronal': 1, 'axial' : 2}[plane.lower()]\n    positions = np.asarray([float(d.ImagePositionPatient[plane]) for d in dicoms])\n    idx = np.argsort(-positions if reverse_sort else positions)\n    ipp = np.asarray([d.ImagePositionPatient for d in dicoms]).astype('float')[idx]\n    array = np.stack([d.pixel_array.astype(float) for d in dicoms])\n    array = array[idx]\n    \n    return {\n        \"array\" : convert_to_8bit(array),\n        'positions' : ipp,\n        'pixel-spacing' : np.asarray(dicoms[0].PixelSpacing).astype('float')\n    }\n    \nimage_dir = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/\"    ","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:24.574656Z","iopub.execute_input":"2024-09-12T04:42:24.575044Z","iopub.status.idle":"2024-09-12T04:42:24.589559Z","shell.execute_reply.started":"2024-09-12T04:42:24.574999Z","shell.execute_reply":"2024-09-12T04:42:24.587859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"resize_transform = A.Compose([\n    A.LongestMaxSize(max_size = 256, interpolation=cv2.INTER_CUBIC, always_apply=True),\n    A.PadIfNeeded(min_height=256, max_height=256, border_mode=cv2.BORDER_CONSTANT, value=(0, 0, 0), always_apply=True)\n])","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:24.591781Z","iopub.execute_input":"2024-09-12T04:42:24.59297Z","iopub.status.idle":"2024-09-12T04:42:24.607882Z","shell.execute_reply.started":"2024-09-12T04:42:24.592906Z","shell.execute_reply":"2024-09-12T04:42:24.6069Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def angle_of_line(x1, y1, x2, y2):\n    return math.degrees(math.atan2(-(y2-y1), x2-x1))\n\n","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:24.609684Z","iopub.execute_input":"2024-09-12T04:42:24.610121Z","iopub.status.idle":"2024-09-12T04:42:24.622835Z","shell.execute_reply.started":"2024-09-12T04:42:24.610079Z","shell.execute_reply":"2024-09-12T04:42:24.621592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_image(img, coord_temp):\n    fig, ax = plt.subplots()\n    ax.imshow(img, cmap='gray')\n    h, w = img.shape\n    \n    p = coord_temp.groupby('level') \\\n        .apply(lambda g : list(zip(g['relative_x'], g['relative_y'])), include_groups=False) \\\n        .reset_index(drop=False, name='vals')\n    \n    for _, row in p.iterrows():\n        level = p['level']\n        x = [_[0] * w for _ in row['vals']]\n        y = [_[1] * h for _ in row['vals']]\n        ax.plot(x, y, marker='o')\n        \n    ax.axis('off')\n    plt.show()\n    \n    ","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:24.625749Z","iopub.execute_input":"2024-09-12T04:42:24.626209Z","iopub.status.idle":"2024-09-12T04:42:24.635537Z","shell.execute_reply.started":"2024-09-12T04:42:24.626156Z","shell.execute_reply":"2024-09-12T04:42:24.634353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_5_crops(img, coord_temp):\n    fig = plt.figure(figsize=(10, 10))\n    gs = gridspec.GridSpec(1, 5, width_ratios=[1] * 5)\n    \n    p = coord_temp.groupby('level') \\\n        .apply(lambda g : list(zip(g['relative_x'], g['relative_y'])), include_groups=False) \\\n        .reset_index(drop=False, name='vals')\n    \n    for idx, (_, rows) in enumerate(p.iterrows()):\n        img_copy = img.copy()\n        h, w = img.shape\n        \n        level = rows['level']\n        vals = sorted(rows['vals'], key=lambda x:x[0])\n        print(vals)\n        a, b = vals\n        a = (a[0] * w, a[1] * h)\n        b = (b[0] * w, b[1] * h)\n        \n        rotation_angle = angle_of_line(a[0], a[1], b[0], b[1])\n        \n        transform = A.Compose([\n            A.Rotate(limit=(-rotation_angle, -rotation_angle), p=1.0)\n        ], keypoint_params = A.KeypointParams(format='xy', remove_invisible=False)\n        )\n        \n        t = transform(image=img_copy, keypoints=[a, b])\n        img_copy = t['image']\n        a, b = t['keypoints']\n        \n        img_copy = crop_between_keypoints(img_copy, a, b)\n        img_copy = resize_transform(image=img_copy)['image']\n        \n        ax = plt.subplot(gs[idx])\n        ax.imshow(img_copy, cmap='gray')\n        ax.set_title(level)\n        ax.axis('off')\n        \n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:24.637113Z","iopub.execute_input":"2024-09-12T04:42:24.63765Z","iopub.status.idle":"2024-09-12T04:42:24.657442Z","shell.execute_reply.started":"2024-09-12T04:42:24.637598Z","shell.execute_reply":"2024-09-12T04:42:24.656391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def crop_between_keypoints(img, keypoint1, keypoint2):\n    h, w = img.shape\n    x1, y1 = int(keypoint1[0]), int(keypoint1[1])\n    x2, y2 = int(keypoint2[0]), int(keypoint2[1])\n    \n    left = int(min(x1, x2))\n    right = int(max(x1, x2))\n    \n    top = int(min(y1, y2) - (h * 0.1))\n    bottom = int(max(y1, y2) + (h * 0.1))\n    \n    return img[top:bottom, left:right]\n","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:24.658707Z","iopub.execute_input":"2024-09-12T04:42:24.65905Z","iopub.status.idle":"2024-09-12T04:42:24.669653Z","shell.execute_reply.started":"2024-09-12T04:42:24.659002Z","shell.execute_reply":"2024-09-12T04:42:24.6686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 10\nN = 2","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:24.67354Z","iopub.execute_input":"2024-09-12T04:42:24.673961Z","iopub.status.idle":"2024-09-12T04:42:24.689936Z","shell.execute_reply.started":"2024-09-12T04:42:24.67391Z","shell.execute_reply":"2024-09-12T04:42:24.688869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dfd= pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv\")\ndfd = dfd[dfd.series_description == 'Sagittal T2/STIR']\ndfd = dfd.sample(frac=1, random_state=SEED).head(N)\n\ndfd","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:24.69169Z","iopub.execute_input":"2024-09-12T04:42:24.692097Z","iopub.status.idle":"2024-09-12T04:42:24.749393Z","shell.execute_reply.started":"2024-09-12T04:42:24.692045Z","shell.execute_reply":"2024-09-12T04:42:24.748177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coords= pd.read_csv(\"/kaggle/input/lumbar-coordinate-pretraining-dataset/coords_rsna_improved.csv\")\ncoords = coords.sort_values(['series_id', 'level', 'side']).reset_index(drop=True)\ncoords = coords[['series_id', 'level', 'side', 'relative_x', 'relative_y']]\n","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:24.751126Z","iopub.execute_input":"2024-09-12T04:42:24.751622Z","iopub.status.idle":"2024-09-12T04:42:24.961204Z","shell.execute_reply.started":"2024-09-12T04:42:24.751569Z","shell.execute_reply":"2024-09-12T04:42:24.959944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for idx, rows in dfd.iterrows():\n    try:\n        print(f\"{'-' * 25} STUDY_ID : {rows.study_id}, SERIES_ID : {rows.series_id} {'-' * 25}\")\n        sag_t2 = load_dicom_stack(os.path.join(image_dir, str(rows.study_id), str(rows.series_id)), plane='sagittal')\n        img = sag_t2['array'][len(sag_t2['array']) // 2]\n        coords_temp = coords[coords['series_id'] == rows.series_id].copy()\n        \n        plot_image(img, coords_temp)\n        plot_5_crops(img, coords_temp)\n    except Exception as e:\n        print(e)\n        pass\n    ","metadata":{"execution":{"iopub.status.busy":"2024-09-12T05:01:24.420227Z","iopub.execute_input":"2024-09-12T05:01:24.421167Z","iopub.status.idle":"2024-09-12T05:01:26.070832Z","shell.execute_reply.started":"2024-09-12T05:01:24.421116Z","shell.execute_reply":"2024-09-12T05:01:26.069506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for files in os.listdir(os.path.join(image_dir, str(rows.study_id), str(rows.series_id))):\n    print(files)","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:27.084445Z","iopub.execute_input":"2024-09-12T04:42:27.084848Z","iopub.status.idle":"2024-09-12T04:42:27.091595Z","shell.execute_reply.started":"2024-09-12T04:42:27.084754Z","shell.execute_reply":"2024-09-12T04:42:27.090558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg = SimpleNamespace(\n    img_dir = \"/kaggle/input/lumbar-coordinate-pretraining-dataset/data/\",\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu'),\n    n_frames = 3,\n    epochs = 10,\n    lr = 0.0005,\n    batch_size=16,\n    backbone='resnet18',\n    seed=0\n)\nset_seed(seed=cfg.seed)","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:27.092796Z","iopub.execute_input":"2024-09-12T04:42:27.093154Z","iopub.status.idle":"2024-09-12T04:42:27.112023Z","shell.execute_reply.started":"2024-09-12T04:42:27.09311Z","shell.execute_reply":"2024-09-12T04:42:27.110858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coords_temp.head(5)","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:27.113681Z","iopub.execute_input":"2024-09-12T04:42:27.114168Z","iopub.status.idle":"2024-09-12T04:42:27.12886Z","shell.execute_reply.started":"2024-09-12T04:42:27.114129Z","shell.execute_reply":"2024-09-12T04:42:27.127821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df= pd.read_csv(\"/kaggle/input/lumbar-coordinate-pretraining-dataset/coords_pretrain.csv\")\ndf = df.sort_values(['source', 'filename', 'level']).reset_index(drop=True)\ndf['filename'] = df['filename'].str.replace('.jpg', '.npy')\ndf['series_id'] = df['source'] + '_'  + df['filename'].str.split('.').str[0]\n\nprint(\"--IMGS per source--\")\ndisplay((df.source.value_counts()/5).astype(int).reset_index())","metadata":{"execution":{"iopub.status.busy":"2024-09-12T04:42:27.130126Z","iopub.execute_input":"2024-09-12T04:42:27.130535Z","iopub.status.idle":"2024-09-12T04:42:27.185092Z","shell.execute_reply.started":"2024-09-12T04:42:27.130495Z","shell.execute_reply":"2024-09-12T04:42:27.183911Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PreTrainDataset(torch.utils.data.Dataset):\n    def __init__(self, df, cfg):\n        self.cfg = cfg\n        self.records = self.load_coords(df)\n        \n    def load_coords(self,df):\n        d = df.groupby(\"series_id\")[['relative_x', 'relative_y']].apply(lambda x : list(x.itertuples(index=False, name=None)))\n        records = {}\n        \n        for i, (k, v) in enumerate(d.items()):\n            records[i] = {'series_id' : k, 'label' : np.array(v).flatten()}\n            assert len(v) == 5\n            \n        return records\n    \n    def pad_image(self, img):\n        n = img.shape[-1]\n        if n >= self.cfg.n_frames:\n            start_idx = (n - self.cfg.n_frames) // 2\n            return img[:, :, start_idx : start_idx + self.cfg.n_frames]\n        else:\n            pad_left = (self.cfg.n_frames - n) // 2\n            pad_right = self.cfg.n_frames - n - pad_left\n            return np.pad(img, ((0, 0), (0, 0),(pad_left, pad_right)), 'constant', constant_values=0 )\n        \n    def load_img(self, source, series_id):\n        fname = os.path.join(self.cfg.img_dir, f\"processed_{source}/{series_id}.npy\")\n        img = np.load(fname).astype(np.float32)\n        img = self.pad_image(img)\n        img = np.transpose(img, (2, 0, 1))\n        img = (img / 255.0)\n        return img\n    \n    def __getitem__(self, idx):\n        d = self.records[idx]\n        label = d['label']\n        source = d['series_id'].split('_')[0]\n        series_id= \"_\".join(d[\"series_id\"].split(\"_\")[1:])\n        \n        image = self.load_img(source, series_id)\n        \n        return {\n            'img' : image, \n            'label' : label\n        }\n    \n    def __len__(self, ):\n        return len(self.records)\n    \nds = PreTrainDataset(df, cfg)\nprint(\"--Sample images--\")\n\nfor k, v in ds[0].items():\n    print(k, v.shape)    ","metadata":{"execution":{"iopub.status.busy":"2024-09-12T05:27:03.231018Z","iopub.execute_input":"2024-09-12T05:27:03.231871Z","iopub.status.idle":"2024-09-12T05:27:03.470234Z","shell.execute_reply.started":"2024-09-12T05:27:03.231819Z","shell.execute_reply":"2024-09-12T05:27:03.468976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def batch_to_device(batch, device, skip_layers=[]):\n    batch_dict = {}\n    for key in batch:\n        if key in skip_layers:\n            batch_dict[key] = batch[key]\n        else:\n            batch_dict[key] = batch[key].to(device)\n            \n    return batch_dict","metadata":{"execution":{"iopub.status.busy":"2024-09-12T05:26:29.643502Z","iopub.execute_input":"2024-09-12T05:26:29.643959Z","iopub.status.idle":"2024-09-12T05:26:29.650135Z","shell.execute_reply.started":"2024-09-12T05:26:29.643913Z","shell.execute_reply":"2024-09-12T05:26:29.648566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_prediction(batch, pred, epochs):\n    mid = cfg.n_frames // 2\n    \n    for idx in range(1):\n        img = batch['img'][idx, mid, :, :].cpu().numpy() * 255\n        cs_true = batch['label'][idx, ...].cpu().numpy() * 256\n        cs = pred[idx, ...].cpu().numpy() * 256\n        \n        coords_list = [(\"TRUE\", \"lightblue\", cs_true), ('PRED', 'orange', cs)]\n        text_labels = [str(x) for x in range(1, 6)]\n        \n        fig, axes = plt.subplots(1, len(coords_list), figsize=(10, 4))\n        fig.suptitle(f'EPOCHS: {epochs}')\n        \n        for ax, (title, color, coords) in zip(axes, coords_list):\n            ax.imshow(img, cmap='gray')\n            ax.scatter(coords[0::2], coords[1::2], c=color, s=50)\n            ax.axis('off')\n            ax.set_title(title)\n            \n            for i, (x, y) in enumerate(zip(coords[0::2], coords[1::2])):\n                if i < len(text_labels):\n                    ax.text(x+10, y, text_labels[i], color='white', fontsize=15, bbox=dict(facecolor='black', alpha=0.5))\n                    \n            fig.suptitle(f\"EPOCHS {epochs}\")\n            plt.show()\n            \n    return ","metadata":{"execution":{"iopub.status.busy":"2024-09-12T05:34:49.789575Z","iopub.execute_input":"2024-09-12T05:34:49.790007Z","iopub.status.idle":"2024-09-12T05:34:49.801118Z","shell.execute_reply.started":"2024-09-12T05:34:49.789964Z","shell.execute_reply":"2024-09-12T05:34:49.800039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_weights_and_skip_mismatch(model, weights_path, device):\n    state_dict = torch.load(weights_path, device)\n    model_dict = model.state_dict()\n\n    params = {}\n\n    for (sdk, sfv), (mdk, mdv) in zip(state_dict.items(), model_dict.items()):\n        if sfv.size() == mdv.size():\n            params[sdk] = sfv\n        else:\n            print(f\"skipping {sdk}, {sfv.size()} != {mdv.size()}\")\n\n\n    model.load_state_dict(params, strict=False)\n    print(f\"Loaded model from {weights_path}\")","metadata":{"execution":{"iopub.status.busy":"2024-09-12T06:15:01.185255Z","iopub.execute_input":"2024-09-12T06:15:01.18574Z","iopub.status.idle":"2024-09-12T06:15:01.193385Z","shell.execute_reply.started":"2024-09-12T06:15:01.185696Z","shell.execute_reply":"2024-09-12T06:15:01.19204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = df[df['source'] != 'spider']\nval_df = df[df['source'] == 'spider']\n\ntrain_ds = PreTrainDataset(train_df, cfg)\nval_ds = PreTrainDataset(val_df, cfg)\n\ntrain_dl = torch.utils.data.DataLoader(\n    train_ds, \n    batch_size = cfg.batch_size,\n    shuffle = True,\n    drop_last = True\n)\n\nval_dl = torch.utils.data.DataLoader(\n    val_ds, \n    batch_size = cfg.batch_size,\n    shuffle=False,\n    \n)\n\nmodel = timm.create_model('resnet18', pretrained=True, num_classes=10)\nmodel = model.to(cfg.device)\n\ncriterion = nn.MSELoss()\noptimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr)\n","metadata":{"execution":{"iopub.status.busy":"2024-09-12T05:27:06.372845Z","iopub.execute_input":"2024-09-12T05:27:06.373806Z","iopub.status.idle":"2024-09-12T05:27:06.971739Z","shell.execute_reply.started":"2024-09-12T05:27:06.373755Z","shell.execute_reply":"2024-09-12T05:27:06.970692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(cfg.epochs + 1):\n    loss = torch.tensor([0.]).float().to(cfg.device)\n    if epoch != 0:\n        model = model.train()\n        for batch in tqdm(train_dl):\n            batch = batch_to_device(batch, cfg.device)\n            optimizer.zero_grad()\n            \n            x_out = model(batch['img'].float())\n            x_out = torch.sigmoid(x_out)\n            \n            loss = criterion(x_out, batch['label'].float())\n            loss.backward()\n            optimizer.step()\n            \n    val_loss = 0\n    with torch.no_grad():\n        model = model.eval()\n        \n        for batch in tqdm(val_dl):\n            batch = batch_to_device(batch, cfg.device)\n            pred = model(batch['img'].float())\n            pred = torch.sigmoid(pred)\n            \n            \n            val_loss += criterion(pred, batch['label'].float()).item()\n        val_loss = val_loss / len(val_dl)\n        \n    \n    visualize_prediction(batch, pred, epoch)\n    \n    print(f\"Epoch {epoch + 1}, Training loss : {loss.item()}, Validation Loss : {val_loss}\")\n    \nprint('Training complete')\n    ","metadata":{"execution":{"iopub.status.busy":"2024-09-12T05:34:52.301496Z","iopub.execute_input":"2024-09-12T05:34:52.301964Z","iopub.status.idle":"2024-09-12T06:00:53.982467Z","shell.execute_reply.started":"2024-09-12T05:34:52.301916Z","shell.execute_reply":"2024-09-12T06:00:53.981124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f = f\"{cfg.backbone}_{cfg.seed}.pt\"\ntorch.save(model.state_dict(), f)\nprint(f\"Saved weights {f}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-09-12T06:05:53.462171Z","iopub.execute_input":"2024-09-12T06:05:53.462622Z","iopub.status.idle":"2024-09-12T06:05:53.53889Z","shell.execute_reply.started":"2024-09-12T06:05:53.462582Z","shell.execute_reply":"2024-09-12T06:05:53.537615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = timm.create_model('resnet18', pretrained=True, num_classes=75)\nmodel = model.to(cfg.device)\nload_weights_and_skip_mismatch(model, f, cfg.device)","metadata":{"execution":{"iopub.status.busy":"2024-09-12T06:15:09.639239Z","iopub.execute_input":"2024-09-12T06:15:09.639691Z","iopub.status.idle":"2024-09-12T06:15:10.083156Z","shell.execute_reply.started":"2024-09-12T06:15:09.639649Z","shell.execute_reply":"2024-09-12T06:15:10.082036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import accuracy_score, confusion_matrix, classification_report\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n","metadata":{"execution":{"iopub.status.busy":"2024-09-12T06:19:19.598664Z","iopub.execute_input":"2024-09-12T06:19:19.599072Z","iopub.status.idle":"2024-09-12T06:19:19.605001Z","shell.execute_reply.started":"2024-09-12T06:19:19.599025Z","shell.execute_reply":"2024-09-12T06:19:19.603276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def calculate_accuracy(model, data_loader, device, threshold=0.5):\n    model.eval()\n    all_preds = []\n    all_labels = []\n    with torch.no_grad():\n        for batch in data_loader:\n            batch = batch_to_device(batch, device)\n            preds = model(batch['img'].float())\n            preds = torch.sigmoid(preds)\n            preds = (preds > threshold).float()\n            print(preds.shape, batch['label'].shape)\n            if preds.shape == (1, 75):\n                continue\n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(batch['label'].cpu().numpy())\n    \n    accuracy = accuracy_score(all_labels, all_preds)\n    return accuracy, all_labels, all_preds\n\nval_accuracy, val_labels, val_preds = calculate_accuracy(model, val_dl, cfg.device)\nprint(f\"Validation accuracy is {val_accuracy}\")\n\n\n    ","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:04:47.76678Z","iopub.execute_input":"2024-09-12T07:04:47.767216Z","iopub.status.idle":"2024-09-12T07:05:00.24015Z","shell.execute_reply.started":"2024-09-12T07:04:47.767174Z","shell.execute_reply":"2024-09-12T07:05:00.238591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}