{"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":"# Welcome to the SPR Age Prediction Challenge\n\nThis challenge aims to train a model capable of predicting the patient's age based on a chest X-ray.\n\nAlthough this is a simple notebook, you can use it as a base to improve it, testing new network architectures, including augmentations, changing learning rates, etc.\n\nGood competition!\n\n**_PS: Before you start, click on the tab Notebook options that is located at the right and select a GPU Accelerator. This is important to accelerate your training._**\n\n#### Acknowledgements\nThis Jupyter Notebook was based on code by Y. Nakama (https://www.kaggle.com/yasufuminakama) in the SETI Challenge","metadata":{}},{"cell_type":"code","source":"!pip install albumentations==0.4.6\n!pip install timm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-03-03T10:48:40.064783Z","iopub.execute_input":"2023-03-03T10:48:40.065533Z","iopub.status.idle":"2023-03-03T10:49:03.067392Z","shell.execute_reply.started":"2023-03-03T10:48:40.065486Z","shell.execute_reply":"2023-03-03T10:49:03.066108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Library\n# ====================================================\n\nimport os\nimport math\nimport time\nimport random\nimport shutil\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\nfrom glob import glob\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn import preprocessing\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold\nfrom sklearn.metrics import mean_absolute_error\n\nfrom tqdm.auto import tqdm\nfrom functools import partial\n\nimport cv2\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD\nimport torchvision.models as models\nfrom torch.nn.parameter import Parameter\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\n\nimport timm\n\nfrom torch.cuda.amp import autocast, GradScaler\n\nimport warnings \nwarnings.filterwarnings('ignore')\n\nfrom matplotlib import pyplot as plt\nimport seaborn as sns\n\n#from preprocess import resize_pad, preprocess\n\n\nVER = '001SPR'\nGPU = '0'\npixels = 224\n\ndevice = torch.device('cuda:' + GPU if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:49:03.071242Z","iopub.execute_input":"2023-03-03T10:49:03.071571Z","iopub.status.idle":"2023-03-03T10:49:07.738599Z","shell.execute_reply.started":"2023-03-03T10:49:03.071534Z","shell.execute_reply":"2023-03-03T10:49:07.737536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/spr-x-ray-age/train_age.csv')\n\n\ndef get_train_file_path(image_id):\n    return \"/kaggle/input/spr-x-ray-age/kaggle/kaggle/train/\" + str(image_id).zfill(6) + \".png\"\n\n\ntrain['file_path'] = train['imageId'].apply(get_train_file_path)\n\nprint(train['age'].isnull().values.any())\n\ndisplay(train.head())\n","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:49:07.740625Z","iopub.execute_input":"2023-03-03T10:49:07.741007Z","iopub.status.idle":"2023-03-03T10:49:07.784572Z","shell.execute_reply.started":"2023-03-03T10:49:07.740964Z","shell.execute_reply":"2023-03-03T10:49:07.783505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(8, 24))\nfor i in range(10):\n    image = cv2.imread(train.loc[i, 'file_path']) # (6, 273, 256)\n    #image = image.astype(np.float32)\n    image = image / 255.\n    image2 = np.zeros((pixels,pixels,3))\n    for j in range(image.shape[2]):\n        image2[..., j] = cv2.resize(image[...,j], (pixels, pixels))\n    plt.subplot(5, 2, i + 1)\n    plt.title(str(train.loc[i, 'age']) + ' years')\n    plt.imshow(image2)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:49:07.787053Z","iopub.execute_input":"2023-03-03T10:49:07.787517Z","iopub.status.idle":"2023-03-03T10:49:09.656608Z","shell.execute_reply.started":"2023-03-03T10:49:07.787479Z","shell.execute_reply":"2023-03-03T10:49:09.653532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train['age'].hist();","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:49:09.657593Z","iopub.execute_input":"2023-03-03T10:49:09.657903Z","iopub.status.idle":"2023-03-03T10:49:09.914165Z","shell.execute_reply.started":"2023-03-03T10:49:09.657872Z","shell.execute_reply":"2023-03-03T10:49:09.913083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\nOUTPUT_DIR = 'model/' + VER + '/'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:49:09.915867Z","iopub.execute_input":"2023-03-03T10:49:09.916234Z","iopub.status.idle":"2023-03-03T10:49:09.922692Z","shell.execute_reply.started":"2023-03-03T10:49:09.916195Z","shell.execute_reply":"2023-03-03T10:49:09.921597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# CFG\n# ====================================================\nclass CFG:\n    apex=False\n    debug=False\n    print_freq=50\n    num_workers=10\n    model_name='vgg11' \n    size_x=pixels\n    size_y=pixels\n    scheduler='CosineAnnealingLR' # ['ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts']\n    epochs=1\n    #factor=0.2 # ReduceLROnPlateau\n    #patience=4 # ReduceLROnPlateau\n    #eps=1e-6 # ReduceLROnPlateau\n    T_max=6 # CosineAnnealingLR\n    #T_0=6 # CosineAnnealingWarmRestarts\n    lr=3e-4\n    min_lr=3e-7\n    batch_size=16\n    weight_decay=1e-6\n    gradient_accumulation_steps=1\n    max_grad_norm=1000\n    seed=42\n    target_size=1\n    target_col='age'\n    n_fold=3\n    trn_fold=[0]\n    tst_fold = 9\n    train=True\n    \nif CFG.debug:\n    CFG.epochs = 1\n    train = train.sample(n=1000, random_state=CFG.seed).reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:50:37.570098Z","iopub.execute_input":"2023-03-03T10:50:37.570836Z","iopub.status.idle":"2023-03-03T10:50:37.578148Z","shell.execute_reply.started":"2023-03-03T10:50:37.570796Z","shell.execute_reply":"2023-03-03T10:50:37.576833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Utils\n# ====================================================\ndef get_score(y_true, y_pred):\n    #score = np.mean(np.abs(y_true - y_pred)) * 200.\n    score = mean_absolute_error(y_true, y_pred) * 200.\n    return score\n\n\ndef init_logger(log_file=OUTPUT_DIR+'train.log'):\n    from logging import getLogger, INFO, FileHandler,  Formatter,  StreamHandler\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\nLOGGER = init_logger()\n\n\ndef seed_torch(seed=42):\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 = True\n\nseed_torch(seed=CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:50:38.335427Z","iopub.execute_input":"2023-03-03T10:50:38.336475Z","iopub.status.idle":"2023-03-03T10:50:38.348167Z","shell.execute_reply.started":"2023-03-03T10:50:38.336433Z","shell.execute_reply":"2023-03-03T10:50:38.347192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Fold = StratifiedKFold(n_splits=CFG.n_fold, shuffle=True, random_state=CFG.seed)\nfor n, (train_index, val_index) in enumerate(Fold.split(train, train[CFG.target_col])):\n    train.loc[val_index, 'fold'] = int(n)\ntrain['fold'] = train['fold'].astype(int)\ndisplay(train.groupby(['fold', 'age']).size())","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:50:38.742166Z","iopub.execute_input":"2023-03-03T10:50:38.742862Z","iopub.status.idle":"2023-03-03T10:50:38.771712Z","shell.execute_reply.started":"2023-03-03T10:50:38.742824Z","shell.execute_reply":"2023-03-03T10:50:38.770515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Dataset\n# ====================================================\nclass TrainDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.file_names = df['file_path'].values\n        self.labels = df[CFG.target_col].values\n        #self.gender = df['male'].values * 1.\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_path = self.file_names[idx]\n        image = cv2.imread(file_path)\n        image = image.astype(np.float32) / 255.\n\n        image2 = np.zeros((CFG.size_x, CFG.size_y,3))\n        \n        for j in range(image.shape[2]):\n            image2[..., j] = cv2.resize(image[...,j], (pixels, pixels))\n        \n        image = image2.astype(np.float32)\n        \n        if self.transform:\n            image = self.transform(image=image)['image']\n        else:\n            image = image[np.newaxis,:,:]\n            image = torch.from_numpy(image).float()\n        \n        #image[:16, :16, ...] = self.gender[idx]\n        \n        label = torch.tensor(self.labels[idx]).float() / 200.\n        \n        return image, label\n    \nclass ValidDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.file_names = df['file_path'].values\n        self.labels = df[CFG.target_col].values\n        #self.gender = df['male'].values * 1.\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_path = self.file_names[idx]\n        image = cv2.imread(file_path)\n        image = image.astype(np.float32) / 255.\n        \n        image2 = np.zeros((CFG.size_x, CFG.size_y,3))\n        \n        for j in range(image.shape[2]):\n            image2[..., j] = cv2.resize(image[...,j], (pixels, pixels))\n        \n        image = image2.astype(np.float32)\n        \n        if self.transform:\n            image = self.transform(image=image)['image']\n        else:\n            image = image[np.newaxis,:,:]\n            image = torch.from_numpy(image).float()\n        \n        #image[:16, :16, ...] = self.gender[idx]\n        \n        label = torch.tensor(self.labels[idx]).float() / 200.\n        return image, label","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:50:38.948211Z","iopub.execute_input":"2023-03-03T10:50:38.948877Z","iopub.status.idle":"2023-03-03T10:50:38.965162Z","shell.execute_reply.started":"2023-03-03T10:50:38.948837Z","shell.execute_reply":"2023-03-03T10:50:38.964237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Transforms\n# ====================================================\ndef get_transforms(*, data):\n    \n    if data == 'train':\n        return A.Compose([\n            A.ShiftScaleRotate(shift_limit=0.8, scale_limit=0.8, rotate_limit=180, p=0.5),\n            A.RandomBrightnessContrast(brightness_limit=0.8, contrast_limit=0.8, p=0.2),\n            A.Resize(CFG.size_x, CFG.size_y),\n            ToTensorV2(),\n        ])\n\n    elif data == 'valid':\n        return A.Compose([\n            A.Resize(CFG.size_x, CFG.size_y),\n            ToTensorV2(),\n        ])\n","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:50:39.142841Z","iopub.execute_input":"2023-03-03T10:50:39.143566Z","iopub.status.idle":"2023-03-03T10:50:39.150311Z","shell.execute_reply.started":"2023-03-03T10:50:39.143522Z","shell.execute_reply":"2023-03-03T10:50:39.149239Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# MODEL\n# ====================================================\nclass CustomModel(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, in_chans=3)\n        self.n_features = self.model.head.fc.in_features\n        self.model.head.fc = nn.Linear(self.n_features, self.cfg.target_size)\n\n    def forward(self, x):\n        output = self.model(x)\n        return output","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:52:42.136396Z","iopub.execute_input":"2023-03-03T10:52:42.137359Z","iopub.status.idle":"2023-03-03T10:52:42.144597Z","shell.execute_reply.started":"2023-03-03T10:52:42.137304Z","shell.execute_reply":"2023-03-03T10:52:42.143323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Helper functions\n# ====================================================\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n\ndef asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\n\ndef timeSince(since, percent):\n    now = time.time()\n    s = now - since\n    es = s / (percent)\n    rs = es - s\n    return '%s (remain %s)' % (asMinutes(s), asMinutes(rs))\n\n\ndef train_fn(train_loader, model, criterion, optimizer, epoch, scheduler, device):\n    if CFG.apex:\n        scaler = GradScaler()\n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    scores = AverageMeter()\n    # switch to train mode\n    model.train()\n    start = end = time.time()\n    global_step = 0\n    for step, (images, labels) in enumerate(train_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        images = images.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n        if CFG.apex:\n            with autocast():\n                y_preds = model(images)\n                loss = criterion(y_preds.view(-1), labels)\n        else:\n            y_preds = model(images)\n            loss = criterion(y_preds.view(-1), labels)\n        # record loss\n        losses.update(loss.item(), batch_size)\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        if CFG.apex:\n            scaler.scale(loss).backward()\n        else:\n            loss.backward()\n        grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.max_grad_norm)\n        if (step + 1) % CFG.gradient_accumulation_steps == 0:\n            if CFG.apex:\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                optimizer.step()\n            optimizer.zero_grad()\n            global_step += 1\n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(train_loader)-1):\n            print('Epoch: [{0}][{1}/{2}] '\n                  'Data {data_time.val:.3f} ({data_time.avg:.3f}) '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  'Grad: {grad_norm:.4f}  '\n                  #'LR: {lr:.6f}  '\n                  .format(\n                   epoch+1, step, len(train_loader), batch_time=batch_time,\n                   data_time=data_time, loss=losses,\n                   remain=timeSince(start, float(step+1)/len(train_loader)),\n                   grad_norm=grad_norm,\n                   #lr=scheduler.get_lr()[0],\n                   ))\n    return losses.avg\n\n\ndef valid_fn(valid_loader, model, criterion, device):\n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    scores = AverageMeter()\n    # switch to evaluation mode\n    model.eval()\n    preds = []\n    start = end = time.time()\n    for step, (images, labels) in enumerate(valid_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        images = images.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n        # compute loss\n        with torch.no_grad():\n            y_preds = model(images)\n        loss = criterion(y_preds.view(-1), labels)\n        losses.update(loss.item(), batch_size)\n        # record accuracy\n        preds.append(y_preds.to('cpu').numpy())\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(valid_loader)-1):\n            print('EVAL: [{0}/{1}] '\n                  'Data {data_time.val:.3f} ({data_time.avg:.3f}) '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  .format(\n                   step, len(valid_loader), batch_time=batch_time,\n                   data_time=data_time, loss=losses,\n                   remain=timeSince(start, float(step+1)/len(valid_loader)),\n                   ))\n    predictions = np.concatenate(preds)\n    return losses.avg, predictions","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:50:39.490988Z","iopub.execute_input":"2023-03-03T10:50:39.493369Z","iopub.status.idle":"2023-03-03T10:50:39.515755Z","shell.execute_reply.started":"2023-03-03T10:50:39.493330Z","shell.execute_reply":"2023-03-03T10:50:39.514557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Train loop\n# ====================================================\ndef train_loop(folds, fold):\n    \n    LOGGER.info(f\"========== fold: {fold} training ==========\")\n\n    # ====================================================\n    # loader\n    # ====================================================\n    trn_idx = folds[(folds['fold'] != fold) & (folds['fold'] != CFG.tst_fold)].index\n    val_idx = folds[folds['fold'] == fold].index\n\n    train_folds = folds.loc[trn_idx].reset_index(drop=True)\n    valid_folds = folds.loc[val_idx].reset_index(drop=True)\n    valid_labels = valid_folds[CFG.target_col].values / 200.\n\n    train_dataset = TrainDataset(train_folds, \n                                 transform=get_transforms(data='train'))\n    valid_dataset = ValidDataset(valid_folds, \n                                 transform=get_transforms(data='valid'))\n\n    train_loader = DataLoader(train_dataset, \n                              batch_size=CFG.batch_size, \n                              shuffle=True, \n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=True)\n    valid_loader = DataLoader(valid_dataset, \n                              batch_size=CFG.batch_size, \n                              shuffle=False, \n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n    \n    # ====================================================\n    # scheduler \n    # ====================================================\n    def get_scheduler(optimizer):\n        if CFG.scheduler=='ReduceLROnPlateau':\n            scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=CFG.factor, patience=CFG.patience, verbose=True, eps=CFG.eps)\n        elif CFG.scheduler=='CosineAnnealingLR':\n            scheduler = CosineAnnealingLR(optimizer, T_max=CFG.T_max, eta_min=CFG.min_lr, last_epoch=-1)\n        elif CFG.scheduler=='CosineAnnealingWarmRestarts':\n            scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=CFG.T_0, T_mult=1, eta_min=CFG.min_lr, last_epoch=-1)\n        return scheduler\n\n    # ====================================================\n    # model & optimizer\n    # ====================================================\n    model = CustomModel(CFG, pretrained=True)\n    model.to(device)\n\n    optimizer = Adam(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay, amsgrad=False)\n    scheduler = get_scheduler(optimizer)\n\n    # ====================================================\n    # loop\n    # ====================================================\n    criterion = nn.L1Loss()\n\n    best_score = np.inf\n    best_loss = np.inf\n    \n    for epoch in range(CFG.epochs):\n        \n        start_time = time.time()\n        \n        # train\n        avg_loss = train_fn(train_loader, model, criterion, optimizer, epoch, scheduler, device)\n\n        # eval\n        avg_val_loss, preds = valid_fn(valid_loader, model, criterion, device)\n        \n        if isinstance(scheduler, ReduceLROnPlateau):\n            scheduler.step(avg_val_loss)\n        elif isinstance(scheduler, CosineAnnealingLR):\n            scheduler.step()\n        elif isinstance(scheduler, CosineAnnealingWarmRestarts):\n            scheduler.step()\n\n        # scoring\n        score = get_score(valid_labels, preds)\n\n        elapsed = time.time() - start_time\n\n        LOGGER.info(f'Epoch {epoch+1} - avg_train_loss: {avg_loss:.4f}  avg_val_loss: {avg_val_loss:.4f}  time: {elapsed:.0f}s')\n        LOGGER.info(f'Epoch {epoch+1} - Score: {score:.4f}')\n\n        if score < best_score:\n            best_score = score\n            LOGGER.info(f'Epoch {epoch+1} - Save Best Score: {best_score:.4f} Model')\n            torch.save({'model': model.state_dict(), \n                        'preds': preds},\n                        OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best_score.pth')\n        \n        if avg_val_loss < best_loss:\n            best_loss = avg_val_loss\n            LOGGER.info(f'Epoch {epoch+1} - Save Best Loss: {best_loss:.4f} Model')\n            torch.save({'model': model.state_dict(), \n                        'preds': preds},\n                        OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best_loss.pth')\n    \n    valid_folds['preds'] = torch.load(OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best_score.pth', \n                                      map_location=torch.device('cpu'))['preds']\n\n    return valid_folds","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:50:39.876566Z","iopub.execute_input":"2023-03-03T10:50:39.877008Z","iopub.status.idle":"2023-03-03T10:50:39.895123Z","shell.execute_reply.started":"2023-03-03T10:50:39.876970Z","shell.execute_reply":"2023-03-03T10:50:39.893975Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# main\n# ====================================================\ndef main():\n\n    \"\"\"\n    Prepare: 1.train \n    \"\"\"\n\n    def get_result(result_df):\n        preds = result_df['preds'].values\n        labels = result_df[CFG.target_col].values / 200.\n        score = get_score(labels, preds)\n        LOGGER.info(f'Score: {score:<.4f}')\n    \n    if CFG.train:\n        # train \n        oof_df = pd.DataFrame()\n\n        for fold in CFG.trn_fold:\n            _oof_df = train_loop(train, fold)\n            oof_df = pd.concat([oof_df, _oof_df])\n            LOGGER.info(f\"========== fold: {fold} result ==========\")\n            get_result(_oof_df)\n        # CV result\n        LOGGER.info(f\"========== CV ==========\")\n        get_result(oof_df)\n        # save result\n        oof_df.to_csv(OUTPUT_DIR+'oof_df.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:50:40.123549Z","iopub.execute_input":"2023-03-03T10:50:40.123963Z","iopub.status.idle":"2023-03-03T10:50:40.132609Z","shell.execute_reply.started":"2023-03-03T10:50:40.123926Z","shell.execute_reply":"2023-03-03T10:50:40.131107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == '__main__':\n    main()","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:52:50.680209Z","iopub.execute_input":"2023-03-03T10:52:50.680673Z","iopub.status.idle":"2023-03-03T10:57:48.776418Z","shell.execute_reply.started":"2023-03-03T10:52:50.680630Z","shell.execute_reply":"2023-03-03T10:57:48.775331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test = pd.read_csv('/kaggle/input/spr-x-ray-age/sample_submission_age.csv')\n\n\ndef get_train_file_path(image_id):\n    return \"/kaggle/input/spr-x-ray-age/kaggle/kaggle/test/\" + str(image_id).zfill(6) + \".png\"\n\n\ntest['file_path'] = test['imageId'].apply(get_train_file_path)\n\nprint(test['age'].isnull().values.any())\n\nprint(len(test))\ndisplay(test.head())\n","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:57:52.581858Z","iopub.execute_input":"2023-03-03T10:57:52.582333Z","iopub.status.idle":"2023-03-03T10:57:52.618655Z","shell.execute_reply.started":"2023-03-03T10:57:52.582284Z","shell.execute_reply":"2023-03-03T10:57:52.617685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_files = glob(\"/kaggle/input/spr-x-ray-age/kaggle/kaggle/test/*.png\")\nlen(test_files)","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:57:54.287499Z","iopub.execute_input":"2023-03-03T10:57:54.288125Z","iopub.status.idle":"2023-03-03T10:57:54.663506Z","shell.execute_reply.started":"2023-03-03T10:57:54.288082Z","shell.execute_reply":"2023-03-03T10:57:54.662416Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TestDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.file_names = df['file_path'].values\n        #self.labels = df[CFG.target_col].values\n        #self.gender = df['male'].values * 1.\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_path = self.file_names[idx]\n        image = cv2.imread(file_path)\n        image = image.astype(np.float32) / 255.\n        \n        image2 = np.zeros((CFG.size_x, CFG.size_y,3))\n        \n        for j in range(image.shape[2]):\n            image2[..., j] = cv2.resize(image[...,j], (pixels, pixels))\n        \n        image = image2.astype(np.float32)\n        \n        if self.transform:\n            image = self.transform(image=image)['image']\n        else:\n            image = image[np.newaxis,:,:]\n            image = torch.from_numpy(image).float()\n        \n        #image[:16, :16, ...] = self.gender[idx]\n        \n\n        \n        #label = torch.tensor(self.labels[idx]).float() / 200.\n        return image #, label","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:57:55.997269Z","iopub.execute_input":"2023-03-03T10:57:55.998012Z","iopub.status.idle":"2023-03-03T10:57:56.008654Z","shell.execute_reply.started":"2023-03-03T10:57:55.997969Z","shell.execute_reply":"2023-03-03T10:57:56.007338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def inference(model, states, test_loader, device):\n    model.to(device)\n    tk0 = tqdm(enumerate(test_loader), total=len(test_loader))\n    probs = []\n    for i, (images) in tk0:\n        #print(images.dtype)\n        images = images.to(device)\n        avg_preds = []\n        for state in states:\n            model.load_state_dict(state['model'])\n            model.eval()\n            with torch.no_grad():\n                y_preds = model(images)\n            avg_preds.append(y_preds.to('cpu').numpy())\n        avg_preds = np.mean(avg_preds, axis=0)\n        probs.append(avg_preds)\n    probs = np.concatenate(probs)\n    return probs","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:57:56.333259Z","iopub.execute_input":"2023-03-03T10:57:56.333633Z","iopub.status.idle":"2023-03-03T10:57:56.343072Z","shell.execute_reply.started":"2023-03-03T10:57:56.333599Z","shell.execute_reply":"2023-03-03T10:57:56.341821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ntest_dataset = TestDataset(test, \n                             transform=get_transforms(data='valid'))\n\ntest_loader = DataLoader(test_dataset, \n                          batch_size=CFG.batch_size * 2, \n                          shuffle=False, \n                          num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n\n\nmodel = CustomModel(CFG, pretrained=False)\nMODEL_DIR = 'model/' + VER + '/'\nstates = [torch.load(MODEL_DIR+f'{CFG.model_name}_fold{fold}_best_score.pth') for fold in CFG.trn_fold]\npreds = inference(model, states, test_loader, device)\n\n","metadata":{"execution":{"iopub.status.busy":"2023-03-03T10:57:56.642519Z","iopub.execute_input":"2023-03-03T10:57:56.643634Z","iopub.status.idle":"2023-03-03T11:02:24.087781Z","shell.execute_reply.started":"2023-03-03T10:57:56.643584Z","shell.execute_reply":"2023-03-03T11:02:24.086632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test['age'] = np.squeeze(preds) * 200.","metadata":{"execution":{"iopub.status.busy":"2023-03-03T11:05:51.272394Z","iopub.execute_input":"2023-03-03T11:05:51.272796Z","iopub.status.idle":"2023-03-03T11:05:51.280131Z","shell.execute_reply.started":"2023-03-03T11:05:51.272760Z","shell.execute_reply":"2023-03-03T11:05:51.279027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test","metadata":{"execution":{"iopub.status.busy":"2023-03-03T11:05:51.767625Z","iopub.execute_input":"2023-03-03T11:05:51.768747Z","iopub.status.idle":"2023-03-03T11:05:51.784528Z","shell.execute_reply.started":"2023-03-03T11:05:51.768689Z","shell.execute_reply":"2023-03-03T11:05:51.783259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test[['imageId', 'age']].to_csv(\"sub2.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-03-03T11:05:52.212248Z","iopub.execute_input":"2023-03-03T11:05:52.213407Z","iopub.status.idle":"2023-03-03T11:05:52.240096Z","shell.execute_reply.started":"2023-03-03T11:05:52.213359Z","shell.execute_reply":"2023-03-03T11:05:52.239049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}