{"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":"# Sumary\n* baseline : train\n* Dataset : 256 png\n* Model : efficientnet\n* GPU : P100","metadata":{}},{"cell_type":"code","source":"!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2022-12-18T13:49:16.389584Z","iopub.execute_input":"2022-12-18T13:49:16.390355Z","iopub.status.idle":"2022-12-18T13:49:17.421541Z","shell.execute_reply.started":"2022-12-18T13:49:16.390315Z","shell.execute_reply":"2022-12-18T13:49:17.420184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#try:\n#    import pylibjpeg\n#except:\n#    !pip install /kaggle/input/rsna-2022-whl/{pydicom-2.3.0-py3-none-any.whl,pylibjpeg-1.4.0-py3-none-any.whl,python_gdcm-3.0.15-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl}","metadata":{"execution":{"iopub.status.busy":"2022-12-18T13:49:17.424339Z","iopub.execute_input":"2022-12-18T13:49:17.424791Z","iopub.status.idle":"2022-12-18T13:49:17.431692Z","shell.execute_reply.started":"2022-12-18T13:49:17.424745Z","shell.execute_reply":"2022-12-18T13:49:17.430408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    folds = 4\n    epochs = 3\n    BATCH = 64\n    model_name = \"efficientnet_b1\"\n    lr = 1e-4\n    weight_decay = 1e-6\n    target_size = 1\n    size = 256\n    seed = 42","metadata":{"execution":{"iopub.status.busy":"2022-12-18T13:49:17.433251Z","iopub.execute_input":"2022-12-18T13:49:17.434317Z","iopub.status.idle":"2022-12-18T13:49:17.442159Z","shell.execute_reply.started":"2022-12-18T13:49:17.434249Z","shell.execute_reply":"2022-12-18T13:49:17.441062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/timm-pytorch-image-models/pytorch-image-models-master/')\nimport timm\n\nimport numpy as np\nimport pandas as pd\nimport glob\nimport os\nimport cv2\nimport random\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom tqdm import tqdm\n\nimport torch\nfrom torch.utils.data import Dataset,DataLoader\nimport torch.nn as nn\nfrom torch.optim import Adam\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n#import pylibjpeg\n#import pydicom\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2022-12-18T13:49:17.4458Z","iopub.execute_input":"2022-12-18T13:49:17.447408Z","iopub.status.idle":"2022-12-18T13:49:17.515653Z","shell.execute_reply.started":"2022-12-18T13:49:17.447381Z","shell.execute_reply":"2022-12-18T13:49:17.514593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# utils","metadata":{}},{"cell_type":"code","source":"#https://www.kaggle.com/code/sohier/probabilistic-f-score\n\ndef pfbeta(labels, predictions, beta):\n    y_true_count = 0\n    ctp = 0\n    cfp = 0\n\n    for idx in range(len(labels)):\n        prediction = min(max(predictions[idx], 0), 1)\n        if (labels[idx]):\n            y_true_count += 1\n            ctp += prediction\n        else:\n            cfp += prediction\n\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp)\n    c_recall = ctp / y_true_count\n    if (c_precision > 0 and c_recall > 0):\n        result = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall)\n        return result\n    else:\n        return 0","metadata":{"execution":{"iopub.status.busy":"2022-12-18T13:49:17.519565Z","iopub.execute_input":"2022-12-18T13:49:17.519868Z","iopub.status.idle":"2022-12-18T13:49:17.528297Z","shell.execute_reply.started":"2022-12-18T13:49:17.519841Z","shell.execute_reply":"2022-12-18T13:49:17.527181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DataFrame","metadata":{}},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/train.csv')\ntest = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/test.csv')\n\ndisplay(train)\ndisplay(test)","metadata":{"execution":{"iopub.status.busy":"2022-12-18T13:49:17.531401Z","iopub.execute_input":"2022-12-18T13:49:17.532447Z","iopub.status.idle":"2022-12-18T13:49:17.684252Z","shell.execute_reply.started":"2022-12-18T13:49:17.532279Z","shell.execute_reply":"2022-12-18T13:49:17.683218Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# image","metadata":{}},{"cell_type":"code","source":"TRAIN_PATH = '/kaggle/input/rsna-breast-cancer-256-pngs/'\nTEST_PATH = '/kaggle/input/rsna-breast-cancer-detection/test_images/'","metadata":{"execution":{"iopub.status.busy":"2022-12-18T13:49:17.685696Z","iopub.execute_input":"2022-12-18T13:49:17.686795Z","iopub.status.idle":"2022-12-18T13:49:17.691259Z","shell.execute_reply.started":"2022-12-18T13:49:17.686753Z","shell.execute_reply":"2022-12-18T13:49:17.690175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# dataset","metadata":{}},{"cell_type":"code","source":"class RSNADataset(Dataset):\n    def __init__(self, df, transform=None): \n        self.df = df\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self,idx):\n        row = self.df.iloc[idx]\n        img_path = TRAIN_PATH + f\"{row.patient_id}_{row.image_id}.png\"\n        image = cv2.imread(img_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        if self.transform:\n            image = self.transform(image=image)['image']\n        target = torch.tensor(row.cancer).float()\n        return image, target","metadata":{"execution":{"iopub.status.busy":"2022-12-18T13:49:17.692764Z","iopub.execute_input":"2022-12-18T13:49:17.69359Z","iopub.status.idle":"2022-12-18T13:49:17.704667Z","shell.execute_reply.started":"2022-12-18T13:49:17.693552Z","shell.execute_reply":"2022-12-18T13:49:17.703774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transforms(data):\n    if data == 'train':\n        return A.Compose([\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])\n    elif data == 'valid':\n        return A.Compose([\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])","metadata":{"execution":{"iopub.status.busy":"2022-12-18T13:49:17.706089Z","iopub.execute_input":"2022-12-18T13:49:17.707202Z","iopub.status.idle":"2022-12-18T13:49:17.715018Z","shell.execute_reply.started":"2022-12-18T13:49:17.707164Z","shell.execute_reply":"2022-12-18T13:49:17.713953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Dataset check\n\ntrain_dataset = RSNADataset(train,transform=get_transforms(data='train'))\nimg,target = train_dataset[0]\nprint(img.shape)\nplt.imshow(img[0])","metadata":{"execution":{"iopub.status.busy":"2022-12-18T13:49:17.718646Z","iopub.execute_input":"2022-12-18T13:49:17.718972Z","iopub.status.idle":"2022-12-18T13:49:18.002669Z","shell.execute_reply.started":"2022-12-18T13:49:17.718946Z","shell.execute_reply":"2022-12-18T13:49:18.001759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# model","metadata":{}},{"cell_type":"code","source":"class RSNAModel(nn.Module):\n    def __init__(self, cfg, pretrained=False):\n        super().__init__()\n        self.cfg = cfg\n        self.model = timm.create_model(self.cfg.model_name, pretrained=pretrained)\n        in_features = self.model.classifier.in_features\n        self.model.classifier = nn.Linear(in_features, self.cfg.target_size)\n                    \n    def forward(self, image):\n        output = self.model(image)\n        return output","metadata":{"execution":{"iopub.status.busy":"2022-12-18T13:49:18.004529Z","iopub.execute_input":"2022-12-18T13:49:18.005563Z","iopub.status.idle":"2022-12-18T13:49:18.012489Z","shell.execute_reply.started":"2022-12-18T13:49:18.005524Z","shell.execute_reply":"2022-12-18T13:49:18.011243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# train","metadata":{}},{"cell_type":"code","source":"#debug\n#train = train.head(1000)","metadata":{"execution":{"iopub.status.busy":"2022-12-18T13:49:18.014015Z","iopub.execute_input":"2022-12-18T13:49:18.014489Z","iopub.status.idle":"2022-12-18T13:49:18.023289Z","shell.execute_reply.started":"2022-12-18T13:49:18.014404Z","shell.execute_reply":"2022-12-18T13:49:18.022386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#CV\nskf = StratifiedGroupKFold(CFG.folds)\n\n#train & valid\nfor fold, (tr_idx, val_idx) in enumerate(skf.split(train, y=train.cancer, groups=train.patient_id)):\n    print(f'========= fold : {fold} =========')\n    \n    train_dataset = RSNADataset(train.loc[tr_idx,:],transform=get_transforms(data='train'))\n    valid_dataset = RSNADataset(train.loc[val_idx,:],transform=get_transforms(data='valid'))\n    \n    train_dataloader = DataLoader(train_dataset, batch_size=CFG.BATCH, shuffle=True, num_workers=4)\n    valid_dataloader = DataLoader(valid_dataset, batch_size=CFG.BATCH, shuffle=False, num_workers=4)\n    \n    model = RSNAModel(CFG,pretrained=True)\n    model.to(device)\n    \n    criterion = nn.BCEWithLogitsLoss()\n    optimizer = Adam(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\n    \n    for epoch in range(CFG.epochs):\n        #train\n        print(f'epoch:{epoch}')\n        model.train()\n        \n        train_losses = 0\n        val_losses = 0\n        train_targets,train_preds = [],[]\n        val_targets,val_preds = [],[]\n\n        for image,target in tqdm(train_dataloader):\n            image = image.to(device)\n            train_target = target.to(device)\n            train_pred = model(image)\n            train_loss = criterion(train_pred.view(-1),train_target)\n            train_losses += train_loss.item()\n            optimizer.zero_grad()\n            train_loss.backward()\n            optimizer.step()\n            \n            train_targets.extend(train_target.cpu().detach().tolist())\n            train_preds.extend(torch.sigmoid(train_pred.view(-1)).cpu().detach().tolist())\n            \n        avg_train_loss = train_losses / len(train_dataloader)\n        train_pfbeta_score = pfbeta(train_targets, train_preds, beta=0.5)\n        \n    \n        #valid\n        model.eval()\n        with torch.no_grad():\n            for img, target in tqdm(valid_dataloader):\n                img = img.to(device)\n                val_target = target.to(device)\n                val_pred = model(img)\n                val_loss = criterion(val_pred.view(-1),val_target)\n                val_losses += val_loss.item()\n                \n                val_targets.extend(val_target.cpu().detach().tolist())\n                val_preds.extend(torch.sigmoid(val_pred.view(-1)).cpu().detach().tolist())\n                \n        avg_val_loss = val_losses / len(valid_dataloader)\n        val_pfbeta_score = pfbeta(val_targets, val_preds, beta=0.5)\n        \n        #score\n        print('avg_train_loss:',avg_train_loss,'avg_val_loss:',avg_val_loss)\n        print('train_pfbeta_score:',train_pfbeta_score,'val_pfbeta_score:',val_pfbeta_score)\n        \n    #model save\n    torch.save(model.state_dict(), f'super_simple_baseline_model{fold}.pth')","metadata":{"execution":{"iopub.status.busy":"2022-12-18T13:49:18.024874Z","iopub.execute_input":"2022-12-18T13:49:18.025303Z","iopub.status.idle":"2022-12-18T14:52:03.460556Z","shell.execute_reply.started":"2022-12-18T13:49:18.025206Z","shell.execute_reply":"2022-12-18T14:52:03.459178Z"},"trusted":true},"execution_count":null,"outputs":[]}]}