{"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"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":71447,"databundleVersionId":8208918,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import hydra\nfrom omegaconf import DictConfig, OmegaConf\nimport sys,gc,os,random,time,math,glob\nimport matplotlib.pyplot as plt\nfrom contextlib import contextmanager\nfrom pathlib import Path\nfrom collections import defaultdict, Counter\nfrom  torch.cuda.amp import autocast, GradScaler \nimport cv2,timm\nfrom sklearn.metrics import roc_auc_score\nfrom PIL import Image\nimport numpy as np\nimport pandas as pd\nimport scipy as sp\nimport sklearn.metrics as metrics\nfrom sklearn.model_selection import StratifiedKFold,GroupKFold\nfrom sklearn.metrics import log_loss\nfrom functools import partial\nfrom tqdm import tqdm\nfrom sklearn.metrics import precision_score,recall_score,f1_score,log_loss,mean_absolute_error,mean_squared_error\nfrom  sklearn.metrics import accuracy_score as acc\nimport torch\nimport torch.nn as nn\nfrom torch.optim import Adam, SGD,AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, ReduceLROnPlateau,CosineAnnealingWarmRestarts\nfrom torch.utils.data import DataLoader, Dataset\nfrom albumentations import Compose, Normalize, HorizontalFlip, VerticalFlip,RandomGamma, RandomRotate90,GaussNoise,RandomBrightnessContrast,Resize\nfrom albumentations.pytorch import ToTensorV2\nimport transformers as T\n\nimport albumentations as A\n#import vision_transformer as vits\n\n### my utils\n# code_factory is from https://github.com/abebe9849/code_factory\nfrom code_factory.pooling import GeM,AdaptiveConcatPool2d\nfrom code_factory.augmix import RandomAugMix\nfrom code_factory.gridmask import GridMask\nfrom code_factory.fmix import *\nfrom code_factory.loss_func import *\n\"\"\"\nage=age/100\nとする　０〜１となったものをbceで\n\n\"\"\"\n###\n\nimport logging\n#from mylib.\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\n\nfrom timm.data.transforms import RandomResizedCropAndInterpolation\ndef select_random_elements(list_a, N):\n    if len(list_a) < N:\n        # 重複を許して要素を選択\n        return random.choices(list_a, k=N)\n    else:\n        # 重複なしで要素を選択\n        return random.sample(list_a, N)\ndef pad_to_square_bbox(x, y, w, h, image_width, image_height):\n    \"\"\"\n    与えられたバウンディングボックス(xywh形式)を、長辺に対してパディングして正方形にします。\n    :param x: バウンディングボックスの左上のX座標\n    :param y: バウンディングボックスの左上のY座標\n    :param w: バウンディングボックスの幅\n    :param h: バウンディングボックスの高さ\n    :param image_width: 画像の幅\n    :param image_height: 画像の高さ\n    :return: パディング後のバウンディングボックス（xywh形式）\n    \"\"\"\n    # 長辺を求める\n    max_side = max(w, h)\n\n    # 新しい幅と高さを長辺に設定\n    new_w = new_h = max_side\n\n    # 新しいX、Y座標を計算\n    new_x = x - (new_w - w) / 2\n    new_y = y - (new_h - h) / 2\n\n    # バウンディングボックスが画像の範囲を超えないように調整\n    new_x = max(0, min(new_x, image_width - new_w))\n    new_y = max(0, min(new_y, image_height - new_h))\n\n    return int(new_x), int(new_y), int(new_w), int(new_h)\n\n# 例: 使用例\nx, y, w, h = 50, 50, 100, 150  # 元のバウンディングボックス\nimage_width, image_height = 500, 500  # 画像のサイズ\nnew_x, new_y, new_w, new_h = pad_to_square_bbox(x, y, w, h, image_width, image_height)\nprint(new_x, new_y, new_w, new_h)\n\n      \ndef crop_func(img,xywh_):\n    x,y,w,h,_ = map(int,xywh_)\n    max_side = max(w, h)\n\n    # 新しい幅と高さを長辺に設定\n    new_w = new_h = max_side\n    image_width,image_height,_ = img.shape\n\n    # 新しいX、Y座標を計算\n    new_x = x - (new_w - w) / 2\n    new_y = y - (new_h - h) / 2\n    new_x = int(max(0, min(new_x, image_width - new_w)))\n    new_y = int(max(0, min(new_y, image_height - new_h)))\n    \n    #return img[y:y+h,x:x+w]\n    return img[new_y:new_y+new_h,new_x:new_x+new_w]\n\n\nclass TrainDataset(Dataset):\n    def __init__(self, df,CFG,train=True,transform1=None):\n        self.df = df\n        self.transform = transform1\n        self.CFG = CFG\n        self.train = train\n        self.image_size_seg = (128, 128, CFG.N_patch)\n        self.ids = self.df[\"StudyID___SeriesID\"].values\n        self.labels = self.df[\"Age\"].values\n        self.root = \"/home/share/dataset/CTage/window_png/25d\"\n        paths_ = [glob.glob(f\"{self.root}/{ct_id}/*.png\") for ct_id in self.ids]\n        self.id_2_paths = dict(zip(self.ids, paths_))\n    def __len__(self):\n        return len(self.ids)\n\n    def __getitem__(self, idx):\n        ct_id = self.ids[idx]\n        img_paths = self.id_2_paths[ct_id]\n        crop_ = [np.load(p.replace(\"25d\",\"25d_xywh\").replace(\".png\",\".npy\")) for p in img_paths]\n\n        img_paths = [[i,j] for i,j in zip(img_paths,crop_) if j[-1]>0.05]\n        crop_ = np.array(crop_)\n        max_xywh = crop_[np.argmax(crop_[:,4])]\n        #print(img_paths)\n        indices = np.quantile(list(range(len(img_paths))), np.linspace(0., 1., self.image_size_seg[2])).round().astype(int)\n        img_paths = [img_paths[i] for i in indices]\n        if self.CFG.precrop:\n            imgs= [crop_func(cv2.imread(i[0]),max_xywh) for i in img_paths]\n        else:\n            imgs = [cv2.imread(i[0]) for i in img_paths]\n        imgs =  np.stack([self.transform(image=img)['image']  for img in imgs])\n        image = torch.from_numpy(imgs.transpose(0,3,1,2)).float()\n\n        label =  self.labels[idx]\n        if self.CFG.loss.name==\"BCE\":\n            label = label/100.\n            label = torch.tensor(label).float()\n        elif self.CFG.loss.name==\"MSE\":\n            label = torch.tensor(label).float()\n        elif self.CFG.loss.name in [\"CE\",\"DLDL\"]:\n            label = torch.tensor(label).long()\n\n        \n        return image, label\n\n\n\n\n#### dataset ==============\n\n#### augmentation ==============\n\n\ndef get_transforms(*, data,CFG):\n    if data == 'train':\n        return Compose([\n            Resize(CFG.preprocess.size,CFG.preprocess.size),\n            #A.augmentations.crops.transforms.CenterCrop(CFG.preprocess.size*0.9,CFG.preprocess.size*0.9),\n            #A.crops.transforms.RandomResizedCrop(CFG.preprocess.size,CFG.preprocess.size,scale=(0.5, 1.0)),\n            #A.crops.transforms.RandomCrop(CFG.preprocess.size,CFG.preprocess.size),\n            A.HorizontalFlip(p=CFG.aug.HorizontalFlip.p),\n            A.VerticalFlip(p=CFG.aug.VerticalFlip.p),\n            A.RandomRotate90(p=CFG.aug.RandomRotate90.p),\n            A.ShiftScaleRotate(\n                shift_limit=CFG.aug.ShiftScaleRotate.shift_limit,\n                scale_limit=CFG.aug.ShiftScaleRotate.scale_limit,\n                rotate_limit=CFG.aug.ShiftScaleRotate.rotate_limit,\n                p=CFG.aug.ShiftScaleRotate.p),\n            A.RandomBrightnessContrast(\n                brightness_limit=CFG.aug.RandomBrightnessContrast.brightness_limit,\n                contrast_limit=CFG.aug.RandomBrightnessContrast.contrast_limit,\n                p=CFG.aug.RandomBrightnessContrast.p),\n            A.CLAHE(\n                clip_limit=(1,4),\n                p=CFG.aug.CLAHE.p),\n            A.OneOf([\n                A.ImageCompression(),\n                A.Downscale(scale_min=0.1, scale_max=0.15),\n                ], p=CFG.aug.compress.p),\n            GridMask(\n                num_grid=CFG.aug.GridMask.num_grid,p=CFG.aug.GridMask.p),\n            #A.CoarseDropout(max_holes=CFG.aug.CoarseDropout.max_holes, max_height=CFG.aug.CoarseDropout.max_height, max_width=CFG.aug.CoarseDropout.max_width, p=CFG.aug.CoarseDropout.p),\n            Normalize(mean=[0.485, 0.456, 0.406],std=[0.229, 0.224, 0.225])\n            ])\n    elif data == 'valid':\n        return Compose([\n            Resize(CFG.preprocess.size,CFG.preprocess.size),\n            #A.augmentations.crops.transforms.CenterCrop(int(CFG.preprocess.size*0.9),int(CFG.preprocess.size*0.9)),\n            Normalize(mean=[0.485, 0.456, 0.406],std=[0.229, 0.224, 0.225])\n            ])\n\n\n#### augmentation ==============\nimport open_clip\nfrom transformers import AutoProcessor, CLIPVisionModel\n#### model ================\nSEQ_POOLING = {\n    'gem': GeM(dim=2),\n    'concat': AdaptiveConcatPool2d(),\n    'avg': nn.AdaptiveAvgPool2d(1),\n    'max': nn.AdaptiveMaxPool2d(1)\n}\ndic_NUM_CLS = {\"BCE\":1,\"CE\":100,\"MSE\":1,\"DLDL2\":100}\n\nclass Model_iafoss(nn.Module):\n    def __init__(self,CFG, base_model='tf_efficientnet_b0_ns',pool=\"avg\",pretrain=True):\n        super(Model_iafoss, self).__init__()\n        self.base_model = base_model \n        NUM_CLS = dic_NUM_CLS[CFG.loss.name]\n        \"\"\"\n        if self.base_model in [\"hipt\",\"plip\",\"qnet\",\"ibot\"]:\n            if self.base_model==\"ibot\":\n                checkpoint_key = \"teacher\"\n                pretrained_weights = \"/home/abe/pandasub/ibot/pandaExp000/checkpoint.pth\"\n                self.model = vits.__dict__[\"vit_base\"](patch_size=16, num_classes=0)\n                state_dict = torch.load(pretrained_weights, map_location=\"cpu\")\n                if checkpoint_key is not None and checkpoint_key in state_dict:\n                    print(f\"Take key {checkpoint_key} in provided checkpoint dict\")\n                    state_dict = state_dict[checkpoint_key]\n                state_dict = {k.replace(\"module.\", \"\"): v for k, v in state_dict.items()}\n                # remove `backbone.` prefix induced by multicrop wrapper\n                state_dict = {k.replace(\"backbone.\", \"\"): v for k, v in state_dict.items()}\n                self.model.load_state_dict(state_dict, strict=False)\n                for _, p in self.model.named_parameters():\n                    p.requires_grad = False\n                for _, p in self.model.head.named_parameters():\n                    p.requires_grad = True\n\n                freeze =9\n                for n, p in self.model.blocks.named_parameters():\n                    if int(n.split(\".\")[0])>=(12-freeze):\n                        p.requires_grad = True\n                        \n                self.n_last_blocks  = 4\n                avgpool_patchtokens = 0\n                \n                nc = self.model.embed_dim * (self.n_last_blocks + int(avgpool_patchtokens))\n            \n            \n            if self.base_model==\"plip\":\n                self.model = CLIPVisionModel.from_pretrained(\"vinid/plip\")\n                nc = 768\n            elif self.base_model==\"qnet\":\n                self.model = open_clip.create_model_and_transforms('hf-hub:wisdomik/QuiltNet-B-32')[0]\n                \n                nc = 512\n            self.gru = nn.GRU(nc, 512, bidirectional=True, batch_first=True, num_layers=2)\n            nc*=CFG.N_patch\n            self.head = nn.Sequential(nn.Linear(nc,512),\n                            nn.ReLU(), nn.Dropout(0.5),nn.Linear(512,3))\n            self.exam_predictor = nn.Linear(512*2, 3)\n            self.pool = nn.AdaptiveAvgPool1d(1)\n        \"\"\"\n        if self.base_model==\"dino\":\n            print(\"not implemet\")\n            exit()\n        else:\n            self.model = timm.create_model(self.base_model, pretrained=True, num_classes=0,in_chans=3)\n            \"\"\"\n            #not work grad_accm not work\n            for module in self.model.modules():\n                \n                if isinstance(module, timm.models.layers.BatchNormAct2d):\n                    #print(module)\n                    if hasattr(module, 'weight'):\n                        module.weight.requires_grad_(False)\n                    if hasattr(module, 'bias'):\n                        module.bias.requires_grad_(False)\n                    module.eval()\n            \"\"\"\n\n            #self.model.conv_stem = nn.Conv2d(2, 32, kernel_size=3, padding=1, stride=1, bias=False)\n            nc = self.model.num_features\n            self.head = nn.Sequential(nn.AdaptiveAvgPool2d(1),nn.Flatten(),nn.Linear(nc,512),\n                            nn.ReLU(), nn.Dropout(0.5),nn.Linear(512,NUM_CLS))\n\n            self.gru = nn.GRU(nc, 512, bidirectional=True, batch_first=True, num_layers=2)\n            self.exam_predictor = nn.Linear(512*2, 1)\n            self.pool = nn.AdaptiveAvgPool1d(1)\n\n        \n    def forward(self, input1):\n        shape = input1.size()\n        batch_size = shape[0]\n        n = shape[1]\n\n        input1 = input1.view(-1,shape[2],shape[3],shape[4])\n\n        if \"ibot\" in self.base_model:\n            intermediate_output = self.model.get_intermediate_layers(input1, self.n_last_blocks)\n            x = torch.cat([x[:, 0] for x in intermediate_output], dim=-1)\n            embeds, _ = self.gru(x.view(batch_size,n,x.shape[1]))\n            embeds = self.pool(embeds.permute(0,2,1))[:,:,0]\n            y = self.exam_predictor(embeds)\n            #x = x.view(batch_size,x.shape[1]*n)\n            #y = self.head(x)\n           \n           #python base.py model.name=\"ibot\" train.lr=0.0001\n            return y\n        elif self.base_model==\"plip\":\n            x = self.model(input1)[\"pooler_output\"]\n            embeds, _ = self.gru(x.view(batch_size,n,x.shape[1]))\n            embeds = self.pool(embeds.permute(0,2,1))[:,:,0]\n            y = self.exam_predictor(embeds)\n           \n            return y\n        elif self.base_model==\"qnet\":\n            x = self.model.encode_image(input1)\n            embeds, _ = self.gru(x.view(batch_size,n,x.shape[1]))\n            embeds = self.pool(embeds.permute(0,2,1))[:,:,0]\n            y = self.exam_predictor(embeds)\n           \n            return y\n            \n        else:\n            \"\"\"\n            x = self.model.forward_features(input1)#bs*num_tile,embed_dim,h,w\n            shape = x.size()\n            x = x.view(-1,n,shape[1],shape[2],shape[3])\n            x = x.permute(0,2,1,3,4).contiguous().view(-1,shape[1],shape[2]*n,shape[3])\n            y = self.head(x)\n            \"\"\"\n            x =  self.model(input1)\n            \n            embeds, _ = self.gru(x.view(batch_size,n,x.shape[1]))\n            embeds = self.pool(embeds.permute(0,2,1))[:,:,0]\n            y = self.exam_predictor(embeds)\n            #\"\"\"\n            return y\n\n#model = Model_iafoss(\"dino_vit_s\")\n\n\n\n#### model ================\n\n\ndef train_fn(CFG,fold,folds,test_pl=0):\n\n\n    torch.cuda.set_device(CFG.general.device)\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    print(f\"### fold: {fold} ###\")\n    trn_idx = folds[folds['fold'] != fold].index\n    val_idx = folds[folds['fold'] == fold].index\n    \n    val_folds = folds.loc[val_idx]\n    tra_folds = folds.loc[trn_idx]\n    if CFG.general.debug:\n        CFG.train.epochs =2\n        tra_folds = tra_folds[tra_folds[\"StudyID___SeriesID\"].isin(tra_folds[\"StudyID___SeriesID\"].unique()[:50])]\n        val_folds = val_folds[val_folds[\"StudyID___SeriesID\"].isin(val_folds[\"StudyID___SeriesID\"].unique()[:50])]\n\n        \n    print(val_folds[\"StudyID___SeriesID\"].nunique(),tra_folds[\"StudyID___SeriesID\"].nunique())\n    if type(test_pl)!=type(0):\n        tra_folds = pd.concat([tra_folds,test_pl]).reset_index(drop=True)\n\n    train_dataset = TrainDataset(tra_folds.reset_index(drop=True),train=True, transform1=get_transforms(data='train',CFG=CFG),CFG=CFG)#get_transforms(data='train',CFG=CFG)\n    valid_dataset = TrainDataset(val_folds.reset_index(drop=True),train=False,transform1=get_transforms(data='valid',CFG=CFG),CFG=CFG)#\n\n\n    train_loader = DataLoader(train_dataset, batch_size=CFG.train.batch_size, shuffle=True, num_workers=8,pin_memory=True)\n    valid_loader = DataLoader(valid_dataset, batch_size=CFG.train.batch_size, shuffle=False, num_workers=8,pin_memory=True)\n\n    ###  model select ============\n    model = Model_iafoss(CFG,base_model=CFG.model.name).to(device)\n    # ============\n\n\n    ###  optim select ============\n    if CFG.train.optim==\"adam\":\n        optimizer = Adam(model.parameters(), lr=CFG.train.lr, amsgrad=False)\n    elif CFG.train.optim==\"adamw\":\n        optimizer = AdamW(model.parameters(), lr=CFG.train.lr,weight_decay=5e-5)\n    # ============\n\n    ###  scheduler select ============\n    if CFG.train.scheduler.name==\"cosine\":\n        scheduler = CosineAnnealingLR(optimizer, T_max=CFG.train.epochs, eta_min=CFG.train.scheduler.min_lr)\n    elif CFG.train.scheduler.name==\"cosine_warmup\":\n        scheduler =T.get_cosine_schedule_with_warmup(optimizer,\n        num_warmup_steps=len(train_loader)*CFG.train.scheduler.warmup,\n        num_training_steps=len(train_loader)*CFG.train.epochs)\n\n    # ============\n\n    ###  loss select ============\n    if CFG.loss.name==\"BCE\":\n        criterion=nn.BCEWithLogitsLoss()\n    elif CFG.loss.name==\"CE\":\n        criterion=nn.CrossEntropyLoss()\n    elif CFG.loss.name==\"DLDL2\":\n        criterion=DLDL2_loss()\n    elif CFG.loss.name==\"MSE\":\n        criterion=nn.HuberLoss() #https://www.kaggle.com/competitions/ventilator-pressure-prediction/discussion/277690\n        \n    print(criterion)\n    ###  loss select ============\n\n    scaler = torch.cuda.amp.GradScaler()\n    best_score = np.inf\n    best_loss = np.inf\n    best_preds = None\n        \n    for epoch in range(CFG.train.epochs):\n        start_time = time.time()\n        model.train()\n        avg_loss = 0.\n\n        tk0 = tqdm(enumerate(train_loader), total=len(train_loader))\n\n        for i, (images, labels) in tk0:\n            optimizer.zero_grad()\n            \n            images = images.to(device)\n            labels = labels.to(device)\n            \n\n            ### mix系のaugumentation=========\n            rand = np.random.rand()\n            ##mixupを終盤のepochでとめる\n            if epoch+1 >=CFG.train.without_hesitate:\n                rand=0\n\n            if CFG.augmentation.mix_p>rand and CFG.augmentation.do_mixup:\n                images, y_a, y_b, lam = mixup_data(images, labels,alpha=CFG.augmentation.mix_alpha)\n            elif CFG.augmentation.mix_p>rand and CFG.augmentation.do_cutmix:\n                images, y_a, y_b, lam = cutmix_data(images, labels,alpha=CFG.augmentation.mix_alpha)\n            elif CFG.augmentation.mix_p>rand and CFG.augmentation.do_resizemix:\n                images, y_a, y_b, lam = resizemix_data(images, labels,alpha=CFG.augmentation.mix_alpha)\n            elif CFG.augmentation.mix_p>rand and CFG.augmentation.do_fmix:\n                images, y_a, y_b, lam = fmix_data(images, labels,alpha=CFG.augmentation.mix_alpha)\n            ### mix系のaugumentation おわり=========\n\n            if CFG.train.amp:\n                with autocast():\n                    y_preds = model(images)\n                    if CFG.augmentation.mix_p>rand:\n                        if CFG.loss.name==\"BCE\" or CFG.loss.name==\"MSE\":\n                            loss_ = mixup_criterion(criterion, y_preds, y_a.view(-1,1), y_b.view(-1,1), lam)\n                        else:\n                            loss_ = mixup_criterion(criterion, y_preds, y_a, y_b, lam)\n                    else:\n                        if CFG.loss.name==\"BCE\" or CFG.loss.name==\"MSE\":\n                           loss_ = criterion(y_preds,labels.view(-1,1))\n                        else:\n                            loss_ = criterion(y_preds,labels)\n\n                    loss=loss_\n\n                scaler.scale(loss).backward()\n\n                if (i+1)%CFG.train.ga_accum==0 or i==-1:\n                    scaler.step(optimizer)\n                    scaler.update()\n                    if CFG.train.scheduler.name==\"cosine_warmup\":\n                        scheduler.step()\n            if CFG.train.scheduler.name==\"cosine\":\n                scheduler.step()\n            avg_loss += loss.item() / len(train_loader)\n        model.eval()\n        avg_val_loss = 0.\n        LOGITS = []\n        valid_labels = []\n        tk1 = tqdm(enumerate(valid_loader), total=len(valid_loader))\n        RANK = torch.Tensor([i for i in range(100)]).to(device)\n        for i, (images, labels) in tk1:\n            images = images.to(device)\n            labels = labels.to(device)\n            with torch.no_grad():\n                with autocast(enabled=False):\n                    logits = model(images)\n                    if CFG.loss.name==\"BCE\" or CFG.loss.name==\"MSE\":\n                        loss_ = criterion(logits,labels.view(-1,1))\n                    else:\n                        loss_ = criterion(logits,labels)\n            valid_labels.append(labels)\n            if CFG.loss.name==\"BCE\":\n                LOGITS.append(logits.detach().sigmoid())\n            elif CFG.loss.name==\"MSE\":\n                LOGITS.append(logits.detach())\n            elif CFG.loss.name in [\"CE\",\"DLDL2\"]:\n                logits = nn.functional.softmax(logits.detach(), dim=1)\n                \n                LOGITS.append(torch.sum(logits * RANK, dim=1))\n                \n            avg_val_loss += loss.item() / len(valid_loader)\n        \n        preds = torch.cat(LOGITS).cpu().numpy().squeeze()\n        valid_labels = torch.cat(valid_labels).cpu().numpy()\n        if CFG.loss.name==\"BCE\":\n            preds*=100\n            valid_labels*=100\n\n\n        #each_auc,score =AUC(true=valid_labels,predict=preds)\n        MAE_score = mean_absolute_error(valid_labels, preds)\n\n\n        elapsed = time.time() - start_time\n        log.info(f\"MAE  {MAE_score}\")\n\n\n        log.info(f'  Epoch {epoch+1} - avg_train_loss: {avg_loss:.6f}  avg_val_loss: {avg_val_loss:.6f}  time: {elapsed:.0f}s')\n\n        #if best_loss>avg_val_loss:#pr_auc best\n        #    best_loss = avg_val_loss\n        #    log.info(f'  Epoch {epoch+1} - Save Best loss: {best_loss:.4f}')\n        #    torch.save(model.state_dict(), f'fold{fold}_{CFG.general.exp_num}_best_loss.pth')\n\n        if best_score>MAE_score:#pr_auc best\n            best_score = MAE_score\n            log.info(f'  Epoch {epoch+1} - Save Best MAE: {best_score:.4f}')\n            best_preds = preds\n            torch.save(model.state_dict(), f'fold{fold}_{CFG.general.exp_num}_best_MAE.pth')\n\n\n    return best_preds, valid_labels\n\n\n\ndef eval_func(model, valid_loader, device,CFG):\n    model.to(device) \n    model.eval()\n\n    valid_labels = []\n    preds = []\n\n    tk1 = tqdm(enumerate(valid_loader), total=len(valid_loader))\n\n    for i, (images, labels) in tk1:\n        images = images.to(device)\n        labels = labels.to(device)\n        with torch.no_grad():\n            with autocast():\n                y_preds = model(images.float())\n                y_preds = y_preds.sigmoid()\n\n        valid_labels.append(labels.to('cpu').numpy())\n        preds.append(y_preds.to('cpu').numpy())\n    preds = np.concatenate(preds)\n    valid_labels = np.concatenate(valid_labels)\n\n    return preds,valid_labels\n\n\n\ndef inf_func(models, valid_loader, device,CFG):\n    for model in models:\n        model.eval()\n\n    preds = []\n    RANK = torch.Tensor([i for i in range(100)]).to(device)\n    \n\n    tk1 = tqdm(enumerate(valid_loader), total=len(valid_loader))\n\n    for i, (images, _) in tk1:\n        images = images.to(device,non_blocking=True)\n        with torch.no_grad():\n            with autocast():\n                if CFG.loss.name==\"BCE\":\n                    y_preds = [m(images.float()).detach().sigmoid() for m  in models]\n                elif CFG.loss.name==\"MSE\":\n                    y_preds = [m(images.float()).detach() for m  in models]\n                elif CFG.loss.name in [\"CE\",\"DLDL2\"]:\n                    y_preds = [torch.sum(nn.functional.softmax(m(images.float()).detach(), dim=1) * RANK, dim=1) for m  in models]                    \n                    \n                y_preds  = torch.stack(y_preds,axis=-1)#.median(0)\n                #https://www.kaggle.com/competitions/ventilator-pressure-prediction/discussion/276138\n\n        preds.append(y_preds)\n        \n    preds = torch.cat(preds).cpu().numpy().squeeze()\n    if CFG.loss.name==\"BCE\":\n        preds*=100\n\n    return preds\n        \n\n    \ndef submit(CFG,num_folds,test,DIR):\n    torch.cuda.set_device(CFG.general.device)\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    models = []\n    for fold in   range(num_folds):\n        model = Model_iafoss(CFG,base_model=CFG.model.name).to(device)\n        model.load_state_dict(torch.load(f\"{DIR}/fold{fold}_{CFG.general.exp_num}_best_MAE.pth\", map_location=\"cpu\"))\n        models.append(model)\n        \n    valid_dataset = TrainDataset(test,train=False,\n                                 transform1=get_transforms(data='valid',CFG=CFG),CFG=CFG)# \n    valid_loader = DataLoader(valid_dataset, batch_size=2, shuffle=False, num_workers=12,pin_memory=True)\n    tets_preds = inf_func(models, valid_loader, device,CFG)\n    print(tets_preds.shape)\n    for i in range(num_folds):\n        col = f\"pred_{i}\"\n        test[col]=tets_preds[:,i]\n    \n    \n    return test\n\n\ndef calculate_median(df):\n# predカラムの値を1つのリストに集約\n    preds = df.filter(like='pred').values.flatten()\n    # 中央値を計算\n    return np.median(preds)\n\nDIR = \"/home/u094724e/CTage/src\"\n\nlog = logging.getLogger(__name__)\n@hydra.main(config_path=f\"{DIR}/\",config_name=\"base\")\ndef main(CFG : DictConfig) -> None:\n\n    seed_torch(seed=CFG.general.seed)\n\n    log.info(f\"===============exp_num{CFG.general.exp_num}============\")\n\n    folds = pd.read_csv(\"/home/share/dataset/CTage/25d_10folds.csv\")\n    num_folds = int(folds[\"fold\"].max())+1\n\n    #\"\"\"\n    preds = []\n    valid_labels = []\n    oof = pd.DataFrame()\n    \n    if CFG.psuedo_label!=0:\n        test_pl=pd.read_csv(CFG.psuedo_label)\n        test_pl[\"fold\"]=999\n    else:\n        test_pl = 0\n    #time.sleep(3600*4)\n\n        \n    \n    \n    for fold in range(num_folds):\n        _preds, _valid_labels = train_fn(CFG,fold,folds,test_pl)\n        preds.append(_preds)\n        valid_labels.append(_valid_labels)\n    preds = np.concatenate(preds)\n    valid_labels = np.concatenate(valid_labels)\n\n    MAE_score = mean_absolute_error(valid_labels, preds)\n    log.info(f\"OOF MAE_score  {MAE_score}\")\n\n\n    #oof.to_csv(f\"oof_{CFG.general.exp_num}.csv\",index=False)\n    test_df = pd.read_csv(\"/home/share/dataset/CTage/25d_test.csv\")\n    test_df[\"Age\"]=0\n    test_df = submit(CFG,num_folds,test_df,DIR=\".\")\n\n    test_df.to_csv(f\"inf_{CFG.general.exp_num}.csv\",index=False)\n    test_df = test_df.groupby('StudyID').apply(calculate_median).reset_index(name='Age')\n    test_df[\"StudyID\"] = [str(i).zfill(6) for i in test_df[\"StudyID\"].values]\n    test_df[\"Age\"] =test_df[\"Age\"].round()\n\n    \n    test_df.to_csv(f\"test_{CFG.general.exp_num}.csv\",index=False)\n\n\nif __name__ == \"__main__\":\n    main()\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]}]}