{"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":"code","source":"!pip install -q timm wandb torch_xla\n!pip -q install torchviz","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-08-22T08:14:14.096049Z","iopub.execute_input":"2023-08-22T08:14:14.096333Z","iopub.status.idle":"2023-08-22T08:14:43.03236Z","shell.execute_reply.started":"2023-08-22T08:14:14.096286Z","shell.execute_reply":"2023-08-22T08:14:43.030996Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport wandb\nfrom tqdm import tqdm\nfrom glob import glob\n\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nimport imageio\n\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.distributed as dist\nimport torchviz\nimport torch.optim as optim\n#import torch_xla.core.xla_model as xm\n#import torch_xla.distributed.xla_multiprocessing as xmp\nimport torchvision\nfrom torchvision import transforms\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn import metrics\n\nfrom PIL import Image\nfrom pathlib import Path","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:43.034852Z","iopub.execute_input":"2023-08-22T08:14:43.035237Z","iopub.status.idle":"2023-08-22T08:14:48.276408Z","shell.execute_reply.started":"2023-08-22T08:14:43.035198Z","shell.execute_reply":"2023-08-22T08:14:48.275326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    api_key = user_secrets.get_secret(\"WANDB\")\n\n    wandb.login(key=api_key)\n    anonymous = None\nexcept:\n    anonymous = \"must\"\n    print('To use your W&B account,\\nGo to Add-ons -> Secrets and provide your W&B access token. Use the Label name as WANDB. \\nGet your W&B access token from here: https://wandb.ai/authorize')","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:48.277903Z","iopub.execute_input":"2023-08-22T08:14:48.278602Z","iopub.status.idle":"2023-08-22T08:14:51.110023Z","shell.execute_reply.started":"2023-08-22T08:14:48.278562Z","shell.execute_reply":"2023-08-22T08:14:51.109017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config:\n    wandb=True\n    competition='rsna-atd'\n    _wandbkernel=\"hemanthh17\"\n    debug=False\n    comment=\"swin_trans_v1\"\n    exp_name=\"baseline_Testing\"\n    verbose=0\n    display_plot=True\n    device='GPU'\n    model_name='swin_base_patch4_window7_224'\n    seed=42\n    folds=4\n    selected_folds=[0,1,2]\n    img_size = [221, 221]\n#     eq_dim = np.prod(img_size)**0.5\n\n    # batch_size and epochs\n    batch_size = 48\n    epochs = 24\n\n    # loss\n    loss      = 'BCE & CCE'  # BCE, Focal\n    \n    # optimizer\n    optimizer = 'Adam'\n\n    \n    # test-time augs\n    tta = 1\n    \n    # target column\n    target_col  = [ \"bowel_injury\", \"extravasation_injury\", \"kidney_healthy\", \"kidney_low\",\n                   \"kidney_high\", \"liver_healthy\", \"liver_low\", \"liver_high\",\n                   \"spleen_healthy\", \"spleen_low\", \"spleen_high\"] # not using \"bowel_healthy\" & \"extravasation_healthy\"","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:51.112274Z","iopub.execute_input":"2023-08-22T08:14:51.112859Z","iopub.status.idle":"2023-08-22T08:14:51.121717Z","shell.execute_reply.started":"2023-08-22T08:14:51.112831Z","shell.execute_reply":"2023-08-22T08:14:51.120369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if \"TPU\" in Config.device:\n    tpu = 'local' if Config.device == 'TPU-VM' else None\n    print(\"Connecting to TPU...\")\n    try:\n        # Connect to TPU\n        # Note: PyTorch does not have a direct equivalent of TPUClusterResolver and TPUStrategy like TensorFlow\n        tpu = xm.xla_device()\n        ngpu = 0\n        dist_backend = 'xla'\n        device = tpu\n    except:\n        # If TPU connection fails, switch to GPU or CPU\n        Config.device = \"GPU\"  # or \"CPU\"\n        ngpu = torch.cuda.device_count()\n        device = torch.device(\"cuda\" if ngpu > 0 else \"cpu\")\n\nif Config.device == \"GPU\":\n    ngpu = torch.cuda.device_count()\n    print(\"Num GPUs Available: \", ngpu)\n    device = torch.device(\"cuda\" if ngpu > 0 else \"cpu\")\n\nif Config.device == \"CPU\":\n    device = torch.device(\"cpu\")","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:51.123191Z","iopub.execute_input":"2023-08-22T08:14:51.123946Z","iopub.status.idle":"2023-08-22T08:14:51.153307Z","shell.execute_reply.started":"2023-08-22T08:14:51.123913Z","shell.execute_reply":"2023-08-22T08:14:51.152344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BASE_PATH = f'/kaggle/input/rsna-atd-512x512-png-v2-dataset'\n#BASE_PATH1 = f'/kaggle/input/rsna-atd-512x512-png-v3-dataset'","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:51.154773Z","iopub.execute_input":"2023-08-22T08:14:51.155121Z","iopub.status.idle":"2023-08-22T08:14:51.159683Z","shell.execute_reply.started":"2023-08-22T08:14:51.155091Z","shell.execute_reply":"2023-08-22T08:14:51.158524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df= pd.read_csv(f\"{BASE_PATH}/train.csv\")\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:51.161566Z","iopub.execute_input":"2023-08-22T08:14:51.161973Z","iopub.status.idle":"2023-08-22T08:14:51.265525Z","shell.execute_reply.started":"2023-08-22T08:14:51.16194Z","shell.execute_reply":"2023-08-22T08:14:51.264369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['image_path'] = f'{BASE_PATH}/train_images'\\\n                    + '/' + df.patient_id.astype(str)\\\n                    + '/' + df.series_id.astype(str)\\\n                    + '/' + df.instance_number.astype(str) +'.png'\ndf=df.drop_duplicates()\n","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:51.267391Z","iopub.execute_input":"2023-08-22T08:14:51.267761Z","iopub.status.idle":"2023-08-22T08:14:51.331103Z","shell.execute_reply.started":"2023-08-22T08:14:51.267728Z","shell.execute_reply":"2023-08-22T08:14:51.330171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.describe()","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:51.332641Z","iopub.execute_input":"2023-08-22T08:14:51.333002Z","iopub.status.idle":"2023-08-22T08:14:51.409155Z","shell.execute_reply.started":"2023-08-22T08:14:51.33297Z","shell.execute_reply":"2023-08-22T08:14:51.408028Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df= pd.read_csv(f\"{BASE_PATH}/test.csv\")\ntest_df['image_path'] = f'{BASE_PATH}/test_images'\\\n                    + '/' + test_df.patient_id.astype(str)\\\n                    + '/' + test_df.series_id.astype(str)\\\n                    + '/' + test_df.instance_number.astype(str) +'.png'\ntest_df = test_df.drop_duplicates()\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:51.41393Z","iopub.execute_input":"2023-08-22T08:14:51.414215Z","iopub.status.idle":"2023-08-22T08:14:51.439679Z","shell.execute_reply.started":"2023-08-22T08:14:51.414189Z","shell.execute_reply":"2023-08-22T08:14:51.43858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df.describe()","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:51.441247Z","iopub.execute_input":"2023-08-22T08:14:51.441704Z","iopub.status.idle":"2023-08-22T08:14:51.476203Z","shell.execute_reply.started":"2023-08-22T08:14:51.441668Z","shell.execute_reply":"2023-08-22T08:14:51.474839Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def is_exist(path):\n    if os.path.exists(path):\n        print(\"it exists\")\n    else:\n        print(\"it does not exist!!\")\n    \nis_exist(df.image_path[0]),is_exist(test_df.image_path[0])","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:51.480218Z","iopub.execute_input":"2023-08-22T08:14:51.480887Z","iopub.status.idle":"2023-08-22T08:14:51.507138Z","shell.execute_reply.started":"2023-08-22T08:14:51.48085Z","shell.execute_reply":"2023-08-22T08:14:51.506093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.shape[0],test_df.shape[0]","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:51.508568Z","iopub.execute_input":"2023-08-22T08:14:51.50964Z","iopub.status.idle":"2023-08-22T08:14:51.516765Z","shell.execute_reply.started":"2023-08-22T08:14:51.509603Z","shell.execute_reply":"2023-08-22T08:14:51.515655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import seaborn as sns\nimport matplotlib.pyplot as plt\n\nsns.set_theme(style=\"white\")\n\n# Compute the correlation matrix\ncorr = df[df.columns[1:14]].corr().round(2)\n\n# Generate a mask for the upper triangle\nmask = np.triu(np.ones_like(corr, dtype=bool))\n\n# Generate a custom diverging colormap\ncmap = sns.diverging_palette(230, 20, as_cmap=True)\n\n# Set up the matplotlib figure\nf, ax = plt.subplots(figsize=(11, 9))\n\n# Draw the heatmap with the mask and correct aspect ratio\nsns.heatmap(corr, mask=mask, annot=True, cmap=cmap, vmax=.3, center=0,\n            square=True, linewidths=.5, cbar_kws={\"shrink\": .5})\nplt.show();\n\n","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:51.518318Z","iopub.execute_input":"2023-08-22T08:14:51.519408Z","iopub.status.idle":"2023-08-22T08:14:52.368711Z","shell.execute_reply.started":"2023-08-22T08:14:51.519374Z","shell.execute_reply":"2023-08-22T08:14:52.367713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['stratify'] = ''\nfor col in Config.target_col:\n    df['stratify'] += df[col].astype(str)\n\ndf = df.reset_index(drop=True)\nskf = StratifiedGroupKFold(n_splits=Config.folds, shuffle=True, random_state=Config.seed)\nfor fold, (train_idx, val_idx) in enumerate(skf.split(df, df['stratify'], df[\"patient_id\"])):\n    df.loc[val_idx, 'fold'] = fold\ndisplay(df.groupby(['fold', 'patient_id']).size())","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:52.369797Z","iopub.execute_input":"2023-08-22T08:14:52.370581Z","iopub.status.idle":"2023-08-22T08:14:52.649054Z","shell.execute_reply.started":"2023-08-22T08:14:52.370544Z","shell.execute_reply":"2023-08-22T08:14:52.647882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.columns","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:52.650927Z","iopub.execute_input":"2023-08-22T08:14:52.651338Z","iopub.status.idle":"2023-08-22T08:14:52.658976Z","shell.execute_reply.started":"2023-08-22T08:14:52.651302Z","shell.execute_reply":"2023-08-22T08:14:52.657976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Prepearation","metadata":{}},{"cell_type":"code","source":"class RSNAATDDataset(torch.utils.data.Dataset):\n    def __init__(self,data,transform=None):\n        self.data=data\n        self.transform=transform\n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self,idx):\n        img_path=self.data.loc[idx,'image_path']\n        image= Image.open(Path(img_path)).convert('RGB')\n        extra=self.data[[col for col in self.data.columns if col in ['extravasation_healthy',\n                                                                         'extravasation_injury']]].values[idx].astype('int')\n        spleen=self.data[[col for col in self.data.columns if col in ['spleen_healthy',\n                                                                          'spleen_low', 'spleen_high']]].values[idx].astype('int')\n        bowel=self.data[[col for col in self.data.columns if col in ['bowel_healthy', 'bowel_injury']]].values[idx].astype('int')\n        liver=self.data[[col for col in self.data.columns if col in ['liver_healthy',\n                                                                         'liver_low', 'liver_high']]].values[idx].astype('int')\n        kidney=self.data[[col for col in self.data.columns if col in ['kidney_healthy',\n                                                                          'kidney_low', 'kidney_high']]].values[idx].astype('int')\n        if self.transform:\n            image = self.transform(image)\n            \n        return {\"image\":image,\n                \"extra\": torch.tensor(extra,dtype=torch.float32),\n                \"spleen\":torch.tensor(spleen,dtype=torch.float32),\n                \"bowel\":torch.tensor(bowel,dtype=torch.float32),\n                \"liver\":torch.tensor(liver,dtype=torch.float32),\n                \"kidney\":torch.tensor(kidney,dtype=torch.float32)}\n        ","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:52.66054Z","iopub.execute_input":"2023-08-22T08:14:52.661358Z","iopub.status.idle":"2023-08-22T08:14:52.675697Z","shell.execute_reply.started":"2023-08-22T08:14:52.661225Z","shell.execute_reply":"2023-08-22T08:14:52.674489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transform = transforms.Compose([\n    transforms.Resize((256,256)),\n    transforms.RandomResizedCrop(224),\n    transforms.RandomAutocontrast(),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1),\n    transforms.RandomRotation(20),\n    transforms.ToTensor()\n    \n])\n\ntest_transform = transforms.Compose([\n    transforms.Resize((256,256)),\n    transforms.ToTensor()    \n])\n\n# Create dataset instances\ndfk= df[df['fold'].isin([0,1,2,3])].reset_index(drop=True)\n\ntrain_dataset = RSNAATDDataset(dfk, transform=train_transform)\nvalid_dataset = RSNAATDDataset(test_df, transform=test_transform)\ndfk.head()\n","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:52.679194Z","iopub.execute_input":"2023-08-22T08:14:52.6799Z","iopub.status.idle":"2023-08-22T08:14:52.7132Z","shell.execute_reply.started":"2023-08-22T08:14:52.679864Z","shell.execute_reply":"2023-08-22T08:14:52.712279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_bs=16\ntest_bs=8\nbatch_size = 32\ntrain_loader = torch.utils.data.DataLoader(train_dataset, batch_size=train_bs, shuffle=True, num_workers=4)\nvalid_loader = torch.utils.data.DataLoader(valid_dataset, batch_size=test_bs,num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:52.714627Z","iopub.execute_input":"2023-08-22T08:14:52.714958Z","iopub.status.idle":"2023-08-22T08:14:52.72315Z","shell.execute_reply.started":"2023-08-22T08:14:52.714933Z","shell.execute_reply":"2023-08-22T08:14:52.722036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_path=\"/kaggle/input/rsna-atd-512x512-png-v2-dataset/test_images/48843/62825/30.png\"\nimport imageio\nimg=imageio.imread(img_path)\nimg.shape","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:52.724623Z","iopub.execute_input":"2023-08-22T08:14:52.725402Z","iopub.status.idle":"2023-08-22T08:14:52.757215Z","shell.execute_reply.started":"2023-08-22T08:14:52.725369Z","shell.execute_reply":"2023-08-22T08:14:52.756272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNAViTTransformer(nn.Module):\n    def __init__(self):\n        super(RSNAViTTransformer,self).__init__()\n        self.prime_model= timm.create_model('swin_base_patch4_window7_224', pretrained=True)\n        self.linear_common1=nn.Linear(1000,32)\n        self.linear_common2=nn.Linear(32,4)\n        self.linear_extra=nn.Linear(4,2)\n        self.linear_bowel=nn.Linear(4,2)\n        self.linear_liver=nn.Linear(4,3)\n        self.linear_kidney=nn.Linear(4,3)\n        self.linear_spleen=nn.Linear(4,3)\n        self.selu= nn.Hardswish()\n        \n        \n    def forward(self, img):\n        \n        out = self.prime_model(img)\n        \n        out = self.selu(self.linear_common2(self.selu(self.linear_common1(self.prime_model(img)))))\n        \n        out_extra = torch.sigmoid(self.linear_extra(out))\n        \n        out_bowel = torch.sigmoid(self.linear_bowel(out))\n        out_liver = torch.softmax(self.linear_liver(out),dim=1)\n        out_kidney = torch.softmax(self.linear_kidney(out),dim=1)\n        \n        out_spleen = torch.softmax(self.linear_spleen(out),dim=1)\n        \n        combined_output = out_extra, out_bowel, out_kidney, out_liver, out_spleen\n    \n        return combined_output\n","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:52.758509Z","iopub.execute_input":"2023-08-22T08:14:52.759408Z","iopub.status.idle":"2023-08-22T08:14:52.769906Z","shell.execute_reply.started":"2023-08-22T08:14:52.759377Z","shell.execute_reply":"2023-08-22T08:14:52.768796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model= RSNAViTTransformer()","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:14:52.771117Z","iopub.execute_input":"2023-08-22T08:14:52.772513Z","iopub.status.idle":"2023-08-22T08:15:00.561189Z","shell.execute_reply.started":"2023-08-22T08:14:52.77248Z","shell.execute_reply":"2023-08-22T08:15:00.560179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataloader = torch.utils.data.DataLoader(train_dataset, batch_size=10, shuffle=True)\n\n# Get the next batch of data using the next() function\nbatch = next(iter(dataloader))\n\n# Extract images and labels from the batch\nd = batch\n\n# Print the shape of the images and labels\n#print(\"Image shape:\", images.shape)\nprint(\"Image Dimension:\",d['image'].shape)\nprint(\"Extravastial Data Dimension:\",d['extra'].shape)\nprint(\"Kidney Data Dimensions:\",d['kidney'].shape)","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:15:00.56261Z","iopub.execute_input":"2023-08-22T08:15:00.563249Z","iopub.status.idle":"2023-08-22T08:15:00.827353Z","shell.execute_reply.started":"2023-08-22T08:15:00.563216Z","shell.execute_reply":"2023-08-22T08:15:00.826287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"optimizer= optim.RMSprop(model.parameters(),lr=2e-4)\ncriterion_extra = nn.BCELoss()\ncriterion_bowel = nn.BCELoss()\ncriterion_liver = nn.CrossEntropyLoss()\ncriterion_kidney = nn.CrossEntropyLoss()\ncriterion_spleen = nn.CrossEntropyLoss()\nscheduler1= optim.lr_scheduler.CosineAnnealingLR(optimizer,T_max=2)\nscheduler2=optim.lr_scheduler.ReduceLROnPlateau(optimizer,patience=3)\ndevice= torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel=model.to(device)\nwandb.init(project=Config.exp_name, name=Config.model_name)\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:15:00.828969Z","iopub.execute_input":"2023-08-22T08:15:00.829366Z","iopub.status.idle":"2023-08-22T08:15:36.273087Z","shell.execute_reply.started":"2023-08-22T08:15:00.829331Z","shell.execute_reply":"2023-08-22T08:15:36.272074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in tqdm(range(Config.epochs),dynamic_ncols=False, bar_format='{l_bar}{bar}| {n_fmt}/{total_fmt}'):\n    model.train()\n    total_loss = 0\n    \n    for batch in tqdm(train_loader,dynamic_ncols=False, bar_format='{l_bar}{bar}| {n_fmt}/{total_fmt}'):\n        images = batch['image'].to(device)\n        labels_extra = batch['extra'].to(device)\n        labels_bowel = batch['bowel'].to(device)\n        labels_liver = batch['liver'].to(device)\n        labels_kidney = batch['kidney'].to(device)\n        labels_spleen = batch['spleen'].to(device)\n        \n        optimizer.zero_grad()\n        \n        # Forward pass\n        outputs = model(images)\n        out_extra, out_bowel, out_kidney, out_liver, out_spleen = outputs\n        \n        # Calculate loss for each component\n        loss_extra = criterion_extra(out_extra, labels_extra)\n        loss_bowel = criterion_bowel(out_bowel, labels_bowel)\n        loss_liver = criterion_liver(out_liver, labels_liver)\n        loss_kidney = criterion_kidney(out_kidney, labels_kidney)\n        loss_spleen = criterion_spleen(out_spleen, labels_spleen)\n        \n        wandb.log({\n            'loss_extra': loss_extra.item(),\n            'loss_bowel': loss_bowel.item(),\n            'loss_kidney': loss_kidney.item(),\n            'loss_liver': loss_liver.item(),\n            'loss_spleen': loss_spleen.item()\n        })\n        \n        # Total loss\n        loss = loss_extra + loss_bowel + loss_liver + loss_kidney + loss_spleen\n        total_loss += loss.item()\n        \n        # Backpropagation and optimization\n        loss.backward()\n        optimizer.step()\n    scheduler1.step()\n    scheduler2.step(total_loss/len(train_loader))\n    \n    print(f\"Epoch [{epoch+1}/{Config.epochs}] - Loss: {total_loss / len(train_loader):.4f}\")\n    wandb.log({'epoch': epoch, 'avg_loss': (total_loss / len(train_loader))})","metadata":{"execution":{"iopub.status.busy":"2023-08-22T08:15:42.089461Z","iopub.execute_input":"2023-08-22T08:15:42.089887Z","iopub.status.idle":"2023-08-22T08:16:10.857507Z","shell.execute_reply.started":"2023-08-22T08:15:42.089854Z","shell.execute_reply":"2023-08-22T08:16:10.85453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wandb.finish()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), 'rsna_atd_swin.pth')","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}