{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9187072,"sourceType":"datasetVersion","datasetId":5504483}],"dockerImageVersionId":30747,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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.nn.functional as F\nimport timm\nimport pydicom\n\ndef 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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-08-17T08:34:36.185834Z","iopub.execute_input":"2024-08-17T08:34:36.186185Z","iopub.status.idle":"2024-08-17T08:34:45.348829Z","shell.execute_reply.started":"2024-08-17T08:34:36.186155Z","shell.execute_reply":"2024-08-17T08:34:45.347894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Lumbar Coordinate Dataset\n\n\nThis notebook shows how the [Lumbar Coordinate Dataset](https://www.kaggle.com/datasets/brendanartley/lumbar-coordinate-pretraining-datasethttps://www.kaggle.com/datasets/brendanartley/lumbar-coordinate-pretraining-dataset) can be used in the RSNA 2024 competition. First I will showcase improved coordinates for the competition data, and then I will showcase using external data to pretrain backbones.","metadata":{}},{"cell_type":"markdown","source":"## 1. Improved RSNA Coordinates\n\n\nWe improve the competition coordinates by labelling the left side of each disc. This gives us the angle of orientation which can be used to improve cropping of each disc.\n\nThanks to [Ian Pan](https://www.kaggle.com/vaillant) for sharing some helper functions for loading the dicom data. See his great notebook [here](https://www.kaggle.com/code/vaillant/cross-reference-images-in-different-mri-planes?scriptVersionId=182551992&cellId=2).","metadata":{}},{"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\")\n\n\ndef load_dicom_stack(dicom_folder, plane, reverse_sort=False):\n    dicom_files = glob.glob(os.path.join(dicom_folder, \"*.dcm\"))\n    dicoms = [pydicom.dcmread(f) for f in 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    # if reverse_sort=False, then increasing array index will be from RIGHT->LEFT and CAUDAL->CRANIAL\n    # thus we do reverse_sort=True for axial so increasing array index is craniocaudal\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(\"float32\") for d in dicoms])\n    array = array[idx]\n    return {\"array\": convert_to_8bit(array), \"positions\": ipp, \"pixel_spacing\": np.asarray(dicoms[0].PixelSpacing).astype(\"float\")}\n\nimage_dir = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/\"","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:34:45.350823Z","iopub.execute_input":"2024-08-17T08:34:45.351193Z","iopub.status.idle":"2024-08-17T08:34:45.364344Z","shell.execute_reply.started":"2024-08-17T08:34:45.351161Z","shell.execute_reply":"2024-08-17T08:34:45.363455Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"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, min_width=256, border_mode=cv2.BORDER_CONSTANT, value=(0, 0, 0), always_apply=True),\n])\n\ndef angle_of_line(x1, y1, x2, y2):\n    return math.degrees(math.atan2(-(y2-y1), x2-x1))\n\ndef plot_img(img, coords_temp):\n    # Plot img\n    fig, ax = plt.subplots()\n    ax.imshow(img, cmap='gray')\n    h, w = img.shape\n    \n    # Kepoints as pairs\n    p= coords_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    # Plot keypoints\n    for _, row in p.iterrows():\n        level = row['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    ax.axis('off')\n    plt.show()\n\ndef plot_5_crops(img, coords_temp):\n    # Create a figure and axis for the grid\n    fig = plt.figure(figsize=(10, 10))\n    gs = gridspec.GridSpec(1, 5, width_ratios=[1]*5)\n    \n    # Plot the crops\n    p= coords_temp.groupby(\"level\").apply(lambda g: list(zip(g['relative_x'], g['relative_y'])), include_groups=False).reset_index(drop=False, name=\"vals\")\n    for idx, (_, row) in enumerate(p.iterrows()):\n        # Copy of img\n        img_copy= img.copy()\n        h, w = img.shape\n\n        # Extract Keypoints\n        level = row['level']\n        vals = sorted(row[\"vals\"], key=lambda x: x[0])\n        a,b= vals\n        a= (a[0]*w, a[1]*h)\n        b= (b[0]*w, b[1]*h)\n\n        # Rotate\n        rotate_angle= angle_of_line(a[0], a[1], b[0], b[1])\n        transform = A.Compose([\n            A.Rotate(limit=(-rotate_angle, -rotate_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        # Crop + Resize\n        img_copy= crop_between_keypoints(img_copy, a, b)\n        img_copy= resize_transform(image=img_copy)[\"image\"]\n        \n        # Plot\n        ax = plt.subplot(gs[idx])\n        ax.imshow(img_copy, cmap='gray')\n        ax.set_title(level)\n        ax.axis('off')\n    plt.show()\n    \n    \ndef 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    # Calculate bounding box around the keypoints\n    left = int(min(x1, x2))\n    right = int(max(x1, x2))\n    top = int(min(y1, y2) - (h * 0.1))\n    bottom = int(max(y1, y2) + (h * 0.1))\n            \n    # Crop the image\n    return img[top:bottom, left:right]","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:34:45.365802Z","iopub.execute_input":"2024-08-17T08:34:45.366147Z","iopub.status.idle":"2024-08-17T08:34:45.394513Z","shell.execute_reply.started":"2024-08-17T08:34:45.366116Z","shell.execute_reply":"2024-08-17T08:34:45.393554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Feel free to change the seed, or increase N to see more samples.","metadata":{}},{"cell_type":"code","source":"SEED= 10\nN= 2\n\n# Load series_ids\ndfd= 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\n# Load coords\ncoords= 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\n# Plot samples\nfor idx, row in dfd.iterrows():\n    try:\n        print(\"-\"*25, \" STUDY_ID: {}, SERIES_ID: {} \".format(row.study_id, row.series_id), \"-\"*25)\n        sag_t2 = load_dicom_stack(os.path.join(image_dir, str(row.study_id), str(row.series_id)), plane=\"sagittal\")\n        \n        # Img + Coords\n        img= sag_t2[\"array\"][len(sag_t2[\"array\"])//2]\n        coords_temp= coords[coords[\"series_id\"] == row.series_id].copy()\n        \n        # Plot\n        plot_img(img, coords_temp)\n        plot_5_crops(img, coords_temp)\n        \n    except Exception as e:\n        print(e)\n        pass","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:34:46.946017Z","iopub.execute_input":"2024-08-17T08:34:46.946991Z","iopub.status.idle":"2024-08-17T08:34:48.918087Z","shell.execute_reply.started":"2024-08-17T08:34:46.946947Z","shell.execute_reply":"2024-08-17T08:34:48.91718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Pretraining\n\nNext, I show a simple pipeline to train a model to predict the x,y coordinates of the 5 lower lumbar vertabrae. The idea is that we first train our image model on a similar task so that the model better suited to tackle our main objective. \n\nThis data was put together in the first version so it does not include left side coordinates. \n\nFor more information on the data, see [here](https://www.kaggle.com/datasets/brendanartley/lumbar-coordinate-pretraining-datasethttps://www.kaggle.com/datasets/brendanartley/lumbar-coordinate-pretraining-dataset).\n\n<h1 align=\"left\">\n<img src=\"https://storage.googleapis.com/kaggle-datasets-images/5464745/9091594/db0b402668602e8a6eb772a162f47eb3/dataset-cover.png?t=2024-08-02-23-33-03\" alt=\"spine_img\" width=\"700\">\n</h1>\n","metadata":{}},{"cell_type":"code","source":"# Config\ncfg= 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) # Makes results reproducable","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:34:52.005777Z","iopub.execute_input":"2024-08-17T08:34:52.006164Z","iopub.status.idle":"2024-08-17T08:34:52.041178Z","shell.execute_reply.started":"2024-08-17T08:34:52.006133Z","shell.execute_reply":"2024-08-17T08:34:52.04027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load metadata\ndf= 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-08-17T08:34:53.105909Z","iopub.execute_input":"2024-08-17T08:34:53.106262Z","iopub.status.idle":"2024-08-17T08:34:53.159643Z","shell.execute_reply.started":"2024-08-17T08:34:53.106233Z","shell.execute_reply":"2024-08-17T08:34:53.158577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset\n\nHere we define the torch dataset that will be used during training.","metadata":{}},{"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        # Convert to dict\n        d = df.groupby(\"series_id\")[[\"relative_x\", \"relative_y\"]].apply(lambda x: list(x.itertuples(index=False, name=None)))\n        records= {}\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, \"processed_{}/{}.npy\".format(source, series_id))\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        \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        img= self.load_img(source, series_id)\n        return {\n            'img': img, \n            'label': label,\n            }\n    \n    def __len__(self,):\n        return len(self.records)\n    \nds= PreTrainDataset(df, cfg)    \n\n# Plot a Single Sample\nprint(\"---- Sample Shapes -----\")\nfor k, v in ds[0].items():\n    print(k, v.shape)","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:34:54.901128Z","iopub.execute_input":"2024-08-17T08:34:54.90148Z","iopub.status.idle":"2024-08-17T08:34:55.160736Z","shell.execute_reply.started":"2024-08-17T08:34:54.90145Z","shell.execute_reply":"2024-08-17T08:34:55.15958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utils\n\n\nHere we have a couple helpers functions. \n\nThe first moves data to the GPU if enabled, the second visualizes predictions during training, and the third loads weights when dealing with mismatched shapes.","metadata":{}},{"cell_type":"code","source":"def batch_to_device(batch, device, skip_keys=[]):\n    batch_dict= {}\n    for key in batch:\n        if key in skip_keys:\n             batch_dict[key]= batch[key]\n        else:    \n            batch_dict[key]= batch[key].to(device)\n    return batch_dict\n\ndef visualize_prediction(batch, pred, epoch):\n    \n    mid= cfg.n_frames//2\n    \n    # Plot\n    for idx in range(1):\n    \n        # Select Data\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        # Plot coords\n        fig, axes = plt.subplots(1, len(coords_list), figsize=(10,4))\n        fig.suptitle(\"EPOCH: {}\".format(epoch))\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            # Add text labels near the coordinates\n            for i, (x, y) in enumerate(zip(coords[0::2], coords[1::2])):\n                if i < len(text_labels):  # Ensure there are enough labels\n                    ax.text(x + 10, y, text_labels[i], color='white', fontsize=15, bbox=dict(facecolor='black', alpha=0.5))\n\n\n        fig.suptitle(\"EPOCH: {}\".format(epoch))\n        plt.show()\n#         plt.close(fig)\n    return\n\ndef load_weights_skip_mismatch(model, weights_path, device):\n    # Load Weights\n    state_dict = torch.load(weights_path, map_location=device)\n    model_dict = model.state_dict()\n    \n    # Iter models\n    params = {}\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(\"Skipping param: {}, {} != {}\".format(sdk, sfv.size(), mdv.size()))\n    \n    # Reload + Skip\n    model.load_state_dict(params, strict=False)\n    print(\"Loaded weights from:\", weights_path)","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:34:56.347729Z","iopub.execute_input":"2024-08-17T08:34:56.348087Z","iopub.status.idle":"2024-08-17T08:34:56.36391Z","shell.execute_reply.started":"2024-08-17T08:34:56.348061Z","shell.execute_reply":"2024-08-17T08:34:56.362853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Training\n\nHere we train on all sources except for the spider dataset, which is used for validation.","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"# Dataframes\ntrain_df= df[df[\"source\"] != \"spider\"]\nval_df= df[df[\"source\"] == \"spider\"]\nprint(\"TRAIN_SIZE: {}, VAL_SIZE: {}\".format(len(train_df)//5, len(val_df)//5))\n\n# Datasets + Dataloaders\ntrain_ds= PreTrainDataset(train_df, cfg)\nval_ds= PreTrainDataset(val_df, cfg)\ntrain_dl = torch.utils.data.DataLoader(train_ds, batch_size=cfg.batch_size, shuffle=True, drop_last=True)\nval_dl = torch.utils.data.DataLoader(val_ds, batch_size=cfg.batch_size, shuffle=False)\n\n# Model\nmodel = timm.create_model('resnet18', pretrained=True, num_classes=75)\nmodel = model.to(cfg.device)\n\n# Loss / Optim\ncriterion = nn.MSELoss()\noptimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr)","metadata":{"execution":{"iopub.status.busy":"2024-08-17T09:09:17.611134Z","iopub.execute_input":"2024-08-17T09:09:17.611533Z","iopub.status.idle":"2024-08-17T09:09:18.27959Z","shell.execute_reply.started":"2024-08-17T09:09:17.611499Z","shell.execute_reply":"2024-08-17T09:09:18.278617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(cfg.epochs+1):\n    \n    # Train Loop\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    # Validation Loop\n    val_loss = 0\n    with torch.no_grad():\n        model = model.eval()\n        for batch in tqdm(val_dl):\n            batch = batch_to_device(batch, cfg.device)\n\n            pred = model(batch[\"img\"].float())\n            pred = torch.sigmoid(pred)\n            \n            val_loss += criterion(pred, batch[\"label\"].float()).item()\n        val_loss /= len(val_dl)\n            \n    # Viz\n    visualize_prediction(batch, pred, epoch)           \n            \n    print(f\"Epoch {epoch+1}, Training Loss: {loss.item()}, Validation Loss: {val_loss}\")\nprint(\"Training complete...\")","metadata":{"execution":{"iopub.status.busy":"2024-08-17T09:09:19.555717Z","iopub.execute_input":"2024-08-17T09:09:19.556547Z","iopub.status.idle":"2024-08-17T09:09:21.109057Z","shell.execute_reply.started":"2024-08-17T09:09:19.556506Z","shell.execute_reply":"2024-08-17T09:09:21.107768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model, 'model.pth')","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:37:50.585202Z","iopub.execute_input":"2024-08-17T08:37:50.585815Z","iopub.status.idle":"2024-08-17T08:37:50.671109Z","shell.execute_reply.started":"2024-08-17T08:37:50.585779Z","shell.execute_reply":"2024-08-17T08:37:50.67012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), 'model.pth')","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:42:17.116772Z","iopub.execute_input":"2024-08-17T08:42:17.117371Z","iopub.status.idle":"2024-08-17T08:42:17.220487Z","shell.execute_reply.started":"2024-08-17T08:42:17.11734Z","shell.execute_reply":"2024-08-17T08:42:17.219498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Save\n\nNext, we save the backbone weights.","metadata":{}},{"cell_type":"code","source":"f= \"{}_{}.pt\".format(cfg.backbone, cfg.seed)\ntorch.save(model.state_dict(), f)\nprint(\"Saved weights: {}\".format(f))","metadata":{"execution":{"iopub.status.busy":"2024-08-17T09:09:24.786904Z","iopub.execute_input":"2024-08-17T09:09:24.78762Z","iopub.status.idle":"2024-08-17T09:09:24.863759Z","shell.execute_reply.started":"2024-08-17T09:09:24.78759Z","shell.execute_reply":"2024-08-17T09:09:24.862816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Finally, the weights can be loaded for a new task (eg. this competition).","metadata":{}},{"cell_type":"code","source":"# Load backbone for RSNA 2024 task\nmodel = timm.create_model('resnet18', pretrained=True, num_classes=75)\nmodel = model.to(cfg.device)\nload_weights_skip_mismatch(model, f, cfg.device)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}